From cf6f0e1817f712648b69d065166897c063af721f Mon Sep 17 00:00:00 2001 From: Raphael Simon Date: Wed, 16 Sep 2026 07:40:30 -0700 Subject: [PATCH] Fix recursive gRPC collection generation --- UPGRADING.md | 4 +- codegen/ARCHITECTURE.md | 8 + ...c_recursive_collection_integration_test.go | 225 ++++++++++++++++++ codegen/go_transform.go | 17 +- codegen/go_transform_test.go | 10 +- codegen/transform_helper_registry.go | 45 ++-- codegen/transform_helper_wrapper_test.go | 45 ++++ codegen/transformer.go | 10 +- codegen/validation_plan.go | 15 +- codegen/validation_plan_test.go | 26 ++ .../oneof_anonymous_user_union_test.go | 5 +- grpc/codegen/proto_hooks.go | 30 ++- .../templates/transform_go_array.go.tpl | 8 + ...t_types_client-result-collection.go.golden | 18 +- ...ection-to-result-type-collection.go.golden | 14 +- ...uf-type_type-array-to-type-array.go.golden | 8 +- ...ection-to-result-type-collection.go.golden | 14 +- ...ce-type_type-array-to-type-array.go.golden | 8 +- ...ixed_view_collection_constructor.go.golden | 6 +- ...r_types_server-result-collection.go.golden | 19 +- ...helper_wrapped_collection_client.go.golden | 17 +- ...helper_wrapped_collection_server.go.golden | 17 +- 22 files changed, 447 insertions(+), 122 deletions(-) create mode 100644 codegen/generator/generate_grpc_recursive_collection_integration_test.go create mode 100644 codegen/transform_helper_wrapper_test.go diff --git a/UPGRADING.md b/UPGRADING.md index ed50927d12..ea79fa2f40 100644 --- a/UPGRADING.md +++ b/UPGRADING.md @@ -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 diff --git a/codegen/ARCHITECTURE.md b/codegen/ARCHITECTURE.md index 735457e756..2df68a9042 100644 --- a/codegen/ARCHITECTURE.md +++ b/codegen/ARCHITECTURE.md @@ -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, diff --git a/codegen/generator/generate_grpc_recursive_collection_integration_test.go b/codegen/generator/generate_grpc_recursive_collection_integration_test.go new file mode 100644 index 0000000000..8f24d3a2d1 --- /dev/null +++ b/codegen/generator/generate_grpc_recursive_collection_integration_test.go @@ -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)) +} diff --git a/codegen/go_transform.go b/codegen/go_transform.go index 3d108cb9b9..12bd92ead6 100644 --- a/codegen/go_transform.go +++ b/codegen/go_transform.go @@ -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, @@ -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 diff --git a/codegen/go_transform_test.go b/codegen/go_transform_test.go index bd05de74ab..91c90505f8 100644 --- a/codegen/go_transform_test.go +++ b/codegen/go_transform_test.go @@ -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 { diff --git a/codegen/transform_helper_registry.go b/codegen/transform_helper_registry.go index 7c63aa4927..c2806cc2dd 100644 --- a/codegen/transform_helper_registry.go +++ b/codegen/transform_helper_registry.go @@ -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 { @@ -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) } @@ -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 { @@ -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 diff --git a/codegen/transform_helper_wrapper_test.go b/codegen/transform_helper_wrapper_test.go new file mode 100644 index 0000000000..02d208137d --- /dev/null +++ b/codegen/transform_helper_wrapper_test.go @@ -0,0 +1,45 @@ +// This file checks that helper sharing uses the generated value inside every +// collection wrapper, including wrappers nested below other collections. +package codegen + +import ( + "fmt" + "testing" + + "github.com/stretchr/testify/require" + + "goa.design/goa/v3/expr" +) + +func TestTransformHelperRegistryFollowsNestedWrappers(t *testing.T) { + for _, wrapTarget := range []bool{false, true} { + t.Run(fmt.Sprintf("wrap-target=%t", wrapTarget), func(t *testing.T) { + rootPlan, source, target := rootWrappedTransformPlan(t, wrapTarget) + registry := NewTransformHelperRegistry() + for depth := range 3 { + plan, err := rootPlan.program.Plan(source, target, "") + require.NoError(t, err) + require.NoError(t, registry.Collect( + plan, + transformTestLayout(t, source, GoLayoutPolicy{UseDefault: true}), + transformTestLayout(t, target, GoLayoutPolicy{UseDefault: true}), + transformTestOrderFactory(fmt.Sprintf("depth-%d", depth)).order, + )) + source = nestedTransformCollection(source) + target = nestedTransformCollection(target) + } + groups, err := registry.Finalize() + require.NoError(t, err) + require.Len(t, groups, 1, "wrapping a value must not change its conversion helper") + }) + } +} + +// nestedTransformCollection puts a value inside both an array and a map, so +// the next helper lookup must cross both collections before opening its wrapper. +func nestedTransformCollection(value *expr.AttributeExpr) *expr.AttributeExpr { + return &expr.AttributeExpr{Type: &expr.Array{ElemType: &expr.AttributeExpr{Type: &expr.Map{ + KeyType: &expr.AttributeExpr{Type: expr.String}, + ElemType: value, + }}}} +} diff --git a/codegen/transformer.go b/codegen/transformer.go index effcdba736..675f24b253 100644 --- a/codegen/transformer.go +++ b/codegen/transformer.go @@ -225,7 +225,7 @@ type ( target *expr.AttributeExpr rootSource *expr.AttributeExpr rootTarget *expr.AttributeExpr - rootWrap *WrapDirective + wrappers map[TransformHelperDefinitionLocation]transformLayoutWrapper sourceBaseline *expr.AttributeExpr targetBaseline *expr.AttributeExpr sourceCopier *expr.AttributeGraphCopier @@ -241,6 +241,14 @@ type ( renders map[transformRenderRequest]transformRenderResult } + // transformLayoutWrapper records the field selected inside a generated + // wrapper, so helper type lookup follows the same path as conversion. + transformLayoutWrapper struct { + wrapper *expr.AttributeExpr + value *expr.AttributeExpr + directive *WrapDirective + } + // transformRenderRequest identifies one Render invocation. Repeating the // same invocation returns its first result instead of invoking hooks again. transformRenderRequest struct { diff --git a/codegen/validation_plan.go b/codegen/validation_plan.go index 3d5460b162..13f9d521fb 100644 --- a/codegen/validation_plan.go +++ b/codegen/validation_plan.go @@ -546,15 +546,12 @@ func attributeNeedsValidation(attribute *expr.AttributeExpr, policy GoLayoutPoli return true } } + if nested, ok := attribute.Type.(expr.UserType); ok && !expr.IsAlias(nested) { + return userTypeNeedsValidation(nested, policy, seen) + } switch { case expr.IsObject(attribute.Type): for _, field := range *expr.AsObject(attribute.Type) { - if nested, ok := field.Attribute.Type.(expr.UserType); ok && !expr.IsAlias(nested) { - if userTypeNeedsValidation(nested, policy, seen) { - return true - } - continue - } if attributeNeedsValidation(field.Attribute, policy, seen) { return true } @@ -576,12 +573,6 @@ func attributeNeedsValidation(attribute *expr.AttributeExpr, policy GoLayoutPoli for _, branch := range expr.AsUnion(attribute.Type).Values { branchPolicy := policy branchPolicy.Pointer = policy.Pointer && expr.IsObject(branch.Attribute.Type) - if nested, ok := branch.Attribute.Type.(expr.UserType); ok && !expr.IsAlias(nested) { - if userTypeNeedsValidation(nested, branchPolicy, seen) { - return true - } - continue - } if attributeNeedsValidation(branch.Attribute, branchPolicy, seen) { return true } diff --git a/codegen/validation_plan_test.go b/codegen/validation_plan_test.go index 748d184282..bcdca10051 100644 --- a/codegen/validation_plan_test.go +++ b/codegen/validation_plan_test.go @@ -345,6 +345,32 @@ func TestNeedsValidation(t *testing.T) { } } +// TestNeedsValidationRecursiveCollections checks that walking a recursive array +// or map terminates and still finds a rule after the recursive field. +func TestNeedsValidationRecursiveCollections(t *testing.T) { + for _, container := range []string{"array", "map"} { + for _, constrained := range []bool{false, true} { + t.Run(fmt.Sprintf("%s/constrained=%t", container, constrained), func(t *testing.T) { + node := goTypeTestUserType("Node", &expr.Object{}) + element := &expr.AttributeExpr{Type: node} + var children expr.DataType = &expr.Array{ElemType: element} + if container == "map" { + children = &expr.Map{KeyType: &expr.AttributeExpr{Type: expr.String}, ElemType: element} + } + label := &expr.AttributeExpr{Type: expr.String} + if constrained { + label.Validation = &expr.ValidationExpr{Pattern: "^valid$"} + } + node.Attribute().Type = &expr.Object{ + {Name: "children", Attribute: &expr.AttributeExpr{Type: children}}, + {Name: "label", Attribute: label}, + } + require.Equal(t, constrained, NeedsValidation(&expr.AttributeExpr{Type: node}, GoLayoutPolicy{SumType: true})) + }) + } + } +} + // TestValidationPlanChecksOnlyRepresentableNullElements verifies that a null // check follows the generated element type instead of the raw DSL flag. func TestValidationPlanChecksOnlyRepresentableNullElements(t *testing.T) { diff --git a/grpc/codegen/oneof_anonymous_user_union_test.go b/grpc/codegen/oneof_anonymous_user_union_test.go index 50ae13f425..522b4be976 100644 --- a/grpc/codegen/oneof_anonymous_user_union_test.go +++ b/grpc/codegen/oneof_anonymous_user_union_test.go @@ -58,8 +58,11 @@ func TestAnonymousUserUnionArrayNoWrappersFromProto(t *testing.T) { freezeProtoBufTransformMessages(t, sd, source) pbCtx := protoBufTypeContext("proto", sd) - code, _, err := protoBufTransform(source, target, pbCtx, svcCtx, false, true) + code, helpers, err := protoBufTransform(source, target, pbCtx, svcCtx, false, true) require.NoError(t, err) + for _, helper := range helpers { + code += "\n" + helper.Code + } out := codegen.FormatTestCode(t, "package foo\nfunc transform(){\n"+code+"}") // Ensure no per-branch wrapper casts (e.g., types.DetailsAlpha(...)). diff --git a/grpc/codegen/proto_hooks.go b/grpc/codegen/proto_hooks.go index b731c83a4d..96f6a97be6 100644 --- a/grpc/codegen/proto_hooks.go +++ b/grpc/codegen/proto_hooks.go @@ -69,7 +69,10 @@ func init() { panic(fmt.Sprintf("create protobuf-to-service transform program: %s", err)) } - fm := template.FuncMap{"transformAttribute": codegen.TransformAttribute} + fm := template.FuncMap{ + "transformAttribute": codegen.TransformAttribute, + "transformHelperName": codegen.TransformHelperName, + } renderGoArrayT = template.Must(template.New("renderGoArray").Funcs(fm).Parse(grpcTemplates.Read(grpcTransformGoArrayT))) renderGoMapT = template.Must(template.New("renderGoMap").Funcs(fm).Parse(grpcTemplates.Read(grpcTransformGoMapT))) renderGoUnionToProtoT = template.Must(template.New("renderGoUnionToProto").Parse(grpcTemplates.Read(grpcTransformGoUnionToProtoT))) @@ -161,7 +164,6 @@ func protoHooks(proto bool) *codegen.TransformHooks { } return "&", true }, - InlineCompositeElems: true, } } @@ -283,8 +285,9 @@ func renderArrayTransform(source, target *expr.Array, sourceVar, targetVar strin } targetRef := ta.TargetCtx.Scope.Ref(elem, ta.TargetCtx.Pkg(elem)) + useHelper := protoCollectionUsesHelper(source.ElemType, target.ElemType) valVar := "val" - if obj := expr.AsObject(source.ElemType.Type); obj != nil && len(*obj) == 0 { + if obj := expr.AsObject(source.ElemType.Type); !useHelper && obj != nil && len(*obj) == 0 { valVar = "" } @@ -303,6 +306,7 @@ func renderArrayTransform(source, target *expr.Array, sourceVar, targetVar strin "TransformAttrs": childAttrs, "LoopVar": loopVar, "ValVar": valVar, + "UseHelper": useHelper, } var buf bytes.Buffer if err := renderGoArrayT.Execute(&buf, data); err != nil { @@ -337,9 +341,14 @@ func renderMapTransform(source, target *expr.Map, sourceVar, targetVar string, n if !proto && isWrappedAttr(source.ElemType) { elemNewVar = false } - elemTransform, err := codegen.TransformAttribute(source.ElemType, target.ElemType, "val", elemTarget, elemNewVar, blockAttrs) - if err != nil { - return "", err + var elemTransform string + if protoCollectionUsesHelper(source.ElemType, target.ElemType) { + elemTransform = fmt.Sprintf("%s := %s(val)\n", elemTarget, codegen.TransformHelperName(source.ElemType, target.ElemType, blockAttrs)) + } else { + elemTransform, err = codegen.TransformAttribute(source.ElemType, target.ElemType, "val", elemTarget, elemNewVar, blockAttrs) + if err != nil { + return "", err + } } if !elemNewVar { elemTransform = fmt.Sprintf("var %s %s\nif val != nil {\n%s}\n", elemTarget, ta.TargetCtx.Scope.Ref(et, ta.TargetCtx.Pkg(et)), elemTransform) @@ -363,6 +372,15 @@ func renderMapTransform(source, target *expr.Map, sourceVar, targetVar string, n return ensureTrailingNewline(buf.String()), nil } +// protoCollectionUsesHelper selects named object conversions, matching the +// shared transform planner. Collection wrappers contain arrays or maps on the +// service side, so they stay inline until their object elements are reached. +func protoCollectionUsesHelper(source, target *expr.AttributeExpr) bool { + _, sourceNamed := source.Type.(expr.UserType) + _, targetNamed := target.Type.(expr.UserType) + return sourceNamed && targetNamed && expr.IsObject(source.Type) && expr.IsObject(target.Type) +} + // renderUnionToProtoTransform writes a service union from sourceVar into the // protobuf oneof in targetVar. func renderUnionToProtoTransform(source, target *expr.AttributeExpr, sourceVar, targetVar string, sourcePtr bool, ta *codegen.TransformAttrs) (string, error) { diff --git a/grpc/codegen/templates/transform_go_array.go.tpl b/grpc/codegen/templates/transform_go_array.go.tpl index 9b96f75939..3465e3915f 100644 --- a/grpc/codegen/templates/transform_go_array.go.tpl +++ b/grpc/codegen/templates/transform_go_array.go.tpl @@ -8,12 +8,20 @@ {{- $arr := printf "arr%s" .LoopVar -}} {{ $arr }} := make([]{{ .ElemTypeRef }}, len({{ if .SourcePtr }}*{{ end }}{{ .SourceVar }})) for {{ .LoopVar }}{{ if .ValVar }}, {{ .ValVar }}{{ end }} := range {{ if .SourcePtr }}*{{ end }}{{ .SourceVar }} { +{{ if .UseHelper -}} + {{ $arr }}[{{ .LoopVar }}] = {{ transformHelperName .SourceElem .TargetElem .TransformAttrs }}(val) +{{ else -}} {{ transformAttribute .SourceElem .TargetElem "val" (printf "%s[%s]" $arr .LoopVar) false .TransformAttrs -}} +{{ end -}} } {{ .TargetVar }} = &{{ $arr }} {{- else -}} {{ .TargetVar }} {{ if .NewVar }}:={{ else }}={{ end }} make([]{{ .ElemTypeRef }}, len({{ if .SourcePtr }}*{{ end }}{{ .SourceVar }})) for {{ .LoopVar }}{{ if .ValVar }}, {{ .ValVar }}{{ end }} := range {{ if .SourcePtr }}*{{ end }}{{ .SourceVar }} { +{{ if .UseHelper -}} + {{ .TargetVar }}[{{ .LoopVar }}] = {{ transformHelperName .SourceElem .TargetElem .TransformAttrs }}(val) +{{ else -}} {{ transformAttribute .SourceElem .TargetElem "val" (printf "%s[%s]" .TargetVar .LoopVar) false .TransformAttrs -}} +{{ end -}} } {{- end -}} diff --git a/grpc/codegen/testdata/golden/client_types_client-result-collection.go.golden b/grpc/codegen/testdata/golden/client_types_client-result-collection.go.golden index b428fb009c..82070d10b5 100644 --- a/grpc/codegen/testdata/golden/client_types_client-result-collection.go.golden +++ b/grpc/codegen/testdata/golden/client_types_client-result-collection.go.golden @@ -25,13 +25,21 @@ func transformProtoResultTToResultT(v *service_result_with_collectionpb.ResultT) if v.CollectionField != nil { res.CollectionField = make([]*serviceresultwithcollection.RT, len(v.CollectionField.Field)) for i, val := range v.CollectionField.Field { - res.CollectionField[i] = &serviceresultwithcollection.RT{} - if val.IntField != nil { - intField := int(*val.IntField) - res.CollectionField[i].IntField = &intField - } + res.CollectionField[i] = transformProtoRTToRT(val) } } return res } + +// transformProtoRTToRT builds a value of type *serviceresultwithcollection.RT +// from a value of type *service_result_with_collectionpb.RT. +func transformProtoRTToRT(v *service_result_with_collectionpb.RT) *serviceresultwithcollection.RT { + res := &serviceresultwithcollection.RT{} + if v.IntField != nil { + intField := int(*v.IntField) + res.IntField = &intField + } + + return res +} diff --git a/grpc/codegen/testdata/golden/protobuf_type_encode_to-protobuf-type_result-type-collection-to-result-type-collection.go.golden b/grpc/codegen/testdata/golden/protobuf_type_encode_to-protobuf-type_result-type-collection-to-result-type-collection.go.golden index 16ec026b91..4204ee3c0b 100644 --- a/grpc/codegen/testdata/golden/protobuf_type_encode_to-protobuf-type_result-type-collection-to-result-type-collection.go.golden +++ b/grpc/codegen/testdata/golden/protobuf_type_encode_to-protobuf-type_result-type-collection-to-result-type-collection.go.golden @@ -4,19 +4,7 @@ func transform() { target.Collection = &proto.ResultTypeCollection2{} target.Collection.Field = make([]*proto.ResultType, len(source.Collection)) for i, val := range source.Collection { - target.Collection.Field[i] = &proto.ResultType{} - if val.Int != nil { - int_ := int32(*val.Int) - target.Collection.Field[i].Int = &int_ - } - if val.Map != nil { - target.Collection.Field[i].Map_ = make(map[int32]string, len(val.Map)) - for key, val := range val.Map { - tk := int32(key) - tv := val - target.Collection.Field[i].Map_[tk] = tv - } - } + target.Collection.Field[i] = svcProtoResultTypeToProtoResultType(val) } } } diff --git a/grpc/codegen/testdata/golden/protobuf_type_encode_to-protobuf-type_type-array-to-type-array.go.golden b/grpc/codegen/testdata/golden/protobuf_type_encode_to-protobuf-type_type-array-to-type-array.go.golden index a34fa583ba..08971033ee 100644 --- a/grpc/codegen/testdata/golden/protobuf_type_encode_to-protobuf-type_type-array-to-type-array.go.golden +++ b/grpc/codegen/testdata/golden/protobuf_type_encode_to-protobuf-type_type-array-to-type-array.go.golden @@ -3,13 +3,7 @@ func transform() { if source.TypeArray != nil { target.TypeArray = make([]*proto.SimpleArray, len(source.TypeArray)) for i, val := range source.TypeArray { - target.TypeArray[i] = &proto.SimpleArray{} - if val.StringArray != nil { - target.TypeArray[i].StringArray = make([]string, len(val.StringArray)) - for j, val := range val.StringArray { - target.TypeArray[i].StringArray[j] = val - } - } + target.TypeArray[i] = svcProtoSimpleArrayToProtoSimpleArray(val) } } } diff --git a/grpc/codegen/testdata/golden/protobuf_type_encode_to-service-type_result-type-collection-to-result-type-collection.go.golden b/grpc/codegen/testdata/golden/protobuf_type_encode_to-service-type_result-type-collection-to-result-type-collection.go.golden index 9db70510f9..8eee949711 100644 --- a/grpc/codegen/testdata/golden/protobuf_type_encode_to-service-type_result-type-collection-to-result-type-collection.go.golden +++ b/grpc/codegen/testdata/golden/protobuf_type_encode_to-service-type_result-type-collection-to-result-type-collection.go.golden @@ -3,19 +3,7 @@ func transform() { if source.Collection != nil { target.Collection = make([]*proto.ResultType, len(source.Collection.Field)) for i, val := range source.Collection.Field { - target.Collection[i] = &proto.ResultType{} - if val.Int != nil { - int_ := int(*val.Int) - target.Collection[i].Int = &int_ - } - if val.Map_ != nil { - target.Collection[i].Map = make(map[int]string, len(val.Map_)) - for key, val := range val.Map_ { - tk := int(key) - tv := val - target.Collection[i].Map[tk] = tv - } - } + target.Collection[i] = protobufProtoResultTypeToProtoResultType(val) } } } diff --git a/grpc/codegen/testdata/golden/protobuf_type_encode_to-service-type_type-array-to-type-array.go.golden b/grpc/codegen/testdata/golden/protobuf_type_encode_to-service-type_type-array-to-type-array.go.golden index a34fa583ba..817b6161a0 100644 --- a/grpc/codegen/testdata/golden/protobuf_type_encode_to-service-type_type-array-to-type-array.go.golden +++ b/grpc/codegen/testdata/golden/protobuf_type_encode_to-service-type_type-array-to-type-array.go.golden @@ -3,13 +3,7 @@ func transform() { if source.TypeArray != nil { target.TypeArray = make([]*proto.SimpleArray, len(source.TypeArray)) for i, val := range source.TypeArray { - target.TypeArray[i] = &proto.SimpleArray{} - if val.StringArray != nil { - target.TypeArray[i].StringArray = make([]string, len(val.StringArray)) - for j, val := range val.StringArray { - target.TypeArray[i].StringArray[j] = val - } - } + target.TypeArray[i] = protobufProtoSimpleArrayToProtoSimpleArray(val) } } } diff --git a/grpc/codegen/testdata/golden/released_fixed_view_collection_constructor.go.golden b/grpc/codegen/testdata/golden/released_fixed_view_collection_constructor.go.golden index c988387ffe..ba6bca8a55 100644 --- a/grpc/codegen/testdata/golden/released_fixed_view_collection_constructor.go.golden +++ b/grpc/codegen/testdata/golden/released_fixed_view_collection_constructor.go.golden @@ -6,11 +6,7 @@ func NewProtoResultTypeCollection(result serviceclientstreamingresulttypecollect message := &service_client_streaming_result_type_collection_with_explicit_viewpb.ResultTypeCollection{} message.Field = make([]*service_client_streaming_result_type_collection_with_explicit_viewpb.ResultType, len(result)) for i, val := range result { - message.Field[i] = &service_client_streaming_result_type_collection_with_explicit_viewpb.ResultType{} - if val.IntField != nil { - intField := int32(*val.IntField) - message.Field[i].IntField = &intField - } + message.Field[i] = transformResultTypeViewToProtoResultType(val) } return message } diff --git a/grpc/codegen/testdata/golden/server_types_server-result-collection.go.golden b/grpc/codegen/testdata/golden/server_types_server-result-collection.go.golden index ace306f00b..7238dd9928 100644 --- a/grpc/codegen/testdata/golden/server_types_server-result-collection.go.golden +++ b/grpc/codegen/testdata/golden/server_types_server-result-collection.go.golden @@ -18,13 +18,22 @@ func transformResultTToProtoResultT(v *serviceresultwithcollection.ResultT) *ser res.CollectionField = &service_result_with_collectionpb.RTCollection{} res.CollectionField.Field = make([]*service_result_with_collectionpb.RT, len(v.CollectionField)) for i, val := range v.CollectionField { - res.CollectionField.Field[i] = &service_result_with_collectionpb.RT{} - if val.IntField != nil { - intField := int32(*val.IntField) - res.CollectionField.Field[i].IntField = &intField - } + res.CollectionField.Field[i] = transformRTToProtoRT(val) } } return res } + +// transformRTToProtoRT builds a value of type +// *service_result_with_collectionpb.RT from a value of type +// *serviceresultwithcollection.RT. +func transformRTToProtoRT(v *serviceresultwithcollection.RT) *service_result_with_collectionpb.RT { + res := &service_result_with_collectionpb.RT{} + if v.IntField != nil { + intField := int32(*v.IntField) + res.IntField = &intField + } + + return res +} diff --git a/grpc/codegen/testdata/golden/transform_helper_wrapped_collection_client.go.golden b/grpc/codegen/testdata/golden/transform_helper_wrapped_collection_client.go.golden index 2d17f7da7e..b3619a7e09 100644 --- a/grpc/codegen/testdata/golden/transform_helper_wrapped_collection_client.go.golden +++ b/grpc/codegen/testdata/golden/transform_helper_wrapped_collection_client.go.golden @@ -11,10 +11,7 @@ func NewProtoListRequest() *wrapped_collectionpb.ListRequest { func NewListResult(message *wrapped_collectionpb.WrappedCollectionItemCollection) wrappedcollectionviews.WrappedCollectionItemCollectionView { result := make([]*wrappedcollectionviews.WrappedCollectionItemView, len(message.Field)) for i, val := range message.Field { - result[i] = &wrappedcollectionviews.WrappedCollectionItemView{} - if val.Child != nil { - result[i].Child = transformProtoWrappedCollectionChildToWrappedCollectionChildView(val.Child) - } + result[i] = transformProtoWrappedCollectionItemToWrappedCollectionItemView(val) } return result } @@ -55,6 +52,18 @@ func validatetest_20_api_WrappedCollection_WrappedCollectionChild_At_child(child return } +// transformProtoWrappedCollectionItemToWrappedCollectionItemView builds a +// value of type *wrappedcollectionviews.WrappedCollectionItemView from a value +// of type *wrapped_collectionpb.WrappedCollectionItem. +func transformProtoWrappedCollectionItemToWrappedCollectionItemView(v *wrapped_collectionpb.WrappedCollectionItem) *wrappedcollectionviews.WrappedCollectionItemView { + res := &wrappedcollectionviews.WrappedCollectionItemView{} + if v.Child != nil { + res.Child = transformProtoWrappedCollectionChildToWrappedCollectionChildView(v.Child) + } + + return res +} + // transformProtoWrappedCollectionChildToWrappedCollectionChildView builds a // value of type *wrappedcollectionviews.WrappedCollectionChildView from a // value of type *wrapped_collectionpb.WrappedCollectionChild. diff --git a/grpc/codegen/testdata/golden/transform_helper_wrapped_collection_server.go.golden b/grpc/codegen/testdata/golden/transform_helper_wrapped_collection_server.go.golden index 8d5a096b2c..998bf7cf22 100644 --- a/grpc/codegen/testdata/golden/transform_helper_wrapped_collection_server.go.golden +++ b/grpc/codegen/testdata/golden/transform_helper_wrapped_collection_server.go.golden @@ -5,14 +5,23 @@ func NewProtoWrappedCollectionItemCollection(result wrappedcollectionviews.Wrapp message := &wrapped_collectionpb.WrappedCollectionItemCollection{} message.Field = make([]*wrapped_collectionpb.WrappedCollectionItem, len(result)) for i, val := range result { - message.Field[i] = &wrapped_collectionpb.WrappedCollectionItem{} - if val.Child != nil { - message.Field[i].Child = transformWrappedCollectionChildViewToProtoWrappedCollectionChild(val.Child) - } + message.Field[i] = transformWrappedCollectionItemViewToProtoWrappedCollectionItem(val) } return message } +// transformWrappedCollectionItemViewToProtoWrappedCollectionItem builds a +// value of type *wrapped_collectionpb.WrappedCollectionItem from a value of +// type *wrappedcollectionviews.WrappedCollectionItemView. +func transformWrappedCollectionItemViewToProtoWrappedCollectionItem(v *wrappedcollectionviews.WrappedCollectionItemView) *wrapped_collectionpb.WrappedCollectionItem { + res := &wrapped_collectionpb.WrappedCollectionItem{} + if v.Child != nil { + res.Child = transformWrappedCollectionChildViewToProtoWrappedCollectionChild(v.Child) + } + + return res +} + // transformWrappedCollectionChildViewToProtoWrappedCollectionChild builds a // value of type *wrapped_collectionpb.WrappedCollectionChild from a value of // type *wrappedcollectionviews.WrappedCollectionChildView.