diff --git a/spring-web/src/main/java/org/springframework/http/client/JdkClientHttpRequest.java b/spring-web/src/main/java/org/springframework/http/client/JdkClientHttpRequest.java index 5ae62cbadaa..2bb7f195420 100644 --- a/spring-web/src/main/java/org/springframework/http/client/JdkClientHttpRequest.java +++ b/spring-web/src/main/java/org/springframework/http/client/JdkClientHttpRequest.java @@ -43,6 +43,7 @@ import java.util.concurrent.Executor; import java.util.concurrent.Flow; import java.util.concurrent.TimeUnit; import java.util.concurrent.atomic.AtomicBoolean; +import java.util.function.Consumer; import java.util.zip.GZIPInputStream; import java.util.zip.InflaterInputStream; @@ -113,16 +114,15 @@ class JdkClientHttpRequest extends AbstractStreamingClientHttpRequest { try { HttpRequest request = buildRequest(headers, body); responseFuture = this.httpClient.sendAsync(request, this.compression ? new DecompressingBodyHandler() : HttpResponse.BodyHandlers.ofInputStream()); - if (this.timeout != null) { timeoutHandler = new TimeoutHandler(responseFuture, this.timeout); HttpResponse response = responseFuture.get(); InputStream inputStream = timeoutHandler.wrapInputStream(response); - return new JdkClientHttpResponse(response, inputStream); + return new JdkClientHttpResponse(response, processResponseHeaders(), inputStream); } else { HttpResponse response = responseFuture.get(); - return new JdkClientHttpResponse(response, response.body()); + return new JdkClientHttpResponse(response, processResponseHeaders(), response.body()); } } catch (InterruptedException ex) { @@ -231,6 +231,19 @@ class JdkClientHttpRequest extends AbstractStreamingClientHttpRequest { return Collections.unmodifiableSet(headers); } + private Consumer processResponseHeaders() { + if (this.compression) { + return headers -> { + String encoding = headers.getFirst(HttpHeaders.CONTENT_ENCODING); + if (encoding != null && SUPPORTED_ENCODINGS.contains(encoding)) { + headers.remove(HttpHeaders.CONTENT_ENCODING); + headers.remove(HttpHeaders.CONTENT_LENGTH); + } + }; + } + return headers -> {}; + } + private static final class ByteBufferMapper implements OutputStreamPublisher.ByteMapper { diff --git a/spring-web/src/main/java/org/springframework/http/client/JdkClientHttpResponse.java b/spring-web/src/main/java/org/springframework/http/client/JdkClientHttpResponse.java index 13143e22abd..ce9ffff1dd8 100644 --- a/spring-web/src/main/java/org/springframework/http/client/JdkClientHttpResponse.java +++ b/spring-web/src/main/java/org/springframework/http/client/JdkClientHttpResponse.java @@ -21,17 +21,14 @@ import java.io.InputStream; import java.net.http.HttpClient; import java.net.http.HttpResponse; import java.util.List; -import java.util.Locale; import java.util.Map; +import java.util.function.Consumer; import org.jspecify.annotations.Nullable; import org.springframework.http.HttpHeaders; import org.springframework.http.HttpStatus; import org.springframework.http.HttpStatusCode; -import org.springframework.util.CollectionUtils; -import org.springframework.util.LinkedCaseInsensitiveMap; -import org.springframework.util.MultiValueMap; import org.springframework.util.StreamUtils; /** @@ -50,18 +47,18 @@ class JdkClientHttpResponse implements ClientHttpResponse { private final InputStream body; - public JdkClientHttpResponse(HttpResponse response, @Nullable InputStream body) { + JdkClientHttpResponse(HttpResponse response, Consumer headersConsumer, @Nullable InputStream body) { this.response = response; - this.headers = adaptHeaders(response); + this.headers = adaptHeaders(response, headersConsumer); this.body = (body != null ? body : InputStream.nullInputStream()); } - private static HttpHeaders adaptHeaders(HttpResponse response) { + private static HttpHeaders adaptHeaders(HttpResponse response, Consumer headersConsumer) { Map> rawHeaders = response.headers().map(); - Map> map = new LinkedCaseInsensitiveMap<>(rawHeaders.size(), Locale.ROOT); - MultiValueMap multiValueMap = CollectionUtils.toMultiValueMap(map); - multiValueMap.putAll(rawHeaders); - return HttpHeaders.readOnlyHttpHeaders(multiValueMap); + HttpHeaders headers = new HttpHeaders(); + rawHeaders.forEach(headers::put); + headersConsumer.accept(headers); + return HttpHeaders.readOnlyHttpHeaders(headers); } diff --git a/spring-web/src/test/java/org/springframework/http/client/AbstractMockWebServerTests.java b/spring-web/src/test/java/org/springframework/http/client/AbstractMockWebServerTests.java index 9fbef5328d4..1376fdc2026 100644 --- a/spring-web/src/test/java/org/springframework/http/client/AbstractMockWebServerTests.java +++ b/spring-web/src/test/java/org/springframework/http/client/AbstractMockWebServerTests.java @@ -136,6 +136,7 @@ public abstract class AbstractMockWebServerTests { .body(buffer) .code(200); builder.setHeader(HttpHeaders.CONTENT_ENCODING, encoding); + builder.setHeader(HttpHeaders.CONTENT_LENGTH, buffer.size()); return builder.build(); } return new MockResponse.Builder().code(404).build(); diff --git a/spring-web/src/test/java/org/springframework/http/client/JdkClientHttpRequestFactoryTests.java b/spring-web/src/test/java/org/springframework/http/client/JdkClientHttpRequestFactoryTests.java index c64782f9f64..391e42e4403 100644 --- a/spring-web/src/test/java/org/springframework/http/client/JdkClientHttpRequestFactoryTests.java +++ b/spring-web/src/test/java/org/springframework/http/client/JdkClientHttpRequestFactoryTests.java @@ -135,7 +135,9 @@ class JdkClientHttpRequestFactoryTests extends AbstractHttpRequestFactoryTests { try (ClientHttpResponse response = request.execute()) { assertThat(response.getStatusCode()).as("Invalid response status").isEqualTo(HttpStatus.OK); assertThat(response.getHeaders().getFirst("Content-Encoding")) - .as("Invalid content encoding").isEqualTo("gzip"); + .as("Content Encoding should be removed").isNull(); + assertThat(response.getHeaders().getFirst("Content-Length")) + .as("Content-Length should be removed").isNull(); assertThat(response.getBody()).as("Invalid request body").hasContent("Payload to compress"); } } @@ -150,7 +152,9 @@ class JdkClientHttpRequestFactoryTests extends AbstractHttpRequestFactoryTests { try (ClientHttpResponse response = request.execute()) { assertThat(response.getStatusCode()).as("Invalid response status").isEqualTo(HttpStatus.OK); assertThat(response.getHeaders().getFirst("Content-Encoding")) - .as("Invalid content encoding").isEqualTo("deflate"); + .as("Content Encoding should be removed").isNull(); + assertThat(response.getHeaders().getFirst("Content-Length")) + .as("Content-Length should be removed").isNull(); assertThat(response.getBody()).as("Invalid request body").hasContent("Payload to compress"); } }