diff --git a/spring-graphql/src/main/java/org/springframework/graphql/data/pagination/ConnectionFieldTypeVisitor.java b/spring-graphql/src/main/java/org/springframework/graphql/data/pagination/ConnectionFieldTypeVisitor.java index d0b260bf..fad3259c 100644 --- a/spring-graphql/src/main/java/org/springframework/graphql/data/pagination/ConnectionFieldTypeVisitor.java +++ b/spring-graphql/src/main/java/org/springframework/graphql/data/pagination/ConnectionFieldTypeVisitor.java @@ -34,8 +34,10 @@ import graphql.schema.DataFetchingEnvironment; import graphql.schema.GraphQLCodeRegistry; import graphql.schema.GraphQLFieldDefinition; import graphql.schema.GraphQLFieldsContainer; +import graphql.schema.GraphQLList; import graphql.schema.GraphQLNonNull; import graphql.schema.GraphQLObjectType; +import graphql.schema.GraphQLOutputType; import graphql.schema.GraphQLSchemaElement; import graphql.schema.GraphQLType; import graphql.schema.GraphQLTypeVisitorStub; @@ -101,14 +103,54 @@ public final class ConnectionFieldTypeVisitor extends GraphQLTypeVisitorStub { return TraversalControl.CONTINUE; } - private static boolean isConnectionField(GraphQLFieldDefinition fieldDefinition) { - GraphQLType returnType = fieldDefinition.getType(); - if (returnType instanceof GraphQLNonNull nonNullType) { - returnType = nonNullType.getWrappedType(); + private static boolean isConnectionField(GraphQLFieldDefinition field) { + GraphQLObjectType type = getAsObjectType(field); + if (type == null || !type.getName().endsWith("Connection")) { + return false; } - return (returnType instanceof GraphQLObjectType objectType && - objectType.getName().endsWith("Connection") && - objectType.getField("pageInfo") != null); + + GraphQLObjectType edgeType = getEdgeType(type.getField("edges")); + if (edgeType == null || !edgeType.getName().endsWith("Edge")) { + return false; + } + if (edgeType.getField("node") == null || edgeType.getField("cursor") == null) { + return false; + } + + GraphQLObjectType pageInfoType = getAsObjectType(type.getField("pageInfo")); + if (pageInfoType == null || !pageInfoType.getName().equals("PageInfo")) { + return false; + } + if (pageInfoType.getField("hasPreviousPage") == null || pageInfoType.getField("hasNextPage") == null || + pageInfoType.getField("startCursor") == null || pageInfoType.getField("endCursor") == null) { + return false; + } + + return true; + } + + @Nullable + private static GraphQLObjectType getAsObjectType(@Nullable GraphQLFieldDefinition field) { + return (getType(field) instanceof GraphQLObjectType type ? type : null); + } + + @Nullable + private static GraphQLObjectType getEdgeType(@Nullable GraphQLFieldDefinition field) { + if (getType(field) instanceof GraphQLList listType) { + if (listType.getWrappedType() instanceof GraphQLObjectType type) { + return type; + } + } + return null; + } + + @Nullable + private static GraphQLType getType(@Nullable GraphQLFieldDefinition field) { + if (field == null) { + return null; + } + GraphQLOutputType type = field.getType(); + return (type instanceof GraphQLNonNull nonNullType ? nonNullType.getWrappedType() : type); } diff --git a/spring-graphql/src/test/java/org/springframework/graphql/data/pagination/ConnectionFieldTypeVisitorTests.java b/spring-graphql/src/test/java/org/springframework/graphql/data/pagination/ConnectionFieldTypeVisitorTests.java index 2e0fc71a..b480c8f0 100644 --- a/spring-graphql/src/test/java/org/springframework/graphql/data/pagination/ConnectionFieldTypeVisitorTests.java +++ b/spring-graphql/src/test/java/org/springframework/graphql/data/pagination/ConnectionFieldTypeVisitorTests.java @@ -19,14 +19,28 @@ package org.springframework.graphql.data.pagination; import java.util.Collection; import java.util.List; +import graphql.schema.DataFetcher; +import graphql.schema.FieldCoordinates; +import graphql.schema.GraphQLFieldDefinition; +import graphql.schema.GraphQLSchema; import graphql.schema.PropertyDataFetcher; +import graphql.schema.SchemaTransformer; +import graphql.schema.idl.RuntimeWiring; +import graphql.schema.idl.SchemaGenerator; +import graphql.schema.idl.SchemaParser; +import graphql.schema.idl.TypeDefinitionRegistry; +import org.junit.jupiter.api.Nested; import org.junit.jupiter.api.Test; import reactor.core.publisher.Mono; +import org.springframework.core.io.Resource; import org.springframework.graphql.BookSource; import org.springframework.graphql.ExecutionGraphQlResponse; import org.springframework.graphql.GraphQlSetup; import org.springframework.graphql.ResponseHelper; +import org.springframework.graphql.execution.ConnectionTypeDefinitionConfigurer; + +import static org.assertj.core.api.Assertions.assertThat; /** * Unit tests for {@link ConnectionFieldTypeVisitor}. @@ -69,18 +83,6 @@ public class ConnectionFieldTypeVisitorTests { ); } - @Test // gh-707 - void trivialDataFetcherIsSkipped() { - - Mono response = GraphQlSetup.schemaResource(BookSource.paginationSchema) - .dataFetcher("Query", "books", new PropertyDataFetcher<>("books")) - .connectionSupport(new ListConnectionAdapter()) - .toGraphQlService() - .execute(BookSource.booksConnectionQuery(null)); - - ResponseHelper.forResponse(response).assertData("{\"books\":null}"); - } - @Test // gh-707 void nullValueTreatedAsEmptyConnection() { @@ -103,6 +105,85 @@ public class ConnectionFieldTypeVisitorTests { } + @Nested + class DecorationTests { + + @Test // gh-707 + void trivialDataFetcherIsNotDecorated() throws Exception { + + FieldCoordinates coordinates = FieldCoordinates.coordinates("Query", "books"); + DataFetcher dataFetcher = new PropertyDataFetcher<>("books"); + + DataFetcher actual = + applyConnectionFieldTypeVisitor(BookSource.paginationSchema, coordinates, dataFetcher); + + assertThat(actual).isSameAs(dataFetcher); + } + + @Test // gh-709 + void connectionTypeWithoutEdgesIsNotDecorated() throws Exception { + + String schemaContent = """ + type Query { + puzzles: PuzzleConnection + } + type PuzzleConnection { + pageInfo: PuzzlePageInfo! + puzzles: [PuzzleEdge] + } + type PuzzlePageInfo { + total: Int! + } + type PuzzleEdge { + puzzle: Puzzle + cursor: String! + } + type Puzzle { + title: String! + } + """; + + FieldCoordinates coordinates = FieldCoordinates.coordinates("Query", "puzzles"); + DataFetcher dataFetcher = env -> null; + + DataFetcher actual = applyConnectionFieldTypeVisitor(schemaContent, coordinates, dataFetcher); + + assertThat(actual).isSameAs(dataFetcher); + } + + private static DataFetcher applyConnectionFieldTypeVisitor( + Object schemaSource, FieldCoordinates coordinates, DataFetcher fetcher) throws Exception { + + TypeDefinitionRegistry registry; + if (schemaSource instanceof Resource resource) { + registry = new SchemaParser().parse(resource.getInputStream()); + } + else if (schemaSource instanceof String schemaContent) { + registry = new SchemaParser().parse(schemaContent); + } + else { + throw new IllegalArgumentException(); + } + + new ConnectionTypeDefinitionConfigurer().configure(registry); + + GraphQLSchema schema = new SchemaGenerator().makeExecutableSchema( + registry, RuntimeWiring.newRuntimeWiring() + .type(coordinates.getTypeName(), b -> b.dataFetcher(coordinates.getFieldName(), fetcher)) + .build()); + + ConnectionFieldTypeVisitor visitor = + ConnectionFieldTypeVisitor.create(List.of(new ListConnectionAdapter())); + + schema = new SchemaTransformer().transform(schema, visitor); + GraphQLFieldDefinition field = schema.getFieldDefinition(coordinates); + return schema.getCodeRegistry().getDataFetcher(coordinates, field); + } + + } + + + private static class ListConnectionAdapter implements ConnectionAdapter { private int initialOffset = 0;