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 1cec047f..8b8a3347 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 @@ -87,8 +87,8 @@ public final class ConnectionFieldTypeVisitor extends GraphQLTypeVisitorStub { GraphQLFieldsContainer parent = (GraphQLFieldsContainer) context.getParentNode(); DataFetcher dataFetcher = codeRegistry.getDataFetcher(parent, fieldDefinition); - if (visitorHelper != null && visitorHelper.isSubscriptionType(parent)) { - return TraversalControl.ABORT; + if (visitorHelper != null && isUnderSubscriptionOperation(visitorHelper, context)) { + return TraversalControl.CONTINUE; } if (isConnectionField(fieldDefinition)) { @@ -108,6 +108,13 @@ public final class ConnectionFieldTypeVisitor extends GraphQLTypeVisitorStub { return TraversalControl.CONTINUE; } + private static boolean isUnderSubscriptionOperation(TypeVisitorHelper visitorHelper, TraverserContext context) { + return context.getBreadcrumbs().stream() + .filter(GraphQLFieldsContainer.class::isInstance) + .map(GraphQLFieldsContainer.class::cast) + .anyMatch(visitorHelper::isSubscriptionType); + } + private static boolean isConnectionField(GraphQLFieldDefinition field) { GraphQLObjectType type = getAsObjectType(field); if (type == null || !type.getName().endsWith("Connection")) { 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 389ccf63..84912e67 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 @@ -16,21 +16,33 @@ package org.springframework.graphql.data.pagination; +import java.util.Arrays; import java.util.Collection; +import java.util.HashMap; +import java.util.HashSet; import java.util.List; +import java.util.Map; +import java.util.Set; import java.util.function.Consumer; import graphql.execution.DataFetcherResult; import graphql.schema.DataFetcher; import graphql.schema.FieldCoordinates; +import graphql.schema.GraphQLCodeRegistry; import graphql.schema.GraphQLFieldDefinition; import graphql.schema.GraphQLSchema; +import graphql.schema.GraphQLSchemaElement; +import graphql.schema.GraphQLTypeVisitor; +import graphql.schema.GraphQLTypeVisitorStub; import graphql.schema.PropertyDataFetcher; import graphql.schema.SchemaTransformer; +import graphql.schema.SchemaTraverser; import graphql.schema.idl.RuntimeWiring; import graphql.schema.idl.SchemaGenerator; import graphql.schema.idl.SchemaParser; import graphql.schema.idl.TypeDefinitionRegistry; +import graphql.util.TraversalControl; +import graphql.util.TraverserContext; import org.junit.jupiter.api.Nested; import org.junit.jupiter.api.Test; import reactor.core.publisher.Mono; @@ -42,6 +54,7 @@ import org.springframework.graphql.ExecutionGraphQlResponse; import org.springframework.graphql.GraphQlSetup; import org.springframework.graphql.ResponseHelper; import org.springframework.graphql.execution.ConnectionTypeDefinitionConfigurer; +import org.springframework.graphql.execution.TypeVisitorHelper; import static org.assertj.core.api.Assertions.assertThat; @@ -223,6 +236,61 @@ public class ConnectionFieldTypeVisitorTests { } + @Nested + class TypeVisitorTests { + + + @Test // gh-861 + void shouldNotAbortTraverseForSubscriptions() { + String schemaContent = """ + type Query { + greeting: String + } + type Subscription { + puzzles: Puzzle + } + type Puzzle { + title: String! + author: Author + } + type Author { + name: String! + } + """; + + ConnectionFieldTypeVisitor connectionFieldTypeVisitor = ConnectionFieldTypeVisitor.create(List.of(new ListConnectionAdapter())); + TrackingTypeVisitor trackingTypeVisitor = new TrackingTypeVisitor(); + visitSchema(schemaContent, connectionFieldTypeVisitor, trackingTypeVisitor); + assertThat(trackingTypeVisitor.visitedDefinitions).contains("puzzles", "title", "author", "name"); + } + + private static void visitSchema(String schemaContent, GraphQLTypeVisitor... typeVisitors) { + TypeDefinitionRegistry registry = new SchemaParser().parse(schemaContent); + new ConnectionTypeDefinitionConfigurer().configure(registry); + GraphQLSchema schema = new SchemaGenerator().makeExecutableSchema(registry, RuntimeWiring.newRuntimeWiring().build()); + + GraphQLCodeRegistry.Builder outputCodeRegistry = + GraphQLCodeRegistry.newCodeRegistry(schema.getCodeRegistry()); + Map, Object> vars = new HashMap<>(2); + vars.put(GraphQLCodeRegistry.Builder.class, outputCodeRegistry); + vars.put(TypeVisitorHelper.class, TypeVisitorHelper.create(schema)); + + List visitorsToUse = Arrays.asList(typeVisitors); + new SchemaTraverser().depthFirstFullSchema(visitorsToUse, schema, vars); + } + + static class TrackingTypeVisitor extends GraphQLTypeVisitorStub { + + Set visitedDefinitions = new HashSet<>(); + + @Override + public TraversalControl visitGraphQLFieldDefinition(GraphQLFieldDefinition node, TraverserContext context) { + this.visitedDefinitions.add(node.getName()); + return TraversalControl.CONTINUE; + } + } + } + private static class ListConnectionAdapter implements ConnectionAdapter {