diff --git a/spring-graphql/src/main/java/org/springframework/graphql/data/method/InvocableHandlerMethodSupport.java b/spring-graphql/src/main/java/org/springframework/graphql/data/method/InvocableHandlerMethodSupport.java index 4b9e21ad..79ac7f96 100644 --- a/spring-graphql/src/main/java/org/springframework/graphql/data/method/InvocableHandlerMethodSupport.java +++ b/spring-graphql/src/main/java/org/springframework/graphql/data/method/InvocableHandlerMethodSupport.java @@ -26,6 +26,7 @@ import java.util.concurrent.CompletableFuture; import java.util.concurrent.Executor; import graphql.GraphQLContext; +import io.micrometer.context.ContextSnapshot; import reactor.core.publisher.Mono; import org.springframework.core.CoroutinesUtils; @@ -110,32 +111,22 @@ public abstract class InvocableHandlerMethodSupport extends HandlerMethod { Object result; if (this.invokeAsync) { Callable callable = () -> method.invoke(getBean(), argValues); - result = adaptCallable(graphQLContext, callable); + result = adaptCallable(graphQLContext, callable, method, argValues); } else { result = method.invoke(getBean(), argValues); if (this.hasCallableReturnValue && result != null) { - result = adaptCallable(graphQLContext, (Callable) result); + result = adaptCallable(graphQLContext, (Callable) result, method, argValues); } } return result; } catch (IllegalArgumentException ex) { - assertTargetBean(method, getBean(), argValues); - String text = (ex.getMessage() != null) ? ex.getMessage() : "Illegal argument"; - return Mono.error(new IllegalStateException(formatInvokeError(text, argValues), ex)); + return Mono.error(processIllegalArgumentException(argValues, ex, method)); } catch (InvocationTargetException ex) { - // Unwrap for DataFetcherExceptionResolvers ... - Throwable targetException = ex.getTargetException(); - if (targetException instanceof Error || targetException instanceof Exception) { - return Mono.error(targetException); - } - else { - return Mono.error(new IllegalStateException( - formatInvokeError("Invocation failure", argValues), targetException)); - } + return Mono.error(processInvocationTargetException(argValues, ex)); } catch (Throwable ex) { return Mono.error(ex); @@ -155,16 +146,46 @@ public abstract class InvocableHandlerMethodSupport extends HandlerMethod { return result; } - private CompletableFuture adaptCallable(GraphQLContext graphQLContext, Callable result) { - return CompletableFuture.supplyAsync(() -> { + @SuppressWarnings("DataFlowIssue") + private CompletableFuture adaptCallable( + GraphQLContext graphQLContext, Callable result, Method method, Object[] argValues) { + + CompletableFuture future = new CompletableFuture<>(); + this.executor.execute(() -> { try { - return ContextSnapshotFactoryHelper.captureFrom(graphQLContext).wrap(result).call(); + ContextSnapshot snapshot = ContextSnapshotFactoryHelper.captureFrom(graphQLContext); + Object value = snapshot.wrap((Callable) result).call(); + future.complete(value); + } + catch (IllegalArgumentException ex) { + future.completeExceptionally(processIllegalArgumentException(argValues, ex, method)); + } + catch (InvocationTargetException ex) { + future.completeExceptionally(processInvocationTargetException(argValues, ex)); } catch (Exception ex) { - String msg = "Failure in Callable returned from " + getBridgedMethod().toGenericString(); - throw new IllegalStateException(msg, ex); + future.completeExceptionally(ex); } - }, this.executor); + }); + return future; + } + + private IllegalStateException processIllegalArgumentException( + Object[] argValues, IllegalArgumentException ex, Method method) { + + assertTargetBean(method, getBean(), argValues); + String text = (ex.getMessage() != null) ? ex.getMessage() : "Illegal argument"; + return new IllegalStateException(formatInvokeError(text, argValues), ex); + } + + private Throwable processInvocationTargetException(Object[] argValues, InvocationTargetException ex) { + // Unwrap for DataFetcherExceptionResolvers ... + Throwable targetException = ex.getTargetException(); + if (targetException instanceof Error || targetException instanceof Exception) { + return targetException; + } + String message = formatInvokeError("Invocation failure", argValues); + return new IllegalStateException(message, targetException); } /** diff --git a/spring-graphql/src/test/java/org/springframework/graphql/data/method/annotation/support/DataFetcherHandlerMethodTests.java b/spring-graphql/src/test/java/org/springframework/graphql/data/method/annotation/support/DataFetcherHandlerMethodTests.java index 650b0c69..0b12bceb 100644 --- a/spring-graphql/src/test/java/org/springframework/graphql/data/method/annotation/support/DataFetcherHandlerMethodTests.java +++ b/spring-graphql/src/test/java/org/springframework/graphql/data/method/annotation/support/DataFetcherHandlerMethodTests.java @@ -1,5 +1,5 @@ /* - * Copyright 2002-2023 the original author or authors. + * Copyright 2002-2024 the original author or authors. * * Licensed under the Apache License, Version 2.0 (the "License"); * you may not use this file except in compliance with the License. @@ -19,6 +19,7 @@ package org.springframework.graphql.data.method.annotation.support; import java.lang.reflect.Method; import java.util.Collections; +import java.util.Map; import java.util.concurrent.Callable; import java.util.concurrent.CompletableFuture; @@ -26,17 +27,15 @@ import graphql.GraphQLContext; import graphql.schema.DataFetchingEnvironment; import graphql.schema.DataFetchingEnvironmentImpl; import org.junit.jupiter.api.Test; -import org.mockito.Mockito; import reactor.core.publisher.Mono; +import reactor.test.StepVerifier; import org.springframework.core.task.SimpleAsyncTaskExecutor; import org.springframework.graphql.data.GraphQlArgumentBinder; import org.springframework.graphql.data.method.HandlerMethod; -import org.springframework.graphql.data.method.HandlerMethodArgumentResolver; import org.springframework.graphql.data.method.HandlerMethodArgumentResolverComposite; import org.springframework.graphql.data.method.annotation.Argument; import org.springframework.graphql.data.method.annotation.QueryMapping; -import org.springframework.lang.Nullable; import org.springframework.security.authentication.TestingAuthenticationToken; import org.springframework.security.core.annotation.AuthenticationPrincipal; import org.springframework.security.core.context.SecurityContextHolder; @@ -72,17 +71,24 @@ public class DataFetcherHandlerMethodTests { @Test void asyncInvocation() throws Exception { - testAsyncInvocation("handleSync", true); + testAsyncInvocation("handleSync", false, true, "A"); } @Test void asyncInvocationWithCallableReturnValue() throws Exception { - testAsyncInvocation("handleAndReturnCallable", false); + testAsyncInvocation("handleAndReturnCallable", false, false, "A"); } - private static void testAsyncInvocation(String methodName, boolean invokeAsync) throws Exception { + @Test + void asyncInvocationWithCallableReturnValueError() throws Exception { + testAsyncInvocation("handleAndReturnCallable", true, false, "simulated exception"); + } + + private static void testAsyncInvocation( + String methodName, boolean raiseError, boolean invokeAsync, String expected) throws Exception { + HandlerMethodArgumentResolverComposite resolvers = new HandlerMethodArgumentResolverComposite(); - resolvers.addResolver(Mockito.mock(HandlerMethodArgumentResolver.class)); + resolvers.addResolver(new ArgumentMethodArgumentResolver(new GraphQlArgumentBinder())); DataFetcherHandlerMethod handlerMethod = new DataFetcherHandlerMethod( handlerMethodFor(new TestController(), methodName), resolvers, null, @@ -90,6 +96,7 @@ public class DataFetcherHandlerMethodTests { DataFetchingEnvironment environment = DataFetchingEnvironmentImpl .newDataFetchingEnvironment() + .arguments(Map.of("raiseError", raiseError)) // gh-973 .graphQLContext(GraphQLContext.newContext().build()) .build(); @@ -97,7 +104,10 @@ public class DataFetcherHandlerMethodTests { assertThat(result).isInstanceOf(CompletableFuture.class); CompletableFuture future = (CompletableFuture) result; - assertThat(future.get()).isEqualTo("A"); + if (raiseError) { + future = future.handle((s, ex) -> ex.getMessage()); + } + assertThat(future.get()).isEqualTo(expected); } @Test @@ -144,14 +154,17 @@ public class DataFetcherHandlerMethodTests { return "Hello, " + name; } - @Nullable public String handleSync() { return "A"; } - @Nullable - public Callable handleAndReturnCallable() { - return () -> "A"; + public Callable handleAndReturnCallable(@Argument boolean raiseError) { + return () -> { + if (raiseError) { + throw new IllegalStateException("simulated exception"); + } + return "A"; + }; } public CompletableFuture handleAndReturnFuture(@AuthenticationPrincipal User user) {