From 5defa7dbf4107ce889b6eca04f3d44598d66eee1 Mon Sep 17 00:00:00 2001 From: Alexander Shparun Date: Mon, 15 Jun 2026 16:05:49 +0200 Subject: [PATCH] [fleet, rpc] make stream consumer announce itself with StreamInit a producer that abandons a stream (e.g. a cancelled client call returning a channel) used to leave the remote consumer waiting forever, since a lone StreamNext is ignored for unregistered streams; the consumer now sends StreamInit too, so the producer replies StreamClosed and the stream is collected GitOrigin-RevId: 7108ee05efc9c8f5811f06bba90f7b35534f5870 --- .../fleet/rpc/client/IRpcClient.kt | 38 +++++++++------- .../fleet/rpc/client/RpcClient.kt | 5 +-- .../fleet/rpc/core/RpcMessage.kt | 45 ++++++++++++------- .../srcCommonMain/fleet/rpc/core/RpcStream.kt | 4 ++ 4 files changed, 57 insertions(+), 35 deletions(-) diff --git a/fleet/rpc/srcCommonMain/fleet/rpc/client/IRpcClient.kt b/fleet/rpc/srcCommonMain/fleet/rpc/client/IRpcClient.kt index 578240d509e0..25a92531db3d 100644 --- a/fleet/rpc/srcCommonMain/fleet/rpc/client/IRpcClient.kt +++ b/fleet/rpc/srcCommonMain/fleet/rpc/client/IRpcClient.kt @@ -12,6 +12,7 @@ import fleet.rpc.core.PrefetchStrategy import fleet.util.UID import kotlinx.coroutines.CoroutineScope import kotlinx.coroutines.Deferred +import kotlinx.coroutines.currentCoroutineContext import kotlinx.coroutines.flow.StateFlow import kotlinx.coroutines.flow.first import kotlinx.coroutines.flow.mapNotNull @@ -20,18 +21,21 @@ import kotlinx.coroutines.withTimeoutOrNull import org.jetbrains.annotations.ApiStatus import kotlin.concurrent.Volatile import kotlin.coroutines.CoroutineContext -import kotlin.coroutines.coroutineContext @ApiStatus.Internal -data class Call(val route: UID, - val service: InstanceId, - val signature: RpcSignature, - val arguments: List) { - fun display(): String = "route=$route, $service/${signature.methodName}(${signature.parameters.map { it.parameterName }.joinToString(", ")})" +data class Call( + val route: UID, + val service: InstanceId, + val signature: RpcSignature, + val arguments: List +) { + fun display(): String = "route=$route, $service/${signature.methodName}(${signature.parameters.joinToString(", ") { it.parameterName }})" } -internal data class RpcStrategyContextElement(val awaitConnection: Boolean = true, - val prefetchStrategy: PrefetchStrategy = PrefetchStrategy.Default) : CoroutineContext.Element { +internal data class RpcStrategyContextElement( + val awaitConnection: Boolean = true, + val prefetchStrategy: PrefetchStrategy = PrefetchStrategy.Default +) : CoroutineContext.Element { companion object : CoroutineContext.Key override val key: CoroutineContext.Key<*> get() = RpcStrategyContextElement @@ -39,7 +43,7 @@ internal data class RpcStrategyContextElement(val awaitConnection: Boolean = tru @ApiStatus.Internal suspend fun withoutAwaitingForReconnect(body: suspend CoroutineScope.() -> T): T { - val strategy = coroutineContext[RpcStrategyContextElement]?.copy(awaitConnection = false) + val strategy = currentCoroutineContext()[RpcStrategyContextElement]?.copy(awaitConnection = false) ?: RpcStrategyContextElement(awaitConnection = false) return withContext(strategy) { body() @@ -48,7 +52,7 @@ suspend fun withoutAwaitingForReconnect(body: suspend CoroutineScope.() -> T @ApiStatus.Internal suspend fun withPrefetchStrategy(prefetchStrategy: PrefetchStrategy, body: suspend CoroutineScope.() -> T): T { - val strategy = coroutineContext[RpcStrategyContextElement]?.copy(prefetchStrategy = prefetchStrategy) + val strategy = currentCoroutineContext()[RpcStrategyContextElement]?.copy(prefetchStrategy = prefetchStrategy) ?: RpcStrategyContextElement(prefetchStrategy = prefetchStrategy) return withContext(strategy) { body() @@ -74,10 +78,12 @@ fun IRpcClient.asHandlerFactory(): InvocationHandlerFactory = object : InvocationHandlerFactory { override fun handler(arg: ProxyClosure): SuspendInvocationHandler { return object : SuspendInvocationHandler { - override suspend fun call(remoteApiDescriptor: RemoteApiDescriptor<*>, - method: String, - args: List, - publish: (SuspendInvocationHandler.CallResult) -> Unit) { + override suspend fun call( + remoteApiDescriptor: RemoteApiDescriptor<*>, + method: String, + args: List, + publish: (SuspendInvocationHandler.CallResult) -> Unit + ) { call( Call( route = arg.route, @@ -103,7 +109,7 @@ internal fun reconnectingRpcClient(attempts: StateFlow Unit) { withTimeoutOrNull(RPC_TIMEOUT) { mayHaveCalls = true - val awaitConnection = coroutineContext[RpcStrategyContextElement]?.awaitConnection ?: true + val awaitConnection = currentCoroutineContext()[RpcStrategyContextElement]?.awaitConnection ?: true val client = if (awaitConnection) { attempts.mapNotNull { (it as? ConnectionStatus.Connected)?.value }.first { it != disconnectedClient } } @@ -123,7 +129,7 @@ internal fun reconnectingRpcClient(attempts: StateFlow { logger.trace { "Got CallResult: requestId = ${message.requestId}" } - // TODO this value can contain send/receive channels that must be closed, we should not drop it on the ground outgoingRpc.remove(message.requestId)?.let { (rpc) -> try { val (returnResult, streams) = run { @@ -425,7 +425,6 @@ private class RpcClient( } } else { - // TODO this value can contain send/receive channels that must be closed logger.trace { "Received StreamData for unregistered stream ${message.streamId}" } } } diff --git a/fleet/rpc/srcCommonMain/fleet/rpc/core/RpcMessage.kt b/fleet/rpc/srcCommonMain/fleet/rpc/core/RpcMessage.kt index 26f9571630cf..501519f0ff05 100644 --- a/fleet/rpc/srcCommonMain/fleet/rpc/core/RpcMessage.kt +++ b/fleet/rpc/srcCommonMain/fleet/rpc/core/RpcMessage.kt @@ -23,11 +23,13 @@ sealed class RpcMessage { @Serializable @SerialName("call") - data class CallRequest(val requestId: UID, - val service: InstanceId, - val method: String, - val args: Map, - val meta: Map = emptyMap()) : RpcMessage() { + data class CallRequest( + val requestId: UID, + val service: InstanceId, + val method: String, + val args: Map, + val meta: Map = emptyMap() + ) : RpcMessage() { val displayName: String get() = "RPC call ${classMethodDisplayName()}[#${requestId}]" fun classMethodDisplayName(): String { @@ -41,14 +43,18 @@ sealed class RpcMessage { @Serializable @SerialName("call_result") - data class CallResult(val requestId: UID, - val result: JsonElement, - val meta: Map = emptyMap()) : RpcMessage() + data class CallResult( + val requestId: UID, + val result: JsonElement, + val meta: Map = emptyMap() + ) : RpcMessage() @Serializable @SerialName("call_failure") - data class CallFailure(val requestId: UID, - val error: FailureInfo) : RpcMessage() + data class CallFailure( + val requestId: UID, + val error: FailureInfo + ) : RpcMessage() @Serializable @SerialName("cancel_call") @@ -56,8 +62,10 @@ sealed class RpcMessage { @Serializable @SerialName("stream_data") - data class StreamData(val streamId: UID, - val data: JsonElement) : RpcMessage() + data class StreamData( + val streamId: UID, + val data: JsonElement + ) : RpcMessage() // optional message sent by a producer to a newly created stream // @@ -67,7 +75,7 @@ sealed class RpcMessage { // when consumer receives StreamInit for a stream that is not registered as in-use, it can respond with StreamClose @Serializable @SerialName("stream_init") - data class StreamInit(val streamId: UID): RpcMessage() + data class StreamInit(val streamId: UID) : RpcMessage() /** * Producer should not send any data unless asked to. Consumer controls the moment and the quantity with this message. @@ -75,12 +83,17 @@ sealed class RpcMessage { */ @Serializable @SerialName("stream_next") - data class StreamNext(val streamId: UID, val count: Int) : RpcMessage() + data class StreamNext( + val streamId: UID, + val count: Int, + ) : RpcMessage() @Serializable @SerialName("stream_closed") - data class StreamClosed(val streamId: UID, - val error: FailureInfo? = null) : RpcMessage() + data class StreamClosed( + val streamId: UID, + val error: FailureInfo? = null + ) : RpcMessage() @Serializable @SerialName("resource_consumed") diff --git a/fleet/rpc/srcCommonMain/fleet/rpc/core/RpcStream.kt b/fleet/rpc/srcCommonMain/fleet/rpc/core/RpcStream.kt index 096ed5e0cd8c..f1d7793b8ac6 100644 --- a/fleet/rpc/srcCommonMain/fleet/rpc/core/RpcStream.kt +++ b/fleet/rpc/srcCommonMain/fleet/rpc/core/RpcStream.kt @@ -337,6 +337,10 @@ fun serveStream( var requested = prefetchStrategy.streamStarted() var remaining = requested require(requested > 0) + // announce the consumer symmetrically to the producer's StreamInit: + // if the producer has already abandoned this stream, it will respond with StreamClosed + // and we won't linger waiting for data that will never arrive (StreamNext alone is ignored for unregistered streams) + sendMessage(RpcMessage.StreamInit(streamId = descriptor.uid)) sendMessage(RpcMessage.StreamNext(descriptor.uid, requested)) descriptor.bufferedChannel.consumeEach { each -> RpcStream.logger.trace { "Channel ${descriptor.uid} ${descriptor.displayName} processes message $each" }