From 7eafec25e300bb8e1de26f1b80623cabfec927e3 Mon Sep 17 00:00:00 2001 From: Rossen Stoyanchev Date: Wed, 7 Oct 2020 21:36:34 +0100 Subject: [PATCH] Support for response headers through WebOutput See gh-6 --- .../graphql/WebFluxGraphQLHandler.java | 8 +++- .../graphql/WebInterceptorExecutionChain.java | 2 +- .../graphql/WebMvcGraphQLHandler.java | 12 ++++-- .../springframework/graphql/WebOutput.java | 41 ++++++++++++++++++- .../WebInterceptorExecutionChainTests.java | 7 +++- 5 files changed, 62 insertions(+), 8 deletions(-) diff --git a/spring-graphql-web/src/main/java/org/springframework/graphql/WebFluxGraphQLHandler.java b/spring-graphql-web/src/main/java/org/springframework/graphql/WebFluxGraphQLHandler.java index 55398c17..4a4cd274 100644 --- a/spring-graphql-web/src/main/java/org/springframework/graphql/WebFluxGraphQLHandler.java +++ b/spring-graphql-web/src/main/java/org/springframework/graphql/WebFluxGraphQLHandler.java @@ -51,7 +51,13 @@ public class WebFluxGraphQLHandler implements HandlerFunction { WebInput webInput = new WebInput(request.uri(), request.headers().asHttpHeaders(), body); return this.executionChain.execute(webInput); }) - .flatMap(output -> ServerResponse.ok().bodyValue(output.toSpecification())); + .flatMap(output -> { + ServerResponse.BodyBuilder builder = ServerResponse.ok(); + if (output.getHeaders() != null) { + builder.headers(headers -> headers.putAll(output.getHeaders())); + } + return builder.bodyValue(output.toSpecification()); + }); } } diff --git a/spring-graphql-web/src/main/java/org/springframework/graphql/WebInterceptorExecutionChain.java b/spring-graphql-web/src/main/java/org/springframework/graphql/WebInterceptorExecutionChain.java index dd5168a2..57a94152 100644 --- a/spring-graphql-web/src/main/java/org/springframework/graphql/WebInterceptorExecutionChain.java +++ b/spring-graphql-web/src/main/java/org/springframework/graphql/WebInterceptorExecutionChain.java @@ -64,7 +64,7 @@ class WebInterceptorExecutionChain { } private Mono createOutputChain(Mono resultMono) { - Mono outputMono = resultMono.map(WebOutput::new); + Mono outputMono = resultMono.map((ExecutionResult executionResult) -> new WebOutput(executionResult, null)); for (int i = this.interceptors.size() - 1 ; i >= 0; i--) { WebInterceptor interceptor = this.interceptors.get(i); outputMono = outputMono.flatMap(interceptor::postHandle); diff --git a/spring-graphql-web/src/main/java/org/springframework/graphql/WebMvcGraphQLHandler.java b/spring-graphql-web/src/main/java/org/springframework/graphql/WebMvcGraphQLHandler.java index 16159031..97454b0f 100644 --- a/spring-graphql-web/src/main/java/org/springframework/graphql/WebMvcGraphQLHandler.java +++ b/spring-graphql-web/src/main/java/org/springframework/graphql/WebMvcGraphQLHandler.java @@ -21,7 +21,6 @@ import java.util.Map; import javax.servlet.ServletException; -import graphql.ExecutionResult; import graphql.GraphQL; import reactor.core.publisher.Mono; @@ -60,8 +59,15 @@ public class WebMvcGraphQLHandler implements HandlerFunction { */ public ServerResponse handle(ServerRequest request) throws ServletException { WebInput webInput = new WebInput(request.uri(), request.headers().asHttpHeaders(), readBody(request)); - Mono outputMono = this.executionChain.execute(webInput); - return ServerResponse.ok().body(outputMono.map(ExecutionResult::toSpecification)); + Mono responseMono = this.executionChain.execute(webInput) + .map(output -> { + ServerResponse.BodyBuilder builder = ServerResponse.ok(); + if (output.getHeaders() != null) { + builder.headers(headers -> headers.putAll(output.getHeaders())); + } + return builder.body(output.toSpecification()); + }); + return ServerResponse.async(responseMono); } private static Map readBody(ServerRequest request) throws ServletException { diff --git a/spring-graphql-web/src/main/java/org/springframework/graphql/WebOutput.java b/spring-graphql-web/src/main/java/org/springframework/graphql/WebOutput.java index b0e9e5a1..5fe937a1 100644 --- a/spring-graphql-web/src/main/java/org/springframework/graphql/WebOutput.java +++ b/spring-graphql-web/src/main/java/org/springframework/graphql/WebOutput.java @@ -24,6 +24,7 @@ import graphql.ExecutionResult; import graphql.ExecutionResultImpl; import graphql.GraphQLError; +import org.springframework.http.HttpHeaders; import org.springframework.lang.Nullable; @@ -35,12 +36,16 @@ public class WebOutput implements ExecutionResult { private final ExecutionResult executionResult; + @Nullable + private final HttpHeaders headers; + /** * Create an instance that wraps the given {@link ExecutionResult}. */ - public WebOutput(ExecutionResult executionResult) { + public WebOutput(ExecutionResult executionResult, @Nullable HttpHeaders headers) { this.executionResult = executionResult; + this.headers = headers; } @@ -69,6 +74,14 @@ public class WebOutput implements ExecutionResult { return this.executionResult.toSpecification(); } + /** + * Return headers to be added to the HTTP response. + */ + @Nullable + public HttpHeaders getHeaders() { + return this.headers; + } + /** * Transform this {@code WebOutput} instance through a {@link Builder} and * return a new instance with the modified values. @@ -90,12 +103,18 @@ public class WebOutput implements ExecutionResult { @Nullable private Map extensions; + @Nullable + private HttpHeaders headers; + + public Builder(WebOutput output) { this.data = output.getData(); this.errors = output.getErrors(); this.extensions = output.getExtensions(); + this.headers = output.getHeaders(); } + /** * Set the execution {@link ExecutionResult#getData() data}. */ @@ -120,8 +139,26 @@ public class WebOutput implements ExecutionResult { return this; } + public Builder header(String name, String... values) { + initHeaders(); + for (String value : values) { + this.headers.add(name, value); + } + return this; + } + + public Builder headers(Consumer consumer) { + initHeaders(); + consumer.accept(this.headers); + return this; + } + + private void initHeaders() { + this.headers = (this.headers != null ? this.headers : new HttpHeaders()); + } + public WebOutput build() { - return new WebOutput(new ExecutionResultImpl(this.data, this.errors, this.extensions)); + return new WebOutput(new ExecutionResultImpl(this.data, this.errors, this.extensions), this.headers); } } diff --git a/spring-graphql-web/src/test/java/org/springframework/graphql/WebInterceptorExecutionChainTests.java b/spring-graphql-web/src/test/java/org/springframework/graphql/WebInterceptorExecutionChainTests.java index b1e1ddef..c90fbea4 100644 --- a/spring-graphql-web/src/test/java/org/springframework/graphql/WebInterceptorExecutionChainTests.java +++ b/spring-graphql-web/src/test/java/org/springframework/graphql/WebInterceptorExecutionChainTests.java @@ -69,8 +69,10 @@ public class WebInterceptorExecutionChainTests { assertThat(sb.toString()).isEqualTo(":pre1:pre2:pre3:post3:post2:post1"); assertThat(webOutput.isDataPresent()).isTrue(); + assertThat(webOutput.getHeaders().get("MyHeader")).containsExactly("MyValue3", "MyValue2", "MyValue1"); } + private static GraphQL createGraphQL() throws Exception { RuntimeWiring runtimeWiring = RuntimeWiring.newRuntimeWiring() .type(newTypeWiring("Query").dataFetcher("bookById", GraphQLDataFetchers.getBookByIdDataFetcher())) @@ -105,7 +107,10 @@ public class WebInterceptorExecutionChainTests { @Override public Mono postHandle(WebOutput webOutput) { this.output.append(":post").append(this.index); - return Mono.delay(Duration.ofMillis(50)).map(aLong -> webOutput); + return Mono.delay(Duration.ofMillis(50)) + .map(aLong -> webOutput.transform(builder -> { + builder.header("myHeader", "MyValue" + this.index); + })); } }