diff --git a/spring-aop/spring-aop.gradle b/spring-aop/spring-aop.gradle index eec30b7bedf..965b231286b 100644 --- a/spring-aop/spring-aop.gradle +++ b/spring-aop/spring-aop.gradle @@ -8,6 +8,7 @@ dependencies { compileOnly("com.google.code.findbugs:jsr305") // for the AOP Alliance fork optional("org.apache.commons:commons-pool2") optional("org.aspectj:aspectjweaver") + optional("org.jetbrains.kotlin:kotlin-reflect") optional("org.jetbrains.kotlinx:kotlinx-coroutines-reactor") testFixturesImplementation(testFixtures(project(":spring-beans"))) testFixturesImplementation(testFixtures(project(":spring-core"))) diff --git a/spring-aop/src/main/java/org/springframework/aop/framework/CglibAopProxy.java b/spring-aop/src/main/java/org/springframework/aop/framework/CglibAopProxy.java index 10edfb43928..07ad2672c99 100644 --- a/spring-aop/src/main/java/org/springframework/aop/framework/CglibAopProxy.java +++ b/spring-aop/src/main/java/org/springframework/aop/framework/CglibAopProxy.java @@ -17,7 +17,6 @@ package org.springframework.aop.framework; import java.io.Serializable; -import java.lang.reflect.InvocationTargetException; import java.lang.reflect.Method; import java.lang.reflect.Modifier; import java.lang.reflect.UndeclaredThrowableException; @@ -52,7 +51,6 @@ import org.springframework.cglib.proxy.MethodProxy; import org.springframework.cglib.proxy.NoOp; import org.springframework.cglib.transform.impl.UndeclaredThrowableStrategy; import org.springframework.core.KotlinDetector; -import org.springframework.core.MethodParameter; import org.springframework.core.SmartClassLoader; import org.springframework.util.Assert; import org.springframework.util.ClassUtils; @@ -97,8 +95,6 @@ class CglibAopProxy implements AopProxy, Serializable { private static final int INVOKE_HASHCODE = 6; - private static final String COROUTINES_FLOW_CLASS_NAME = "kotlinx.coroutines.flow.Flow"; - private static final boolean COROUTINES_REACTOR_PRESENT = ClassUtils.isPresent( "kotlinx.coroutines.reactor.MonoKt", CglibAopProxy.class.getClassLoader()); @@ -422,8 +418,7 @@ class CglibAopProxy implements AopProxy, Serializable { * Also takes care of the conversion from {@code Mono} to Kotlin Coroutines if needed. */ private static @Nullable Object processReturnType( - Object proxy, @Nullable Object target, Method method, Object[] arguments, @Nullable Object returnValue) throws - NoSuchMethodException, InvocationTargetException, IllegalAccessException { + Object proxy, @Nullable Object target, Method method, Object[] arguments, @Nullable Object returnValue) { // Massage return value if necessary if (returnValue != null && returnValue == target && @@ -438,14 +433,7 @@ class CglibAopProxy implements AopProxy, Serializable { "Null return value from advice does not match primitive return type for: " + method); } if (COROUTINES_REACTOR_PRESENT && KotlinDetector.isSuspendingFunction(method)) { - Class returnParameterType = new MethodParameter(method, -1).getParameterType(); - if (COROUTINES_FLOW_CLASS_NAME.equals(returnParameterType.getName())) { - return CoroutinesUtils.asFlow(returnValue); - } - Object awaitResult = CoroutinesUtils.awaitSingleOrNull(returnValue, arguments[arguments.length - 1]); - return KotlinDetector.isInlineClass(returnParameterType) && - awaitResult != null && KotlinDetector.isInlineClass(awaitResult.getClass()) ? - awaitResult.getClass().getDeclaredMethod("unbox-impl").invoke(awaitResult) : awaitResult; + return CoroutinesUtils.adaptReturnValue(method, returnValue, arguments[arguments.length - 1]); } return returnValue; } diff --git a/spring-aop/src/main/java/org/springframework/aop/framework/CoroutinesUtils.java b/spring-aop/src/main/java/org/springframework/aop/framework/CoroutinesUtils.java index f0e56cab0fa..e1b279fe572 100644 --- a/spring-aop/src/main/java/org/springframework/aop/framework/CoroutinesUtils.java +++ b/spring-aop/src/main/java/org/springframework/aop/framework/CoroutinesUtils.java @@ -16,13 +16,32 @@ package org.springframework.aop.framework; +import java.lang.reflect.Method; +import java.util.List; +import java.util.Map; +import java.util.Objects; + import kotlin.coroutines.Continuation; +import kotlin.jvm.JvmClassMappingKt; +import kotlin.reflect.KClass; +import kotlin.reflect.KFunction; +import kotlin.reflect.KParameter; +import kotlin.reflect.KType; +import kotlin.reflect.KTypeParameter; +import kotlin.reflect.full.KClasses; +import kotlin.reflect.jvm.KTypesJvm; +import kotlin.reflect.jvm.ReflectJvmMapping; import kotlinx.coroutines.reactive.ReactiveFlowKt; import kotlinx.coroutines.reactor.MonoKt; import org.jspecify.annotations.Nullable; import org.reactivestreams.Publisher; import reactor.core.publisher.Mono; +import org.springframework.core.KotlinDetector; +import org.springframework.core.MethodParameter; +import org.springframework.util.ConcurrentReferenceHashMap; +import org.springframework.util.ReflectionUtils; + /** * Package-visible class designed to avoid a hard dependency on Kotlin and Coroutines dependency at runtime. * @@ -31,6 +50,21 @@ import reactor.core.publisher.Mono; */ abstract class CoroutinesUtils { + private static final String COROUTINES_FLOW_CLASS_NAME = "kotlinx.coroutines.flow.Flow"; + + private static final Method NO_UNBOX_METHOD; + + private static final Map unboxMethodCache = new ConcurrentReferenceHashMap<>(); + + static { + try { + NO_UNBOX_METHOD = CoroutinesUtils.class.getDeclaredMethod("noUnboxMethod"); + } + catch (NoSuchMethodException ex) { + throw new IllegalStateException("Expected method not found: " + ex); + } + } + static Object asFlow(@Nullable Object publisher) { if (publisher instanceof Publisher rsPublisher) { return ReactiveFlowKt.asFlow(rsPublisher); @@ -46,4 +80,109 @@ abstract class CoroutinesUtils { (Continuation) continuation); } + /** + * Adapt the return value of a proxied suspending function: convert it to a + * {@code Flow} or await its completion, unboxing Kotlin value classes when the + * caller expects their unboxed representation. + * @param method the suspending function + * @param returnValue the return value, typically a {@code Publisher} + * @param continuation the {@code Continuation} argument of the suspending function + * @return the adapted return value + */ + static @Nullable Object adaptReturnValue(Method method, @Nullable Object returnValue, Object continuation) { + MethodParameter returnParameter = new MethodParameter(method, -1); + Class returnParameterType = returnParameter.getParameterType(); + if (COROUTINES_FLOW_CLASS_NAME.equals(returnParameterType.getName())) { + return asFlow(returnValue); + } + Object result = awaitSingleOrNull(returnValue, continuation); + if (KotlinDetector.isInlineClass(returnParameterType) && returnParameterType.isInstance(result)) { + Method unboxMethod = unboxMethodCache.computeIfAbsent(method, key -> + findUnboxMethod(key, returnParameterType, returnParameter.isOptional())); + if (unboxMethod != NO_UNBOX_METHOD) { + return ReflectionUtils.invokeMethod(unboxMethod, result); + } + } + return result; + } + + private static Method findUnboxMethod(Method method, Class valueClass, boolean nullableReturnType) { + Method unboxMethod = ReflectionUtils.findMethod(valueClass, "unbox-impl"); + if (unboxMethod == null || unboxMethod.getReturnType().isPrimitive() || + (nullableReturnType && isUnderlyingTypeNullable(valueClass)) || + overridesFunctionWithDifferentReturnType(method, valueClass)) { + return NO_UNBOX_METHOD; + } + ReflectionUtils.makeAccessible(unboxMethod); + return unboxMethod; + } + + private static boolean overridesFunctionWithDifferentReturnType(Method method, Class valueClass) { + KFunction function = ReflectJvmMapping.getKotlinFunction(method); + if (function == null) { + return false; + } + KClass valueKClass = JvmClassMappingKt.getKotlinClass(valueClass); + for (KType superType : KClasses.getAllSupertypes(JvmClassMappingKt.getKotlinClass(method.getDeclaringClass()))) { + if (superType.getClassifier() instanceof KClass superClass) { + for (KFunction candidate : KClasses.getDeclaredMemberFunctions(superClass)) { + if (candidate.getName().equals(function.getName()) && candidate.isSuspend() && + hasSameParameterTypes(function, candidate, superType) && + candidate.getReturnType().getClassifier() != valueKClass) { + return true; + } + } + } + } + return false; + } + + private static boolean hasSameParameterTypes(KFunction function, KFunction candidate, KType superType) { + List parameters = function.getParameters(); + List candidateParameters = candidate.getParameters(); + if (parameters.size() != candidateParameters.size()) { + return false; + } + for (int i = 1; i < parameters.size(); i++) { + KType candidateType = candidateParameters.get(i).getType(); + if (candidateType.getClassifier() instanceof KTypeParameter typeParameter && + superType.getClassifier() instanceof KClass superClass) { + int index = superClass.getTypeParameters().indexOf(typeParameter); + KType argumentType = (index != -1 ? superType.getArguments().get(index).getType() : null); + if (argumentType != null) { + candidateType = argumentType; + } + } + if (!KTypesJvm.getJvmErasure(candidateType).equals(KTypesJvm.getJvmErasure(parameters.get(i).getType()))) { + return false; + } + } + return true; + } + + private static boolean isUnderlyingTypeNullable(Class valueClass) { + KFunction constructor = Objects.requireNonNull( + KClasses.getPrimaryConstructor(JvmClassMappingKt.getKotlinClass(valueClass))); + return isNullable(constructor.getParameters().get(0).getType()); + } + + private static boolean isNullable(KType type) { + if (type.isMarkedNullable()) { + return true; + } + if (type.getClassifier() instanceof KTypeParameter typeParameter) { + return typeParameter.getUpperBounds().stream().anyMatch(CoroutinesUtils::isNullable); + } + return (type.getClassifier() instanceof KClass kClass && + KotlinDetector.isInlineClass(JvmClassMappingKt.getJavaClass(kClass)) && + isUnderlyingTypeNullable(JvmClassMappingKt.getJavaClass(kClass))); + } + + /** + * For the {@link #NO_UNBOX_METHOD} constant. + */ + @SuppressWarnings("unused") + private static void noUnboxMethod() { + } + } diff --git a/spring-aop/src/main/java/org/springframework/aop/framework/JdkDynamicAopProxy.java b/spring-aop/src/main/java/org/springframework/aop/framework/JdkDynamicAopProxy.java index 23484340775..daf05614dbc 100644 --- a/spring-aop/src/main/java/org/springframework/aop/framework/JdkDynamicAopProxy.java +++ b/spring-aop/src/main/java/org/springframework/aop/framework/JdkDynamicAopProxy.java @@ -35,7 +35,6 @@ import org.springframework.aop.TargetSource; import org.springframework.aop.support.AopUtils; import org.springframework.core.DecoratingProxy; import org.springframework.core.KotlinDetector; -import org.springframework.core.MethodParameter; import org.springframework.util.Assert; import org.springframework.util.ClassUtils; @@ -73,8 +72,6 @@ final class JdkDynamicAopProxy implements AopProxy, InvocationHandler, Serializa private static final long serialVersionUID = 5531744639992436476L; - private static final String COROUTINES_FLOW_CLASS_NAME = "kotlinx.coroutines.flow.Flow"; - private static final boolean COROUTINES_REACTOR_PRESENT = ClassUtils.isPresent( "kotlinx.coroutines.reactor.MonoKt", JdkDynamicAopProxy.class.getClassLoader()); @@ -237,14 +234,7 @@ final class JdkDynamicAopProxy implements AopProxy, InvocationHandler, Serializa "Null return value from advice does not match primitive return type for: " + method); } if (COROUTINES_REACTOR_PRESENT && KotlinDetector.isSuspendingFunction(method)) { - Class returnParameterType = new MethodParameter(method, -1).getParameterType(); - if (COROUTINES_FLOW_CLASS_NAME.equals(returnParameterType.getName())) { - return CoroutinesUtils.asFlow(retVal); - } - Object awaitResult = CoroutinesUtils.awaitSingleOrNull(retVal, args[args.length - 1]); - return KotlinDetector.isInlineClass(returnParameterType) && - awaitResult != null && KotlinDetector.isInlineClass(awaitResult.getClass()) ? - awaitResult.getClass().getDeclaredMethod("unbox-impl").invoke(awaitResult) : awaitResult; + return CoroutinesUtils.adaptReturnValue(method, retVal, args[args.length - 1]); } return retVal; } diff --git a/spring-aop/src/test/kotlin/org/springframework/aop/framework/CglibAopProxyKotlinTests.kt b/spring-aop/src/test/kotlin/org/springframework/aop/framework/CglibAopProxyKotlinTests.kt index 55f8a1944ef..5b6477c16f1 100644 --- a/spring-aop/src/test/kotlin/org/springframework/aop/framework/CglibAopProxyKotlinTests.kt +++ b/spring-aop/src/test/kotlin/org/springframework/aop/framework/CglibAopProxyKotlinTests.kt @@ -149,6 +149,172 @@ class CglibAopProxyKotlinTests { assertThat(proxy.returnAny()).isEqualTo(ValueClass("bar")) } + @Test + suspend fun proxiedSuspendedInvocationValueClassPrimitiveValue() { + val proxyFactory = ProxyFactory(TestBean()) + proxyFactory.addAdvice(MethodInterceptor { + ValueClassPrimitiveValue(1) + }) + val proxy = proxyFactory.proxy as TestBean + assertThat(proxy.returnValueClassPrimitiveValue()).isEqualTo(ValueClassPrimitiveValue(1)) + } + + @Test + suspend fun proxiedSuspendedInvocationValueClassPrimitiveValueProceed() { + val proxyFactory = ProxyFactory(TestBean()) + proxyFactory.addAdvice(MethodInterceptor { + it.proceed() + }) + val proxy = proxyFactory.proxy as TestBean + assertThat(proxy.returnValueClassPrimitiveValue()).isEqualTo(ValueClassPrimitiveValue(0)) + } + + @Test + suspend fun proxiedSuspendedInvocationNullableValueClassPrimitiveValue() { + val proxyFactory = ProxyFactory(TestBean()) + proxyFactory.addAdvice(MethodInterceptor { + ValueClassPrimitiveValue(1) + }) + val proxy = proxyFactory.proxy as TestBean + assertThat(proxy.returnNullableValueClassPrimitiveValue()).isEqualTo(ValueClassPrimitiveValue(1)) + } + + @Test + suspend fun proxiedSuspendedInvocationNullableValueClassPrimitiveValueNull() { + val proxyFactory = ProxyFactory(TestBean()) + proxyFactory.addAdvice(MethodInterceptor { + null + }) + val proxy = proxyFactory.proxy as TestBean + assertThat(proxy.returnNullableValueClassPrimitiveValue()).isNull() + } + + @Test + suspend fun proxiedSuspendedInvocationNullableValueClassNullableValue() { + val proxyFactory = ProxyFactory(TestBean()) + proxyFactory.addAdvice(MethodInterceptor { + ValueClassNullableValue("bar") + }) + val proxy = proxyFactory.proxy as TestBean + assertThat(proxy.returnNullableValueClassNullableValue()).isEqualTo(ValueClassNullableValue("bar")) + } + + @Test + suspend fun proxiedSuspendedInvocationNullableValueClassNullableValueNull() { + val proxyFactory = ProxyFactory(TestBean()) + proxyFactory.addAdvice(MethodInterceptor { + ValueClassNullableValue(null) + }) + val proxy = proxyFactory.proxy as TestBean + assertThat(proxy.returnNullableValueClassNullableValue()).isEqualTo(ValueClassNullableValue(null)) + } + + @Test + suspend fun proxiedSuspendedInvocationNullableGenericValueClass() { + val proxyFactory = ProxyFactory(TestBean()) + proxyFactory.addAdvice(MethodInterceptor { + GenericValueClass("bar") + }) + val proxy = proxyFactory.proxy as TestBean + assertThat(proxy.returnNullableGenericValueClass()).isEqualTo(GenericValueClass("bar")) + } + + @Test + suspend fun proxiedSuspendedInvocationNestedValueClass() { + val proxyFactory = ProxyFactory(TestBean()) + proxyFactory.addAdvice(MethodInterceptor { + NestedValueClass(ValueClass("bar")) + }) + val proxy = proxyFactory.proxy as TestBean + assertThat(proxy.returnNestedValueClass()).isEqualTo(NestedValueClass(ValueClass("bar"))) + } + + @Test + suspend fun proxiedSuspendedInvocationNullableNestedValueClassPrimitiveValue() { + val proxyFactory = ProxyFactory(TestBean()) + proxyFactory.addAdvice(MethodInterceptor { + NestedValueClassPrimitiveValue(ValueClassPrimitiveValue(1)) + }) + val proxy = proxyFactory.proxy as TestBean + assertThat(proxy.returnNullableNestedValueClassPrimitiveValue()) + .isEqualTo(NestedValueClassPrimitiveValue(ValueClassPrimitiveValue(1))) + } + + @Test + suspend fun proxiedSuspendedInvocationNullableNestedValueClassNullableValue() { + val proxyFactory = ProxyFactory(TestBean()) + proxyFactory.addAdvice(MethodInterceptor { + NestedValueClassNullableValue(ValueClassNullableValue("bar")) + }) + val proxy = proxyFactory.proxy as TestBean + assertThat(proxy.returnNullableNestedValueClassNullableValue()) + .isEqualTo(NestedValueClassNullableValue(ValueClassNullableValue("bar"))) + } + + @Test + suspend fun proxiedSuspendedInvocationNullableNonNullGenericValueClass() { + val proxyFactory = ProxyFactory(TestBean()) + proxyFactory.addAdvice(MethodInterceptor { + NonNullGenericValueClass("bar") + }) + val proxy = proxyFactory.proxy as TestBean + assertThat(proxy.returnNullableNonNullGenericValueClass()).isEqualTo(NonNullGenericValueClass("bar")) + } + + @Test + suspend fun proxiedSuspendedInvocationNullableResult() { + val proxyFactory = ProxyFactory(TestBean()) + proxyFactory.addAdvice(MethodInterceptor { + Result.success("bar") + }) + val proxy = proxyFactory.proxy as TestBean + assertThat(proxy.returnNullableResult()?.getOrNull()).isEqualTo("bar") + } + + @Test + suspend fun proxiedSuspendedInvocationValueClassOverridingAny() { + val proxyFactory = ProxyFactory(ValueClassOverridingAnyBean()) + proxyFactory.isProxyTargetClass = true + proxyFactory.addAdvice(MethodInterceptor { + ValueClass("bar") + }) + val proxy = proxyFactory.proxy as ValueClassOverridingAnyBean + assertThat(proxy.returnValue()).isEqualTo(ValueClass("bar")) + } + + @Test + suspend fun proxiedSuspendedInvocationValueClassOverridingGeneric() { + val proxyFactory = ProxyFactory(ValueClassOverridingGenericBean()) + proxyFactory.isProxyTargetClass = true + proxyFactory.addAdvice(MethodInterceptor { + ValueClass("bar") + }) + val proxy = proxyFactory.proxy as ValueClassOverridingGenericBean + assertThat(proxy.returnValue()).isEqualTo(ValueClass("bar")) + } + + @Test + suspend fun proxiedSuspendedInvocationValueClassOverridingGenericParameter() { + val proxyFactory = ProxyFactory(ValueClassOverridingGenericParameterBean()) + proxyFactory.isProxyTargetClass = true + proxyFactory.addAdvice(MethodInterceptor { + ValueClass("bar") + }) + val proxy = proxyFactory.proxy as ValueClassOverridingGenericParameterBean + assertThat(proxy.returnValue("foo")).isEqualTo(ValueClass("bar")) + } + + @Test + suspend fun proxiedSuspendedInvocationValueClassOverloadingGenericParameter() { + val proxyFactory = ProxyFactory(ValueClassOverloadingGenericParameterBean()) + proxyFactory.isProxyTargetClass = true + proxyFactory.addAdvice(MethodInterceptor { + ValueClass("bar") + }) + val proxy = proxyFactory.proxy as ValueClassOverloadingGenericParameterBean + assertThat(proxy.returnValue("foo").value).isEqualTo("bar") + } + open class MyKotlinBean { open fun capitalize(value: String) = value.uppercase() @@ -183,42 +349,172 @@ class CglibAopProxyKotlinTests { val updatedAt: LocalDateTime? = null, ) + @Test + suspend fun proxiedSuspendedInvocationPrivateValueClass() { + val proxyFactory = ProxyFactory(PrivateValueClassBean()) + proxyFactory.isProxyTargetClass = true + proxyFactory.addAdvice(MethodInterceptor { + PrivateValueClass("bar") + }) + val proxy = proxyFactory.proxy as PrivateValueClassBean + assertThat(proxy.returnValue()).isEqualTo(PrivateValueClass("bar")) + } + @JvmInline value class ValueClass(val value: String) @JvmInline value class ValueClassNullableValue(val value: String?) + @JvmInline + value class ValueClassPrimitiveValue(val value: Int) + + @JvmInline + value class GenericValueClass(val value: T) + + @JvmInline + value class NestedValueClass(val value: ValueClass) + + @JvmInline + value class NestedValueClassPrimitiveValue(val value: ValueClassPrimitiveValue) + + @JvmInline + value class NestedValueClassNullableValue(val value: ValueClassNullableValue) + + @JvmInline + value class NonNullGenericValueClass(val value: T) + open class TestBean { open suspend fun returnValueClass(): ValueClass { - delay(1000.milliseconds) + delay(10.milliseconds) return ValueClass("foo") } open suspend fun returnNullableValueClass(): ValueClass? { - delay(1000.milliseconds) + delay(10.milliseconds) return null } open suspend fun returnValueClassNullableValue(): ValueClassNullableValue { - delay(1000.milliseconds) + delay(10.milliseconds) return ValueClassNullableValue(null) } open suspend fun returnResult(): Result { - delay(1000.milliseconds) + delay(10.milliseconds) return Result.success("foo") } open suspend fun returnString(): String { - delay(1000.milliseconds) + delay(10.milliseconds) return "foo" } open suspend fun returnAny(): Any { - delay(1000.milliseconds) + delay(10.milliseconds) + return ValueClass("foo") + } + + open suspend fun returnValueClassPrimitiveValue(): ValueClassPrimitiveValue { + delay(10.milliseconds) + return ValueClassPrimitiveValue(0) + } + + open suspend fun returnNullableValueClassPrimitiveValue(): ValueClassPrimitiveValue? { + delay(10.milliseconds) + return null + } + + open suspend fun returnNullableValueClassNullableValue(): ValueClassNullableValue? { + delay(10.milliseconds) + return null + } + + open suspend fun returnNullableGenericValueClass(): GenericValueClass? { + delay(10.milliseconds) + return null + } + + open suspend fun returnNestedValueClass(): NestedValueClass { + delay(10.milliseconds) + return NestedValueClass(ValueClass("foo")) + } + + open suspend fun returnNullableNestedValueClassPrimitiveValue(): NestedValueClassPrimitiveValue? { + delay(10.milliseconds) + return null + } + + open suspend fun returnNullableNestedValueClassNullableValue(): NestedValueClassNullableValue? { + delay(10.milliseconds) + return null + } + + open suspend fun returnNullableNonNullGenericValueClass(): NonNullGenericValueClass? { + delay(10.milliseconds) + return null + } + + open suspend fun returnNullableResult(): Result? { + delay(10.milliseconds) + return null + } + } + + + interface AnyBean { + suspend fun returnValue(): Any + } + + open class ValueClassOverridingAnyBean : AnyBean { + override suspend fun returnValue(): ValueClass { + delay(10.milliseconds) return ValueClass("foo") } } + interface GenericBean { + suspend fun returnValue(): T + } + + open class ValueClassOverridingGenericBean : GenericBean { + override suspend fun returnValue(): ValueClass { + delay(10.milliseconds) + return ValueClass("foo") + } + } + + interface GenericParameterBean { + suspend fun returnValue(value: T): Any + } + + open class ValueClassOverridingGenericParameterBean : GenericParameterBean { + override suspend fun returnValue(value: String): ValueClass { + delay(10.milliseconds) + return ValueClass(value) + } + } + + open class ValueClassOverloadingGenericParameterBean : GenericParameterBean { + override suspend fun returnValue(value: Int): Any { + delay(10.milliseconds) + return value + } + + open suspend fun returnValue(value: String): ValueClass { + delay(10.milliseconds) + return ValueClass(value) + } + } + + @JvmInline + private value class PrivateValueClass(val value: String) + + private open class PrivateValueClassBean { + open suspend fun returnValue(): PrivateValueClass { + delay(10.milliseconds) + return PrivateValueClass("foo") + } + } + } diff --git a/spring-aop/src/test/kotlin/org/springframework/aop/framework/JdkDynamicAopProxyKotlinTests.kt b/spring-aop/src/test/kotlin/org/springframework/aop/framework/JdkDynamicAopProxyKotlinTests.kt index 75d4e69eb4d..9c4122638b8 100644 --- a/spring-aop/src/test/kotlin/org/springframework/aop/framework/JdkDynamicAopProxyKotlinTests.kt +++ b/spring-aop/src/test/kotlin/org/springframework/aop/framework/JdkDynamicAopProxyKotlinTests.kt @@ -26,6 +26,7 @@ import kotlin.time.Duration.Companion.milliseconds * Tests for Kotlin support in [JdkDynamicAopProxy]. * * @author Dmitry Sulman + * @author Sebastien Deleuze */ class JdkDynamicAopProxyKotlinTests { @@ -119,12 +120,182 @@ class JdkDynamicAopProxyKotlinTests { assertThat(proxy.returnAny()).isEqualTo(ValueClass("bar")) } + @Test + suspend fun proxiedSuspendedInvocationValueClassPrimitiveValue() { + val proxyFactory = ProxyFactory(TestBeanImpl()) + proxyFactory.addAdvice(MethodInterceptor { + ValueClassPrimitiveValue(1) + }) + val proxy = proxyFactory.proxy as TestBean + assertThat(proxy.returnValueClassPrimitiveValue()).isEqualTo(ValueClassPrimitiveValue(1)) + } + + @Test + suspend fun proxiedSuspendedInvocationValueClassPrimitiveValueProceed() { + val proxyFactory = ProxyFactory(TestBeanImpl()) + proxyFactory.addAdvice(MethodInterceptor { + it.proceed() + }) + val proxy = proxyFactory.proxy as TestBean + assertThat(proxy.returnValueClassPrimitiveValue()).isEqualTo(ValueClassPrimitiveValue(0)) + } + + @Test + suspend fun proxiedSuspendedInvocationNullableValueClassPrimitiveValue() { + val proxyFactory = ProxyFactory(TestBeanImpl()) + proxyFactory.addAdvice(MethodInterceptor { + ValueClassPrimitiveValue(1) + }) + val proxy = proxyFactory.proxy as TestBean + assertThat(proxy.returnNullableValueClassPrimitiveValue()).isEqualTo(ValueClassPrimitiveValue(1)) + } + + @Test + suspend fun proxiedSuspendedInvocationNullableValueClassPrimitiveValueNull() { + val proxyFactory = ProxyFactory(TestBeanImpl()) + proxyFactory.addAdvice(MethodInterceptor { + null + }) + val proxy = proxyFactory.proxy as TestBean + assertThat(proxy.returnNullableValueClassPrimitiveValue()).isNull() + } + + @Test + suspend fun proxiedSuspendedInvocationNullableValueClassNullableValue() { + val proxyFactory = ProxyFactory(TestBeanImpl()) + proxyFactory.addAdvice(MethodInterceptor { + ValueClassNullableValue("bar") + }) + val proxy = proxyFactory.proxy as TestBean + assertThat(proxy.returnNullableValueClassNullableValue()).isEqualTo(ValueClassNullableValue("bar")) + } + + @Test + suspend fun proxiedSuspendedInvocationNullableValueClassNullableValueNull() { + val proxyFactory = ProxyFactory(TestBeanImpl()) + proxyFactory.addAdvice(MethodInterceptor { + ValueClassNullableValue(null) + }) + val proxy = proxyFactory.proxy as TestBean + assertThat(proxy.returnNullableValueClassNullableValue()).isEqualTo(ValueClassNullableValue(null)) + } + + @Test + suspend fun proxiedSuspendedInvocationNullableGenericValueClass() { + val proxyFactory = ProxyFactory(TestBeanImpl()) + proxyFactory.addAdvice(MethodInterceptor { + GenericValueClass("bar") + }) + val proxy = proxyFactory.proxy as TestBean + assertThat(proxy.returnNullableGenericValueClass()).isEqualTo(GenericValueClass("bar")) + } + + @Test + suspend fun proxiedSuspendedInvocationNestedValueClass() { + val proxyFactory = ProxyFactory(TestBeanImpl()) + proxyFactory.addAdvice(MethodInterceptor { + NestedValueClass(ValueClass("bar")) + }) + val proxy = proxyFactory.proxy as TestBean + assertThat(proxy.returnNestedValueClass()).isEqualTo(NestedValueClass(ValueClass("bar"))) + } + + @Test + suspend fun proxiedSuspendedInvocationNullableNestedValueClassPrimitiveValue() { + val proxyFactory = ProxyFactory(TestBeanImpl()) + proxyFactory.addAdvice(MethodInterceptor { + NestedValueClassPrimitiveValue(ValueClassPrimitiveValue(1)) + }) + val proxy = proxyFactory.proxy as TestBean + assertThat(proxy.returnNullableNestedValueClassPrimitiveValue()) + .isEqualTo(NestedValueClassPrimitiveValue(ValueClassPrimitiveValue(1))) + } + + @Test + suspend fun proxiedSuspendedInvocationNullableNestedValueClassNullableValue() { + val proxyFactory = ProxyFactory(TestBeanImpl()) + proxyFactory.addAdvice(MethodInterceptor { + NestedValueClassNullableValue(ValueClassNullableValue("bar")) + }) + val proxy = proxyFactory.proxy as TestBean + assertThat(proxy.returnNullableNestedValueClassNullableValue()) + .isEqualTo(NestedValueClassNullableValue(ValueClassNullableValue("bar"))) + } + + @Test + suspend fun proxiedSuspendedInvocationNullableNonNullGenericValueClass() { + val proxyFactory = ProxyFactory(TestBeanImpl()) + proxyFactory.addAdvice(MethodInterceptor { + NonNullGenericValueClass("bar") + }) + val proxy = proxyFactory.proxy as TestBean + assertThat(proxy.returnNullableNonNullGenericValueClass()).isEqualTo(NonNullGenericValueClass("bar")) + } + + @Test + suspend fun proxiedSuspendedInvocationNullableResult() { + val proxyFactory = ProxyFactory(TestBeanImpl()) + proxyFactory.addAdvice(MethodInterceptor { + Result.success("bar") + }) + val proxy = proxyFactory.proxy as TestBean + assertThat(proxy.returnNullableResult()?.getOrNull()).isEqualTo("bar") + } + + @Test + suspend fun proxiedSuspendedInvocationValueClassOverridingAny() { + val proxyFactory = ProxyFactory(ValueClassOverridingAnyBeanImpl()) + proxyFactory.addAdvice(MethodInterceptor { + ValueClass("bar") + }) + val proxy = proxyFactory.proxy as ValueClassOverridingAnyBean + assertThat(proxy.returnValue()).isEqualTo(ValueClass("bar")) + } + + @Test + suspend fun proxiedSuspendedInvocationValueClassOverridingGeneric() { + val proxyFactory = ProxyFactory(ValueClassOverridingGenericBeanImpl()) + proxyFactory.addAdvice(MethodInterceptor { + ValueClass("bar") + }) + val proxy = proxyFactory.proxy as ValueClassOverridingGenericBean + assertThat(proxy.returnValue()).isEqualTo(ValueClass("bar")) + } + + @Test + suspend fun proxiedSuspendedInvocationPrivateValueClass() { + val proxyFactory = ProxyFactory(PrivateValueClassBeanImpl()) + proxyFactory.addAdvice(MethodInterceptor { + PrivateValueClass("bar") + }) + val proxy = proxyFactory.proxy as PrivateValueClassBean + assertThat(proxy.returnValue()).isEqualTo(PrivateValueClass("bar")) + } + @JvmInline value class ValueClass(val value: String) @JvmInline value class ValueClassNullableValue(val value: String?) + @JvmInline + value class ValueClassPrimitiveValue(val value: Int) + + @JvmInline + value class GenericValueClass(val value: T) + + @JvmInline + value class NestedValueClass(val value: ValueClass) + + @JvmInline + value class NestedValueClassPrimitiveValue(val value: ValueClassPrimitiveValue) + + @JvmInline + value class NestedValueClassNullableValue(val value: ValueClassNullableValue) + + @JvmInline + value class NonNullGenericValueClass(val value: T) + interface TestBean { suspend fun returnValueClass(): ValueClass @@ -137,38 +308,146 @@ class JdkDynamicAopProxyKotlinTests { suspend fun returnString(): String suspend fun returnAny(): Any + + suspend fun returnValueClassPrimitiveValue(): ValueClassPrimitiveValue + + suspend fun returnNullableValueClassPrimitiveValue(): ValueClassPrimitiveValue? + + suspend fun returnNullableValueClassNullableValue(): ValueClassNullableValue? + + suspend fun returnNullableGenericValueClass(): GenericValueClass? + + suspend fun returnNestedValueClass(): NestedValueClass + + suspend fun returnNullableNestedValueClassPrimitiveValue(): NestedValueClassPrimitiveValue? + + suspend fun returnNullableNestedValueClassNullableValue(): NestedValueClassNullableValue? + + suspend fun returnNullableNonNullGenericValueClass(): NonNullGenericValueClass? + + suspend fun returnNullableResult(): Result? } class TestBeanImpl : TestBean { override suspend fun returnValueClass(): ValueClass { - delay(1000.milliseconds) + delay(10.milliseconds) return ValueClass("foo") } override suspend fun returnNullableValueClass(): ValueClass? { - delay(1000.milliseconds) + delay(10.milliseconds) return null } override suspend fun returnValueClassNullableValue(): ValueClassNullableValue { - delay(1000.milliseconds) + delay(10.milliseconds) return ValueClassNullableValue(null) } override suspend fun returnResult(): Result { - delay(1000.milliseconds) + delay(10.milliseconds) return Result.success("foo") } override suspend fun returnString(): String { - delay(1000.milliseconds) + delay(10.milliseconds) return "foo" } override suspend fun returnAny(): Any { - delay(1000.milliseconds) + delay(10.milliseconds) + return ValueClass("foo") + } + + override suspend fun returnValueClassPrimitiveValue(): ValueClassPrimitiveValue { + delay(10.milliseconds) + return ValueClassPrimitiveValue(0) + } + + override suspend fun returnNullableValueClassPrimitiveValue(): ValueClassPrimitiveValue? { + delay(10.milliseconds) + return null + } + + override suspend fun returnNullableValueClassNullableValue(): ValueClassNullableValue? { + delay(10.milliseconds) + return null + } + + override suspend fun returnNullableGenericValueClass(): GenericValueClass? { + delay(10.milliseconds) + return null + } + + override suspend fun returnNestedValueClass(): NestedValueClass { + delay(10.milliseconds) + return NestedValueClass(ValueClass("foo")) + } + + override suspend fun returnNullableNestedValueClassPrimitiveValue(): NestedValueClassPrimitiveValue? { + delay(10.milliseconds) + return null + } + + override suspend fun returnNullableNestedValueClassNullableValue(): NestedValueClassNullableValue? { + delay(10.milliseconds) + return null + } + + override suspend fun returnNullableNonNullGenericValueClass(): NonNullGenericValueClass? { + delay(10.milliseconds) + return null + } + + override suspend fun returnNullableResult(): Result? { + delay(10.milliseconds) + return null + } + } + + + interface AnyBean { + suspend fun returnValue(): Any + } + + interface ValueClassOverridingAnyBean : AnyBean { + override suspend fun returnValue(): ValueClass + } + + class ValueClassOverridingAnyBeanImpl : ValueClassOverridingAnyBean { + override suspend fun returnValue(): ValueClass { + delay(10.milliseconds) return ValueClass("foo") } } -} \ No newline at end of file + interface GenericBean { + suspend fun returnValue(): T + } + + interface ValueClassOverridingGenericBean : GenericBean { + override suspend fun returnValue(): ValueClass + } + + class ValueClassOverridingGenericBeanImpl : ValueClassOverridingGenericBean { + override suspend fun returnValue(): ValueClass { + delay(10.milliseconds) + return ValueClass("foo") + } + } + + @JvmInline + private value class PrivateValueClass(val value: String) + + private interface PrivateValueClassBean { + suspend fun returnValue(): PrivateValueClass + } + + private class PrivateValueClassBeanImpl : PrivateValueClassBean { + override suspend fun returnValue(): PrivateValueClass { + delay(10.milliseconds) + return PrivateValueClass("foo") + } + } + +}