diff --git a/spring-websocket/src/main/java/org/springframework/web/socket/WebSocketHttpHeaders.java b/spring-websocket/src/main/java/org/springframework/web/socket/WebSocketHttpHeaders.java index be90695e7b9..84d81fa7b44 100644 --- a/spring-websocket/src/main/java/org/springframework/web/socket/WebSocketHttpHeaders.java +++ b/spring-websocket/src/main/java/org/springframework/web/socket/WebSocketHttpHeaders.java @@ -19,18 +19,15 @@ package org.springframework.web.socket; import java.util.ArrayList; import java.util.Collections; import java.util.List; -import java.util.Map; -import java.util.Set; import org.jspecify.annotations.Nullable; import org.springframework.http.HttpHeaders; import org.springframework.util.CollectionUtils; -import org.springframework.util.MultiValueMap; /** - * An {@link org.springframework.http.HttpHeaders} variant that adds support for - * the HTTP headers defined by the WebSocket specification RFC 6455. + * An {@link HttpHeaders} variant that adds support for the HTTP headers defined + * by the WebSocket specification RFC 6455. * * @author Rossen Stoyanchev * @author Sam Brannen @@ -51,23 +48,32 @@ public class WebSocketHttpHeaders extends HttpHeaders { private static final long serialVersionUID = -6644521016187828916L; - private final HttpHeaders headers; - - /** - * Create a new instance. + * Construct a new, empty {@code WebSocketHttpHeaders} instance. */ public WebSocketHttpHeaders() { - this(new HttpHeaders()); + super(); } /** - * Create an instance that wraps the given pre-existing HttpHeaders and also - * propagate all changes to it. - * @param headers the HTTP headers to wrap + * Construct a new {@code WebSocketHttpHeaders} instance backed by the supplied + * {@code HttpHeaders}. + *

Changes to the {@code WebSocketHttpHeaders} created by this constructor + * will write through to the supplied {@code HttpHeaders}. If you wish to copy + * an existing {@code HttpHeaders} or {@code WebSocketHttpHeaders} instance, + * use {@link #copyOf(HttpHeaders)} instead. Note, however, that {@code copyOf()} + * does not create an instance of {@code WebSocketHttpHeaders}. + *

If the supplied {@code HttpHeaders} instance is a + * {@linkplain #readOnlyHttpHeaders(HttpHeaders) read-only} + * {@code HttpHeaders} wrapper, it will be unwrapped to ensure that the + * {@code WebSocketHttpHeaders} instance created by this constructor is mutable. + * Once the writable instance is mutated, the read-only instance is likely to + * be out of sync and should be discarded. + * @param httpHeaders the headers to expose + * @see #copyOf(HttpHeaders) */ - public WebSocketHttpHeaders(HttpHeaders headers) { - this.headers = headers; + public WebSocketHttpHeaders(HttpHeaders httpHeaders) { + super(httpHeaders); } @@ -182,132 +188,4 @@ public class WebSocketHttpHeaders extends HttpHeaders { return getFirst(SEC_WEBSOCKET_VERSION); } - @Override - public @Nullable List get(String headerName) { - return this.headers.get(headerName); - } - - @Override - public @Nullable String getFirst(String headerName) { - return this.headers.getFirst(headerName); - } - - @Override - public @Nullable List put(String key, List value) { - return this.headers.put(key, value); - } - - @Override - public @Nullable List putIfAbsent(String headerName, List headerValues) { - return this.headers.putIfAbsent(headerName, headerValues); - } - - @Override - public void add(String headerName, @Nullable String headerValue) { - this.headers.add(headerName, headerValue); - } - - /** - * {@inheritDoc} - * @since 7.0 - */ - @Override - public void addAll(String headerName, List headerValues) { - this.headers.addAll(headerName, headerValues); - } - - @Override - public void set(String headerName, @Nullable String headerValue) { - this.headers.set(headerName, headerValue); - } - - @Override - public void setAll(Map values) { - this.headers.setAll(values); - } - - @Override - public Map toSingleValueMap() { - return this.headers.toSingleValueMap(); - } - - /** - * {@inheritDoc} - * @since 7.0 - * @deprecated in favor of {@link #toSingleValueMap()} which performs a copy but - * ensures that collection-iterating methods like {@code entrySet()} are - * case-insensitive - */ - @Override - @Deprecated(since = "7.0", forRemoval = true) - @SuppressWarnings("removal") - public Map asSingleValueMap() { - return this.headers.asSingleValueMap(); - } - - /** - * {@inheritDoc} - * @since 7.0 - * @deprecated This method is provided for backward compatibility with APIs - * that would only accept maps. Generally avoid using HttpHeaders as a Map - * or MultiValueMap. - */ - @Override - @Deprecated(since = "7.0", forRemoval = true) - @SuppressWarnings("removal") - public MultiValueMap asMultiValueMap() { - return this.headers.asMultiValueMap(); - } - - @Override - public boolean containsHeader(String key) { - return this.headers.containsHeader(key); - } - - @Override - public boolean isEmpty() { - return this.headers.isEmpty(); - } - - @Override - public int size() { - return this.headers.size(); - } - - @Override - public @Nullable List remove(String key) { - return this.headers.remove(key); - } - - @Override - public void clear() { - this.headers.clear(); - } - - @Override - public Set headerNames() { - return this.headers.headerNames(); - } - - @Override - public Set>> headerSet() { - return this.headers.headerSet(); - } - - @Override - public boolean equals(@Nullable Object other) { - return (this == other || (other instanceof WebSocketHttpHeaders that && - this.headers.equals(that.headers))); - } - - @Override - public int hashCode() { - return this.headers.hashCode(); - } - - @Override - public String toString() { - return this.headers.toString(); - } - } diff --git a/spring-websocket/src/test/java/org/springframework/web/socket/handler/WebSocketHttpHeadersTests.java b/spring-websocket/src/test/java/org/springframework/web/socket/handler/WebSocketHttpHeadersTests.java index f5c47083e77..b0ec0244d78 100644 --- a/spring-websocket/src/test/java/org/springframework/web/socket/handler/WebSocketHttpHeadersTests.java +++ b/spring-websocket/src/test/java/org/springframework/web/socket/handler/WebSocketHttpHeadersTests.java @@ -25,7 +25,14 @@ import org.springframework.http.HttpHeaders; import org.springframework.web.socket.WebSocketHttpHeaders; import static org.assertj.core.api.Assertions.assertThat; +import static org.assertj.core.api.Assertions.assertThatExceptionOfType; import static org.assertj.core.api.Assertions.entry; +import static org.springframework.http.HttpHeaders.CONTENT_TYPE; +import static org.springframework.http.MediaType.APPLICATION_JSON; +import static org.springframework.http.MediaType.APPLICATION_JSON_VALUE; +import static org.springframework.http.MediaType.TEXT_PLAIN; +import static org.springframework.http.MediaType.TEXT_PLAIN_VALUE; +import static org.springframework.web.socket.WebSocketHttpHeaders.SEC_WEBSOCKET_EXTENSIONS; /** * Tests for {@link WebSocketHttpHeaders}. @@ -35,13 +42,101 @@ import static org.assertj.core.api.Assertions.entry; */ class WebSocketHttpHeadersTests { - private WebSocketHttpHeaders headers = new WebSocketHttpHeaders(); + private final WebSocketHttpHeaders headers = new WebSocketHttpHeaders(); + @Test // gh-35792 + void constructorCopiesHeaders() { + headers.setContentType(APPLICATION_JSON); + assertThat(headers.getContentType()).isEqualTo(APPLICATION_JSON); + + var copy = new WebSocketHttpHeaders(headers); + assertThat(copy.getContentType()).isEqualTo(APPLICATION_JSON); + + copy.setContentType(TEXT_PLAIN); + + assertThat(copy.getContentType()).isEqualTo(TEXT_PLAIN); + // The WebSocketHttpHeaders "copy constructor" creates an WebSocketHttpHeaders + // instance that mutates the state of the original WebSocketHttpHeaders instance. + assertThat(headers.getContentType()).isEqualTo(TEXT_PLAIN); + } + + @Test // gh-35792 + void constructorUnwrapsReadonly() { + headers.setContentType(APPLICATION_JSON); + + var readOnly = HttpHeaders.readOnlyHttpHeaders(headers); + assertThat(readOnly.getContentType()).isEqualTo(APPLICATION_JSON); + + var writable = new WebSocketHttpHeaders(readOnly); + writable.setContentType(TEXT_PLAIN); + + // content-type value is cached by ReadOnlyHttpHeaders + assertThat(readOnly.getContentType()).isEqualTo(APPLICATION_JSON); + assertThat(writable.getContentType()).isEqualTo(TEXT_PLAIN); + assertThat(headers.getContentType()).isEqualTo(TEXT_PLAIN); + } + + @Test // gh-35792 + void copyOf() { + headers.setContentType(APPLICATION_JSON); + headers.set("X-Project", "Spring Framework"); + + assertThat(headers.getContentType()).isEqualTo(APPLICATION_JSON); + assertThat(headers.toSingleValueMap()).containsOnly( + entry(CONTENT_TYPE, APPLICATION_JSON_VALUE), + entry("X-Project", "Spring Framework") + ); + + var copy = HttpHeaders.copyOf(headers); + + assertThat(copy.getContentType()).isEqualTo(APPLICATION_JSON); + assertThat(copy.toSingleValueMap()).containsOnly( + entry(CONTENT_TYPE, APPLICATION_JSON_VALUE), + entry("X-Project", "Spring Framework") + ); + + copy.setContentType(TEXT_PLAIN); + copy.set("X-Project", "Project X"); + + assertThat(headers.getContentType()).isEqualTo(APPLICATION_JSON); + assertThat(headers.toSingleValueMap()).containsOnly( + entry(CONTENT_TYPE, APPLICATION_JSON_VALUE), + entry("X-Project", "Spring Framework") + ); + assertThat(copy.getContentType()).isEqualTo(TEXT_PLAIN); + assertThat(copy.toSingleValueMap()).containsOnly( + entry(CONTENT_TYPE, TEXT_PLAIN_VALUE), + entry("X-Project", "Project X") + ); + } + + @Test // gh-35792 + void readOnlyHttpHeaders() { + headers.setContentType(APPLICATION_JSON); + headers.add("X-Project", "Spring Framework"); + + assertThat(headers.getContentType()).isEqualTo(APPLICATION_JSON); + assertThat(headers.toSingleValueMap()).containsOnly( + entry(CONTENT_TYPE, APPLICATION_JSON_VALUE), + entry("X-Project", "Spring Framework") + ); + + var readOnly = HttpHeaders.readOnlyHttpHeaders(headers); + assertThat(readOnly.getContentType()).isEqualTo(APPLICATION_JSON); + assertThat(readOnly.toSingleValueMap()).containsOnly( + entry(CONTENT_TYPE, APPLICATION_JSON_VALUE), + entry("X-Project", "Spring Framework") + ); + + assertThatExceptionOfType(UnsupportedOperationException.class) + .isThrownBy(() -> readOnly.setContentType(TEXT_PLAIN)); + } + @Test void parseWebSocketExtensions() { var extensions = List.of("x-foo-extension, x-bar-extension", "x-test-extension"); - this.headers.put(WebSocketHttpHeaders.SEC_WEBSOCKET_EXTENSIONS, extensions); + this.headers.put(SEC_WEBSOCKET_EXTENSIONS, extensions); var parsedExtensions = this.headers.getSecWebSocketExtensions(); assertThat(parsedExtensions).hasSize(3);