diff --git a/spring-graphql/src/main/java/org/springframework/graphql/client/HttpGraphQlTransport.java b/spring-graphql/src/main/java/org/springframework/graphql/client/HttpGraphQlTransport.java index 400b3be6..a59b9e93 100644 --- a/spring-graphql/src/main/java/org/springframework/graphql/client/HttpGraphQlTransport.java +++ b/spring-graphql/src/main/java/org/springframework/graphql/client/HttpGraphQlTransport.java @@ -24,6 +24,7 @@ import reactor.core.publisher.Mono; import org.springframework.core.ParameterizedTypeReference; import org.springframework.graphql.GraphQlRequest; import org.springframework.graphql.GraphQlResponse; +import org.springframework.http.HttpHeaders; import org.springframework.http.MediaType; import org.springframework.util.Assert; import org.springframework.web.reactive.function.client.WebClient; @@ -46,17 +47,27 @@ final class HttpGraphQlTransport implements GraphQlTransport { private final WebClient webClient; + private final MediaType contentType; + HttpGraphQlTransport(WebClient webClient) { Assert.notNull(webClient, "WebClient is required"); this.webClient = webClient; + this.contentType = initContentType(webClient); + } + + private static MediaType initContentType(WebClient webClient) { + HttpHeaders headers = new HttpHeaders(); + webClient.mutate().defaultHeaders(headers::putAll); + MediaType contentType = headers.getContentType(); + return (contentType != null ? contentType : MediaType.APPLICATION_GRAPHQL); } @Override public Mono execute(GraphQlRequest request) { return this.webClient.post() - .contentType(MediaType.APPLICATION_GRAPHQL) + .contentType(this.contentType) .accept(MediaType.APPLICATION_GRAPHQL, MediaType.APPLICATION_JSON) .bodyValue(request.toMap()) .retrieve() diff --git a/spring-graphql/src/test/java/org/springframework/graphql/client/WebGraphQlClientBuilderTests.java b/spring-graphql/src/test/java/org/springframework/graphql/client/WebGraphQlClientBuilderTests.java index 414d01c9..8e9331a3 100644 --- a/spring-graphql/src/test/java/org/springframework/graphql/client/WebGraphQlClientBuilderTests.java +++ b/spring-graphql/src/test/java/org/springframework/graphql/client/WebGraphQlClientBuilderTests.java @@ -38,6 +38,8 @@ import org.springframework.graphql.server.WebGraphQlRequest; import org.springframework.graphql.server.webflux.GraphQlHttpHandler; import org.springframework.graphql.server.webflux.GraphQlWebSocketHandler; import org.springframework.graphql.support.DocumentSource; +import org.springframework.http.HttpHeaders; +import org.springframework.http.MediaType; import org.springframework.http.codec.ClientCodecConfigurer; import org.springframework.http.codec.json.Jackson2JsonDecoder; import org.springframework.http.server.reactive.HttpHandler; @@ -174,6 +176,28 @@ public class WebGraphQlClientBuilderTests { assertThat(builderSetup.getActualRequest().getUri().toString()).isEqualTo("/graphql%20one"); } + @Test + void contentTypeDefault() { + + HttpBuilderSetup setup = new HttpBuilderSetup(); + setup.initBuilder().build().document(DOCUMENT).execute().block(TIMEOUT); + + WebGraphQlRequest request = setup.getActualRequest(); + assertThat(request.getHeaders().getContentType()).isEqualTo(MediaType.APPLICATION_GRAPHQL); + } + + @Test + void contentTypeOverride() { + + HttpBuilderSetup setup = new HttpBuilderSetup(); + setup.initBuilder().header(HttpHeaders.CONTENT_TYPE, MediaType.APPLICATION_JSON_VALUE).build() + .document(DOCUMENT).execute().block(TIMEOUT); + + WebGraphQlRequest request = setup.getActualRequest(); + assertThat(request.getHeaders().getContentType()).isEqualTo(MediaType.APPLICATION_JSON); + + } + @ParameterizedTest @MethodSource("argumentSource") void codecConfigurerRegistersJsonPathMappingProvider(ClientBuilderSetup builderSetup) {