Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
91 changes: 70 additions & 21 deletions pkg/generators/openapi.go
Original file line number Diff line number Diff line change
Expand Up @@ -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
}
}
Expand Down Expand Up @@ -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
}
Expand Down
110 changes: 110 additions & 0 deletions pkg/generators/openapi_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down