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 cf4b8c8633e..cee16887c10 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,9 +16,12 @@ package org.springframework.http.server; +import java.util.AbstractSet; import java.util.ArrayList; import java.util.Collection; +import java.util.Collections; import java.util.Enumeration; +import java.util.Iterator; import java.util.LinkedHashMap; import java.util.LinkedHashSet; import java.util.List; @@ -46,7 +49,7 @@ final class ServletRequestHeadersAdapter implements MultiValueMap values = this.request.getHeaders(headerName); if (values.hasMoreElements()) { - List result = new ArrayList<>(); + String value = values.nextElement(); + if (!values.hasMoreElements()) { + return Collections.singletonList(value); + } + List result = new ArrayList<>(4); + result.add(value); while (values.hasMoreElements()) { result.add(values.nextElement()); } @@ -157,12 +165,7 @@ final class ServletRequestHeadersAdapter implements MultiValueMap keySet() { - Set set = new LinkedHashSet<>(); - Enumeration names = this.request.getHeaderNames(); - while (names.hasMoreElements()) { - set.add(names.nextElement()); - } - return set; + return new HeaderNames(); } @Override @@ -183,30 +186,6 @@ final class ServletRequestHeadersAdapter implements 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 public String toString() { return HttpHeaders.formatHeaders(this); @@ -214,13 +193,57 @@ final class ServletRequestHeadersAdapter implements MultiValueMap requestHeaderOverrideWrapper(MultiValueMap headers) { - return new RequestHeaderOverrideWrapper(headers); + static MultiValueMap create(HttpServletRequest request) { + return new RequestHeaderOverrideWrapper(new ServletRequestHeadersAdapter(request)); + } + + + private class HeaderNames extends AbstractSet { + + @Override + public Iterator iterator() { + return new HeaderNamesIterator(request.getHeaderNames()); + } + + @Override + public int size() { + Enumeration names = request.getHeaderNames(); + int size = 0; + while (names.hasMoreElements()) { + names.nextElement(); + size++; + } + return size; + } + } + + + private static final class HeaderNamesIterator implements Iterator { + + private final Enumeration enumeration; + + private HeaderNamesIterator(Enumeration enumeration) { + this.enumeration = enumeration; + } + + @Override + public boolean hasNext() { + return this.enumeration.hasMoreElements(); + } + + @Override + public String next() { + return this.enumeration.nextElement(); + } + + @Override + public void remove() { + throw immutableRequestException(); + } } @@ -351,11 +374,12 @@ final class ServletRequestHeadersAdapter implements MultiValueMap keySet() { - Set set = this.delegate.keySet(); if (this.overrideMap != null) { + Set set = new LinkedHashSet<>(this.delegate.keySet()); set.addAll(this.overrideMap.keySet()); + return set; } - return set; + return this.delegate.keySet(); } @Override @@ -375,7 +399,6 @@ final class ServletRequestHeadersAdapter implements MultiValueMap { public @Nullable List get(Object key) { if (key instanceof String headerName) { Collection values = this.response.getHeaders(headerName); - if (values.isEmpty()) { - return (this.response.containsHeader(headerName) ? Collections.emptyList() : null); + if (!values.isEmpty()) { + return new ArrayList<>(values); } - return new ArrayList<>(values); } return null; } @@ -162,7 +161,7 @@ class ServletResponseHeadersAdapter implements MultiValueMap { @Override public Set keySet() { - return new LinkedHashSet<>(this.response.getHeaderNames()); + return new HeaderNames(); } @Override @@ -205,4 +204,43 @@ class ServletResponseHeadersAdapter implements MultiValueMap { return HttpHeaders.formatHeaders(this); } + + private class HeaderNames extends AbstractSet { + + @Override + public Iterator iterator() { + return new HeaderNamesIterator(response.getHeaderNames()); + } + + @Override + public int size() { + return ServletResponseHeadersAdapter.this.size(); + } + } + + + private static final class HeaderNamesIterator implements Iterator { + + private final Iterator values; + + private HeaderNamesIterator(Collection values) { + this.values = values.iterator(); + } + + @Override + public boolean hasNext() { + return this.values.hasNext(); + } + + @Override + public String next() { + return this.values.next(); + } + + @Override + public void remove() { + throw new UnsupportedOperationException(); + } + } + } 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 149d92f108f..f7567c1f509 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 @@ -22,7 +22,6 @@ import java.io.IOException; import java.io.InputStream; import java.io.OutputStreamWriter; import java.io.Writer; -import java.lang.reflect.Field; import java.net.InetSocketAddress; import java.net.URI; import java.net.URISyntaxException; @@ -42,9 +41,6 @@ import java.util.Map; import java.util.Set; import jakarta.servlet.http.HttpServletRequest; -import jakarta.servlet.http.HttpServletRequestWrapper; -import org.apache.catalina.connector.RequestFacade; -import org.apache.coyote.Request; import org.jspecify.annotations.Nullable; import org.springframework.http.HttpHeaders; @@ -52,10 +48,7 @@ import org.springframework.http.HttpMethod; import org.springframework.http.InvalidMediaTypeException; import org.springframework.http.MediaType; import org.springframework.util.Assert; -import org.springframework.util.ClassUtils; import org.springframework.util.LinkedCaseInsensitiveMap; -import org.springframework.util.MultiValueMap; -import org.springframework.util.ReflectionUtils; import org.springframework.util.StringUtils; /** @@ -70,9 +63,6 @@ public class ServletServerHttpRequest implements ServerHttpRequest { protected static final Charset FORM_CHARSET = StandardCharsets.UTF_8; - private static final boolean TOMCAT_PRESENT = ClassUtils.isPresent( - "org.apache.tomcat.util.http.MimeHeaders", ServletServerHttpRequest.class.getClassLoader()); - private final HttpServletRequest servletRequest; @@ -166,8 +156,7 @@ public class ServletServerHttpRequest implements ServerHttpRequest { @Override public HttpHeaders getHeaders() { if (this.headers == null) { - MultiValueMap headersAdapter = initHeadersMultiValueMap(); - this.headers = new HttpHeaders(headersAdapter); + this.headers = new HttpHeaders(ServletRequestHeadersAdapter.create(this.servletRequest)); // HttpServletRequest exposes some headers as properties: // we should include those if not already present @@ -208,19 +197,6 @@ public class ServletServerHttpRequest implements ServerHttpRequest { return this.headers; } - private MultiValueMap initHeadersMultiValueMap() { - MultiValueMap nativeHeaders = null; - if (TOMCAT_PRESENT) { - nativeHeaders = TomcatInitializer.createTomcatHttpHeaders(this.servletRequest); - } - if (nativeHeaders == null) { - nativeHeaders = new ServletRequestHeadersAdapter(this.servletRequest); - } - return ServletRequestHeadersAdapter.requestHeaderOverrideWrapper(nativeHeaders); - } - - - @Override public @Nullable Principal getPrincipal() { return this.servletRequest.getUserPrincipal(); } @@ -330,41 +306,6 @@ public class ServletServerHttpRequest implements ServerHttpRequest { } - private static final class TomcatInitializer { - - private static final Field COYOTE_REQUEST_FIELD; - - static { - Field field = ReflectionUtils.findField(RequestFacade.class, "request"); - Assert.state(field != null, "Incompatible Tomcat implementation"); - ReflectionUtils.makeAccessible(field); - COYOTE_REQUEST_FIELD = field; - } - - public static @Nullable MultiValueMap createTomcatHttpHeaders(HttpServletRequest servletRequest) { - RequestFacade requestFacade = getRequestFacade(servletRequest); - if (requestFacade == null) { - return null; - } - Object field = ReflectionUtils.getField(COYOTE_REQUEST_FIELD, requestFacade); - Assert.state(field != null, "No Tomcat connector request"); - Request coyoteRequest = ((org.apache.catalina.connector.Request) field).getCoyoteRequest(); - return new TomcatHeadersAdapter(coyoteRequest.getMimeHeaders()); - } - - private static @Nullable RequestFacade getRequestFacade(HttpServletRequest request) { - if (request instanceof RequestFacade facade) { - return facade; - } - else if (request instanceof HttpServletRequestWrapper wrapper) { - HttpServletRequest wrappedRequest = (HttpServletRequest) wrapper.getRequest(); - return getRequestFacade(wrappedRequest); - } - return null; - } - } - - private final class AttributesMap extends AbstractMap { private @Nullable transient Set keySet; diff --git a/spring-web/src/main/java/org/springframework/http/server/ServletServerHttpResponse.java b/spring-web/src/main/java/org/springframework/http/server/ServletServerHttpResponse.java index 843f603de3c..f7bda8df86d 100644 --- a/spring-web/src/main/java/org/springframework/http/server/ServletServerHttpResponse.java +++ b/spring-web/src/main/java/org/springframework/http/server/ServletServerHttpResponse.java @@ -87,13 +87,13 @@ public class ServletServerHttpResponse implements ServerHttpResponse { @Override public OutputStream getBody() throws IOException { this.bodyUsed = true; - writeHeaders(); + this.headersWritten = true; return this.servletResponse.getOutputStream(); } @Override public void flush() throws IOException { - writeHeaders(); + this.headersWritten = true; if (this.bodyUsed) { this.servletResponse.flushBuffer(); } @@ -101,13 +101,7 @@ public class ServletServerHttpResponse implements ServerHttpResponse { @Override public void close() { - writeHeaders(); - } - - private void writeHeaders() { - if (!this.headersWritten) { - this.headersWritten = true; - } + this.headersWritten = true; } }