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 55454b00f20..9b37dc433a2 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 @@ -216,7 +216,7 @@ class JdkClientHttpRequest extends AbstractStreamingClientHttpRequest { * {@code jdk.httpclient.allowRestrictedHeaders} system property. * @see jdk.internal.net.http.common.Utils#getDisallowedHeaders() */ - private static Set disallowedHeaders() { + static Set disallowedHeaders() { TreeSet headers = new TreeSet<>(String.CASE_INSENSITIVE_ORDER); headers.addAll(Set.of("connection", "content-length", "expect", "host", "upgrade")); 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 f44246b3d8a..e41ec8177b9 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 @@ -21,9 +21,6 @@ import java.net.URI; import java.nio.charset.StandardCharsets; import java.time.Duration; -import org.jspecify.annotations.Nullable; -import org.junit.jupiter.api.AfterAll; -import org.junit.jupiter.api.BeforeAll; import org.junit.jupiter.api.Test; import org.junit.jupiter.api.condition.EnabledForJreRange; import org.junit.jupiter.api.condition.JRE; @@ -32,7 +29,6 @@ import org.junit.jupiter.params.provider.ValueSource; import org.springframework.http.HttpMethod; import org.springframework.http.HttpStatus; -import org.springframework.http.HttpStatusCode; import org.springframework.util.StreamUtils; import static org.assertj.core.api.Assertions.assertThat; @@ -45,26 +41,6 @@ import static org.assertj.core.api.Assertions.assertThat; */ class JdkClientHttpRequestFactoryTests extends AbstractHttpRequestFactoryTests { - private static @Nullable String originalPropertyValue; - - - @BeforeAll - static void setProperty() { - originalPropertyValue = System.getProperty("jdk.httpclient.allowRestrictedHeaders"); - System.setProperty("jdk.httpclient.allowRestrictedHeaders", "expect"); - } - - @AfterAll - static void restoreProperty() { - if (originalPropertyValue != null) { - System.setProperty("jdk.httpclient.allowRestrictedHeaders", originalPropertyValue); - } - else { - System.clearProperty("jdk.httpclient.allowRestrictedHeaders"); - } - } - - @Override protected ClientHttpRequestFactory createRequestFactory() { return new JdkClientHttpRequestFactory(); @@ -77,17 +53,6 @@ class JdkClientHttpRequestFactoryTests extends AbstractHttpRequestFactoryTests { assertHttpMethod("patch", HttpMethod.PATCH); } - @Test - void customizeDisallowedHeaders() throws IOException { - URI uri = URI.create(this.baseUrl + "/status/299"); - ClientHttpRequest request = this.factory.createRequest(uri, HttpMethod.PUT); - request.getHeaders().set("Expect", "299"); - - try (ClientHttpResponse response = request.execute()) { - assertThat(response.getStatusCode()).as("Invalid status code").isEqualTo(HttpStatusCode.valueOf(299)); - } - } - @Test // gh-31451 void contentLength0() throws IOException { URI uri = URI.create(this.baseUrl + "/methods/get"); diff --git a/spring-web/src/test/java/org/springframework/http/client/JdkClientHttpRequestTests.java b/spring-web/src/test/java/org/springframework/http/client/JdkClientHttpRequestTests.java index 5b2f0bdc42b..e0e58f9c120 100644 --- a/spring-web/src/test/java/org/springframework/http/client/JdkClientHttpRequestTests.java +++ b/spring-web/src/test/java/org/springframework/http/client/JdkClientHttpRequestTests.java @@ -24,16 +24,19 @@ import java.net.http.HttpRequest; import java.net.http.HttpResponse; import java.net.http.HttpTimeoutException; import java.time.Duration; +import java.util.Set; import java.util.concurrent.CompletableFuture; import java.util.concurrent.ExecutorService; import java.util.concurrent.Executors; +import org.jspecify.annotations.Nullable; import org.junit.jupiter.api.AutoClose; import org.junit.jupiter.api.Test; import org.springframework.http.HttpHeaders; import org.springframework.http.HttpMethod; +import static org.assertj.core.api.Assertions.assertThat; import static org.assertj.core.api.Assertions.assertThatThrownBy; import static org.mockito.ArgumentMatchers.any; import static org.mockito.Mockito.mock; @@ -71,6 +74,53 @@ class JdkClientHttpRequestTests { .isExactlyInstanceOf(IOException.class); } + @Test + void disallowedHeadersByDefault() { + assertThat(disallowedHeaders(null)) + .containsExactlyInAnyOrder("connection", "content-length", "expect", "host", "upgrade"); + } + + @Test + void disallowedHeadersWithSingleHeaderAllowed() { + assertThat(disallowedHeaders("expect")) + .containsExactlyInAnyOrder("connection", "content-length", "host", "upgrade"); + } + + @Test + void disallowedHeadersWithMultipleHeadersAllowed() { + assertThat(disallowedHeaders("expect,host")) + .containsExactlyInAnyOrder("connection", "content-length", "upgrade"); + } + + @Test + void disallowedHeadersAreCaseInsensitive() { + // AssertJ's contains() uses equals(), so go through Set.contains() instead. + assertThat(disallowedHeaders(null).contains("EXPECT")).isTrue(); + assertThat(disallowedHeaders("Expect").contains("expect")).isFalse(); + } + + private static Set disallowedHeaders(@Nullable String allowRestrictedHeaders) { + String key = "jdk.httpclient.allowRestrictedHeaders"; + String original = System.getProperty(key); + try { + if (allowRestrictedHeaders != null) { + System.setProperty(key, allowRestrictedHeaders); + } + else { + System.clearProperty(key); + } + return JdkClientHttpRequest.disallowedHeaders(); + } + finally { + if (original != null) { + System.setProperty(key, original); + } + else { + System.clearProperty(key); + } + } + } + private JdkClientHttpRequest createRequest(Duration timeout) { return new JdkClientHttpRequest(client, URI.create("https://abc.com"), HttpMethod.GET, executor, timeout, false); }