diff --git a/pkg/generators/openapi.go b/pkg/generators/openapi.go index 195620b35..665d7f9c8 100644 --- a/pkg/generators/openapi.go +++ b/pkg/generators/openapi.go @@ -340,34 +340,83 @@ func typeShortName(t *types.Type) string { return path.Base(t.Name.Package) + "." + t.Name.Name } -func (g openAPITypeWriter) generateMembers(t *types.Type, required []string) ([]string, error) { - var err error - for t.Kind == types.Pointer { // fast-forward to effective type containing members - t = t.Elem +type memberCandidate struct { + member types.Member + parent *types.Type + depth int + ambiguous bool +} + +func (g openAPITypeWriter) generateMembers(t *types.Type) ([]string, error) { + type typeAtDepth struct { + typeToVisit *types.Type + depth int } - for _, m := range t.Members { - if hasOpenAPITagValue(m.CommentLines, tagValueFalse) { - continue + + queue := []typeAtDepth{{typeToVisit: t}} + candidatesByName := map[string]*memberCandidate{} + var candidates []*memberCandidate + + // Visit embedded types breadth-first so fields at the shallowest depth are + // selected. Multiple fields with the same name at that depth are ambiguous. + for i := 0; i < len(queue); i++ { + current := queue[i] + for current.typeToVisit.Kind == types.Pointer { // fast-forward to effective type containing members + current.typeToVisit = current.typeToVisit.Elem } - if shouldInlineMembers(&m) { - required, err = g.generateMembers(m.Type, required) - if err != nil { - return required, err + for _, m := range current.typeToVisit.Members { + if hasOpenAPITagValue(m.CommentLines, tagValueFalse) { + continue + } + if shouldInlineMembers(&m) { + queue = append(queue, typeAtDepth{typeToVisit: m.Type, depth: current.depth + 1}) + continue + } + name := getReferableName(&m) + if name == "" { + continue + } + + candidate, found := candidatesByName[name] + switch { + case !found: + candidate = &memberCandidate{ + member: m, + parent: current.typeToVisit, + depth: current.depth, + } + candidatesByName[name] = candidate + candidates = append(candidates, candidate) + case current.depth < candidate.depth: + candidate.member = m + candidate.parent = current.typeToVisit + candidate.depth = current.depth + candidate.ambiguous = false + case current.depth == candidate.depth: + candidate.ambiguous = true } - continue } - name := getReferableName(&m) - if name == "" { + } + + required := []string{} + requiredNames := map[string]struct{}{} + for _, candidate := range candidates { + if candidate.ambiguous { continue } - if isOptional, err := isOptional(&m); err != nil { - klog.Errorf("Error when generating: %v, %v\n", name, m) + + name := getReferableName(&candidate.member) + if optional, err := isOptional(&candidate.member); err != nil { + klog.Errorf("Error when generating: %v, %v\n", name, candidate.member) return required, err - } else if !isOptional { - required = append(required, name) + } else if !optional { + if _, found := requiredNames[name]; !found { + required = append(required, name) + requiredNames[name] = struct{}{} + } } - if err = g.generateProperty(&m, t); err != nil { - klog.Errorf("Error when generating: %v, %v\n", name, m) + if err := g.generateProperty(&candidate.member, candidate.parent); err != nil { + klog.Errorf("Error when generating: %v, %v\n", name, candidate.member) return required, err } } @@ -678,7 +727,7 @@ func (g openAPITypeWriter) generate(t *types.Type) error { propertiesBuf := bytes.Buffer{} bsw := g bsw.SnippetWriter = generator.NewSnippetWriter(&propertiesBuf, g.context, "$", "$") - required, err := bsw.generateMembers(t, []string{}) + required, err := bsw.generateMembers(t) if err != nil { return err } diff --git a/pkg/generators/openapi_test.go b/pkg/generators/openapi_test.go index c2c064796..1962a56bd 100644 --- a/pkg/generators/openapi_test.go +++ b/pkg/generators/openapi_test.go @@ -867,6 +867,116 @@ Required: []string{"String"}, }) } +func TestEmbeddedInlineFieldResolution(t *testing.T) { + tests := []struct { + name string + inputFile string + propertyCounts map[string]int + required string + includedFragment string + excludedFragment string + }{ + { + name: "same depth is ambiguous", + inputFile: ` + package foo + + type Common struct { + Shared string ` + "`json:\"shared\"`" + ` + } + + type Left struct { + Common ` + "`json:\",inline\"`" + ` + Left string ` + "`json:\"left\"`" + ` + } + + type Right struct { + Common ` + "`json:\",inline\"`" + ` + Right string ` + "`json:\"right\"`" + ` + } + + type Blah struct { + Left ` + "`json:\",inline\"`" + ` + Right ` + "`json:\",inline\"`" + ` + }`, + propertyCounts: map[string]int{ + "left": 1, + "right": 1, + "shared": 0, + }, + required: `Required: []string{"left","right"},`, + }, + { + name: "shallower field wins", + inputFile: ` + package foo + + type Intermediate struct { + // Deep shared field. + Shared int ` + "`json:\"shared\"`" + ` + } + + type Deep struct { + Intermediate ` + "`json:\",inline\"`" + ` + } + + type Shallow struct { + // Shallow shared field. + Shared string ` + "`json:\"shared\"`" + ` + } + + type Blah struct { + Deep ` + "`json:\",inline\"`" + ` + Shallow ` + "`json:\",inline\"`" + ` + }`, + propertyCounts: map[string]int{ + "shared": 1, + }, + required: `Required: []string{"shared"},`, + includedFragment: `Description: "Shallow shared field."`, + excludedFragment: `Description: "Deep shared field."`, + }, + } + + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + packagestest.TestAll(t, func(t *testing.T, x packagestest.Exporter) { + e := packagestest.Export(t, x, []packagestest.Module{{ + Name: "example.com/base/foo", + Files: map[string]interface{}{ + "foo.go": test.inputFile, + }, + }}) + defer e.Cleanup() + + callErr, funcErr, _, funcBuffer, _ := testOpenAPITypeWriter(t, e.Config) + if callErr != nil { + t.Fatal(callErr) + } + if funcErr != nil { + t.Fatal(funcErr) + } + + generated := funcBuffer.String() + for property, expectedCount := range test.propertyCounts { + if count := strings.Count(generated, fmt.Sprintf("%q: {", property)); count != expectedCount { + t.Errorf("property %q emitted %d times, want %d\n%s", property, count, expectedCount, generated) + } + } + if count := strings.Count(generated, test.required); count != 1 { + t.Errorf("required fields emitted %d times, want 1 occurrence of %q\n%s", count, test.required, generated) + } + if test.includedFragment != "" && !strings.Contains(generated, test.includedFragment) { + t.Errorf("generated output does not contain winning field schema %q\n%s", test.includedFragment, generated) + } + if test.excludedFragment != "" && strings.Contains(generated, test.excludedFragment) { + t.Errorf("generated output contains shadowed field schema %q\n%s", test.excludedFragment, generated) + } + }) + }) + } +} + func TestNestedMapString(t *testing.T) { inputFile := ` package foo