diff --git a/spring-cloud-netflix-core/src/main/java/org/springframework/cloud/netflix/zuul/filters/pre/FormBodyWrapperFilter.java b/spring-cloud-netflix-core/src/main/java/org/springframework/cloud/netflix/zuul/filters/pre/FormBodyWrapperFilter.java index 5d597986..60103eaa 100644 --- a/spring-cloud-netflix-core/src/main/java/org/springframework/cloud/netflix/zuul/filters/pre/FormBodyWrapperFilter.java +++ b/spring-cloud-netflix-core/src/main/java/org/springframework/cloud/netflix/zuul/filters/pre/FormBodyWrapperFilter.java @@ -30,21 +30,15 @@ import javax.servlet.ServletRequest; import javax.servlet.ServletRequestWrapper; import javax.servlet.http.HttpServletRequest; -import org.springframework.core.io.InputStreamResource; -import org.springframework.core.io.Resource; -import org.springframework.http.HttpEntity; +import org.springframework.cloud.netflix.zuul.util.RequestContentDataExtractor; import org.springframework.http.HttpHeaders; import org.springframework.http.HttpOutputMessage; import org.springframework.http.InvalidMediaTypeException; import org.springframework.http.MediaType; import org.springframework.http.converter.support.AllEncompassingFormHttpMessageConverter; import org.springframework.util.Assert; -import org.springframework.util.LinkedMultiValueMap; import org.springframework.util.MultiValueMap; import org.springframework.util.ReflectionUtils; -import org.springframework.util.StringUtils; -import org.springframework.web.multipart.MultipartFile; -import org.springframework.web.multipart.MultipartRequest; import org.springframework.web.servlet.DispatcherServlet; import com.netflix.zuul.ZuulFilter; @@ -184,36 +178,9 @@ public class FormBodyWrapperFilter extends ZuulFilter { private synchronized void buildContentData() { try { - MultiValueMap builder = new LinkedMultiValueMap(); - Set queryParams = findQueryParams(); - for (Entry entry : this.request.getParameterMap() - .entrySet()) { - if (!queryParams.contains(entry.getKey())) { - for (String value : entry.getValue()) { - builder.add(entry.getKey(), value); - } - } - } - if (this.request instanceof MultipartRequest) { - MultipartRequest multi = (MultipartRequest) this.request; - for (Entry> parts : multi - .getMultiFileMap().entrySet()) { - for (MultipartFile file : parts.getValue()) { - HttpHeaders headers = new HttpHeaders(); - headers.setContentDispositionFormData(file.getName(), - file.getOriginalFilename()); - if (file.getContentType() != null) { - headers.setContentType( - MediaType.valueOf(file.getContentType())); - } - HttpEntity entity = new HttpEntity( - new InputStreamResource(file.getInputStream()), - headers); - builder.add(parts.getKey(), entity); - } - } - } + MultiValueMap builder = RequestContentDataExtractor.extract(this.request); FormHttpOutputMessage data = new FormHttpOutputMessage(); + this.contentType = MediaType.valueOf(this.request.getContentType()); data.getHeaders().setContentType(this.contentType); this.converter.write(builder, this.contentType, data); @@ -227,20 +194,6 @@ public class FormBodyWrapperFilter extends ZuulFilter { } } - private Set findQueryParams() { - Set result = new HashSet<>(); - String query = this.request.getQueryString(); - if (query != null) { - for (String value : StringUtils.tokenizeToStringArray(query, "&")) { - if (value.contains("=")) { - value = value.substring(0, value.indexOf("=")); - } - result.add(value); - } - } - return result; - } - private class FormHttpOutputMessage implements HttpOutputMessage { private HttpHeaders headers = new HttpHeaders(); diff --git a/spring-cloud-netflix-core/src/main/java/org/springframework/cloud/netflix/zuul/util/RequestContentDataExtractor.java b/spring-cloud-netflix-core/src/main/java/org/springframework/cloud/netflix/zuul/util/RequestContentDataExtractor.java new file mode 100644 index 00000000..b2aad040 --- /dev/null +++ b/spring-cloud-netflix-core/src/main/java/org/springframework/cloud/netflix/zuul/util/RequestContentDataExtractor.java @@ -0,0 +1,97 @@ +package org.springframework.cloud.netflix.zuul.util; + +import org.springframework.core.io.InputStreamResource; +import org.springframework.http.HttpEntity; +import org.springframework.http.HttpHeaders; +import org.springframework.http.MediaType; +import org.springframework.util.LinkedMultiValueMap; +import org.springframework.util.MultiValueMap; +import org.springframework.util.StringUtils; +import org.springframework.web.multipart.MultipartFile; +import org.springframework.web.multipart.MultipartHttpServletRequest; + +import javax.servlet.http.HttpServletRequest; +import java.io.IOException; +import java.util.HashSet; +import java.util.List; +import java.util.Map.Entry; +import java.util.Set; + +public class RequestContentDataExtractor { + public static MultiValueMap extract(HttpServletRequest request) throws IOException { + return (request instanceof MultipartHttpServletRequest) ? + extractFromMultipartRequest((MultipartHttpServletRequest) request) : + extractFromRequest(request); + } + + private static MultiValueMap extractFromRequest(HttpServletRequest request) throws IOException { + MultiValueMap builder = new LinkedMultiValueMap<>(); + Set queryParams = findQueryParams(request); + + for (Entry entry : request.getParameterMap().entrySet()) { + String key = entry.getKey(); + + if (!queryParams.contains(key)) { + for (String value : entry.getValue()) { + builder.add(key, value); + } + } + } + + return builder; + } + + private static MultiValueMap extractFromMultipartRequest(MultipartHttpServletRequest request) + throws IOException { + MultiValueMap builder = new LinkedMultiValueMap<>(); + Set queryParams = findQueryParams(request); + + for (Entry entry : request.getParameterMap().entrySet()) { + String key = entry.getKey(); + + if (!queryParams.contains(key)) { + for (String value : entry.getValue()) { + HttpHeaders headers = new HttpHeaders(); + String type = request.getMultipartContentType(key); + + if (type != null) { + headers.setContentType(MediaType.valueOf(type)); + } + + builder.add(key, new HttpEntity<>(value, headers)); + } + } + } + + for (Entry> parts : request.getMultiFileMap().entrySet()) { + for (MultipartFile file : parts.getValue()) { + HttpHeaders headers = new HttpHeaders(); + headers.setContentDispositionFormData(file.getName(), file.getOriginalFilename()); + if (file.getContentType() != null) { + headers.setContentType(MediaType.valueOf(file.getContentType())); + } + + HttpEntity entity = new HttpEntity<>(new InputStreamResource(file.getInputStream()), headers); + builder.add(parts.getKey(), entity); + } + } + + return builder; + } + + private static Set findQueryParams(HttpServletRequest request) { + Set result = new HashSet<>(); + String query = request.getQueryString(); + + if (query != null) { + for (String value : StringUtils.tokenizeToStringArray(query, "&")) { + if (value.contains("=")) { + value = value.substring(0, value.indexOf("=")); + } + result.add(value); + } + } + + return result; + } +} diff --git a/spring-cloud-netflix-core/src/test/java/org/springframework/cloud/netflix/zuul/FormZuulProxyApplicationTests.java b/spring-cloud-netflix-core/src/test/java/org/springframework/cloud/netflix/zuul/FormZuulProxyApplicationTests.java index 75cd7b5a..3d513415 100644 --- a/spring-cloud-netflix-core/src/test/java/org/springframework/cloud/netflix/zuul/FormZuulProxyApplicationTests.java +++ b/spring-cloud-netflix-core/src/test/java/org/springframework/cloud/netflix/zuul/FormZuulProxyApplicationTests.java @@ -22,11 +22,11 @@ import java.util.Map; import org.junit.Before; import org.junit.Test; import org.junit.runner.RunWith; -import org.springframework.beans.factory.annotation.Value; import org.springframework.boot.actuate.trace.InMemoryTraceRepository; import org.springframework.boot.actuate.trace.TraceRepository; import org.springframework.boot.autoconfigure.EnableAutoConfiguration; import org.springframework.boot.builder.SpringApplicationBuilder; +import org.springframework.boot.context.embedded.LocalServerPort; import org.springframework.boot.test.context.SpringBootTest; import org.springframework.boot.test.context.SpringBootTest.WebEnvironment; import org.springframework.boot.test.web.client.TestRestTemplate; @@ -37,7 +37,6 @@ import org.springframework.context.annotation.Bean; import org.springframework.context.annotation.Configuration; import org.springframework.http.HttpEntity; import org.springframework.http.HttpHeaders; -import org.springframework.http.HttpMethod; import org.springframework.http.HttpStatus; import org.springframework.http.MediaType; import org.springframework.http.ResponseEntity; @@ -48,6 +47,7 @@ import org.springframework.util.MultiValueMap; import org.springframework.web.bind.annotation.RequestMapping; import org.springframework.web.bind.annotation.RequestMethod; import org.springframework.web.bind.annotation.RequestParam; +import org.springframework.web.bind.annotation.RequestPart; import org.springframework.web.bind.annotation.RestController; import org.springframework.web.multipart.MultipartFile; @@ -62,124 +62,150 @@ import static org.springframework.util.StreamUtils.copyToString; import lombok.extern.slf4j.Slf4j; +import javax.inject.Inject; +import javax.servlet.http.Part; + @RunWith(SpringJUnit4ClassRunner.class) @SpringBootTest(classes = FormZuulProxyApplication.class, webEnvironment = WebEnvironment.RANDOM_PORT, value = { "zuul.routes.simple:/simple/**" }) @DirtiesContext public class FormZuulProxyApplicationTests { - @Value("${local.server.port}") - private int port; + @Inject + private TestRestTemplate restTemplate; @Before - public void setTestRequestcontext() { - RequestContext context = new RequestContext(); - RequestContext.testSetCurrentContext(context); + public void setTestRequestContext() { + RequestContext.testSetCurrentContext(new RequestContext()); } @Test public void postWithForm() { - MultiValueMap form = new LinkedMultiValueMap(); + MultiValueMap form = new LinkedMultiValueMap<>(); form.set("foo", "bar"); HttpHeaders headers = new HttpHeaders(); headers.setContentType(MediaType.APPLICATION_FORM_URLENCODED); - ResponseEntity result = new TestRestTemplate().exchange( - "http://localhost:" + this.port + "/simple/form", HttpMethod.POST, - new HttpEntity>(form, headers), - String.class); + + ResponseEntity result = sendPost("/simple/form", form, headers); + assertEquals(HttpStatus.OK, result.getStatusCode()); assertEquals("Posted! {foo=[bar]}", result.getBody()); } @Test public void postWithMultipartForm() { - MultiValueMap form = new LinkedMultiValueMap(); + MultiValueMap form = new LinkedMultiValueMap<>(); form.set("foo", "bar"); HttpHeaders headers = new HttpHeaders(); headers.setContentType(MediaType.MULTIPART_FORM_DATA); - ResponseEntity result = new TestRestTemplate().exchange( - "http://localhost:" + this.port + "/simple/form", HttpMethod.POST, - new HttpEntity>(form, headers), - String.class); + + ResponseEntity result = sendPost("/simple/form", form, headers); + assertEquals(HttpStatus.OK, result.getStatusCode()); assertEquals("Posted! {foo=[bar]}", result.getBody()); } @Test public void postWithMultipartFile() { - MultiValueMap form = new LinkedMultiValueMap(); + MultiValueMap form = new LinkedMultiValueMap<>(); + HttpHeaders part = new HttpHeaders(); part.setContentType(MediaType.TEXT_PLAIN); part.setContentDispositionFormData("file", "foo.txt"); - form.set("foo", new HttpEntity("bar".getBytes(), part)); + + form.set("foo", new HttpEntity<>("bar".getBytes(), part)); + HttpHeaders headers = new HttpHeaders(); headers.setContentType(MediaType.MULTIPART_FORM_DATA); - ResponseEntity result = new TestRestTemplate().exchange( - "http://localhost:" + this.port + "/simple/file", HttpMethod.POST, - new HttpEntity>(form, headers), - String.class); + + ResponseEntity result = sendPost("/simple/file", form, headers); + assertEquals(HttpStatus.OK, result.getStatusCode()); assertEquals("Posted! bar", result.getBody()); } @Test public void postWithMultipartFileAndForm() { - MultiValueMap form = new LinkedMultiValueMap(); + MultiValueMap form = new LinkedMultiValueMap<>(); + HttpHeaders part = new HttpHeaders(); part.setContentType(MediaType.TEXT_PLAIN); part.setContentDispositionFormData("file", "foo.txt"); - form.set("foo", new HttpEntity("bar".getBytes(), part)); + form.set("foo", new HttpEntity<>("bar".getBytes(), part)); + form.set("field", "data"); + HttpHeaders headers = new HttpHeaders(); headers.setContentType(MediaType.MULTIPART_FORM_DATA); - ResponseEntity result = new TestRestTemplate().exchange( - "http://localhost:" + this.port + "/simple/fileandform", HttpMethod.POST, - new HttpEntity>(form, headers), - String.class); + + ResponseEntity result = sendPost("/simple/fileandform", form, headers); + assertEquals(HttpStatus.OK, result.getStatusCode()); assertEquals("Posted! bar!field!data", result.getBody()); } @Test - public void postWithUTF8Form() { - MultiValueMap form = new LinkedMultiValueMap(); - form.set("foo", "bar"); + public void postWithMultipartApplicationJson() { + MultiValueMap form = new LinkedMultiValueMap<>(); + + HttpHeaders partHeaders = new HttpHeaders(); + partHeaders.setContentType(MediaType.APPLICATION_JSON); + form.set("field", new HttpEntity<>("{foo=[bar]}", partHeaders)); + HttpHeaders headers = new HttpHeaders(); - headers.setContentType(MediaType.valueOf( - MediaType.APPLICATION_FORM_URLENCODED_VALUE + "; charset=UTF-8")); - ResponseEntity result = new TestRestTemplate().exchange( - "http://localhost:" + this.port + "/simple/form", HttpMethod.POST, - new HttpEntity>(form, headers), - String.class); + headers.setContentType(MediaType.MULTIPART_FORM_DATA); + + ResponseEntity result = sendPost("/simple/json", form, headers); + + assertEquals(HttpStatus.OK, result.getStatusCode()); + assertEquals("Posted! {foo=[bar]} as application/json", result.getBody()); + } + + @Test + public void postWithUTF8Form() { + MultiValueMap form = new LinkedMultiValueMap<>(); + + form.set("foo", "bar"); + + HttpHeaders headers = new HttpHeaders(); + headers.setContentType(MediaType.valueOf(MediaType.APPLICATION_FORM_URLENCODED_VALUE + "; charset=UTF-8")); + + ResponseEntity result = sendPost("/simple/form", form, headers); + assertEquals(HttpStatus.OK, result.getStatusCode()); assertEquals("Posted! {foo=[bar]}", result.getBody()); } @Test public void postWithUrlParams() throws Exception { - MultiValueMap form = new LinkedMultiValueMap(); + MultiValueMap form = new LinkedMultiValueMap<>(); + form.set("foo", "bar"); + HttpHeaders headers = new HttpHeaders(); - headers.setContentType(MediaType.valueOf( - MediaType.APPLICATION_FORM_URLENCODED_VALUE + "; charset=UTF-8")); - ResponseEntity result = new TestRestTemplate().exchange( - "http://localhost:" + this.port + "/simple/form?uriParam=uriValue", - HttpMethod.POST, - new HttpEntity>(form, headers), - String.class); + headers.setContentType(MediaType.valueOf(MediaType.APPLICATION_FORM_URLENCODED_VALUE + "; charset=UTF-8")); + + ResponseEntity result = sendPost("/simple/form?uriParam=uriValue", form, headers); + assertEquals(HttpStatus.OK, result.getStatusCode()); assertEquals("Posted! {uriParam=[uriValue], foo=[bar]}", result.getBody()); } @Test public void getWithUrlParams() throws Exception { - ResponseEntity result = new TestRestTemplate().exchange( - "http://localhost:" + this.port + "/simple/form?uriParam=uriValue", - HttpMethod.GET, null, String.class); + ResponseEntity result = sendGet("/simple/form?uriParam=uriValue"); + assertEquals(HttpStatus.OK, result.getStatusCode()); assertEquals("Posted! {uriParam=[uriValue]}", result.getBody()); } + private ResponseEntity sendPost(String url, MultiValueMap form, HttpHeaders headers) { + return restTemplate.postForEntity(url, new HttpEntity<>(form, headers), String.class); + } + + private ResponseEntity sendGet(String url) { + return restTemplate.getForEntity(url, String.class); + } } // Don't use @SpringBootApplication because we don't want to component scan @@ -220,6 +246,13 @@ class FormZuulProxyApplication { return "Posted! " + copyToString(file.getInputStream(), defaultCharset()) + "!field!" + field; } + @RequestMapping(value = "/json", method = RequestMethod.POST) + public String fileAndJson(@RequestPart Part field) + throws IOException { + + return "Posted! " + copyToString(field.getInputStream(), defaultCharset()) + " as " + field.getContentType(); + } + @Bean public ZuulFilter sampleFilter() { return new ZuulFilter() { @@ -274,7 +307,7 @@ class FormZuulProxyApplication { @Configuration class FormRibbonClientConfiguration { - @Value("${local.server.port}") + @LocalServerPort private int port; @Bean