From ff9daa93775961ea3c6561c760d40f43103fdbba Mon Sep 17 00:00:00 2001 From: Rossen Stoyanchev Date: Sat, 20 Jun 2020 08:33:26 +0100 Subject: [PATCH] Adjust WebFlux behavior for @RequestPart List List support was added relatively late, incorrectly decoding each part to T which means no way to decode a single part to List and thatis the most common case (vs multipart parts with the same name). This behavior was further misaligned with Spring MVC as well as with the behavior for T[]. Closes gh-22973 --- .../RequestPartMethodArgumentResolver.java | 119 ++++++++++-------- ...equestPartMethodArgumentResolverTests.java | 59 +++++---- 2 files changed, 94 insertions(+), 84 deletions(-) diff --git a/spring-webflux/src/main/java/org/springframework/web/reactive/result/method/annotation/RequestPartMethodArgumentResolver.java b/spring-webflux/src/main/java/org/springframework/web/reactive/result/method/annotation/RequestPartMethodArgumentResolver.java index c14f7a9300..555d5b076c 100644 --- a/spring-webflux/src/main/java/org/springframework/web/reactive/result/method/annotation/RequestPartMethodArgumentResolver.java +++ b/spring-webflux/src/main/java/org/springframework/web/reactive/result/method/annotation/RequestPartMethodArgumentResolver.java @@ -1,5 +1,5 @@ /* - * Copyright 2002-2019 the original author or authors. + * Copyright 2002-2020 the original author or authors. * * Licensed under the Apache License, Version 2.0 (the "License"); * you may not use this file except in compliance with the License. @@ -16,6 +16,7 @@ package org.springframework.web.reactive.result.method.annotation; +import java.util.Collection; import java.util.Collections; import java.util.List; @@ -33,6 +34,7 @@ import org.springframework.http.server.reactive.ServerHttpRequest; import org.springframework.http.server.reactive.ServerHttpRequestDecorator; import org.springframework.lang.Nullable; import org.springframework.util.CollectionUtils; +import org.springframework.util.StringUtils; import org.springframework.web.bind.annotation.RequestPart; import org.springframework.web.reactive.BindingContext; import org.springframework.web.server.ServerWebExchange; @@ -69,77 +71,86 @@ public class RequestPartMethodArgumentResolver extends AbstractMessageReaderArgu RequestPart requestPart = parameter.getParameterAnnotation(RequestPart.class); boolean isRequired = (requestPart == null || requestPart.required()); - String name = getPartName(parameter, requestPart); + Class paramType = parameter.getParameterType(); + Flux partValues = getPartValues(parameter, requestPart, isRequired, exchange); - Flux parts = exchange.getMultipartData() + if (Part.class.isAssignableFrom(paramType)) { + return partValues.next().cast(Object.class); + } + + if (Collection.class.isAssignableFrom(paramType) || List.class.isAssignableFrom(paramType)) { + MethodParameter elementType = parameter.nested(); + if (Part.class.isAssignableFrom(elementType.getNestedParameterType())) { + return partValues.collectList().cast(Object.class); + } + else { + return partValues.next() + .flatMap(part -> decode(part, parameter, bindingContext, exchange, isRequired)) + .defaultIfEmpty(Collections.emptyList()); + } + } + + ReactiveAdapter adapter = getAdapterRegistry().getAdapter(paramType); + if (adapter == null) { + return partValues.next().flatMap(part -> + decode(part, parameter, bindingContext, exchange, isRequired)); + } + + MethodParameter elementType = parameter.nested(); + if (Part.class.isAssignableFrom(elementType.getNestedParameterType())) { + return Mono.just(adapter.fromPublisher(partValues)); + } + + Flux flux = partValues.flatMap(part -> decode(part, elementType, bindingContext, exchange, isRequired)); + return Mono.just(adapter.fromPublisher(flux)); + } + + public Flux getPartValues( + MethodParameter parameter, @Nullable RequestPart requestPart, boolean isRequired, + ServerWebExchange exchange) { + + String name = getPartName(parameter, requestPart); + return exchange.getMultipartData() .flatMapIterable(map -> { List list = map.get(name); if (CollectionUtils.isEmpty(list)) { if (isRequired) { - throw getMissingPartException(name, parameter); + String reason = "Required request part '" + name + "' is not present"; + throw new ServerWebInputException(reason, parameter); } return Collections.emptyList(); } return list; }); - - if (Part.class.isAssignableFrom(parameter.getParameterType())) { - return parts.next().cast(Object.class); - } - - if (List.class.isAssignableFrom(parameter.getParameterType())) { - MethodParameter elementType = parameter.nested(); - if (Part.class.isAssignableFrom(elementType.getNestedParameterType())) { - return parts.collectList().cast(Object.class); - } - else { - return decodePartValues(parts, elementType, bindingContext, exchange, isRequired) - .collectList().cast(Object.class); - } - } - - ReactiveAdapter adapter = getAdapterRegistry().getAdapter(parameter.getParameterType()); - if (adapter != null) { - MethodParameter elementType = parameter.nested(); - return Mono.just(adapter.fromPublisher( - Part.class.isAssignableFrom(elementType.getNestedParameterType()) ? - parts : decodePartValues(parts, elementType, bindingContext, exchange, isRequired))); - } - - return decodePartValues(parts, parameter, bindingContext, exchange, isRequired) - .next().cast(Object.class); } private String getPartName(MethodParameter methodParam, @Nullable RequestPart requestPart) { - String partName = (requestPart != null ? requestPart.name() : ""); - if (partName.isEmpty()) { - partName = methodParam.getParameterName(); - if (partName == null) { - throw new IllegalArgumentException("Request part name for argument type [" + - methodParam.getNestedParameterType().getName() + - "] not specified, and parameter name information not found in class file either."); - } + String name = null; + if (requestPart != null) { + name = requestPart.name(); } - return partName; + if (StringUtils.isEmpty(name)) { + name = methodParam.getParameterName(); + } + if (StringUtils.isEmpty(name)) { + throw new IllegalArgumentException("Request part name for argument type [" + + methodParam.getNestedParameterType().getName() + + "] not specified, and parameter name information not found in class file either."); + } + return name; } - private ServerWebInputException getMissingPartException(String name, MethodParameter param) { - String reason = "Required request part '" + name + "' is not present"; - return new ServerWebInputException(reason, param); - } - - - private Flux decodePartValues(Flux parts, MethodParameter elementType, BindingContext bindingContext, + @SuppressWarnings("unchecked") + private Mono decode( + Part part, MethodParameter elementType, BindingContext bindingContext, ServerWebExchange exchange, boolean isRequired) { - return parts.flatMap(part -> { - ServerHttpRequest partRequest = new PartServerHttpRequest(exchange.getRequest(), part); - ServerWebExchange partExchange = exchange.mutate().request(partRequest).build(); - if (logger.isDebugEnabled()) { - logger.debug(exchange.getLogPrefix() + "Decoding part '" + part.name() + "'"); - } - return readBody(elementType, isRequired, bindingContext, partExchange); - }); + ServerHttpRequest partRequest = new PartServerHttpRequest(exchange.getRequest(), part); + ServerWebExchange partExchange = exchange.mutate().request(partRequest).build(); + if (logger.isDebugEnabled()) { + logger.debug(exchange.getLogPrefix() + "Decoding part '" + part.name() + "'"); + } + return (Mono) readBody(elementType, isRequired, bindingContext, partExchange); } diff --git a/spring-webflux/src/test/java/org/springframework/web/reactive/result/method/annotation/RequestPartMethodArgumentResolverTests.java b/spring-webflux/src/test/java/org/springframework/web/reactive/result/method/annotation/RequestPartMethodArgumentResolverTests.java index 875f355c2c..76f5fafe63 100644 --- a/spring-webflux/src/test/java/org/springframework/web/reactive/result/method/annotation/RequestPartMethodArgumentResolverTests.java +++ b/spring-webflux/src/test/java/org/springframework/web/reactive/result/method/annotation/RequestPartMethodArgumentResolverTests.java @@ -17,6 +17,7 @@ package org.springframework.web.reactive.result.method.annotation; import java.time.Duration; +import java.util.Arrays; import java.util.Collections; import java.util.List; @@ -61,17 +62,17 @@ import static org.springframework.web.testfixture.method.MvcAnnotationPredicates * @author Rossen Stoyanchev * @author Ilya Lukyanovich */ -public class RequestPartMethodArgumentResolverTests { +class RequestPartMethodArgumentResolverTests { private RequestPartMethodArgumentResolver resolver; - private ResolvableMethod testMethod = ResolvableMethod.on(getClass()).named("handle").build(); + private final ResolvableMethod testMethod = ResolvableMethod.on(getClass()).named("handle").build(); private MultipartHttpMessageWriter writer; @BeforeEach - public void setup() throws Exception { + void setup() { List> readers = ServerCodecConfigurer.create().getReaders(); ReactiveAdapterRegistry registry = ReactiveAdapterRegistry.getSharedInstance(); this.resolver = new RequestPartMethodArgumentResolver(readers, registry); @@ -82,7 +83,7 @@ public class RequestPartMethodArgumentResolverTests { @Test - public void supportsParameter() { + void supportsParameter() { MethodParameter param; @@ -110,7 +111,7 @@ public class RequestPartMethodArgumentResolverTests { @Test - public void person() { + void person() { MethodParameter param = this.testMethod.annot(requestPart()).arg(Person.class); MultipartBodyBuilder bodyBuilder = new MultipartBodyBuilder(); bodyBuilder.part("name", new Person("Jones")); @@ -120,11 +121,10 @@ public class RequestPartMethodArgumentResolverTests { } @Test - public void listPerson() { + void listPerson() { MethodParameter param = this.testMethod.annot(requestPart()).arg(List.class, Person.class); MultipartBodyBuilder bodyBuilder = new MultipartBodyBuilder(); - bodyBuilder.part("name", new Person("Jones")); - bodyBuilder.part("name", new Person("James")); + bodyBuilder.part("name", Arrays.asList(new Person("Jones"), new Person("James"))); List actual = resolveArgument(param, bodyBuilder); assertThat(actual.get(0).getName()).isEqualTo("Jones"); @@ -132,7 +132,7 @@ public class RequestPartMethodArgumentResolverTests { } @Test // gh-23060 - public void listPersonNotRequired() { + void listPersonNotRequired() { MethodParameter param = this.testMethod.annot(requestPart().notRequired()).arg(List.class, Person.class); MultipartBodyBuilder bodyBuilder = new MultipartBodyBuilder(); List actual = resolveArgument(param, bodyBuilder); @@ -141,7 +141,7 @@ public class RequestPartMethodArgumentResolverTests { } @Test - public void monoPerson() { + void monoPerson() { MethodParameter param = this.testMethod.annot(requestPart()).arg(Mono.class, Person.class); MultipartBodyBuilder bodyBuilder = new MultipartBodyBuilder(); bodyBuilder.part("name", new Person("Jones")); @@ -151,7 +151,7 @@ public class RequestPartMethodArgumentResolverTests { } @Test // gh-23060 - public void monoPersonNotRequired() { + void monoPersonNotRequired() { MethodParameter param = this.testMethod.annot(requestPart().notRequired()).arg(Mono.class, Person.class); MultipartBodyBuilder bodyBuilder = new MultipartBodyBuilder(); Mono actual = resolveArgument(param, bodyBuilder); @@ -160,7 +160,7 @@ public class RequestPartMethodArgumentResolverTests { } @Test - public void fluxPerson() { + void fluxPerson() { MethodParameter param = this.testMethod.annot(requestPart()).arg(Flux.class, Person.class); MultipartBodyBuilder bodyBuilder = new MultipartBodyBuilder(); bodyBuilder.part("name", new Person("Jones")); @@ -173,7 +173,7 @@ public class RequestPartMethodArgumentResolverTests { } @Test // gh-23060 - public void fluxPersonNotRequired() { + void fluxPersonNotRequired() { MethodParameter param = this.testMethod.annot(requestPart().notRequired()).arg(Flux.class, Person.class); MultipartBodyBuilder bodyBuilder = new MultipartBodyBuilder(); Flux actual = resolveArgument(param, bodyBuilder); @@ -182,7 +182,7 @@ public class RequestPartMethodArgumentResolverTests { } @Test - public void part() { + void part() { MethodParameter param = this.testMethod.annot(requestPart()).arg(Part.class); MultipartBodyBuilder bodyBuilder = new MultipartBodyBuilder(); bodyBuilder.part("name", new Person("Jones")); @@ -193,7 +193,7 @@ public class RequestPartMethodArgumentResolverTests { } @Test - public void listPart() { + void listPart() { MethodParameter param = this.testMethod.annot(requestPart()).arg(List.class, Part.class); MultipartBodyBuilder bodyBuilder = new MultipartBodyBuilder(); bodyBuilder.part("name", new Person("Jones")); @@ -205,7 +205,7 @@ public class RequestPartMethodArgumentResolverTests { } @Test // gh-23060 - public void listPartNotRequired() { + void listPartNotRequired() { MethodParameter param = this.testMethod.annot(requestPart().notRequired()).arg(List.class, Part.class); MultipartBodyBuilder bodyBuilder = new MultipartBodyBuilder(); List actual = resolveArgument(param, bodyBuilder); @@ -214,7 +214,7 @@ public class RequestPartMethodArgumentResolverTests { } @Test - public void monoPart() { + void monoPart() { MethodParameter param = this.testMethod.annot(requestPart()).arg(Mono.class, Part.class); MultipartBodyBuilder bodyBuilder = new MultipartBodyBuilder(); bodyBuilder.part("name", new Person("Jones")); @@ -225,7 +225,7 @@ public class RequestPartMethodArgumentResolverTests { } @Test // gh-23060 - public void monoPartNotRequired() { + void monoPartNotRequired() { MethodParameter param = this.testMethod.annot(requestPart().notRequired()).arg(Mono.class, Part.class); MultipartBodyBuilder bodyBuilder = new MultipartBodyBuilder(); Mono actual = resolveArgument(param, bodyBuilder); @@ -234,7 +234,7 @@ public class RequestPartMethodArgumentResolverTests { } @Test - public void fluxPart() { + void fluxPart() { MethodParameter param = this.testMethod.annot(requestPart()).arg(Flux.class, Part.class); MultipartBodyBuilder bodyBuilder = new MultipartBodyBuilder(); bodyBuilder.part("name", new Person("Jones")); @@ -247,7 +247,7 @@ public class RequestPartMethodArgumentResolverTests { } @Test // gh-23060 - public void fluxPartNotRequired() { + void fluxPartNotRequired() { MethodParameter param = this.testMethod.annot(requestPart().notRequired()).arg(Flux.class, Part.class); MultipartBodyBuilder bodyBuilder = new MultipartBodyBuilder(); Flux actual = resolveArgument(param, bodyBuilder); @@ -256,7 +256,7 @@ public class RequestPartMethodArgumentResolverTests { } @Test - public void personRequired() { + void personRequired() { MethodParameter param = this.testMethod.annot(requestPart()).arg(Person.class); ServerWebExchange exchange = createExchange(new MultipartBodyBuilder()); Mono result = this.resolver.resolveArgument(param, new BindingContext(), exchange); @@ -265,7 +265,7 @@ public class RequestPartMethodArgumentResolverTests { } @Test - public void personNotRequired() { + void personNotRequired() { MethodParameter param = this.testMethod.annot(requestPart().notRequired()).arg(Person.class); ServerWebExchange exchange = createExchange(new MultipartBodyBuilder()); Mono result = this.resolver.resolveArgument(param, new BindingContext(), exchange); @@ -274,7 +274,7 @@ public class RequestPartMethodArgumentResolverTests { } @Test - public void partRequired() { + void partRequired() { MethodParameter param = this.testMethod.annot(requestPart()).arg(Part.class); ServerWebExchange exchange = createExchange(new MultipartBodyBuilder()); Mono result = this.resolver.resolveArgument(param, new BindingContext(), exchange); @@ -283,7 +283,7 @@ public class RequestPartMethodArgumentResolverTests { } @Test - public void partNotRequired() { + void partNotRequired() { MethodParameter param = this.testMethod.annot(requestPart().notRequired()).arg(Part.class); ServerWebExchange exchange = createExchange(new MultipartBodyBuilder()); Mono result = this.resolver.resolveArgument(param, new BindingContext(), exchange); @@ -308,16 +308,15 @@ public class RequestPartMethodArgumentResolverTests { this.writer.write(Mono.just(builder.build()), forClass(MultiValueMap.class), MediaType.MULTIPART_FORM_DATA, clientRequest, Collections.emptyMap()).block(); - MockServerHttpRequest serverRequest = MockServerHttpRequest.post("/") - .contentType(clientRequest.getHeaders().getContentType()) - .body(clientRequest.getBody()); + MediaType contentType = clientRequest.getHeaders().getContentType(); + Flux body = clientRequest.getBody(); + MockServerHttpRequest serverRequest = MockServerHttpRequest.post("/").contentType(contentType).body(body); return MockServerWebExchange.from(serverRequest); } private String partToUtf8String(Part part) { - DataBuffer buffer = DataBufferUtils.join(part.content()).block(); - return buffer.toString(UTF_8); + return DataBufferUtils.join(part.content()).block().toString(UTF_8); } @@ -344,7 +343,7 @@ public class RequestPartMethodArgumentResolverTests { private static class Person { - private String name; + private final String name; @JsonCreator public Person(@JsonProperty("name") String name) {