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) {