From 8870cd16a85e8a66ab29bd19f32dd66f3b36b357 Mon Sep 17 00:00:00 2001 From: rstoyanchev Date: Fri, 16 Jun 2023 10:10:24 +0100 Subject: [PATCH] Polishing in ContextDataFetcherDecorator See gh-722 --- .../ContextDataFetcherDecorator.java | 57 +++++++++++-------- 1 file changed, 33 insertions(+), 24 deletions(-) diff --git a/spring-graphql/src/main/java/org/springframework/graphql/execution/ContextDataFetcherDecorator.java b/spring-graphql/src/main/java/org/springframework/graphql/execution/ContextDataFetcherDecorator.java index 20ad46e0..825087ff 100644 --- a/spring-graphql/src/main/java/org/springframework/graphql/execution/ContextDataFetcherDecorator.java +++ b/spring-graphql/src/main/java/org/springframework/graphql/execution/ContextDataFetcherDecorator.java @@ -113,40 +113,49 @@ final class ContextDataFetcherDecorator implements DataFetcher { * data fetchers with the {@link ContextDataFetcherDecorator}. */ static GraphQLTypeVisitor createVisitor(List resolvers) { + return new ContextTypeVisitor(resolvers); + } - SubscriptionExceptionResolver exceptionResolver = new CompositeSubscriptionExceptionResolver(resolvers); - return new GraphQLTypeVisitorStub() { + /** + * Type visitor to apply {@link ContextDataFetcherDecorator}. + */ + private static class ContextTypeVisitor extends GraphQLTypeVisitorStub { - @Override - public TraversalControl visitGraphQLFieldDefinition( - GraphQLFieldDefinition fieldDefinition, TraverserContext context) { + private final SubscriptionExceptionResolver exceptionResolver; - TypeVisitorHelper visitorHelper = context.getVarFromParents(TypeVisitorHelper.class); - GraphQLCodeRegistry.Builder codeRegistry = context.getVarFromParents(GraphQLCodeRegistry.Builder.class); + private ContextTypeVisitor(List resolvers) { + this.exceptionResolver = new CompositeSubscriptionExceptionResolver(resolvers); + } - GraphQLFieldsContainer parent = (GraphQLFieldsContainer) context.getParentNode(); - DataFetcher dataFetcher = codeRegistry.getDataFetcher(parent, fieldDefinition); + @Override + public TraversalControl visitGraphQLFieldDefinition( + GraphQLFieldDefinition fieldDefinition, TraverserContext context) { - if (applyDecorator(dataFetcher)) { - boolean handlesSubscription = visitorHelper.isSubscriptionType(parent); - dataFetcher = new ContextDataFetcherDecorator(dataFetcher, handlesSubscription, exceptionResolver); - codeRegistry.dataFetcher(parent, fieldDefinition, dataFetcher); - } + TypeVisitorHelper visitorHelper = context.getVarFromParents(TypeVisitorHelper.class); + GraphQLCodeRegistry.Builder codeRegistry = context.getVarFromParents(GraphQLCodeRegistry.Builder.class); - return TraversalControl.CONTINUE; + GraphQLFieldsContainer parent = (GraphQLFieldsContainer) context.getParentNode(); + DataFetcher dataFetcher = codeRegistry.getDataFetcher(parent, fieldDefinition); + + if (applyDecorator(dataFetcher)) { + boolean handlesSubscription = visitorHelper.isSubscriptionType(parent); + dataFetcher = new ContextDataFetcherDecorator(dataFetcher, handlesSubscription, exceptionResolver); + codeRegistry.dataFetcher(parent, fieldDefinition, dataFetcher); } - private boolean applyDecorator(DataFetcher dataFetcher) { - Class type = dataFetcher.getClass(); - String packageName = type.getPackage().getName(); - if (packageName.startsWith("graphql.")) { - return (type.getSimpleName().startsWith("DataFetcherFactories") || - packageName.startsWith("graphql.validation")); - } - return true; + return TraversalControl.CONTINUE; + } + + private boolean applyDecorator(DataFetcher dataFetcher) { + Class type = dataFetcher.getClass(); + String packageName = type.getPackage().getName(); + if (packageName.startsWith("graphql.")) { + return (type.getSimpleName().startsWith("DataFetcherFactories") || + packageName.startsWith("graphql.validation")); } - }; + return true; + } } }