diff --git a/spring-core/src/main/java/org/springframework/core/CoroutinesUtils.java b/spring-core/src/main/java/org/springframework/core/CoroutinesUtils.java index 568d94b6c03..020a00064c9 100644 --- a/spring-core/src/main/java/org/springframework/core/CoroutinesUtils.java +++ b/spring-core/src/main/java/org/springframework/core/CoroutinesUtils.java @@ -134,9 +134,10 @@ public abstract class CoroutinesUtils { Object arg = args[index]; if (!(parameter.isOptional() && arg == null)) { KType type = parameter.getType(); - if (!type.isMarkedNullable() && + if (!(type.isMarkedNullable() && arg == null) && type.getClassifier() instanceof KClass kClass && - KotlinDetector.isInlineClass(JvmClassMappingKt.getJavaClass(kClass))) { + KotlinDetector.isInlineClass(JvmClassMappingKt.getJavaClass(kClass)) && + !JvmClassMappingKt.getJavaClass(kClass).isInstance(arg)) { arg = box(kClass, arg); } argMap.put(parameter, arg); @@ -166,9 +167,10 @@ public abstract class CoroutinesUtils { private static Object box(KClass kClass, @Nullable Object arg) { KFunction constructor = Objects.requireNonNull(KClasses.getPrimaryConstructor(kClass)); KType type = constructor.getParameters().get(0).getType(); - if (!type.isMarkedNullable() && + if (!(type.isMarkedNullable() && arg == null) && type.getClassifier() instanceof KClass parameterClass && - KotlinDetector.isInlineClass(JvmClassMappingKt.getJavaClass(parameterClass))) { + KotlinDetector.isInlineClass(JvmClassMappingKt.getJavaClass(parameterClass)) && + !JvmClassMappingKt.getJavaClass(parameterClass).isInstance(arg)) { arg = box(parameterClass, arg); } if (!KCallablesJvm.isAccessible(constructor)) { diff --git a/spring-core/src/main/java/org/springframework/core/ReactiveAdapterRegistry.java b/spring-core/src/main/java/org/springframework/core/ReactiveAdapterRegistry.java index 7d66979afb2..5be957ecc2b 100644 --- a/spring-core/src/main/java/org/springframework/core/ReactiveAdapterRegistry.java +++ b/spring-core/src/main/java/org/springframework/core/ReactiveAdapterRegistry.java @@ -71,6 +71,8 @@ public class ReactiveAdapterRegistry { private static final boolean MUTINY_PRESENT; + private static final boolean CONTEXT_PROPAGATION_PRESENT; + static { ClassLoader classLoader = ReactiveAdapterRegistry.class.getClassLoader(); REACTIVE_STREAMS_PRESENT = ClassUtils.isPresent("org.reactivestreams.Publisher", classLoader); @@ -78,6 +80,7 @@ public class ReactiveAdapterRegistry { RXJAVA_3_PRESENT = ClassUtils.isPresent("io.reactivex.rxjava3.core.Flowable", classLoader); COROUTINES_REACTOR_PRESENT = ClassUtils.isPresent("kotlinx.coroutines.reactor.MonoKt", classLoader); MUTINY_PRESENT = ClassUtils.isPresent("io.smallrye.mutiny.Multi", classLoader); + CONTEXT_PROPAGATION_PRESENT = ClassUtils.isPresent("io.micrometer.context.ContextSnapshotFactory", classLoader); } private final List adapters = new ArrayList<>(); @@ -356,7 +359,9 @@ public class ReactiveAdapterRegistry { registry.registerReactiveType( ReactiveTypeDescriptor.multiValue(kotlinx.coroutines.flow.Flow.class, kotlinx.coroutines.flow.FlowKt::emptyFlow), - source -> kotlinx.coroutines.reactor.ReactorFlowKt.asFlux((kotlinx.coroutines.flow.Flow) source), + CONTEXT_PROPAGATION_PRESENT ? + source -> kotlinx.coroutines.reactor.ReactorFlowKt.asFlux((kotlinx.coroutines.flow.Flow) source, new PropagationContextElement()) : + source -> kotlinx.coroutines.reactor.ReactorFlowKt.asFlux((kotlinx.coroutines.flow.Flow) source), kotlinx.coroutines.reactive.ReactiveFlowKt::asFlow); } } diff --git a/spring-core/src/test/kotlin/org/springframework/core/CoroutinesUtilsTests.kt b/spring-core/src/test/kotlin/org/springframework/core/CoroutinesUtilsTests.kt index 4cc2fdeefdb..0f3f4312d2c 100644 --- a/spring-core/src/test/kotlin/org/springframework/core/CoroutinesUtilsTests.kt +++ b/spring-core/src/test/kotlin/org/springframework/core/CoroutinesUtilsTests.kt @@ -236,6 +236,13 @@ class CoroutinesUtilsTests { Assertions.assertThat(mono.awaitSingleOrNull()).isEqualTo("foo") } + @Test + suspend fun invokeSuspendingFunctionWithNullableValueClassParameterAndUnderlyingValue() { + val method = CoroutinesUtilsTests::class.java.declaredMethods.first { it.name.startsWith("suspendingFunctionWithNullableValueClass") } + val mono = CoroutinesUtils.invokeSuspendingFunction(method, this, "foo", null) as Mono + Assertions.assertThat(mono.awaitSingleOrNull()).isEqualTo("foo") + } + @Test suspend fun invokeSuspendingFunctionWithNullableValueClassParameter() { val method = CoroutinesUtilsTests::class.java.declaredMethods.first { it.name.startsWith("suspendingFunctionWithNullableValueClass") } diff --git a/spring-core/src/test/kotlin/org/springframework/core/ReactiveAdapterRegistryKotlinTests.kt b/spring-core/src/test/kotlin/org/springframework/core/ReactiveAdapterRegistryKotlinTests.kt index 817fecbed5b..26b244256f6 100644 --- a/spring-core/src/test/kotlin/org/springframework/core/ReactiveAdapterRegistryKotlinTests.kt +++ b/spring-core/src/test/kotlin/org/springframework/core/ReactiveAdapterRegistryKotlinTests.kt @@ -16,13 +16,18 @@ package org.springframework.core +import io.micrometer.observation.Observation +import io.micrometer.observation.tck.TestObservationRegistry import kotlinx.coroutines.Deferred import kotlinx.coroutines.DelicateCoroutinesApi +import kotlinx.coroutines.Dispatchers import kotlinx.coroutines.GlobalScope import kotlinx.coroutines.async import kotlinx.coroutines.flow.Flow import kotlinx.coroutines.flow.flow import kotlinx.coroutines.flow.toList +import kotlinx.coroutines.reactive.awaitSingle +import kotlinx.coroutines.runBlocking import org.assertj.core.api.Assertions.assertThat import org.junit.jupiter.api.Test import org.reactivestreams.Publisher @@ -40,6 +45,8 @@ import kotlin.reflect.KClass @OptIn(DelicateCoroutinesApi::class) class ReactiveAdapterRegistryKotlinTests { + private val observationRegistry = TestObservationRegistry.create() + private val registry = ReactiveAdapterRegistry.getSharedInstance() @Test @@ -82,6 +89,23 @@ class ReactiveAdapterRegistryKotlinTests { assertThat((target as Flow<*>).toList()).contains(1, 2, 3) } + @Test + fun propagateMicrometerContextToFlow() { + val source = flow { + val currentObservation = observationRegistry.currentObservation + assertThat(currentObservation).isNotNull + emit(currentObservation?.context?.name) + } + val observation = Observation.createNotStarted("coroutine", observationRegistry) + observation.observe { + val target: Publisher = getAdapter(Flow::class).toPublisher(source) + val result = runBlocking(Dispatchers.IO) { + target.awaitSingle() + } + assertThat(result).isEqualTo("coroutine") + } + } + private fun getAdapter(reactiveType: KClass<*>): ReactiveAdapter { return this.registry.getAdapter(reactiveType.java)!! } diff --git a/spring-web/src/main/java/org/springframework/web/method/support/InvocableHandlerMethod.java b/spring-web/src/main/java/org/springframework/web/method/support/InvocableHandlerMethod.java index 6dca729698e..1a1ef3a011a 100644 --- a/spring-web/src/main/java/org/springframework/web/method/support/InvocableHandlerMethod.java +++ b/spring-web/src/main/java/org/springframework/web/method/support/InvocableHandlerMethod.java @@ -316,9 +316,10 @@ public class InvocableHandlerMethod extends HandlerMethod { Object arg = args[index]; if (!(parameter.isOptional() && arg == null)) { KType type = parameter.getType(); - if (!type.isMarkedNullable() && + if (!(type.isMarkedNullable() && arg == null) && type.getClassifier() instanceof KClass kClass && - KotlinDetector.isInlineClass(JvmClassMappingKt.getJavaClass(kClass))) { + KotlinDetector.isInlineClass(JvmClassMappingKt.getJavaClass(kClass)) && + !JvmClassMappingKt.getJavaClass(kClass).isInstance(arg)) { arg = box(kClass, arg); } argMap.put(parameter, arg); @@ -337,9 +338,10 @@ public class InvocableHandlerMethod extends HandlerMethod { private static Object box(KClass kClass, @Nullable Object arg) { KFunction constructor = Objects.requireNonNull(KClasses.getPrimaryConstructor(kClass)); KType type = constructor.getParameters().get(0).getType(); - if (!type.isMarkedNullable() && + if (!(type.isMarkedNullable() && arg == null) && type.getClassifier() instanceof KClass parameterClass && - KotlinDetector.isInlineClass(JvmClassMappingKt.getJavaClass(parameterClass))) { + KotlinDetector.isInlineClass(JvmClassMappingKt.getJavaClass(parameterClass)) && + !JvmClassMappingKt.getJavaClass(parameterClass).isInstance(arg)) { arg = box(parameterClass, arg); } if (!KCallablesJvm.isAccessible(constructor)) { diff --git a/spring-web/src/test/kotlin/org/springframework/web/method/support/InvocableHandlerMethodKotlinTests.kt b/spring-web/src/test/kotlin/org/springframework/web/method/support/InvocableHandlerMethodKotlinTests.kt index b7da2b669f2..a3dbd8f91a4 100644 --- a/spring-web/src/test/kotlin/org/springframework/web/method/support/InvocableHandlerMethodKotlinTests.kt +++ b/spring-web/src/test/kotlin/org/springframework/web/method/support/InvocableHandlerMethodKotlinTests.kt @@ -155,6 +155,13 @@ class InvocableHandlerMethodKotlinTests { Assertions.assertThat(value).isEqualTo(1L) } + @Test + fun valueClassWithNullableAndUnderlyingValue() { + composite.addResolver(StubArgumentResolver(LongValueClass::class.java, 1L)) + val value = getInvocable(ValueClassHandler::valueClassWithNullable.javaMethod!!).invokeForRequest(request, null) + Assertions.assertThat(value).isEqualTo(1L) + } + @Test fun valueClassWithNullable() { composite.addResolver(StubArgumentResolver(LongValueClass::class.java, null)) @@ -215,6 +222,14 @@ class InvocableHandlerMethodKotlinTests { StepVerifier.create(value as Mono).verifyComplete() } + @Test + fun suspendingValueClassWithNullableAndUnderlyingValue() { + composite.addResolver(ContinuationHandlerMethodArgumentResolver()) + composite.addResolver(StubArgumentResolver(LongValueClass::class.java, 1L)) + val value = getInvocable(SuspendingValueClassHandler::valueClassWithNullable.javaMethod!!).invokeForRequest(request, null) + StepVerifier.create(value as Mono).expectNext(1L).verifyComplete() + } + @Test fun suspendingValueClassWithPrivateConstructor() { composite.addResolver(ContinuationHandlerMethodArgumentResolver()) diff --git a/spring-webflux/src/main/java/org/springframework/web/reactive/result/method/InvocableHandlerMethod.java b/spring-webflux/src/main/java/org/springframework/web/reactive/result/method/InvocableHandlerMethod.java index c06ea62a2aa..11e312bbe8b 100644 --- a/spring-webflux/src/main/java/org/springframework/web/reactive/result/method/InvocableHandlerMethod.java +++ b/spring-webflux/src/main/java/org/springframework/web/reactive/result/method/InvocableHandlerMethod.java @@ -356,9 +356,10 @@ public class InvocableHandlerMethod extends HandlerMethod { Object arg = args[index]; if (!(parameter.isOptional() && arg == null)) { KType type = parameter.getType(); - if (!type.isMarkedNullable() && + if (!(type.isMarkedNullable() && arg == null) && type.getClassifier() instanceof KClass kClass && - KotlinDetector.isInlineClass(JvmClassMappingKt.getJavaClass(kClass))) { + KotlinDetector.isInlineClass(JvmClassMappingKt.getJavaClass(kClass)) && + !JvmClassMappingKt.getJavaClass(kClass).isInstance(arg)) { arg = box(kClass, arg); } argMap.put(parameter, arg); @@ -378,9 +379,10 @@ public class InvocableHandlerMethod extends HandlerMethod { private static Object box(KClass kClass, @Nullable Object arg) { KFunction constructor = Objects.requireNonNull(KClasses.getPrimaryConstructor(kClass)); KType type = constructor.getParameters().get(0).getType(); - if (!type.isMarkedNullable() && + if (!(type.isMarkedNullable() && arg == null) && type.getClassifier() instanceof KClass parameterClass && - KotlinDetector.isInlineClass(JvmClassMappingKt.getJavaClass(parameterClass))) { + KotlinDetector.isInlineClass(JvmClassMappingKt.getJavaClass(parameterClass)) && + !JvmClassMappingKt.getJavaClass(parameterClass).isInstance(arg)) { arg = box(parameterClass, arg); } if (!KCallablesJvm.isAccessible(constructor)) { diff --git a/spring-webflux/src/test/kotlin/org/springframework/web/reactive/result/InvocableHandlerMethodKotlinTests.kt b/spring-webflux/src/test/kotlin/org/springframework/web/reactive/result/InvocableHandlerMethodKotlinTests.kt index 9e41a222399..f8041148619 100644 --- a/spring-webflux/src/test/kotlin/org/springframework/web/reactive/result/InvocableHandlerMethodKotlinTests.kt +++ b/spring-webflux/src/test/kotlin/org/springframework/web/reactive/result/InvocableHandlerMethodKotlinTests.kt @@ -258,6 +258,14 @@ class InvocableHandlerMethodKotlinTests { assertHandlerResultValue(result, "1") } + @Test + fun valueClassWithNullableAndUnderlyingValue() { + this.resolvers.add(stubResolver(1L, LongValueClass::class.java)) + val method = ValueClassController::valueClassWithNullable.javaMethod!! + val result = invoke(ValueClassController(), method) + assertHandlerResultValue(result, "1") + } + @Test fun valueClassWithNullable() { this.resolvers.add(stubResolver(null, LongValueClass::class.java)) @@ -320,6 +328,14 @@ class InvocableHandlerMethodKotlinTests { assertHandlerResultValue(result, "null") } + @Test + fun suspendingValueClassWithNullableAndUnderlyingValue() { + this.resolvers.add(stubResolver(1L, LongValueClass::class.java)) + val method = SuspendingValueClassController::valueClassWithNullable.javaMethod!! + val result = invoke(SuspendingValueClassController(), method) + assertHandlerResultValue(result, "1") + } + @Test fun suspendingValueClassWithPrivateConstructor() { this.resolvers.add(stubResolver(1L, Long::class.java)) @@ -590,4 +606,4 @@ class InvocableHandlerMethodKotlinTests { } class CustomException(message: String) : Throwable(message) -} \ No newline at end of file +} diff --git a/spring-webflux/src/test/kotlin/org/springframework/web/reactive/result/method/annotation/CoroutinesIntegrationTests.kt b/spring-webflux/src/test/kotlin/org/springframework/web/reactive/result/method/annotation/CoroutinesIntegrationTests.kt index b932e5204b4..6850468bb22 100644 --- a/spring-webflux/src/test/kotlin/org/springframework/web/reactive/result/method/annotation/CoroutinesIntegrationTests.kt +++ b/spring-webflux/src/test/kotlin/org/springframework/web/reactive/result/method/annotation/CoroutinesIntegrationTests.kt @@ -16,18 +16,16 @@ package org.springframework.web.reactive.result.method.annotation -import kotlinx.coroutines.Deferred -import kotlinx.coroutines.DelicateCoroutinesApi -import kotlinx.coroutines.GlobalScope -import kotlinx.coroutines.async -import kotlinx.coroutines.delay +import io.micrometer.observation.ObservationRegistry +import io.micrometer.observation.tck.TestObservationRegistry +import kotlinx.coroutines.* import kotlinx.coroutines.flow.Flow import kotlinx.coroutines.flow.flow import org.assertj.core.api.Assertions.assertThat import org.assertj.core.api.Assertions.assertThatExceptionOfType -import org.junit.jupiter.api.Assumptions.assumeFalse import org.springframework.context.ApplicationContext import org.springframework.context.annotation.AnnotationConfigApplicationContext +import org.springframework.context.annotation.Bean import org.springframework.context.annotation.ComponentScan import org.springframework.context.annotation.Configuration import org.springframework.http.HttpHeaders @@ -86,6 +84,15 @@ class CoroutinesIntegrationTests : AbstractRequestMappingIntegrationTests() { assertThat(entity.body).isEqualTo("foobar") } + @ParameterizedHttpServerTest + fun `Handler method returning Flow with observation`(httpServer: HttpServer) { + startServer(httpServer) + + val entity = performGet("/flow-observation", HttpHeaders.EMPTY, String::class.java) + assertThat(entity.statusCode).isEqualTo(HttpStatus.OK) + assertThat(entity.body).isEqualTo("http.server.requests") + } + @ParameterizedHttpServerTest fun `Suspending handler method returning Flow`(httpServer: HttpServer) { startServer(httpServer) @@ -135,11 +142,16 @@ class CoroutinesIntegrationTests : AbstractRequestMappingIntegrationTests() { @Configuration @EnableWebFlux @ComponentScan(resourcePattern = "**/CoroutinesIntegrationTests*") - open class WebConfig + open class WebConfig { + + @Bean + open fun observationRegistry(): ObservationRegistry = TestObservationRegistry.create() + + } @OptIn(DelicateCoroutinesApi::class) @RestController - class CoroutinesController { + class CoroutinesController(private val observationRegistry: ObservationRegistry) { @GetMapping("/suspend") suspend fun suspendingEndpoint(): String { @@ -167,6 +179,15 @@ class CoroutinesIntegrationTests : AbstractRequestMappingIntegrationTests() { delay(1) } + @GetMapping("/flow-observation") + fun flowObservationEndpoint(): Flow { + return flow { + val currentObservation = observationRegistry.currentObservation + assertThat(currentObservation).isNotNull + emit(currentObservation?.context?.name) + } + } + @GetMapping("/suspending-flow") suspend fun suspendingFlowEndpoint(): Flow { delay(1)