[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
This commit is contained in:
Alexander Shparun
2026-06-15 18:51:50 +00:00
committed by intellij-monorepo-bot
parent 44112b97bb
commit 5defa7dbf4
4 changed files with 57 additions and 35 deletions
@@ -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<Any?>) {
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<Any?>
) {
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<RpcStrategyContextElement>
override val key: CoroutineContext.Key<*> get() = RpcStrategyContextElement
@@ -39,7 +43,7 @@ internal data class RpcStrategyContextElement(val awaitConnection: Boolean = tru
@ApiStatus.Internal
suspend fun <T> 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 <T> withoutAwaitingForReconnect(body: suspend CoroutineScope.() -> T
@ApiStatus.Internal
suspend fun <T> 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<ProxyClosure> =
object : InvocationHandlerFactory<ProxyClosure> {
override fun handler(arg: ProxyClosure): SuspendInvocationHandler {
return object : SuspendInvocationHandler {
override suspend fun call(remoteApiDescriptor: RemoteApiDescriptor<*>,
method: String,
args: List<Any?>,
publish: (SuspendInvocationHandler.CallResult) -> Unit) {
override suspend fun call(
remoteApiDescriptor: RemoteApiDescriptor<*>,
method: String,
args: List<Any?>,
publish: (SuspendInvocationHandler.CallResult) -> Unit
) {
call(
Call(
route = arg.route,
@@ -103,7 +109,7 @@ internal fun reconnectingRpcClient(attempts: StateFlow<ConnectionStatus<IRpcClie
override suspend fun call(call: Call, publish: (SuspendInvocationHandler.CallResult) -> 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<ConnectionStatus<IRpcClie
throw x
}
Unit
} ?: throw RpcTimeoutException("Client did not connect after ${RPC_TIMEOUT}", null)
} ?: throw RpcTimeoutException("Client did not connect after $RPC_TIMEOUT", null)
}
}
}
@@ -75,8 +75,9 @@ import kotlin.coroutines.Continuation
import kotlin.coroutines.coroutineContext
import kotlin.coroutines.resumeWithException
import kotlin.coroutines.startCoroutine
import kotlin.time.Duration.Companion.milliseconds
internal const val RPC_TIMEOUT = 60_000L
internal val RPC_TIMEOUT = 60_000.milliseconds
private data class OutgoingRequest(
val route: UID,
@@ -347,7 +348,6 @@ private class RpcClient(
when (message) {
is RpcMessage.CallResult -> {
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}" }
}
}
@@ -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<String, JsonElement>,
val meta: Map<String, JsonElement> = emptyMap()) : RpcMessage() {
data class CallRequest(
val requestId: UID,
val service: InstanceId,
val method: String,
val args: Map<String, JsonElement>,
val meta: Map<String, JsonElement> = 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<String, JsonElement> = emptyMap()) : RpcMessage()
data class CallResult(
val requestId: UID,
val result: JsonElement,
val meta: Map<String, JsonElement> = 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")
@@ -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" }