From 94e2f49e9fbb0679e4053f76367f4fdd0f7fb931 Mon Sep 17 00:00:00 2001 From: Brian Clozel Date: Tue, 31 Mar 2026 19:13:06 +0200 Subject: [PATCH] Read multipart requests from RestTestClient Prior to this commit, the `RestTestClient` MockMvc integration would support transforming client requests into MockMvc requests. `RestTestClient` can serialize `MultiValueMap` request bodies as multipart requests. In this case, the `MockMvcClientHttpRequestFactory` would only read the body as byte stream and would not turn this into a proper `MockMultipartHttpServletRequestBuilder`. This commit uses the new `MultipartHttpMessageConverter` to parse the request payload as `MockPart` instances to be added to the MockMvc requests. Closes gh-35569 --- .../MockMvcClientHttpRequestFactory.java | 83 ++++++++++++++++--- .../client/MockMvcRestTestClientTests.java | 40 ++++++++- 2 files changed, 110 insertions(+), 13 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 acb23a32444..218b4ebcb8d 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 @@ -16,28 +16,40 @@ package org.springframework.test.web.servlet.client; +import java.io.IOException; +import java.io.InputStream; import java.net.URI; import java.nio.charset.StandardCharsets; import java.util.List; import jakarta.servlet.http.Cookie; +import org.springframework.core.ResolvableType; import org.springframework.http.HttpHeaders; import org.springframework.http.HttpMethod; import org.springframework.http.HttpStatus; import org.springframework.http.HttpStatusCode; +import org.springframework.http.MediaType; import org.springframework.http.client.ClientHttpRequest; import org.springframework.http.client.ClientHttpRequestFactory; import org.springframework.http.client.ClientHttpResponse; +import org.springframework.http.converter.multipart.FilePart; +import org.springframework.http.converter.multipart.MultipartHttpMessageConverter; +import org.springframework.http.converter.multipart.Part; +import org.springframework.mock.http.MockHttpInputMessage; import org.springframework.mock.http.client.MockClientHttpRequest; import org.springframework.mock.http.client.MockClientHttpResponse; import org.springframework.mock.web.MockHttpServletResponse; +import org.springframework.mock.web.MockPart; import org.springframework.test.web.servlet.MockMvc; -import org.springframework.test.web.servlet.request.MockHttpServletRequestBuilder; +import org.springframework.test.web.servlet.request.AbstractMockHttpServletRequestBuilder; +import org.springframework.test.web.servlet.request.MockMultipartHttpServletRequestBuilder; import org.springframework.util.Assert; +import org.springframework.util.MultiValueMap; import org.springframework.util.ObjectUtils; import org.springframework.util.StringUtils; +import static org.springframework.test.web.servlet.request.MockMvcRequestBuilders.multipart; import static org.springframework.test.web.servlet.request.MockMvcRequestBuilders.request; /** @@ -45,10 +57,16 @@ import static org.springframework.test.web.servlet.request.MockMvcRequestBuilder * * @author Rossen Stoyanchev * @author Rob Worsnop + * @author Brian Clozel * @since 7.0 */ public class MockMvcClientHttpRequestFactory implements ClientHttpRequestFactory { + private static final ResolvableType MULTIVALUEMAP_TYPE = ResolvableType.forClassWithGenerics(MultiValueMap.class, + String.class, Part.class); + + private static final MultipartHttpMessageConverter MULTIPART_CONVERTER = new MultipartHttpMessageConverter(); + private final MockMvc mockMvc; @@ -79,11 +97,8 @@ public class MockMvcClientHttpRequestFactory implements ClientHttpRequestFactory @Override public ClientHttpResponse executeInternal() { try { - MockHttpServletRequestBuilder servletRequestBuilder = request(getMethod(), getURI()) - .headers(getHeaders()) - .content(getBodyAsBytes()); - - addCookies(servletRequestBuilder); + AbstractMockHttpServletRequestBuilder servletRequestBuilder = + adaptRequest(getMethod(), getURI(), getHeaders(), getBodyAsBytes()); MockHttpServletResponse servletResponse = MockMvcClientHttpRequestFactory.this.mockMvc .perform(servletRequestBuilder) @@ -104,15 +119,61 @@ public class MockMvcClientHttpRequestFactory implements ClientHttpRequestFactory } } - private void addCookies(MockHttpServletRequestBuilder requestBuilder) { - List values = getHeaders().get(HttpHeaders.COOKIE); + private AbstractMockHttpServletRequestBuilder adaptRequest( + HttpMethod httpMethod, URI uri, HttpHeaders headers, byte[] bytes) throws IOException { + + String contentType = headers.getFirst(HttpHeaders.CONTENT_TYPE); + AbstractMockHttpServletRequestBuilder requestBuilder; + + if (StringUtils.hasLength(contentType) && + MediaType.MULTIPART_FORM_DATA.includes(MediaType.parseMediaType(contentType))) { + + MockMultipartHttpServletRequestBuilder multipartRequestBuilder = multipart(httpMethod, uri); + Assert.notNull(bytes, "No multipart content"); + MockHttpInputMessage inputMessage = new MockHttpInputMessage(bytes); + inputMessage.getHeaders().putAll(headers); + + MultiValueMap parts = MULTIPART_CONVERTER.read(MULTIVALUEMAP_TYPE, inputMessage, null); + for (List partValues : parts.values()) { + for (Part part : partValues) { + parsePart(part, multipartRequestBuilder); + } + } + requestBuilder = multipartRequestBuilder; + } + else { + requestBuilder = request(httpMethod, uri); + if (!ObjectUtils.isEmpty(bytes)) { + requestBuilder.content(bytes); + } + } + + requestBuilder.headers(headers); + addCookies(headers, requestBuilder); + + return requestBuilder; + } + + private void parsePart(Part part, MockMultipartHttpServletRequestBuilder multipartRequestBuilder) throws IOException { + try (InputStream content = part.content()) { + byte[] partBytes = content.readAllBytes(); + MockPart mockPart = (part instanceof FilePart filePart ? + new MockPart(part.name(), filePart.filename(), partBytes) : + new MockPart(part.name(), partBytes)); + mockPart.getHeaders().putAll(part.headers()); + multipartRequestBuilder.part(mockPart); + } + } + + private void addCookies(HttpHeaders headers, AbstractMockHttpServletRequestBuilder requestBuilder) { + List values = headers.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]); + String[] cookieParts = StringUtils.split(value, "="); + Assert.isTrue(cookieParts != null && cookieParts.length == 2, "Invalid cookie: '" + value + "'"); + return new Cookie(cookieParts[0], cookieParts[1]); }) .forEach(requestBuilder::cookie); } diff --git a/spring-test/src/test/java/org/springframework/test/web/servlet/client/MockMvcRestTestClientTests.java b/spring-test/src/test/java/org/springframework/test/web/servlet/client/MockMvcRestTestClientTests.java index c1cc6aa0c4e..5a88f1b147e 100644 --- a/spring-test/src/test/java/org/springframework/test/web/servlet/client/MockMvcRestTestClientTests.java +++ b/spring-test/src/test/java/org/springframework/test/web/servlet/client/MockMvcRestTestClientTests.java @@ -17,6 +17,7 @@ package org.springframework.test.web.servlet.client; import java.io.IOException; +import java.nio.charset.StandardCharsets; import jakarta.servlet.http.Cookie; import jakarta.servlet.http.HttpServletResponse; @@ -26,6 +27,7 @@ import org.springframework.core.ParameterizedTypeReference; import org.springframework.core.io.ByteArrayResource; import org.springframework.core.io.Resource; import org.springframework.http.MediaType; +import org.springframework.http.ResponseEntity; import org.springframework.http.converter.multipart.FilePart; import org.springframework.http.converter.multipart.FormFieldPart; import org.springframework.http.converter.multipart.Part; @@ -35,6 +37,8 @@ import org.springframework.util.LinkedMultiValueMap; import org.springframework.util.MultiValueMap; import org.springframework.web.bind.annotation.CookieValue; import org.springframework.web.bind.annotation.GetMapping; +import org.springframework.web.bind.annotation.PostMapping; +import org.springframework.web.bind.annotation.RequestParam; import org.springframework.web.bind.annotation.RestController; import static org.assertj.core.api.Assertions.assertThat; @@ -106,6 +110,25 @@ class MockMvcRestTestClientTests { }); } + @Test + void writeMultipart() { + MultiValueMap parts = new LinkedMultiValueMap<>(); + parts.add("text1", "value1"); + parts.add("file1", new ByteArrayResource("filecontent1".getBytes()) { + @Override + public String getFilename() { + return "spring.txt"; + } + }); + + client.post() + .uri("/multipart") + .contentType(MediaType.MULTIPART_FORM_DATA) + .body(parts) + .exchange() + .expectStatus().isOk(); + } + @RestController static class TestController { @@ -126,8 +149,8 @@ class MockMvcRestTestClientTests { response.getWriter().write("some really bad request"); } - @GetMapping(value = "/multipart", produces = MediaType.MULTIPART_FORM_DATA_VALUE) - MultiValueMap multipart() { + @GetMapping(path = "/multipart", produces = MediaType.MULTIPART_FORM_DATA_VALUE) + MultiValueMap writeMultipart() { MultiValueMap parts = new LinkedMultiValueMap<>(); parts.add("text1", "a"); parts.add("text2", "b"); @@ -140,6 +163,19 @@ class MockMvcRestTestClientTests { parts.add("file1", resource); return parts; } + + @PostMapping(path = "/multipart", consumes = MediaType.MULTIPART_FORM_DATA_VALUE) + ResponseEntity readMultipart(@RequestParam MultiValueMap parts) throws Exception { + assertThat(parts.keySet()).containsOnly("text1", "file1"); + jakarta.servlet.http.Part text1 = parts.get("text1").get(0); + assertThat(text1.getName()).isEqualTo("text1"); + assertThat(text1.getInputStream()).asString(StandardCharsets.UTF_8).isEqualTo("value1"); + jakarta.servlet.http.Part file1 = parts.get("file1").get(0); + assertThat(file1.getName()).isEqualTo("file1"); + assertThat(file1.getSubmittedFileName()).isEqualTo("spring.txt"); + assertThat(file1.getInputStream()).asString(StandardCharsets.UTF_8).isEqualTo("filecontent1"); + return ResponseEntity.ok().build(); + } } }