Support Kotlin parameter default values in handler methods

This commit adds support for Kotlin parameter default values
in handler methods. It allows to write:
@RequestParam value: String = "default"
as an alternative to:
@RequestParam(defaultValue = "default") value: String

Both Spring MVC and WebFlux are supported, including on
suspending functions.

Closes gh-21139
This commit is contained in:
Sébastien Deleuze
2023-06-21 18:49:11 +02:00
parent 254fb39567
commit f06cf21341
12 changed files with 679 additions and 41 deletions

View File

@@ -18,6 +18,7 @@ package org.springframework.core;
import java.lang.reflect.InvocationTargetException;
import java.lang.reflect.Method;
import java.util.Map;
import java.util.Objects;
import kotlin.Unit;
@@ -26,6 +27,7 @@ import kotlin.jvm.JvmClassMappingKt;
import kotlin.reflect.KClass;
import kotlin.reflect.KClassifier;
import kotlin.reflect.KFunction;
import kotlin.reflect.KParameter;
import kotlin.reflect.full.KCallables;
import kotlin.reflect.jvm.KCallablesJvm;
import kotlin.reflect.jvm.ReflectJvmMapping;
@@ -42,6 +44,7 @@ import reactor.core.publisher.Flux;
import reactor.core.publisher.Mono;
import org.springframework.util.Assert;
import org.springframework.util.CollectionUtils;
/**
* Utilities for working with Kotlin Coroutines.
@@ -104,8 +107,22 @@ public abstract class CoroutinesUtils {
if (method.isAccessible() && !KCallablesJvm.isAccessible(function)) {
KCallablesJvm.setAccessible(function, true);
}
Mono<Object> mono = MonoKt.mono(context, (scope, continuation) ->
KCallables.callSuspend(function, getSuspendedFunctionArgs(method, target, args), continuation))
Mono<Object> mono = MonoKt.mono(context, (scope, continuation) -> {
Map<KParameter, Object> argMap = CollectionUtils.newHashMap(args.length + 1);
int index = 0;
for (KParameter parameter : function.getParameters()) {
switch (parameter.getKind()) {
case INSTANCE -> argMap.put(parameter, target);
case VALUE -> {
if (!parameter.isOptional() || args[index] != null) {
argMap.put(parameter, args[index]);
}
index++;
}
}
}
return KCallables.callSuspendBy(function, argMap, continuation);
})
.filter(result -> !Objects.equals(result, Unit.INSTANCE))
.onErrorMap(InvocationTargetException.class, InvocationTargetException::getTargetException);
@@ -125,14 +142,6 @@ public abstract class CoroutinesUtils {
return mono;
}
private static Object[] getSuspendedFunctionArgs(Method method, Object target, Object... args) {
int length = (args.length == method.getParameterCount() - 1 ? args.length + 1 : args.length);
Object[] functionArgs = new Object[length];
functionArgs[0] = target;
System.arraycopy(args, 0, functionArgs, 1, length - 1);
return functionArgs;
}
private static Flux<?> asFlux(Object flow) {
return ReactorFlowKt.asFlux(((Flow<?>) flow));
}

View File

@@ -1,5 +1,5 @@
/*
* Copyright 2002-2019 the original author or authors.
* Copyright 2002-2023 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.
@@ -38,6 +38,8 @@ class KotlinMethodParameterTests {
private val nonNullableMethod = javaClass.getMethod("nonNullable", String::class.java)
private val withDefaultValueMethod: Method = javaClass.getMethod("withDefaultValue", String::class.java)
private val innerClassConstructor = InnerClass::class.java.getConstructor(KotlinMethodParameterTests::class.java)
private val innerClassWithParametersConstructor = InnerClassWithParameter::class.java
@@ -52,6 +54,16 @@ class KotlinMethodParameterTests {
assertThat(MethodParameter(nonNullableMethod, 0).isOptional).isFalse()
}
@Test
fun `Method parameter with default value`() {
assertThat(MethodParameter(withDefaultValueMethod, 0).isOptional).isTrue()
}
@Test
fun `Method parameter without default value`() {
assertThat(MethodParameter(nonNullableMethod, 0).isOptional).isFalse()
}
@Test
fun `Method return type nullability`() {
assertThat(MethodParameter(nullableMethod, -1).isOptional).isTrue()
@@ -123,6 +135,8 @@ class KotlinMethodParameterTests {
@Suppress("unused_parameter")
fun nonNullable(nonNullable: String): Int = 42
fun withDefaultValue(withDefaultValue: String = "default") = withDefaultValue
inner class InnerClass
@Suppress("unused_parameter")

View File

@@ -16,16 +16,22 @@
package org.springframework.web.method.annotation;
import java.lang.reflect.Method;
import java.util.Map;
import java.util.Objects;
import java.util.concurrent.ConcurrentHashMap;
import jakarta.servlet.ServletException;
import kotlin.reflect.KFunction;
import kotlin.reflect.KParameter;
import kotlin.reflect.jvm.ReflectJvmMapping;
import org.springframework.beans.ConversionNotSupportedException;
import org.springframework.beans.TypeMismatchException;
import org.springframework.beans.factory.config.BeanExpressionContext;
import org.springframework.beans.factory.config.BeanExpressionResolver;
import org.springframework.beans.factory.config.ConfigurableBeanFactory;
import org.springframework.core.KotlinDetector;
import org.springframework.core.MethodParameter;
import org.springframework.lang.Nullable;
import org.springframework.web.bind.ServletRequestBindingException;
@@ -60,6 +66,7 @@ import org.springframework.web.method.support.ModelAndViewContainer;
* @author Arjen Poutsma
* @author Rossen Stoyanchev
* @author Juergen Hoeller
* @author Sebastien Deleuze
* @since 3.1
*/
public abstract class AbstractNamedValueMethodArgumentResolver implements HandlerMethodArgumentResolver {
@@ -98,6 +105,9 @@ public abstract class AbstractNamedValueMethodArgumentResolver implements Handle
NamedValueInfo namedValueInfo = getNamedValueInfo(parameter);
MethodParameter nestedParameter = parameter.nestedIfOptional();
boolean hasDefaultValue = KotlinDetector.isKotlinReflectPresent()
&& KotlinDetector.isKotlinType(parameter.getDeclaringClass())
&& KotlinDelegate.hasDefaultValue(nestedParameter);
Object resolvedName = resolveEmbeddedValuesAndExpressions(namedValueInfo.name);
if (resolvedName == null) {
@@ -113,13 +123,15 @@ public abstract class AbstractNamedValueMethodArgumentResolver implements Handle
else if (namedValueInfo.required && !nestedParameter.isOptional()) {
handleMissingValue(namedValueInfo.name, nestedParameter, webRequest);
}
arg = handleNullValue(namedValueInfo.name, arg, nestedParameter.getNestedParameterType());
if (!hasDefaultValue) {
arg = handleNullValue(namedValueInfo.name, arg, nestedParameter.getNestedParameterType());
}
}
else if ("".equals(arg) && namedValueInfo.defaultValue != null) {
arg = resolveEmbeddedValuesAndExpressions(namedValueInfo.defaultValue);
}
if (binderFactory != null) {
if (binderFactory != null && (arg != null || !hasDefaultValue)) {
WebDataBinder binder = binderFactory.createBinder(webRequest, null, namedValueInfo.name);
try {
arg = binder.convertIfNecessary(arg, parameter.getParameterType(), parameter);
@@ -304,4 +316,27 @@ public abstract class AbstractNamedValueMethodArgumentResolver implements Handle
}
}
/**
* Inner class to avoid a hard dependency on Kotlin at runtime.
*/
private static class KotlinDelegate {
/**
* Check whether the specified {@link MethodParameter} represents a nullable Kotlin type
* or an optional parameter (with a default value in the Kotlin declaration).
*/
public static boolean hasDefaultValue(MethodParameter parameter) {
Method method = Objects.requireNonNull(parameter.getMethod());
KFunction<?> function = ReflectJvmMapping.getKotlinFunction(method);
if (function != null) {
int index = 0;
for (KParameter kParameter : function.getParameters()) {
if (KParameter.Kind.VALUE.equals(kParameter.getKind()) && parameter.getParameterIndex() == index++) {
return kParameter.isOptional();
}
}
}
return false;
}
}
}

View File

@@ -19,7 +19,13 @@ package org.springframework.web.method.support;
import java.lang.reflect.InvocationTargetException;
import java.lang.reflect.Method;
import java.util.Arrays;
import java.util.Map;
import java.util.Objects;
import kotlin.reflect.KFunction;
import kotlin.reflect.KParameter;
import kotlin.reflect.jvm.KCallablesJvm;
import kotlin.reflect.jvm.ReflectJvmMapping;
import org.reactivestreams.Publisher;
import org.springframework.context.MessageSource;
@@ -29,6 +35,7 @@ import org.springframework.core.KotlinDetector;
import org.springframework.core.MethodParameter;
import org.springframework.core.ParameterNameDiscoverer;
import org.springframework.lang.Nullable;
import org.springframework.util.CollectionUtils;
import org.springframework.util.ObjectUtils;
import org.springframework.validation.beanvalidation.MethodValidator;
import org.springframework.web.bind.WebDataBinder;
@@ -236,8 +243,13 @@ public class InvocableHandlerMethod extends HandlerMethod {
protected Object doInvoke(Object... args) throws Exception {
Method method = getBridgedMethod();
try {
if (KotlinDetector.isSuspendingFunction(method)) {
return invokeSuspendingFunction(method, getBean(), args);
if (KotlinDetector.isKotlinReflectPresent()) {
if (KotlinDetector.isSuspendingFunction(method)) {
return invokeSuspendingFunction(method, getBean(), args);
}
else if (KotlinDetector.isKotlinType(method.getDeclaringClass())) {
return KotlinDelegate.invokeFunction(method, getBean(), args);
}
}
return method.invoke(getBean(), args);
}
@@ -279,4 +291,33 @@ public class InvocableHandlerMethod extends HandlerMethod {
return CoroutinesUtils.invokeSuspendingFunction(method, target, args);
}
/**
* Inner class to avoid a hard dependency on Kotlin at runtime.
*/
private static class KotlinDelegate {
@Nullable
@SuppressWarnings("deprecation")
public static Object invokeFunction(Method method, Object target, Object[] args) {
KFunction<?> function = Objects.requireNonNull(ReflectJvmMapping.getKotlinFunction(method));
if (method.isAccessible() && !KCallablesJvm.isAccessible(function)) {
KCallablesJvm.setAccessible(function, true);
}
Map<KParameter, Object> argMap = CollectionUtils.newHashMap(args.length + 1);
int index = 0;
for (KParameter parameter : function.getParameters()) {
switch (parameter.getKind()) {
case INSTANCE -> argMap.put(parameter, target);
case VALUE -> {
if (!parameter.isOptional() || args[index] != null) {
argMap.put(parameter, args[index]);
}
index++;
}
}
}
return function.callBy(argMap);
}
}
}

View File

@@ -1,5 +1,5 @@
/*
* Copyright 2002-2018 the original author or authors.
* Copyright 2002-2023 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.
@@ -47,7 +47,7 @@ public class StubArgumentResolver implements HandlerMethodArgumentResolver {
this(valueType, null);
}
public StubArgumentResolver(Class<?> valueType, Object value) {
public StubArgumentResolver(Class<?> valueType, @Nullable Object value) {
this.valueType = valueType;
this.value = value;
}

View File

@@ -1,5 +1,5 @@
/*
* Copyright 2002-2020 the original author or authors.
* Copyright 2002-2023 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.
@@ -58,6 +58,13 @@ class RequestParamMethodArgumentResolverKotlinTests {
lateinit var nonNullableParamRequired: MethodParameter
lateinit var nonNullableParamNotRequired: MethodParameter
lateinit var defaultValueBooleanParamRequired: MethodParameter
lateinit var defaultValueBooleanParamNotRequired: MethodParameter
lateinit var defaultValueIntParamRequired: MethodParameter
lateinit var defaultValueIntParamNotRequired: MethodParameter
lateinit var defaultValueStringParamRequired: MethodParameter
lateinit var defaultValueStringParamNotRequired: MethodParameter
lateinit var nullableMultipartParamRequired: MethodParameter
lateinit var nullableMultipartParamNotRequired: MethodParameter
lateinit var nonNullableMultipartParamRequired: MethodParameter
@@ -73,20 +80,27 @@ class RequestParamMethodArgumentResolverKotlinTests {
binderFactory = DefaultDataBinderFactory(initializer)
webRequest = ServletWebRequest(request, MockHttpServletResponse())
val method = ReflectionUtils.findMethod(javaClass, "handle", String::class.java,
String::class.java, String::class.java, String::class.java,
MultipartFile::class.java, MultipartFile::class.java,
MultipartFile::class.java, MultipartFile::class.java)!!
val method = ReflectionUtils.findMethod(javaClass, "handle",
String::class.java, String::class.java, String::class.java, String::class.java,
Boolean::class.java, Boolean::class.java, Int::class.java, Int::class.java, String::class.java, String::class.java,
MultipartFile::class.java, MultipartFile::class.java, MultipartFile::class.java, MultipartFile::class.java)!!
nullableParamRequired = SynthesizingMethodParameter(method, 0)
nullableParamNotRequired = SynthesizingMethodParameter(method, 1)
nonNullableParamRequired = SynthesizingMethodParameter(method, 2)
nonNullableParamNotRequired = SynthesizingMethodParameter(method, 3)
nullableMultipartParamRequired = SynthesizingMethodParameter(method, 4)
nullableMultipartParamNotRequired = SynthesizingMethodParameter(method, 5)
nonNullableMultipartParamRequired = SynthesizingMethodParameter(method, 6)
nonNullableMultipartParamNotRequired = SynthesizingMethodParameter(method, 7)
defaultValueBooleanParamRequired = SynthesizingMethodParameter(method, 4)
defaultValueBooleanParamNotRequired = SynthesizingMethodParameter(method, 5)
defaultValueIntParamRequired = SynthesizingMethodParameter(method, 6)
defaultValueIntParamNotRequired = SynthesizingMethodParameter(method, 7)
defaultValueStringParamRequired = SynthesizingMethodParameter(method, 8)
defaultValueStringParamNotRequired = SynthesizingMethodParameter(method, 9)
nullableMultipartParamRequired = SynthesizingMethodParameter(method, 10)
nullableMultipartParamNotRequired = SynthesizingMethodParameter(method, 11)
nonNullableMultipartParamRequired = SynthesizingMethodParameter(method, 12)
nonNullableMultipartParamNotRequired = SynthesizingMethodParameter(method, 13)
}
@Test
@@ -143,6 +157,84 @@ class RequestParamMethodArgumentResolverKotlinTests {
}
}
@Test
fun resolveDefaultValueRequiredWithBooleanParameter() {
request.addParameter("value", "false")
val result = resolver.resolveArgument(defaultValueBooleanParamRequired, null, webRequest, binderFactory)
assertThat(result).isEqualTo(false)
}
@Test
fun resolveDefaultValueRequiredWithoutBooleanParameter() {
val result = resolver.resolveArgument(defaultValueBooleanParamRequired, null, webRequest, binderFactory)
assertThat(result).isEqualTo(null)
}
@Test
fun resolveDefaultValueNotRequiredWithBooleanParameter() {
request.addParameter("value", "false")
val result = resolver.resolveArgument(defaultValueBooleanParamNotRequired, null, webRequest, binderFactory)
assertThat(result).isEqualTo(false)
}
@Test
fun resolveDefaultValueNotRequiredWithoutBooleanParameter() {
val result = resolver.resolveArgument(defaultValueBooleanParamNotRequired, null, webRequest, binderFactory)
assertThat(result).isEqualTo(null)
}
@Test
fun resolveDefaultValueRequiredWithIntParameter() {
request.addParameter("value", "123")
val result = resolver.resolveArgument(defaultValueIntParamRequired, null, webRequest, binderFactory)
assertThat(result).isEqualTo(123)
}
@Test
fun resolveDefaultValueRequiredWithoutIntParameter() {
val result = resolver.resolveArgument(defaultValueIntParamRequired, null, webRequest, binderFactory)
assertThat(result).isEqualTo(null)
}
@Test
fun resolveDefaultValueNotRequiredWithIntParameter() {
request.addParameter("value", "123")
val result = resolver.resolveArgument(defaultValueIntParamNotRequired, null, webRequest, binderFactory)
assertThat(result).isEqualTo(123)
}
@Test
fun resolveDefaultValueNotRequiredWithoutIntParameter() {
val result = resolver.resolveArgument(defaultValueIntParamNotRequired, null, webRequest, binderFactory)
assertThat(result).isEqualTo(null)
}
@Test
fun resolveDefaultValueRequiredWithStringParameter() {
request.addParameter("value", "123")
val result = resolver.resolveArgument(defaultValueStringParamRequired, null, webRequest, binderFactory)
assertThat(result).isEqualTo("123")
}
@Test
fun resolveDefaultValueRequiredWithoutStringParameter() {
val result = resolver.resolveArgument(defaultValueStringParamRequired, null, webRequest, binderFactory)
assertThat(result).isEqualTo(null)
}
@Test
fun resolveDefaultValueNotRequiredWithStringParameter() {
request.addParameter("value", "123")
val result = resolver.resolveArgument(defaultValueStringParamNotRequired, null, webRequest, binderFactory)
assertThat(result).isEqualTo("123")
}
@Test
fun resolveDefaultValueNotRequiredWithoutStringParameter() {
val result = resolver.resolveArgument(defaultValueStringParamNotRequired, null, webRequest, binderFactory)
assertThat(result).isEqualTo(null)
}
@Test
fun resolveNullableRequiredWithMultipartParameter() {
val request = MockMultipartHttpServletRequest()
@@ -233,6 +325,13 @@ class RequestParamMethodArgumentResolverKotlinTests {
@RequestParam("name") nonNullableParamRequired: String,
@RequestParam("name", required = false) nonNullableParamNotRequired: String,
@RequestParam("value") withDefaultValueBooleanParamRequired: Boolean = true,
@RequestParam("value", required = false) withDefaultValueBooleanParamNotRequired: Boolean = true,
@RequestParam("value") withDefaultValueIntParamRequired: Int = 20,
@RequestParam("value", required = false) withDefaultValueIntParamNotRequired: Int = 20,
@RequestParam("value") withDefaultValueStringParamRequired: String = "default",
@RequestParam("value", required = false) withDefaultValueStringParamNotRequired: String = "default",
@RequestParam("mfile") nullableMultipartParamRequired: MultipartFile?,
@RequestParam("mfile", required = false) nullableMultipartParamNotRequired: MultipartFile?,
@RequestParam("mfile") nonNullableMultipartParamRequired: MultipartFile,

View File

@@ -0,0 +1,100 @@
/*
* Copyright 2002-2023 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.
* You may obtain a copy of the License at
*
* https://www.apache.org/licenses/LICENSE-2.0
*
* Unless required by applicable law or agreed to in writing, software
* distributed under the License is distributed on an "AS IS" BASIS,
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
* See the License for the specific language governing permissions and
* limitations under the License.
*/
package org.springframework.web.method.support
import org.assertj.core.api.Assertions
import org.junit.jupiter.api.Test
import org.springframework.web.context.request.NativeWebRequest
import org.springframework.web.context.request.ServletWebRequest
import org.springframework.web.testfixture.method.ResolvableMethod
import org.springframework.web.testfixture.servlet.MockHttpServletRequest
import org.springframework.web.testfixture.servlet.MockHttpServletResponse
/**
* Kotlin unit tests for {@link InvocableHandlerMethod}.
*
* @author Sebastien Deleuze
*/
class InvocableHandlerMethodKotlinTests {
private val request: NativeWebRequest = ServletWebRequest(MockHttpServletRequest(), MockHttpServletResponse())
private val composite = HandlerMethodArgumentResolverComposite()
@Test
fun intDefaultValue() {
composite.addResolver(StubArgumentResolver(Int::class.java, null))
val value = getInvocable(Int::class.java).invokeForRequest(request, null)
Assertions.assertThat(getStubResolver(0).resolvedParameters).hasSize(1)
Assertions.assertThat(value).isEqualTo("20")
}
@Test
fun booleanDefaultValue() {
composite.addResolver(StubArgumentResolver(Boolean::class.java, null))
val value = getInvocable(Boolean::class.java).invokeForRequest(request, null)
Assertions.assertThat(getStubResolver(0).resolvedParameters).hasSize(1)
Assertions.assertThat(value).isEqualTo("true")
}
@Test
fun nullableIntDefaultValue() {
composite.addResolver(StubArgumentResolver(Int::class.javaObjectType, null))
val value = getInvocable(Int::class.javaObjectType).invokeForRequest(request, null)
Assertions.assertThat(getStubResolver(0).resolvedParameters).hasSize(1)
Assertions.assertThat(value).isEqualTo("20")
}
@Test
fun nullableBooleanDefaultValue() {
composite.addResolver(StubArgumentResolver(Boolean::class.javaObjectType, null))
val value = getInvocable(Boolean::class.javaObjectType).invokeForRequest(request, null)
Assertions.assertThat(getStubResolver(0).resolvedParameters).hasSize(1)
Assertions.assertThat(value).isEqualTo("true")
}
private fun getInvocable(vararg argTypes: Class<*>): InvocableHandlerMethod {
val method = ResolvableMethod.on(Handler::class.java).argTypes(*argTypes).resolveMethod()
val handlerMethod = InvocableHandlerMethod(Handler(), method)
handlerMethod.setHandlerMethodArgumentResolvers(composite)
return handlerMethod
}
private fun getStubResolver(index: Int): StubArgumentResolver {
return composite.resolvers[index] as StubArgumentResolver
}
private class Handler {
fun intDefaultValue(limit: Int = 20) =
limit.toString()
fun nullableIntDefaultValue(limit: Int? = 20) =
limit.toString()
fun booleanDefaultValue(status: Boolean = true) =
status.toString()
fun nullableBooleanDefaultValue(status: Boolean? = true) =
status.toString()
}
}

View File

@@ -22,8 +22,14 @@ import java.lang.reflect.ParameterizedType;
import java.lang.reflect.Type;
import java.util.ArrayList;
import java.util.List;
import java.util.Map;
import java.util.Objects;
import java.util.stream.Stream;
import kotlin.reflect.KFunction;
import kotlin.reflect.KParameter;
import kotlin.reflect.jvm.KCallablesJvm;
import kotlin.reflect.jvm.ReflectJvmMapping;
import reactor.core.publisher.Mono;
import org.springframework.core.CoroutinesUtils;
@@ -36,6 +42,7 @@ import org.springframework.core.ReactiveAdapterRegistry;
import org.springframework.http.HttpStatusCode;
import org.springframework.http.server.reactive.ServerHttpResponse;
import org.springframework.lang.Nullable;
import org.springframework.util.CollectionUtils;
import org.springframework.util.ObjectUtils;
import org.springframework.validation.beanvalidation.MethodValidator;
import org.springframework.web.method.HandlerMethod;
@@ -159,8 +166,13 @@ public class InvocableHandlerMethod extends HandlerMethod {
Method method = getBridgedMethod();
boolean isSuspendingFunction = KotlinDetector.isSuspendingFunction(method);
try {
if (isSuspendingFunction) {
value = CoroutinesUtils.invokeSuspendingFunction(method, getBean(), args);
if (KotlinDetector.isKotlinReflectPresent() && KotlinDetector.isKotlinType(method.getDeclaringClass())) {
if (isSuspendingFunction) {
value = CoroutinesUtils.invokeSuspendingFunction(method, getBean(), args);
}
else {
value = KotlinDelegate.invokeFunction(method, getBean(), args);
}
}
else {
value = method.invoke(getBean(), args);
@@ -278,4 +290,33 @@ public class InvocableHandlerMethod extends HandlerMethod {
return false;
}
/**
* Inner class to avoid a hard dependency on Kotlin at runtime.
*/
private static class KotlinDelegate {
@Nullable
@SuppressWarnings("deprecation")
public static Object invokeFunction(Method method, Object target, Object[] args) {
KFunction<?> function = Objects.requireNonNull(ReflectJvmMapping.getKotlinFunction(method));
if (method.isAccessible() && !KCallablesJvm.isAccessible(function)) {
KCallablesJvm.setAccessible(function, true);
}
Map<KParameter, Object> argMap = CollectionUtils.newHashMap(args.length + 1);
int index = 0;
for (KParameter parameter : function.getParameters()) {
switch (parameter.getKind()) {
case INSTANCE -> argMap.put(parameter, target);
case VALUE -> {
if (!parameter.isOptional() || args[index] != null) {
argMap.put(parameter, args[index]);
}
index++;
}
}
}
return function.callBy(argMap);
}
}
}

View File

@@ -16,9 +16,14 @@
package org.springframework.web.reactive.result.method.annotation;
import java.lang.reflect.Method;
import java.util.Map;
import java.util.Objects;
import java.util.concurrent.ConcurrentHashMap;
import kotlin.reflect.KFunction;
import kotlin.reflect.KParameter;
import kotlin.reflect.jvm.ReflectJvmMapping;
import reactor.core.publisher.Mono;
import org.springframework.beans.ConversionNotSupportedException;
@@ -26,6 +31,7 @@ import org.springframework.beans.TypeMismatchException;
import org.springframework.beans.factory.config.BeanExpressionContext;
import org.springframework.beans.factory.config.BeanExpressionResolver;
import org.springframework.beans.factory.config.ConfigurableBeanFactory;
import org.springframework.core.KotlinDetector;
import org.springframework.core.MethodParameter;
import org.springframework.core.ReactiveAdapterRegistry;
import org.springframework.lang.Nullable;
@@ -57,6 +63,7 @@ import org.springframework.web.server.ServerWebInputException;
* {@link ConfigurableBeanFactory} must be supplied to the class constructor.
*
* @author Rossen Stoyanchev
* @author Sebastien Deleuze
* @since 5.0
*/
public abstract class AbstractNamedValueArgumentResolver extends HandlerMethodArgumentResolverSupport {
@@ -210,14 +217,21 @@ public abstract class AbstractNamedValueArgumentResolver extends HandlerMethodAr
return Mono.fromSupplier(() -> {
Object value = null;
boolean hasDefaultValue = KotlinDetector.isKotlinReflectPresent()
&& KotlinDetector.isKotlinType(parameter.getDeclaringClass())
&& KotlinDelegate.hasDefaultValue(parameter);
if (namedValueInfo.defaultValue != null) {
value = resolveEmbeddedValuesAndExpressions(namedValueInfo.defaultValue);
}
else if (namedValueInfo.required && !parameter.isOptional()) {
handleMissingValue(namedValueInfo.name, parameter, exchange);
}
value = handleNullValue(namedValueInfo.name, value, parameter.getNestedParameterType());
value = applyConversion(value, namedValueInfo, parameter, bindingContext, exchange);
if (!hasDefaultValue) {
value = handleNullValue(namedValueInfo.name, value, parameter.getNestedParameterType());
}
if (value != null || !hasDefaultValue) {
value = applyConversion(value, namedValueInfo, parameter, bindingContext, exchange);
}
handleResolvedValue(value, namedValueInfo.name, parameter, model, exchange);
return value;
});
@@ -304,4 +318,28 @@ public abstract class AbstractNamedValueArgumentResolver extends HandlerMethodAr
}
}
/**
* Inner class to avoid a hard dependency on Kotlin at runtime.
*/
private static class KotlinDelegate {
/**
* Check whether the specified {@link MethodParameter} represents a nullable Kotlin type
* or an optional parameter (with a default value in the Kotlin declaration).
*/
public static boolean hasDefaultValue(MethodParameter parameter) {
Method method = Objects.requireNonNull(parameter.getMethod());
KFunction<?> function = ReflectJvmMapping.getKotlinFunction(method);
if (function != null) {
int index = 0;
for (KParameter kParameter : function.getParameters()) {
if (KParameter.Kind.VALUE.equals(kParameter.getKind()) && parameter.getParameterIndex() == index++) {
return kParameter.isOptional();
}
}
}
return false;
}
}
}

View File

@@ -1,5 +1,5 @@
/*
* Copyright 2002-2022 the original author or authors.
* Copyright 2002-2023 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.
@@ -21,16 +21,21 @@ import io.mockk.mockk
import kotlinx.coroutines.delay
import org.assertj.core.api.Assertions.assertThat
import org.junit.jupiter.api.Test
import org.springframework.core.ReactiveAdapterRegistry
import org.springframework.http.HttpStatus
import org.springframework.http.server.reactive.ServerHttpResponse
import org.springframework.web.testfixture.http.server.reactive.MockServerHttpRequest.get
import org.springframework.web.testfixture.server.MockServerWebExchange
import org.springframework.web.bind.annotation.RequestMapping
import org.springframework.web.bind.annotation.RequestParam
import org.springframework.web.bind.annotation.ResponseStatus
import org.springframework.web.bind.annotation.RestController
import org.springframework.web.reactive.BindingContext
import org.springframework.web.reactive.HandlerResult
import org.springframework.web.reactive.result.method.HandlerMethodArgumentResolver
import org.springframework.web.reactive.result.method.InvocableHandlerMethod
import org.springframework.web.reactive.result.method.annotation.ContinuationHandlerMethodArgumentResolver
import org.springframework.web.reactive.result.method.annotation.RequestParamMethodArgumentResolver
import org.springframework.web.testfixture.http.server.reactive.MockServerHttpRequest.get
import org.springframework.web.testfixture.server.MockServerWebExchange
import reactor.core.publisher.Mono
import reactor.test.StepVerifier
import java.lang.reflect.Method
@@ -39,9 +44,10 @@ import kotlin.reflect.jvm.javaMethod
class KotlinInvocableHandlerMethodTests {
private val exchange = MockServerWebExchange.from(get("http://localhost:8080/path"))
private var exchange = MockServerWebExchange.from(get("http://localhost:8080/path"))
private val resolvers = mutableListOf<HandlerMethodArgumentResolver>(ContinuationHandlerMethodArgumentResolver())
private val resolvers = mutableListOf<HandlerMethodArgumentResolver>(ContinuationHandlerMethodArgumentResolver(),
RequestParamMethodArgumentResolver(null, ReactiveAdapterRegistry.getSharedInstance(), false))
@Test
fun resolveNoArg() {
@@ -104,6 +110,58 @@ class KotlinInvocableHandlerMethodTests {
assertHandlerResultValue(result, "success:foo")
}
@Test
fun defaultValue() {
this.resolvers.add(stubResolver(Mono.empty()))
val method = DefaultValueController::handle.javaMethod!!
val result = invoke(DefaultValueController(), method)
assertHandlerResultValue(result, "default")
}
@Test
fun defaultValueOverridden() {
this.resolvers.add(stubResolver(Mono.empty()))
val method = DefaultValueController::handle.javaMethod!!
exchange = MockServerWebExchange.from(get("http://localhost:8080/path").queryParam("value", "override"))
val result = invoke(DefaultValueController(), method)
assertHandlerResultValue(result, "override")
}
@Test
fun defaultValues() {
this.resolvers.add(stubResolver(Mono.empty()))
val method = DefaultValueController::handleMultiple.javaMethod!!
val result = invoke(DefaultValueController(), method)
assertHandlerResultValue(result, "10-20")
}
@Test
fun defaultValuesOverridden() {
this.resolvers.add(stubResolver(Mono.empty()))
val method = DefaultValueController::handleMultiple.javaMethod!!
exchange = MockServerWebExchange.from(get("http://localhost:8080/path").queryParam("limit2", "40"))
val result = invoke(DefaultValueController(), method)
assertHandlerResultValue(result, "10-40")
}
@Test
fun suspendingDefaultValue() {
this.resolvers.add(stubResolver(Mono.empty()))
val method = DefaultValueController::handleSuspending.javaMethod!!
val result = invoke(DefaultValueController(), method)
assertHandlerResultValue(result, "default")
}
@Test
fun suspendingDefaultValueOverridden() {
this.resolvers.add(stubResolver(Mono.empty()))
val method = DefaultValueController::handleSuspending.javaMethod!!
exchange = MockServerWebExchange.from(get("http://localhost:8080/path").queryParam("value", "override"))
val result = invoke(DefaultValueController(), method)
assertHandlerResultValue(result, "override")
}
private fun invokeForResult(handler: Any, method: Method, vararg providedArgs: Any): HandlerResult? {
return invoke(handler, method, *providedArgs).block(Duration.ofSeconds(5))
}
@@ -127,8 +185,13 @@ class KotlinInvocableHandlerMethodTests {
private fun assertHandlerResultValue(mono: Mono<HandlerResult>, expected: String) {
StepVerifier.create(mono)
.consumeNextWith { StepVerifier.create(it.returnValue as Mono<*>).expectNext(expected).verifyComplete() }
.verifyComplete()
.consumeNextWith {
if (it.returnValue is Mono<*>) {
StepVerifier.create(it.returnValue as Mono<*>).expectNext(expected).verifyComplete()
} else {
assertThat(it.returnValue).isEqualTo(expected)
}
}.verifyComplete()
}
class CoroutinesController {
@@ -166,4 +229,16 @@ class KotlinInvocableHandlerMethodTests {
return "success:$q"
}
}
@RestController
class DefaultValueController {
fun handle(@RequestParam value: String = "default") = value
fun handleMultiple(@RequestParam(defaultValue = "10") limit1: Int, @RequestParam limit2: Int = 20) = "${limit1}-${limit2}"
@Suppress("RedundantSuspendModifier")
suspend fun handleSuspending(@RequestParam value: String = "default") = value
}
}

View File

@@ -46,6 +46,13 @@ class RequestParamMethodArgumentResolverKotlinTests {
lateinit var nonNullableParamRequired: MethodParameter
lateinit var nonNullableParamNotRequired: MethodParameter
lateinit var defaultValueBooleanParamRequired: MethodParameter
lateinit var defaultValueBooleanParamNotRequired: MethodParameter
lateinit var defaultValueIntParamRequired: MethodParameter
lateinit var defaultValueIntParamNotRequired: MethodParameter
lateinit var defaultValueStringParamRequired: MethodParameter
lateinit var defaultValueStringParamNotRequired: MethodParameter
@BeforeEach
fun setup() {
@@ -54,13 +61,22 @@ class RequestParamMethodArgumentResolverKotlinTests {
initializer.conversionService = DefaultFormattingConversionService()
bindingContext = BindingContext(initializer)
val method = ReflectionUtils.findMethod(javaClass, "handle", String::class.java,
String::class.java, String::class.java, String::class.java)!!
val method = ReflectionUtils.findMethod(javaClass, "handle",
String::class.java, String::class.java, String::class.java, String::class.java,
Boolean::class.java, Boolean::class.java, Int::class.java, Int::class.java,
String::class.java, String::class.java)!!
nullableParamRequired = SynthesizingMethodParameter(method, 0)
nullableParamNotRequired = SynthesizingMethodParameter(method, 1)
nonNullableParamRequired = SynthesizingMethodParameter(method, 2)
nonNullableParamNotRequired = SynthesizingMethodParameter(method, 3)
defaultValueBooleanParamRequired = SynthesizingMethodParameter(method, 4)
defaultValueBooleanParamNotRequired = SynthesizingMethodParameter(method, 5)
defaultValueIntParamRequired = SynthesizingMethodParameter(method, 6)
defaultValueIntParamNotRequired = SynthesizingMethodParameter(method, 7)
defaultValueStringParamRequired = SynthesizingMethodParameter(method, 8)
defaultValueStringParamNotRequired = SynthesizingMethodParameter(method, 9)
}
@Test
@@ -119,13 +135,104 @@ class RequestParamMethodArgumentResolverKotlinTests {
StepVerifier.create(result).expectComplete().verify()
}
@Test
fun resolveDefaultValueRequiredWithBooleanParameter() {
val exchange = MockServerWebExchange.from(MockServerHttpRequest.get("/path?value=false"))
val result = resolver.resolveArgument(defaultValueBooleanParamRequired, bindingContext, exchange)
StepVerifier.create(result).expectNext(false).expectComplete().verify()
}
@Test
fun resolveDefaultValueRequiredWithoutBooleanParameter() {
val exchange = MockServerWebExchange.from(MockServerHttpRequest.get("/"))
val result = resolver.resolveArgument(defaultValueBooleanParamRequired, bindingContext, exchange)
StepVerifier.create(result).expectComplete().verify()
}
@Test
fun resolveDefaultValueNotRequiredWithBooleanParameter() {
val exchange = MockServerWebExchange.from(MockServerHttpRequest.get("/path?value=false"))
val result = resolver.resolveArgument(defaultValueBooleanParamNotRequired, bindingContext, exchange)
StepVerifier.create(result).expectNext(false).expectComplete().verify()
}
@Test
fun resolveDefaultValueNotRequiredWithoutBooleanParameter() {
val exchange = MockServerWebExchange.from(MockServerHttpRequest.get("/"))
val result = resolver.resolveArgument(defaultValueBooleanParamNotRequired, bindingContext, exchange)
StepVerifier.create(result).expectComplete().verify()
}
@Test
fun resolveDefaultValueRequiredWithIntParameter() {
val exchange = MockServerWebExchange.from(MockServerHttpRequest.get("/path?value=123"))
val result = resolver.resolveArgument(defaultValueIntParamRequired, bindingContext, exchange)
StepVerifier.create(result).expectNext(123).expectComplete().verify()
}
@Test
fun resolveDefaultValueRequiredWithoutIntParameter() {
val exchange = MockServerWebExchange.from(MockServerHttpRequest.get("/"))
val result = resolver.resolveArgument(defaultValueIntParamRequired, bindingContext, exchange)
StepVerifier.create(result).expectComplete().verify()
}
@Test
fun resolveDefaultValueNotRequiredWithIntParameter() {
val exchange = MockServerWebExchange.from(MockServerHttpRequest.get("/path?value=123"))
val result = resolver.resolveArgument(defaultValueIntParamNotRequired, bindingContext, exchange)
StepVerifier.create(result).expectNext(123).expectComplete().verify()
}
@Test
fun resolveDefaultValueNotRequiredWithoutIntParameter() {
val exchange = MockServerWebExchange.from(MockServerHttpRequest.get("/"))
val result = resolver.resolveArgument(defaultValueIntParamNotRequired, bindingContext, exchange)
StepVerifier.create(result).expectComplete().verify()
}
@Test
fun resolveDefaultValueRequiredWithStringParameter() {
val exchange = MockServerWebExchange.from(MockServerHttpRequest.get("/path?value=123"))
val result = resolver.resolveArgument(defaultValueStringParamRequired, bindingContext, exchange)
StepVerifier.create(result).expectNext("123").expectComplete().verify()
}
@Test
fun resolveDefaultValueRequiredWithoutStringParameter() {
val exchange = MockServerWebExchange.from(MockServerHttpRequest.get("/"))
val result = resolver.resolveArgument(defaultValueStringParamRequired, bindingContext, exchange)
StepVerifier.create(result).expectComplete().verify()
}
@Test
fun resolveDefaultValueNotRequiredWithStringParameter() {
val exchange = MockServerWebExchange.from(MockServerHttpRequest.get("/path?value=123"))
val result = resolver.resolveArgument(defaultValueStringParamNotRequired, bindingContext, exchange)
StepVerifier.create(result).expectNext("123").expectComplete().verify()
}
@Test
fun resolveDefaultValueNotRequiredWithoutStringParameter() {
val exchange = MockServerWebExchange.from(MockServerHttpRequest.get("/"))
val result = resolver.resolveArgument(defaultValueStringParamNotRequired, bindingContext, exchange)
StepVerifier.create(result).expectComplete().verify()
}
@Suppress("unused_parameter")
fun handle(
@RequestParam("name") nullableParamRequired: String?,
@RequestParam("name", required = false) nullableParamNotRequired: String?,
@RequestParam("name") nonNullableParamRequired: String,
@RequestParam("name", required = false) nonNullableParamNotRequired: String) {
@RequestParam("name", required = false) nonNullableParamNotRequired: String,
@RequestParam("value") withDefaultValueBooleanParamRequired: Boolean = true,
@RequestParam("value", required = false) withDefaultValueBooleanParamNotRequired: Boolean = true,
@RequestParam("value") withDefaultValueIntParamRequired: Int = 20,
@RequestParam("value", required = false) withDefaultValueIntParamNotRequired: Int = 20,
@RequestParam("value") withDefaultValueStringParamRequired: String = "default",
@RequestParam("value", required = false) withDefaultValueStringParamNotRequired: String = "default") {
}
}

View File

@@ -16,9 +16,9 @@
package org.springframework.web.servlet.mvc.method.annotation
import kotlinx.coroutines.delay
import org.assertj.core.api.Assertions.assertThat
import org.springframework.web.bind.annotation.RequestMapping
import org.springframework.web.bind.annotation.RequestParam
import org.springframework.web.bind.annotation.RestController
import org.springframework.web.context.request.async.WebAsyncUtils
import org.springframework.web.servlet.handler.PathPatternsParameterizedTest
@@ -84,6 +84,71 @@ class ServletAnnotationControllerHandlerMethodKotlinTests : AbstractServletHandl
assertThat(WebAsyncUtils.getAsyncManager(request).concurrentResult).isEqualTo("foo")
}
@PathPatternsParameterizedTest
fun defaultValue(usePathPatterns: Boolean) {
initDispatcherServlet(DefaultValueController::class.java, usePathPatterns)
val request = MockHttpServletRequest("GET", "/default-value")
val response = MockHttpServletResponse()
servlet.service(request, response)
assertThat(response.contentAsString).isEqualTo("default")
}
@PathPatternsParameterizedTest
fun defaultValueOverridden(usePathPatterns: Boolean) {
initDispatcherServlet(DefaultValueController::class.java, usePathPatterns)
val request = MockHttpServletRequest("GET", "/default-value")
request.addParameter("value", "override")
val response = MockHttpServletResponse()
servlet.service(request, response)
assertThat(response.contentAsString).isEqualTo("override")
}
@PathPatternsParameterizedTest
fun defaultValues(usePathPatterns: Boolean) {
initDispatcherServlet(DefaultValueController::class.java, usePathPatterns)
val request = MockHttpServletRequest("GET", "/default-values")
val response = MockHttpServletResponse()
servlet.service(request, response)
assertThat(response.contentAsString).isEqualTo("10-20")
}
@PathPatternsParameterizedTest
fun defaultValuesOverridden(usePathPatterns: Boolean) {
initDispatcherServlet(DefaultValueController::class.java, usePathPatterns)
val request = MockHttpServletRequest("GET", "/default-values")
request.addParameter("limit2", "40")
val response = MockHttpServletResponse()
servlet.service(request, response)
assertThat(response.contentAsString).isEqualTo("10-40")
}
@PathPatternsParameterizedTest
fun suspendingDefaultValue(usePathPatterns: Boolean) {
initDispatcherServlet(DefaultValueController::class.java, usePathPatterns)
val request = MockHttpServletRequest("GET", "/suspending-default-value")
request.isAsyncSupported = true
val response = MockHttpServletResponse()
servlet.service(request, response)
assertThat(WebAsyncUtils.getAsyncManager(request).concurrentResult).isEqualTo("default")
}
@PathPatternsParameterizedTest
fun suspendingDefaultValueOverridden(usePathPatterns: Boolean) {
initDispatcherServlet(DefaultValueController::class.java, usePathPatterns)
val request = MockHttpServletRequest("GET", "/suspending-default-value")
request.isAsyncSupported = true
request.addParameter("value", "override")
val response = MockHttpServletResponse()
servlet.service(request, response)
assertThat(WebAsyncUtils.getAsyncManager(request).concurrentResult).isEqualTo("override")
}
data class DataClass(val param1: String, val param2: Int)
@@ -107,6 +172,20 @@ class ServletAnnotationControllerHandlerMethodKotlinTests : AbstractServletHandl
suspend fun handle(): String {
return "foo"
}
}
@RestController
class DefaultValueController {
@RequestMapping("/default-value")
fun handle(@RequestParam value: String = "default") = value
@RequestMapping("/default-values")
fun handleMultiple(@RequestParam(defaultValue = "10") limit1: Int, @RequestParam limit2: Int = 20) = "${limit1}-${limit2}"
@Suppress("RedundantSuspendModifier")
@RequestMapping("/suspending-default-value")
suspend fun handleSuspending(@RequestParam value: String = "default") = value
}