Skip to content
Merged
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
4 changes: 3 additions & 1 deletion UPGRADING.md
Original file line number Diff line number Diff line change
Expand Up @@ -559,7 +559,9 @@ The issue review reproduced two problems that also affect v3.30.0:
implementation when a UUID value is needed.
- A recursive gRPC result containing an array of itself can make generation
recurse indefinitely ([#2515](https://github.com/goadesign/goa/issues/2515)).
This release does not fix that case.
v3.31.0 does not fix that case. The fix for recursive arrays and maps is on
`v3` after v3.31.0; regenerate with a version containing the fix. The design
and protobuf wire format do not need to change.

## Report a problem

Expand Down
8 changes: 8 additions & 0 deletions codegen/ARCHITECTURE.md
Original file line number Diff line number Diff line change
Expand Up @@ -424,6 +424,14 @@ example, identical recursive values reached through a direct field and a map
value may share one function even though their generated layouts must first be
resolved at two different paths.

Array elements and map values that are named objects use planned conversion
helpers in gRPC as well as HTTP. A recursive collection calls the helper already
being planned instead of expanding the same object again. When a transport
stores a collection in a wrapper message, the transform plan retains the chosen
wrapper field at every location, including nested collections. Helper type
lookup follows those saved fields, so function sharing compares the actual
parameter and result types used by the generated conversion.

`Helpers` and `HelperDefinitions` return detached type descriptions with
plan-owned IDs. A caller
may inspect or change those descriptions while choosing a function declaration,
Expand Down
Original file line number Diff line number Diff line change
@@ -0,0 +1,225 @@
// This file checks that recursive collection designs generate working gRPC code.
package generator

import (
"os"
"path/filepath"
"testing"

"github.com/stretchr/testify/require"

"goa.design/goa/v3/codegen"
d "goa.design/goa/v3/dsl"
"goa.design/goa/v3/eval"
)

func TestGenerateRecursiveGRPCResult(t *testing.T) {
root := codegen.RunDSL(t, grpcRecursiveCollectionsDSL)
plan := mustTestPlan(t, "generated.local/gen", []eval.Root{root}, planTransportData)
files, err := testServiceFiles(plan)
require.NoError(t, err)
transport, err := testTransportFiles(plan)
require.NoError(t, err)
dir := t.TempDir()
writeGeneratedModule(t, dir, "generated.local")
for _, file := range append(files, transport...) {
_, err := file.Render(dir)
require.NoError(t, err)
}
writeGRPCRecursiveCollectionsTest(t, dir)
runGeneratedTests(t, dir)
}

// grpcRecursiveCollectionsDSL keeps the unconstrained result from #2515 and
// adds request and response types with validation and nested collection wrappers.
func grpcRecursiveCollectionsDSL() {
category := d.ResultType("application/vnd.category", func() {
d.TypeName("CategoryResult")
d.Attributes(func() {
d.Field(1, "id", d.Int)
d.Field(3, "children_category", d.ArrayOf("CategoryResult"))
d.Field(4, "name", d.String)
})
})
node := d.Type("Node", func() {
d.Field(1, "children", d.ArrayOf("Node"))
d.Field(2, "by_name", d.MapOf(d.String, "Node"))
d.Field(3, "groups", d.MapOf(d.String, d.ArrayOf("Node")))
d.Field(4, "matrix", d.ArrayOf(d.ArrayOf("Node")))
d.Field(5, "branches", d.ArrayOf("Branch"))
d.Field(6, "name", d.String, func() {
d.Pattern("^[a-z]+$")
})
d.Field(7, "count", d.Int)
d.Field(8, "enabled", d.Boolean, func() {
d.Default(true)
})
d.Required("name")
})
d.Type("Branch", func() {
d.Field(1, "nodes", d.MapOf(d.String, node))
})
d.Service("categories", func() {
d.Method("list", func() {
d.Result(category)
d.GRPC(func() {})
})
d.Method("exchange", func() {
d.Payload(node)
d.Result(node)
d.GRPC(func() {})
})
})
}

// writeGRPCRecursiveCollectionsTest exercises all four generated conversion
// directions and verifies that validators inspect values below recursive fields.
func writeGRPCRecursiveCollectionsTest(t *testing.T, moduleDir string) {
t.Helper()
dir := filepath.Join(moduleDir, "recursivetest")
require.NoError(t, os.MkdirAll(dir, 0o750))
const source = `package recursivetest_test

import (
"context"
"testing"

"github.com/stretchr/testify/require"
"google.golang.org/grpc/metadata"
"google.golang.org/protobuf/proto"

gencategories "generated.local/gen/categories"
genclient "generated.local/gen/grpc/categories/client"
genpb "generated.local/gen/grpc/categories/pb"
genserver "generated.local/gen/grpc/categories/server"
)

func TestRecursiveResultRoundTrip(t *testing.T) {
id, name := 42, "category"
result := &gencategories.CategoryResult{ID: &id, Name: &name, ChildrenCategory: []*gencategories.CategoryResult{
{ID: &id, Name: &name, ChildrenCategory: []*gencategories.CategoryResult{{ID: &id, Name: &name}}},
}}
viewed := gencategories.NewViewedCategoryResult(result, "default")
headers, trailers := metadata.MD{}, metadata.MD{}
message, err := genserver.EncodeListResponse(context.Background(), viewed, &headers, &trailers)
require.NoError(t, err)
wire, ok := message.(proto.Message)
require.True(t, ok)
data, err := proto.Marshal(wire)
require.NoError(t, err)
received := wire.ProtoReflect().New().Interface()
require.NoError(t, proto.Unmarshal(data, received))
converted, err := genclient.DecodeListResponse(context.Background(), received, headers, trailers)
require.NoError(t, err)
require.Equal(t, result, converted)
}

func TestRecursiveCollectionsRoundTrip(t *testing.T) {
tests := []struct {
name string
value *gencategories.Node
}{
{
name: "array",
value: &gencategories.Node{Name: "root", Children: []*gencategories.Node{
{Name: "child", Children: []*gencategories.Node{{Name: "leaf"}}},
}},
},
{
name: "map",
value: &gencategories.Node{Name: "root", ByName: map[string]*gencategories.Node{
"child": {Name: "child", ByName: map[string]*gencategories.Node{"leaf": {Name: "leaf"}}},
}},
},
{
name: "map of arrays",
value: &gencategories.Node{Name: "root", Groups: map[string][]*gencategories.Node{
"group": {{Name: "child", Groups: map[string][]*gencategories.Node{"group": {{Name: "leaf"}}}}},
}},
},
{
name: "array of arrays",
value: &gencategories.Node{Name: "root", Matrix: [][]*gencategories.Node{
{{Name: "child", Matrix: [][]*gencategories.Node{{{Name: "leaf"}}}}},
}},
},
{
name: "mutual recursion",
value: &gencategories.Node{Name: "root", Branches: []*gencategories.Branch{
{Nodes: map[string]*gencategories.Node{"child": {Name: "child", Branches: []*gencategories.Branch{
{Nodes: map[string]*gencategories.Node{"leaf": {Name: "leaf"}}},
}}}},
}},
},
}
for _, test := range tests {
t.Run(test.name, func(t *testing.T) {
request := genclient.NewProtoExchangeRequest(test.value)
data, err := proto.Marshal(request)
require.NoError(t, err)
received := new(genpb.ExchangeRequest)
require.NoError(t, proto.Unmarshal(data, received))
require.NoError(t, genserver.ValidateExchangeRequest(received))
payload := genserver.NewExchangePayload(received)
require.Equal(t, test.value, payload)

response := genserver.NewProtoExchangeResponse(payload)
data, err = proto.Marshal(response)
require.NoError(t, err)
result := new(genpb.ExchangeResponse)
require.NoError(t, proto.Unmarshal(data, result))
require.NoError(t, genclient.ValidateExchangeResponse(result))
require.Equal(t, test.value, genclient.NewExchangeResult(result))
})
}
}

func TestRecursiveCollectionsValidateDescendants(t *testing.T) {
valid := "valid"
invalid := "INVALID"
for _, name := range []string{"array", "map", "map of arrays", "array of arrays", "mutual recursion"} {
t.Run(name, func(t *testing.T) {
leaf := &gencategories.Node{Name: valid}
middle := &gencategories.Node{Name: valid}
root := &gencategories.Node{Name: valid}
for parent, child := range map[*gencategories.Node]*gencategories.Node{root: middle, middle: leaf} {
switch name {
case "array":
parent.Children = []*gencategories.Node{child}
case "map":
parent.ByName = map[string]*gencategories.Node{"child": child}
case "map of arrays":
parent.Groups = map[string][]*gencategories.Node{"group": {child}}
case "array of arrays":
parent.Matrix = [][]*gencategories.Node{{child}}
case "mutual recursion":
parent.Branches = []*gencategories.Branch{{Nodes: map[string]*gencategories.Node{"child": child}}}
}
}
require.NoError(t, genserver.ValidateExchangeRequest(genclient.NewProtoExchangeRequest(root)))
leaf.Name = invalid
require.Error(t, genserver.ValidateExchangeRequest(genclient.NewProtoExchangeRequest(root)))
require.Error(t, genclient.ValidateExchangeResponse(genserver.NewProtoExchangeResponse(root)))
})
}
}

func TestRecursiveCollectionPresenceAndDefaults(t *testing.T) {
name, zero, disabled := "valid", int32(0), false
message := &genpb.ExchangeRequest{Name: &name, Children: []*genpb.Node{
{Name: &name},
{Name: &name, Count: &zero, Enabled: &disabled},
}}
require.NoError(t, genserver.ValidateExchangeRequest(message))
payload := genserver.NewExchangePayload(message)
require.Nil(t, payload.Children[0].Count)
require.True(t, payload.Children[0].Enabled)
require.NotNil(t, payload.Children[1].Count)
require.Zero(t, *payload.Children[1].Count)
require.False(t, payload.Children[1].Enabled)
response := genserver.NewProtoExchangeResponse(payload)
require.Equal(t, payload, genclient.NewExchangeResult(response))
}
`
require.NoError(t, os.WriteFile(filepath.Join(dir, "recursive_test.go"), []byte(source), 0o600))
}
17 changes: 14 additions & 3 deletions codegen/go_transform.go
Original file line number Diff line number Diff line change
Expand Up @@ -390,6 +390,7 @@ func newTransformPlan(source, target *expr.AttributeExpr, prefix string, program
target: target,
rootSource: source,
rootTarget: target,
wrappers: make(map[TransformHelperDefinitionLocation]transformLayoutWrapper),
sourceBaseline: baselineSource.Copy(source),
targetBaseline: baselineTarget.Copy(target),
sourceCopier: sourceCopier,
Expand Down Expand Up @@ -1450,14 +1451,24 @@ func planTransformOperationWithHelper(source, target *expr.AttributeExpr, requir
return fmt.Errorf("custom union transform helper requires a named source or target type")
}
}
var rootWrap *WrapDirective
var wrapper *WrapDirective
if plan.hooks != nil && plan.hooks.UnwrapPair != nil {
source, target, rootWrap = plan.hooks.UnwrapPair(source, target)
source, target, wrapper = plan.hooks.UnwrapPair(source, target)
}
if location.encoded == "" {
plan.rootSource = source
plan.rootTarget = target
plan.rootWrap = rootWrap
}
if wrapper != nil {
selected := transformLayoutWrapper{
wrapper: helperSource,
value: source,
directive: wrapper,
}
if wrapper.WrapTarget {
selected.wrapper, selected.value = helperTarget, target
}
plan.wrappers[location] = selected
}
if !forceHelper {
helperSource, helperTarget = source, target
Expand Down
10 changes: 5 additions & 5 deletions codegen/go_transform_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -1457,23 +1457,23 @@ func TestTransformHelperRegistryRejectsInvalidRootWrapper(t *testing.T) {
{
name: "missing generated field",
change: func(plan *TransformPlan, _ *GoTypePlan) {
plan.rootWrap.FieldName = "Missing"
plan.wrappers[TransformHelperDefinitionLocation{}].directive.FieldName = "Missing"
},
err: `select target root wrapper: wrapper field "Missing" is missing`,
err: `find target layout for transform helper occurrence 1: select wrapper field: wrapper field "Missing" is missing`,
},
{
name: "ambiguous generated field",
change: func(_ *TransformPlan, layout *GoTypePlan) {
layout.fields[1].fieldNameUpper = "Field"
},
err: `select target root wrapper: wrapper field "Field" is ambiguous`,
err: `find target layout for transform helper occurrence 1: select wrapper field: wrapper field "Field" is ambiguous`,
},
{
name: "different design field",
change: func(plan *TransformPlan, _ *GoTypePlan) {
plan.rootWrap.FieldName = "Other"
plan.wrappers[TransformHelperDefinitionLocation{}].directive.FieldName = "Other"
},
err: `select target root wrapper: wrapper field "Other" does not hold the selected value`,
err: `find target layout for transform helper occurrence 1: select wrapper field: wrapper field "Other" does not hold the selected value`,
},
}
for _, test := range tests {
Expand Down
45 changes: 15 additions & 30 deletions codegen/transform_helper_registry.go
Original file line number Diff line number Diff line change
Expand Up @@ -87,10 +87,6 @@ func (r *TransformHelperRegistry) Collect(plan *TransformPlan, sourceLayout, tar
if order == nil {
return fmt.Errorf("transform helper order must not be nil")
}
sourceLayout, targetLayout, err := transformRootLayouts(plan, sourceLayout, targetLayout)
if err != nil {
return err
}
definitions := make(map[int]TransformHelperDefinition, len(plan.helpers))
for _, definition := range plan.definitions {
for _, index := range definition.helpers {
Expand All @@ -102,11 +98,11 @@ func (r *TransformHelperRegistry) Collect(plan *TransformPlan, sourceLayout, tar
if !ok {
return fmt.Errorf("transform helper occurrence %d has no definition", helper.Occurrence)
}
source, err := transformLayoutAtLocation(sourceLayout, plan.rootSource, helper.location)
source, err := transformLayoutAtLocation(sourceLayout, plan.rootSource, helper.location, plan.wrappers, false)
if err != nil {
return fmt.Errorf("find source layout for transform helper occurrence %d: %w", helper.Occurrence, err)
}
target, err := transformLayoutAtLocation(targetLayout, plan.rootTarget, helper.location)
target, err := transformLayoutAtLocation(targetLayout, plan.rootTarget, helper.location, plan.wrappers, true)
if err != nil {
return fmt.Errorf("find target layout for transform helper occurrence %d: %w", helper.Occurrence, err)
}
Expand All @@ -127,30 +123,10 @@ func (r *TransformHelperRegistry) Collect(plan *TransformPlan, sourceLayout, tar
return nil
}

// transformRootLayouts selects the generated wrapper field that the transform
// reads or writes before it calls any nested conversion functions.
func transformRootLayouts(plan *TransformPlan, sourceLayout, targetLayout *GoTypePlan) (*GoTypePlan, *GoTypePlan, error) {
if plan.rootWrap == nil {
return sourceLayout, targetLayout, nil
}
if plan.rootWrap.WrapTarget {
selected, err := transformRootWrapperField(targetLayout, plan.target, plan.rootTarget, plan.rootWrap.FieldName)
if err != nil {
return nil, nil, fmt.Errorf("select target root wrapper: %w", err)
}
return sourceLayout, selected, nil
}
selected, err := transformRootWrapperField(sourceLayout, plan.source, plan.rootSource, plan.rootWrap.FieldName)
if err != nil {
return nil, nil, fmt.Errorf("select source root wrapper: %w", err)
}
return selected, targetLayout, nil
}

// transformRootWrapperField returns the generated field chosen by the saved
// transformWrapperField returns the generated field chosen by the saved
// wrapper instruction. The generated field and design field must describe the
// same value that the transform planned to read or write.
func transformRootWrapperField(layout *GoTypePlan, wrapper, selected *expr.AttributeExpr, fieldName string) (*GoTypePlan, error) {
func transformWrapperField(layout *GoTypePlan, wrapper, selected *expr.AttributeExpr, fieldName string) (*GoTypePlan, error) {
layout, wrapper = transformLayoutValue(layout, wrapper)
object := expr.AsObject(wrapper.Type)
if object == nil || layout.kind != GoStruct {
Expand Down Expand Up @@ -481,10 +457,19 @@ func transformSemanticDataTypesEqual(left, right expr.DataType, seen map[transfo
}

// transformLayoutAtLocation follows the authored field, collection, and union
// path saved by TransformPlan and returns the generated layout at that point.
func transformLayoutAtLocation(layout *GoTypePlan, attribute *expr.AttributeExpr, location TransformHelperDefinitionLocation) (*GoTypePlan, error) {
// path saved by TransformPlan, entering wrapper fields on the chosen side, and
// returns the generated layout at that point.
func transformLayoutAtLocation(layout *GoTypePlan, attribute *expr.AttributeExpr, location TransformHelperDefinitionLocation, wrappers map[TransformHelperDefinitionLocation]transformLayoutWrapper, target bool) (*GoTypePlan, error) {
remaining := location.encoded
for len(remaining) > 0 {
parent := TransformHelperDefinitionLocation{encoded: location.encoded[:len(location.encoded)-len(remaining)]}
if wrapper, ok := wrappers[parent]; ok && wrapper.directive.WrapTarget == target {
selected, err := transformWrapperField(layout, wrapper.wrapper, wrapper.value, wrapper.directive.FieldName)
if err != nil {
return nil, fmt.Errorf("select wrapper field: %w", err)
}
layout, attribute = selected, wrapper.value
}
kind := remaining[0]
remaining = remaining[1:]
var name strings.Builder
Expand Down
Loading