From 451f2b5a906b90218c66105627bf3016679b3d93 Mon Sep 17 00:00:00 2001 From: rstoyanchev Date: Fri, 3 Oct 2025 11:49:16 +0100 Subject: [PATCH] Polishing in MockMvcClientHttpRequestFactory Replace anonymous subclass with inner. Simplify methods to adapt client-server request and response. --- .../MockMvcClientHttpRequestFactory.java | 113 +++++++++--------- 1 file changed, 59 insertions(+), 54 deletions(-) diff --git a/spring-test/src/main/java/org/springframework/test/web/servlet/client/MockMvcClientHttpRequestFactory.java b/spring-test/src/main/java/org/springframework/test/web/servlet/client/MockMvcClientHttpRequestFactory.java index 509443698eb..3c0edb6975d 100644 --- a/spring-test/src/main/java/org/springframework/test/web/servlet/client/MockMvcClientHttpRequestFactory.java +++ b/spring-test/src/main/java/org/springframework/test/web/servlet/client/MockMvcClientHttpRequestFactory.java @@ -21,7 +21,6 @@ import java.nio.charset.StandardCharsets; import java.util.List; import jakarta.servlet.http.Cookie; -import org.jspecify.annotations.Nullable; import org.springframework.http.HttpHeaders; import org.springframework.http.HttpMethod; @@ -36,12 +35,13 @@ import org.springframework.mock.web.MockHttpServletResponse; import org.springframework.test.web.servlet.MockMvc; import org.springframework.test.web.servlet.request.MockHttpServletRequestBuilder; import org.springframework.util.Assert; +import org.springframework.util.ObjectUtils; import org.springframework.util.StringUtils; import static org.springframework.test.web.servlet.request.MockMvcRequestBuilders.request; /** - * A {@link ClientHttpRequestFactory} for requests executed via {@link MockMvc}. + * {@link ClientHttpRequestFactory} for requests executed via {@link MockMvc}. * * @author Rossen Stoyanchev * @author Rob Worsnop @@ -60,74 +60,79 @@ class MockMvcClientHttpRequestFactory implements ClientHttpRequestFactory { @Override public ClientHttpRequest createRequest(URI uri, HttpMethod httpMethod) { - return new MockClientHttpRequest(httpMethod, uri) { - @Override - public ClientHttpResponse executeInternal() { - return getClientHttpResponse(httpMethod, uri, getHeaders(), getBodyAsBytes()); - } - }; + return new MockMvcClientHttpRequest(httpMethod, uri); } - private ClientHttpResponse getClientHttpResponse( - HttpMethod httpMethod, URI uri, HttpHeaders requestHeaders, byte[] requestBody) { - try { - Cookie[] cookies = parseCookies(requestHeaders.get(HttpHeaders.COOKIE)); - MockHttpServletRequestBuilder requestBuilder = request(httpMethod, uri) - .content(requestBody).headers(requestHeaders); - if (cookies.length > 0) { - requestBuilder.cookie(cookies); + /** + * {@link ClientHttpRequest} that executes via MockMvc. + */ + private class MockMvcClientHttpRequest extends MockClientHttpRequest { + + MockMvcClientHttpRequest(HttpMethod httpMethod, URI uri) { + super(httpMethod, uri); + } + + @Override + public ClientHttpResponse executeInternal() { + try { + MockHttpServletRequestBuilder servletRequestBuilder = request(getMethod(), getURI()) + .headers(getHeaders()) + .content(getBodyAsBytes()); + + addCookies(servletRequestBuilder); + + MockHttpServletResponse servletResponse = MockMvcClientHttpRequestFactory.this.mockMvc + .perform(servletRequestBuilder) + .andReturn() + .getResponse(); + + MockClientHttpResponse clientResponse = new MockClientHttpResponse( + getResponseBody(servletResponse), + HttpStatusCode.valueOf(servletResponse.getStatus())); + + copyHeaders(servletResponse, clientResponse); + + return clientResponse; } - MockHttpServletResponse servletResponse = this.mockMvc - .perform(requestBuilder) - .andReturn() - .getResponse(); + catch (Exception ex) { + byte[] body = ex.toString().getBytes(StandardCharsets.UTF_8); + return new MockClientHttpResponse(body, HttpStatus.INTERNAL_SERVER_ERROR); + } + } - HttpStatusCode status = HttpStatusCode.valueOf(servletResponse.getStatus()); + private void addCookies(MockHttpServletRequestBuilder requestBuilder) { + List values = getHeaders().get(HttpHeaders.COOKIE); + if (!ObjectUtils.isEmpty(values)) { + values.stream() + .flatMap(header -> StringUtils.commaDelimitedListToSet(header).stream()) + .map(value -> { + String[] parts = StringUtils.split(value, "="); + Assert.isTrue(parts != null && parts.length == 2, "Invalid cookie: '" + value + "'"); + return new Cookie(parts[0], parts[1]); + }) + .forEach(requestBuilder::cookie); + } + } + + private static byte[] getResponseBody(MockHttpServletResponse servletResponse) { byte[] body = servletResponse.getContentAsByteArray(); if (body.length == 0) { String error = servletResponse.getErrorMessage(); if (StringUtils.hasLength(error)) { - // sendError message as default body body = error.getBytes(StandardCharsets.UTF_8); } } - - MockClientHttpResponse clientResponse = new MockClientHttpResponse(body, status); - clientResponse.getHeaders().putAll(getResponseHeaders(servletResponse)); - return clientResponse; + return body; } - catch (Exception ex) { - byte[] body = ex.toString().getBytes(StandardCharsets.UTF_8); - return new MockClientHttpResponse(body, HttpStatus.INTERNAL_SERVER_ERROR); - } - } - private static Cookie[] parseCookies(@Nullable List headerValues) { - if (headerValues == null) { - return new Cookie[0]; - } - return headerValues.stream() - .flatMap(header -> StringUtils.commaDelimitedListToSet(header).stream()) - .map(MockMvcClientHttpRequestFactory::parseCookie) - .toArray(Cookie[]::new); - } + private static void copyHeaders( + MockHttpServletResponse servletResponse, MockClientHttpResponse clientResponse) { - private static Cookie parseCookie(String cookie) { - String[] parts = StringUtils.split(cookie, "="); - Assert.isTrue(parts != null && parts.length == 2, "Invalid cookie: '" + cookie + "'"); - return new Cookie(parts[0], parts[1]); - } - - private HttpHeaders getResponseHeaders(MockHttpServletResponse response) { - HttpHeaders headers = new HttpHeaders(); - for (String name : response.getHeaderNames()) { - List values = response.getHeaders(name); - for (String value : values) { - headers.add(name, value); - } + servletResponse.getHeaderNames() + .forEach(name -> servletResponse.getHeaders(name) + .forEach(value -> clientResponse.getHeaders().add(name, value))); } - return headers; } }