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

@@ -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()
}
}