diff --git a/spring-web/src/main/java/org/springframework/http/HttpHeaders.java b/spring-web/src/main/java/org/springframework/http/HttpHeaders.java index 24d811e867f..fa37ec81f5f 100644 --- a/spring-web/src/main/java/org/springframework/http/HttpHeaders.java +++ b/spring-web/src/main/java/org/springframework/http/HttpHeaders.java @@ -491,7 +491,12 @@ public class HttpHeaders implements Serializable { */ public static HttpHeaders copyOf(MultiValueMap headers) { HttpHeaders httpHeadersCopy = new HttpHeaders(); - headers.forEach((key, values) -> httpHeadersCopy.put(key, new ArrayList<>(values))); + for (String name : headers.keySet()) { + List values = headers.get(name); + if (values != null) { + httpHeadersCopy.put(name, new ArrayList<>(values)); + } + } return httpHeadersCopy; } @@ -1984,7 +1989,9 @@ public class HttpHeaders implements Serializable { * @see #put(String, List) */ public void putAll(Map> headers) { - headers.forEach(this::put); + for (String name : headers.keySet()) { + put(name, headers.get(name)); + } } /** diff --git a/spring-web/src/main/java/org/springframework/http/ReadOnlyHttpHeaders.java b/spring-web/src/main/java/org/springframework/http/ReadOnlyHttpHeaders.java index 87abd29ae56..d6f6bebfa30 100644 --- a/spring-web/src/main/java/org/springframework/http/ReadOnlyHttpHeaders.java +++ b/spring-web/src/main/java/org/springframework/http/ReadOnlyHttpHeaders.java @@ -181,7 +181,10 @@ class ReadOnlyHttpHeaders extends HttpHeaders { @Override public void forEach(BiConsumer> action) { - this.headers.forEach((k, vs) -> action.accept(k, Collections.unmodifiableList(vs))); + for (String name : this.headers.keySet()) { + List values = this.headers.get(name); + action.accept(name, (values != null ? Collections.unmodifiableList(values) : Collections.emptyList())); + } } } diff --git a/spring-web/src/main/java/org/springframework/http/server/ServletRequestHeadersAdapter.java b/spring-web/src/main/java/org/springframework/http/server/ServletRequestHeadersAdapter.java index 7c1b49bf678..cf4b8c8633e 100644 --- a/spring-web/src/main/java/org/springframework/http/server/ServletRequestHeadersAdapter.java +++ b/spring-web/src/main/java/org/springframework/http/server/ServletRequestHeadersAdapter.java @@ -16,11 +16,9 @@ package org.springframework.http.server; -import java.util.AbstractSet; import java.util.ArrayList; import java.util.Collection; import java.util.Enumeration; -import java.util.Iterator; import java.util.LinkedHashMap; import java.util.LinkedHashSet; import java.util.List; @@ -34,6 +32,7 @@ import org.jspecify.annotations.Nullable; import org.springframework.http.HttpHeaders; import org.springframework.util.CollectionUtils; import org.springframework.util.LinkedCaseInsensitiveMap; +import org.springframework.util.LinkedMultiValueMap; import org.springframework.util.MultiValueMap; /** @@ -59,27 +58,27 @@ final class ServletRequestHeadersAdapter implements MultiValueMap values) { - throw new UnsupportedOperationException(); + throw immutableRequestException(); } @Override public void addAll(MultiValueMap map) { - throw new UnsupportedOperationException(); + throw httpHeadersMapException(); } @Override public void set(String key, @Nullable String value) { - throw new UnsupportedOperationException(); + throw immutableRequestException(); } @Override public void setAll(Map map) { - throw new UnsupportedOperationException(); + throw immutableRequestException(); } @Override @@ -95,12 +94,7 @@ final class ServletRequestHeadersAdapter implements MultiValueMap names = this.request.getHeaderNames(); - Set set = new LinkedHashSet<>(); - while (names.hasMoreElements()) { - set.add(names.nextElement().toLowerCase(Locale.ROOT)); - } - return set.size(); + return keySet().size(); } @Override @@ -123,18 +117,7 @@ final class ServletRequestHeadersAdapter implements MultiValueMap names = this.request.getHeaderNames(); - while (names.hasMoreElements()) { - Enumeration values = this.request.getHeaders(names.nextElement()); - while (values.hasMoreElements()) { - if (text.equals(values.nextElement())) { - return true; - } - } - } - } - return false; + throw httpHeadersMapException(); } @Override @@ -154,22 +137,22 @@ final class ServletRequestHeadersAdapter implements MultiValueMap put(String key, List value) { - throw new UnsupportedOperationException(); + throw immutableRequestException(); } @Override public @Nullable List remove(Object key) { - throw new UnsupportedOperationException(); + throw immutableRequestException(); } @Override public void putAll(Map> map) { - throw new UnsupportedOperationException(); + throw httpHeadersMapException(); } @Override public void clear() { - throw new UnsupportedOperationException(); + throw immutableRequestException(); } @Override @@ -184,45 +167,44 @@ final class ServletRequestHeadersAdapter implements MultiValueMap> values() { - List> allValues = new ArrayList<>(); - Enumeration names = this.request.getHeaderNames(); - while (names.hasMoreElements()) { - String name = names.nextElement(); - List currentValues = new ArrayList<>(); - Enumeration values = this.request.getHeaders(name); - while (values.hasMoreElements()) { - currentValues.add(values.nextElement()); - } - allValues.add(currentValues); - } - return allValues; + throw httpHeadersMapException(); } @Override public Set>> entrySet() { - return new AbstractSet<>() { - @Override - public Iterator>> iterator() { - return new EntryIterator(); - } + throw httpHeadersMapException(); + } - @Override - public int size() { - return ServletRequestHeadersAdapter.this.size(); - } - }; + private static UnsupportedOperationException immutableRequestException() { + return new UnsupportedOperationException("Request headers are immutable"); + } + + private static UnsupportedOperationException httpHeadersMapException() { + return new UnsupportedOperationException("HttpHeaders does not support all Map operations"); } @Override public int hashCode() { - return Map.copyOf(this).hashCode(); + return toMultiValueMap().hashCode(); } @Override public boolean equals(@Nullable Object other) { - return (this == other || - (other instanceof MultiValueMap that && Map.copyOf(this).equals(that))); + return (this == other || (other instanceof MultiValueMap that && toMultiValueMap().equals(that))); + } + + private MultiValueMap toMultiValueMap() { + MultiValueMap map = new LinkedMultiValueMap<>(); + Enumeration names = this.request.getHeaderNames(); + while (names.hasMoreElements()) { + String name = names.nextElement(); + Enumeration values = this.request.getHeaders(name); + while (values.hasMoreElements()) { + map.add(name, values.nextElement()); + } + } + return map; } @Override @@ -237,124 +219,80 @@ final class ServletRequestHeadersAdapter implements MultiValueMap overrideHeadersWrapper(MultiValueMap headers) { - return new OverrideHeaderWrapper(headers); - } - - - private class EntryIterator implements Iterator>> { - - private final Iterator names = ServletRequestHeadersAdapter.this.keySet().iterator(); - - @Override - public boolean hasNext() { - return this.names.hasNext(); - } - - @Override - public Entry> next() { - return new HeaderEntry(this.names.next()); - } - } - - - private final class HeaderEntry implements Entry> { - - private final String key; - - HeaderEntry(String key) { - this.key = key; - } - - @Override - public String getKey() { - return this.key; - } - - @Override - public @Nullable List getValue() { - return get(this.key); - } - - @Override - public @Nullable List setValue(List value) { - List previous = getValue(); - remove(this.key); - addAll(this.key, value); - return previous; - } + static MultiValueMap requestHeaderOverrideWrapper(MultiValueMap headers) { + return new RequestHeaderOverrideWrapper(headers); } /** - * Wrapper that supports optional override values. + * Wrapper that holds override values. */ - private static class OverrideHeaderWrapper implements MultiValueMap { + private static class RequestHeaderOverrideWrapper implements MultiValueMap { private final MultiValueMap delegate; - private @Nullable MultiValueMap overrideHeaders; + private @Nullable MultiValueMap overrideMap; - OverrideHeaderWrapper(MultiValueMap delegate) { + RequestHeaderOverrideWrapper(MultiValueMap delegate) { this.delegate = delegate; } @Override public @Nullable String getFirst(String key) { - String value = (this.overrideHeaders != null ? this.overrideHeaders.getFirst(key) : null); + String value = (this.overrideMap != null ? this.overrideMap.getFirst(key) : null); return (value != null ? value : this.delegate.getFirst(key)); } @Override public void add(String key, @Nullable String value) { - initOverrideHeaders().add(key, value); + initOverrideMap().add(key, value); } @Override public void addAll(String key, List values) { - initOverrideHeaders().addAll(key, values); + initOverrideMap().addAll(key, values); } @Override public void addAll(MultiValueMap map) { - initOverrideHeaders().addAll(map); + throw httpHeadersMapException(); } @Override public void set(String key, @Nullable String value) { - initOverrideHeaders().set(key, value); + initOverrideMap().set(key, value); } @Override public void setAll(Map map) { - initOverrideHeaders().setAll(map); + initOverrideMap().setAll(map); } @Override public Map toSingleValueMap() { Map map = this.delegate.toSingleValueMap(); - if (this.overrideHeaders != null) { - this.overrideHeaders.forEach((key, values) -> map.put(key, values.get(0))); + if (this.overrideMap != null) { + this.overrideMap.forEach((key, values) -> map.put(key, values.get(0))); } return map; } @Override public int size() { - if (this.overrideHeaders == null) { + if (this.overrideMap == null) { return this.delegate.size(); } Set set = new LinkedHashSet<>(); for (String name : this.delegate.keySet()) { set.add(name.toLowerCase(Locale.ROOT)); } - this.overrideHeaders.keySet().forEach(key -> set.add(key.toLowerCase(Locale.ROOT))); + this.overrideMap.keySet().forEach(key -> set.add(key.toLowerCase(Locale.ROOT))); return set.size(); } @Override public boolean isEmpty() { - return (this.delegate.isEmpty() && (this.overrideHeaders == null || this.overrideHeaders.isEmpty())); + return (this.delegate.isEmpty() && (this.overrideMap == null || this.overrideMap.isEmpty())); } @Override @@ -363,8 +301,8 @@ final class ServletRequestHeadersAdapter implements MultiValueMap get(Object key) { if (key instanceof String headerName) { - if (this.overrideHeaders != null) { - List values = this.overrideHeaders.get(headerName); + if (this.overrideMap != null) { + List values = this.overrideMap.get(headerName); if (values != null) { return values; } @@ -399,129 +329,80 @@ final class ServletRequestHeadersAdapter implements MultiValueMap put(String key, List value) { - return initOverrideHeaders().put(key, value); + return initOverrideMap().put(key, value); } @Override public @Nullable List remove(Object key) { - return initOverrideHeaders().remove(key); + return initOverrideMap().remove(key); } @Override public void putAll(Map> map) { - initOverrideHeaders().putAll(map); + throw httpHeadersMapException(); } @Override public void clear() { - initOverrideHeaders().clear(); + if (this.overrideMap != null) { + this.overrideMap.clear(); + } } @Override public Set keySet() { Set set = this.delegate.keySet(); - if (this.overrideHeaders != null) { - set.addAll(this.overrideHeaders.keySet()); + if (this.overrideMap != null) { + set.addAll(this.overrideMap.keySet()); } return set; } @Override public Collection> values() { - List> allValues = new ArrayList<>(); - for (String name : keySet()) { - if (this.overrideHeaders != null && this.overrideHeaders.containsKey(name)) { - allValues.add(this.overrideHeaders.get(name)); - } - else { - allValues.add(this.delegate.get(name)); - } - } - return allValues; + throw httpHeadersMapException(); } @Override public Set>> entrySet() { - return new AbstractSet<>() { - @Override - public Iterator>> iterator() { - return new OverrideHeaderWrapper.EntryIterator(); - } - - @Override - public int size() { - return OverrideHeaderWrapper.this.size(); - } - }; + throw httpHeadersMapException(); } - private MultiValueMap initOverrideHeaders() { - if (this.overrideHeaders == null) { - this.overrideHeaders = CollectionUtils.toMultiValueMap(new LinkedCaseInsensitiveMap<>(8, Locale.ROOT)); + private MultiValueMap initOverrideMap() { + if (this.overrideMap == null) { + this.overrideMap = CollectionUtils.toMultiValueMap(new LinkedCaseInsensitiveMap<>(8, Locale.ROOT)); } - return this.overrideHeaders; + return this.overrideMap; } @Override public int hashCode() { - return Map.copyOf(this).hashCode(); + return toMultiValueMap().hashCode(); } @Override public boolean equals(@Nullable Object other) { - return (this == other || - (other instanceof MultiValueMap that && Map.copyOf(this).equals(that))); + return (this == other || (other instanceof MultiValueMap that && toMultiValueMap().equals(that))); + } + + private MultiValueMap toMultiValueMap() { + MultiValueMap map = new LinkedMultiValueMap<>(); + for (String name : keySet()) { + List values = get(name); + if (values != null) { + for (String value : values) { + map.add(name, value); + } + } + } + return map; } @Override public String toString() { return HttpHeaders.formatHeaders(this); } - - - private class EntryIterator implements Iterator>> { - - private final Iterator names = OverrideHeaderWrapper.this.keySet().iterator(); - - @Override - public boolean hasNext() { - return this.names.hasNext(); - } - - @Override - public Entry> next() { - return new OverrideHeaderWrapper.HeaderEntry(this.names.next()); - } - } - - - private final class HeaderEntry implements Entry> { - - private final String key; - - HeaderEntry(String key) { - this.key = key; - } - - @Override - public String getKey() { - return this.key; - } - - @Override - public @Nullable List getValue() { - return get(this.key); - } - - @Override - public @Nullable List setValue(List value) { - List previous = getValue(); - remove(this.key); - addAll(this.key, value); - return previous; - } - } } } diff --git a/spring-web/src/main/java/org/springframework/http/server/ServletResponseHeadersAdapter.java b/spring-web/src/main/java/org/springframework/http/server/ServletResponseHeadersAdapter.java index d9456bd415e..78c25a9218e 100644 --- a/spring-web/src/main/java/org/springframework/http/server/ServletResponseHeadersAdapter.java +++ b/spring-web/src/main/java/org/springframework/http/server/ServletResponseHeadersAdapter.java @@ -16,11 +16,9 @@ package org.springframework.http.server; -import java.util.AbstractSet; import java.util.ArrayList; import java.util.Collection; import java.util.Collections; -import java.util.Iterator; import java.util.LinkedHashMap; import java.util.LinkedHashSet; import java.util.List; @@ -31,6 +29,7 @@ import jakarta.servlet.http.HttpServletResponse; import org.jspecify.annotations.Nullable; import org.springframework.http.HttpHeaders; +import org.springframework.util.LinkedMultiValueMap; import org.springframework.util.MultiValueMap; /** @@ -61,16 +60,14 @@ class ServletResponseHeadersAdapter implements MultiValueMap { @Override public void addAll(String key, List values) { - values.forEach(value -> this.response.addHeader(key, value)); + for (String value : values) { + this.response.addHeader(key, value); + } } @Override public void addAll(MultiValueMap map) { - for (Entry> entry : map.entrySet()) { - for (String value : entry.getValue()) { - this.response.addHeader(entry.getKey(), value); - } - } + throw httpHeadersUnsupportedOperationException(); } @Override @@ -115,27 +112,17 @@ class ServletResponseHeadersAdapter implements MultiValueMap { @Override public boolean containsValue(Object rawValue) { - if (rawValue instanceof String text) { - for (String name : this.response.getHeaderNames()) { - Collection values = this.response.getHeaders(name); - for (String value : values) { - if (text.equals(value)) { - return true; - } - } - } - } - return false; + throw httpHeadersUnsupportedOperationException(); } @Override public @Nullable List get(Object key) { if (key instanceof String headerName) { Collection values = this.response.getHeaders(headerName); - if (!values.isEmpty()) { - return (values instanceof List ? (List) values : new ArrayList<>(values)); + if (values.isEmpty()) { + return (this.response.containsHeader(headerName) ? Collections.emptyList() : null); } - return (this.response.containsHeader(headerName) ? Collections.emptyList() : null); + return new ArrayList<>(values); } return null; } @@ -163,12 +150,7 @@ class ServletResponseHeadersAdapter implements MultiValueMap { @Override public void putAll(Map> map) { - for (Entry> entry : map.entrySet()) { - this.response.setHeader(entry.getKey(), null); - for (String value : entry.getValue()) { - this.response.addHeader(entry.getKey(), value); - } - } + throw httpHeadersUnsupportedOperationException(); } @Override @@ -185,38 +167,37 @@ class ServletResponseHeadersAdapter implements MultiValueMap { @Override public Collection> values() { - List> allValues = new ArrayList<>(); - for (String name : this.response.getHeaderNames()) { - allValues.add(new ArrayList<>(this.response.getHeaders(name))); - } - return allValues; + throw httpHeadersUnsupportedOperationException(); } @Override public Set>> entrySet() { - return new AbstractSet<>() { - @Override - public Iterator>> iterator() { - return new EntryIterator(); - } + throw httpHeadersUnsupportedOperationException(); + } - @Override - public int size() { - return ServletResponseHeadersAdapter.this.size(); - } - }; + private static UnsupportedOperationException httpHeadersUnsupportedOperationException() { + return new UnsupportedOperationException("HttpHeaders does not support all Map operations"); } @Override public int hashCode() { - return Map.copyOf(this).hashCode(); + return toMultiValueMap().hashCode(); } @Override public boolean equals(@Nullable Object other) { - return (this == other || - (other instanceof MultiValueMap that && Map.copyOf(this).equals(that))); + return (this == other || (other instanceof MultiValueMap that && toMultiValueMap().equals(that))); + } + + private MultiValueMap toMultiValueMap() { + MultiValueMap map = new LinkedMultiValueMap<>(); + for (String name : this.response.getHeaderNames()) { + for (String value : this.response.getHeaders(name)) { + map.add(name, value); + } + } + return map; } @Override @@ -224,46 +205,4 @@ class ServletResponseHeadersAdapter implements MultiValueMap { return HttpHeaders.formatHeaders(this); } - - private class EntryIterator implements Iterator>> { - - private final Iterator names = - ServletResponseHeadersAdapter.this.response.getHeaderNames().iterator(); - - @Override - public boolean hasNext() { - return this.names.hasNext(); - } - - @Override - public Entry> next() { - return new HeaderEntry(this.names.next()); - } - } - - - private final class HeaderEntry implements Entry> { - - private final String key; - - HeaderEntry(String key) { - this.key = key; - } - - @Override - public String getKey() { - return this.key; - } - - @Override - public @Nullable List getValue() { - return ServletResponseHeadersAdapter.this.get(this.key); - } - - @Override - public @Nullable List setValue(List values) { - return ServletResponseHeadersAdapter.this.put(this.key, values); - } - } - } diff --git a/spring-web/src/main/java/org/springframework/http/server/ServletServerHttpRequest.java b/spring-web/src/main/java/org/springframework/http/server/ServletServerHttpRequest.java index a2724f51b06..149d92f108f 100644 --- a/spring-web/src/main/java/org/springframework/http/server/ServletServerHttpRequest.java +++ b/spring-web/src/main/java/org/springframework/http/server/ServletServerHttpRequest.java @@ -216,7 +216,7 @@ public class ServletServerHttpRequest implements ServerHttpRequest { if (nativeHeaders == null) { nativeHeaders = new ServletRequestHeadersAdapter(this.servletRequest); } - return ServletRequestHeadersAdapter.overrideHeadersWrapper(nativeHeaders); + return ServletRequestHeadersAdapter.requestHeaderOverrideWrapper(nativeHeaders); }