diff --git a/graphql/schema/request.go b/graphql/schema/request.go index 15c90299afc..bc5616dcffc 100644 --- a/graphql/schema/request.go +++ b/graphql/schema/request.go @@ -85,10 +85,27 @@ func (s *schema) Operation(req *Request) (Operation, error) { interfaceImplFragFields: map[*ast.Field]string{}, } - // recursively expand fragments in operation as selection set fields - for _, s := range op.SelectionSet { - recursivelyExpandFragmentSelections(s.(*ast.Field), operation) + rootType := s.schema.Query + switch op.Operation { + case ast.Mutation: + rootType = s.schema.Mutation + case ast.Subscription: + rootType = s.schema.Subscription } + if rootType == nil { + return nil, errors.Errorf("Not resolving operation because schema doesn't have a root type defined for %s.", + op.Operation) + } + // Normalize root fragments with the same type-aware collector used for nested selections. + root := &ast.Field{ + Definition: &ast.FieldDefinition{Type: ast.NamedType(rootType.Name, nil)}, + SelectionSet: op.SelectionSet, + } + recursivelyExpandFragmentSelections(root, operation) + if op.Operation == ast.Subscription && len(root.SelectionSet) != 1 { + return nil, errors.New("Subscription must select exactly one top level field.") + } + op.SelectionSet = root.SelectionSet return operation, nil } diff --git a/graphql/schema/wrappers_test.go b/graphql/schema/wrappers_test.go index 7cc0fe4cd0f..8713cf72305 100644 --- a/graphql/schema/wrappers_test.go +++ b/graphql/schema/wrappers_test.go @@ -21,6 +21,130 @@ import ( "github.com/dgraph-io/gqlparser/v2/ast" ) +// TestOperationRootFragments checks root normalization across all supported operation kinds. +func TestOperationRootFragments(t *testing.T) { + handler, err := NewHandler(`type Book @withSubscription { id: ID! name: String }`, false) + require.NoError(t, err) + sch, err := FromString(handler.GQLSchema(), x.RootNamespace) + require.NoError(t, err) + for _, testCase := range []struct { + name, document, operation string + variables map[string]interface{} + aliases []string + mutation bool + }{ + {"named_query", `query { ...Books } fragment Books on Query { first: queryBook { id } }`, "", nil, []string{"first"}, false}, + {"inline_query", `query { ... on Query { first: queryBook { id } } }`, "", nil, []string{"first"}, false}, + {"multiple_query_fields", `query { first: queryBook { id } second: queryBook { name } }`, "", nil, []string{"first", "second"}, false}, + {"named_mutation", `mutation { ...Books } fragment Books on Mutation { first: addBook(input: [{name: "first"}]) { numUids } second: addBook(input: [{name: "second"}]) { numUids } }`, "", nil, []string{"first", "second"}, true}, + {"named_subscription", `subscription { ...Books } fragment Books on Subscription { first: queryBook { id } }`, "", nil, []string{"first"}, false}, + {"included_fragment", `query($enabled: Boolean!) { ...Books @include(if: $enabled) } fragment Books on Query { first: queryBook { id } }`, "", map[string]interface{}{"enabled": true}, []string{"first"}, false}, + {"excluded_fragment", `query($enabled: Boolean!) { ...Books @include(if: $enabled) } fragment Books on Query { first: queryBook { id } }`, "", map[string]interface{}{"enabled": false}, nil, false}, + {"shared_named_fragment", `query First { ...Books } query Second { ...Books } fragment Books on Query { first: queryBook { id } }`, "Second", nil, []string{"first"}, false}, + {"merged_alias", `query { first: queryBook { id } ...Books } fragment Books on Query { first: queryBook { name } }`, "", nil, []string{"first"}, false}, + {"skipped_inline", `query { ... on Query @skip(if: true) { first: queryBook { id } } }`, "", nil, nil, false}, + } { + t.Run(testCase.name, func(t *testing.T) { + var op Operation + require.NotPanics(t, func() { + op, err = sch.Operation(&Request{Query: testCase.document, OperationName: testCase.operation, Variables: testCase.variables}) + }) + require.NoError(t, err) + var actual []string + if testCase.mutation { + for _, field := range op.Mutations() { + actual = append(actual, field.ResponseName()) + } + } else { + for _, field := range op.Queries() { + actual = append(actual, field.ResponseName()) + } + } + require.Equal(t, testCase.aliases, actual) + if testCase.name == "merged_alias" { + fields := op.Queries()[0].SelectionSet() + require.Len(t, fields, 2) + require.Equal(t, "id", fields[0].Name()) + require.Equal(t, "name", fields[1].Name()) + } + }) + } +} + +// TestOperationMissingRootType rejects unsupported selected operations without panicking. +func TestOperationMissingRootType(t *testing.T) { + sch, err := FromString(`type Query { ping: String }`, x.RootNamespace) + require.NoError(t, err) + for _, testCase := range []struct { + name, document, selected, wantError string + }{ + {"mutation", `mutation { __typename }`, "", "mutation"}, + {"selected_mutation", `query Read { __typename } mutation Write { __typename }`, "Write", "mutation"}, + {"selected_subscription", `query Read { __typename } subscription Watch { __typename }`, "Watch", "subscription"}, + {"subscription", `subscription { __typename }`, "", "subscription"}, + {"read_with_unused_mutation", `query Read { __typename } mutation Write { __typename }`, "Read", ""}, + {"read_with_unused_subscription", `query Read { __typename } subscription Watch { __typename }`, "Read", ""}, + } { + t.Run(testCase.name, func(t *testing.T) { + var op Operation + var operationErr error + require.NotPanics(t, func() { + op, operationErr = sch.Operation(&Request{Query: testCase.document, OperationName: testCase.selected}) + }) + if testCase.wantError != "" { + require.ErrorContains(t, operationErr, testCase.wantError) + require.Nil(t, op) + return + } + require.NoError(t, operationErr) + require.Len(t, op.Queries(), 1) + require.Equal(t, "__typename", op.Queries()[0].Name()) + }) + } +} + +// TestOperationSubscriptionRootFields checks cardinality after fragment collection and alias merging. +func TestOperationSubscriptionRootFields(t *testing.T) { + handler, err := NewHandler(`type Book @withSubscription { id: ID! name: String }`, false) + require.NoError(t, err) + sch, err := FromString(handler.GQLSchema(), x.RootNamespace) + require.NoError(t, err) + for _, testCase := range []struct { + name, document, selected string + wantError bool + }{ + {"multiple_named_fields", `subscription { ...Books } fragment Books on Subscription { first: queryBook { id } second: queryBook { name } }`, "", true}, + {"multiple_nested_fields", `subscription { ...Outer } fragment Outer on Subscription { first: queryBook { id } ...Inner } fragment Inner on Subscription { second: queryBook { name } }`, "", true}, + {"multiple_inline_fields", `subscription { ... on Subscription { first: queryBook { id } second: queryBook { name } } }`, "", true}, + {"excluded_root_fragment", `subscription { ...Books @skip(if: true) } fragment Books on Subscription { first: queryBook { id } }`, "", true}, + {"single_nested_field", `subscription { ...Outer } fragment Outer on Subscription { ...Inner } fragment Inner on Subscription { first: queryBook { id } }`, "", false}, + {"merged_alias", `subscription { ...Books } fragment Books on Subscription { first: queryBook { id } first: queryBook { name } }`, "", false}, + {"selected_subscription", `query Read { first: queryBook { id } second: queryBook { name } } subscription Watch { ...Books } fragment Books on Subscription { first: queryBook { id } }`, "Watch", false}, + {"selected_query", `query Read { first: queryBook { id } second: queryBook { name } } subscription Watch { ...Books } fragment Books on Subscription { first: queryBook { id } second: queryBook { name } }`, "Read", false}, + } { + t.Run(testCase.name, func(t *testing.T) { + op, operationErr := sch.Operation(&Request{Query: testCase.document, OperationName: testCase.selected}) + if testCase.wantError { + require.ErrorContains(t, operationErr, "exactly one top level field") + require.Nil(t, op) + return + } + require.NoError(t, operationErr) + if testCase.selected == "Read" { + require.False(t, op.IsSubscription()) + require.Len(t, op.Queries(), 2) + return + } + require.True(t, op.IsSubscription()) + require.Len(t, op.Queries(), 1) + require.Equal(t, "first", op.Queries()[0].ResponseName()) + if testCase.name == "merged_alias" { + require.Len(t, op.Queries()[0].SelectionSet(), 2) + } + }) + } +} + func TestDgraphMapping_WithoutDirectives(t *testing.T) { schemaStr := ` type Author {