From aae19d388a9a13c1e8aed5778105df4bf917254b Mon Sep 17 00:00:00 2001 From: rstoyanchev Date: Wed, 12 Jun 2024 08:35:42 +0100 Subject: [PATCH] Refactoring in EntitiesDataFetcher See gh-991 --- .../data/federation/EntitiesDataFetcher.java | 144 ++++++++---------- 1 file changed, 63 insertions(+), 81 deletions(-) diff --git a/spring-graphql/src/main/java/org/springframework/graphql/data/federation/EntitiesDataFetcher.java b/spring-graphql/src/main/java/org/springframework/graphql/data/federation/EntitiesDataFetcher.java index c51ccd25..b6fb28f3 100644 --- a/spring-graphql/src/main/java/org/springframework/graphql/data/federation/EntitiesDataFetcher.java +++ b/spring-graphql/src/main/java/org/springframework/graphql/data/federation/EntitiesDataFetcher.java @@ -40,7 +40,6 @@ import reactor.core.publisher.Mono; import org.springframework.graphql.data.method.annotation.support.HandlerDataFetcherExceptionResolver; import org.springframework.graphql.execution.ErrorType; import org.springframework.lang.Nullable; -import org.springframework.util.Assert; /** * DataFetcher that handles the "_entities" query by invoking @@ -93,12 +92,11 @@ final class EntitiesDataFetcher implements DataFetcher invokeEntityMethod( - DataFetchingEnvironment env, EntityHandlerMethod handlerMethod, Map map, int index) { + DataFetchingEnvironment environment, EntityHandlerMethod handlerMethod, + Map representation, int index) { - return handlerMethod.getEntity(env, map) - .switchIfEmpty(Mono.error(new RepresentationNotResolvedException(map, handlerMethod))) - .onErrorResume((ex) -> resolveException(ex, env, handlerMethod, index)); + return handlerMethod.getEntity(environment, representation) + .switchIfEmpty(Mono.error(new RepresentationNotResolvedException(representation, handlerMethod))) + .onErrorResume((ex) -> resolveException(ex, environment, handlerMethod, index)); + } + + private Mono invokeEntitiesMethod( + DataFetchingEnvironment environment, EntityHandlerMethod handlerMethod, + List> representations, String type) { + + List> typeRepresentations = new ArrayList<>(); + List originalIndexes = new ArrayList<>(); + + for (int i = 0; i < representations.size(); i++) { + Map map = representations.get(i); + if (type.equals(map.get("__typename"))) { + typeRepresentations.add(map); + originalIndexes.add(i); + } + } + + return handlerMethod.getEntities(environment, typeRepresentations) + .mapNotNull((result) -> (((List) result).isEmpty()) ? null : result) + .switchIfEmpty(Mono.defer(() -> { + List> exceptions = new ArrayList<>(originalIndexes.size()); + for (int i = 0; i < originalIndexes.size(); i++) { + exceptions.add(resolveException( + new RepresentationNotResolvedException(typeRepresentations.get(i), handlerMethod), + environment, handlerMethod, originalIndexes.get(i))); + } + return Mono.zip(exceptions, Arrays::asList); + })) + .onErrorResume((ex) -> { + List> list = new ArrayList<>(); + for (Integer index : originalIndexes) { + list.add(resolveException(ex, environment, handlerMethod, index)); + } + return Mono.zip(list, Arrays::asList); + }) + .map((result) -> new EntitiesResultContainer((List) result, originalIndexes)); } private Mono resolveException( @@ -141,8 +176,8 @@ final class EntitiesDataFetcher implements DataFetcher errors = new ArrayList<>(); for (int i = 0; i < entities.size(); i++) { Object entity = entities.get(i); - if (entity instanceof EntityBatchDelegate delegate) { - delegate.processResults(entities, errors); + if (entity instanceof EntitiesResultContainer resultHandler) { + resultHandler.applyResults(entities, errors); } if (entity instanceof ErrorContainer errorContainer) { errors.addAll(errorContainer.errors()); @@ -153,77 +188,6 @@ final class EntitiesDataFetcher implements DataFetcher> filteredRepresentations = new ArrayList<>(); - - private final List indexes = new ArrayList<>(); - - @Nullable - private List resultList; - - EntityBatchDelegate( - DataFetchingEnvironment env, List> allRepresentations, - EntityHandlerMethod handlerMethod, String type) { - - this.environment = env; - this.handlerMethod = handlerMethod; - for (int i = 0; i < allRepresentations.size(); i++) { - Map map = allRepresentations.get(i); - if (type.equals(map.get("__typename"))) { - this.filteredRepresentations.add(map); - this.indexes.add(i); - } - } - } - - Mono invokeEntityBatchMethod() { - return this.handlerMethod.getEntities(this.environment, this.filteredRepresentations) - .mapNotNull((result) -> (((List) result).isEmpty()) ? null : result) - .switchIfEmpty(Mono.defer(this::handleEmptyResult)) - .onErrorResume(this::handleErrorResult) - .map((result) -> { - this.resultList = (List) result; - return this; - }); - } - - Mono handleEmptyResult() { - List> exceptions = new ArrayList<>(this.indexes.size()); - for (int i = 0; i < this.indexes.size(); i++) { - Map map = this.filteredRepresentations.get(i); - Exception ex = new RepresentationNotResolvedException(map, this.handlerMethod); - exceptions.add(resolveException(ex, this.environment, this.handlerMethod, this.indexes.get(i))); - } - return Mono.zip(exceptions, Arrays::asList); - } - - Mono> handleErrorResult(Throwable ex) { - List> list = new ArrayList<>(); - for (Integer index : this.indexes) { - list.add(resolveException(ex, this.environment, this.handlerMethod, index)); - } - return Mono.zip(list, Arrays::asList); - } - - void processResults(List entities, List errors) { - Assert.state(this.resultList != null, "Expected resultList"); - for (int i = 0; i < this.resultList.size(); i++) { - Object entity = this.resultList.get(i); - if (entity instanceof ErrorContainer errorContainer) { - errors.addAll(errorContainer.errors()); - entity = null; - } - entities.set(this.indexes.get(i), entity); - } - } - } - - private static class IndexedDataFetchingEnvironment extends DelegatingDataFetchingEnvironment { private final ExecutionStepInfo executionStepInfo; @@ -242,6 +206,24 @@ final class EntitiesDataFetcher implements DataFetcher results, List originalIndexes) { + + public void applyResults(List entities, List errors) { + for (int i = 0; i < this.results.size(); i++) { + Object result = this.results.get(i); + Integer index = this.originalIndexes.get(i); + if (result instanceof ErrorContainer container) { + errors.addAll(container.errors()); + entities.set(index, null); + } + else { + entities.set(index, result); + } + } + } + } + + private record ErrorContainer(List errors) { ErrorContainer(GraphQLError error) {