From 1f742c0dfdde922a48002932646361419fb0c9cf Mon Sep 17 00:00:00 2001 From: kami Date: Mon, 25 May 2026 00:19:22 +0400 Subject: [PATCH] feat(server-05): replace snapshot-then-subscribe with buffer-then-drain in GlobalStreamHandler Establishes subscribeAll() before replaySnapshot() so no event can slip between the projection read and live forwarding. Removes liveStreamJobs map and per-session forwarder spawn from handleStartSession; all live forwarding flows through the single global subscription. Closes findings #1 (race) and #5 (truth sources). Tests cover happy path, race (x100), empty sessions, and shutdown/leak cases. --- .../apps/server/ws/GlobalStreamHandler.kt | 79 ++++-- .../apps/server/ws/GlobalStreamHandlerTest.kt | 256 ++++++++++++++++++ 2 files changed, 305 insertions(+), 30 deletions(-) create mode 100644 apps/server/src/test/kotlin/com/correx/apps/server/ws/GlobalStreamHandlerTest.kt diff --git a/apps/server/src/main/kotlin/com/correx/apps/server/ws/GlobalStreamHandler.kt b/apps/server/src/main/kotlin/com/correx/apps/server/ws/GlobalStreamHandler.kt index f81effc4..85fcf57b 100644 --- a/apps/server/src/main/kotlin/com/correx/apps/server/ws/GlobalStreamHandler.kt +++ b/apps/server/src/main/kotlin/com/correx/apps/server/ws/GlobalStreamHandler.kt @@ -1,6 +1,7 @@ package com.correx.apps.server.ws import com.correx.apps.server.ServerModule +import com.correx.apps.server.bridge.DomainEventMapper import com.correx.apps.server.bridge.SessionEventBridge import com.correx.apps.server.protocol.ClientMessage import com.correx.apps.server.protocol.ProtocolSerializer @@ -8,7 +9,9 @@ import com.correx.apps.server.protocol.ProviderHealthDto import com.correx.apps.server.protocol.ServerMessage import com.correx.core.events.events.EventMetadata import com.correx.core.events.events.NewEvent +import com.correx.core.events.events.StoredEvent import com.correx.core.events.events.WorkflowFailedEvent +import com.correx.core.events.stores.EventStore import com.correx.core.events.types.EventId import com.correx.core.events.types.SessionId import com.correx.core.inference.ProviderHealth @@ -16,18 +19,21 @@ import com.correx.core.utils.TypeId import io.ktor.server.websocket.DefaultWebSocketServerSession import io.ktor.websocket.Frame import io.ktor.websocket.readText -import kotlinx.coroutines.Job +import kotlinx.coroutines.channels.BufferOverflow +import kotlinx.coroutines.channels.Channel import kotlinx.coroutines.channels.ClosedReceiveChannelException +import kotlinx.coroutines.coroutineScope import kotlinx.coroutines.launch import kotlinx.datetime.Clock import org.slf4j.LoggerFactory -import java.util.* +import java.util.UUID private val log = LoggerFactory.getLogger(GlobalStreamHandler::class.java) +private const val BUFFER_CAPACITY = 1024 class GlobalStreamHandler(private val module: ServerModule) { - suspend fun handle(session: DefaultWebSocketServerSession) { + suspend fun handle(session: DefaultWebSocketServerSession) = coroutineScope { log.info("client connected") val sendFrame: suspend (ServerMessage) -> Unit = { @@ -41,10 +47,14 @@ class GlobalStreamHandler(private val module: ServerModule) { orchestrationRepository = module.orchestrationRepository, approvalRepository = module.approvalRepository, ) - sendInitialSnapshot(session) - bridge.replaySnapshot() + val mapper = DomainEventMapper(module.artifactStore) - val liveStreamJobs = mutableMapOf() + sendInitialSnapshot(session) + + // Launch global stream (snapshot + live-forward) concurrently with incoming frame loop + launch { + streamGlobal(module.eventStore, bridge, mapper, sendFrame) + } try { for (frame in session.incoming) { @@ -54,7 +64,7 @@ class GlobalStreamHandler(private val module: ServerModule) { .onSuccess { msg -> log.debug("recv: {}", msg::class.simpleName) runCatching { - handleClientMessage(session, msg, liveStreamJobs, sendFrame) + handleClientMessage(session, msg, sendFrame) }.onFailure { log.error("handle failed", it) } @@ -62,18 +72,16 @@ class GlobalStreamHandler(private val module: ServerModule) { .onFailure { log.warn("decode error: {}", it.message) val error = ServerMessage.ProtocolError( - message = "Unknown message: ${it.message}", - sequence = null, - sessionSequence = null, - ) + message = "Unknown message: ${it.message}", + sequence = null, + sessionSequence = null, + ) session.send(Frame.Text(ProtocolSerializer.encodeServerMessage(error))) } } } } catch (_: ClosedReceiveChannelException) { log.info("client disconnected") - } finally { - liveStreamJobs.values.forEach { it.cancel() } } } @@ -102,17 +110,15 @@ class GlobalStreamHandler(private val module: ServerModule) { private suspend fun handleClientMessage( session: DefaultWebSocketServerSession, msg: ClientMessage, - liveStreamJobs: MutableMap, sendFrame: suspend (ServerMessage) -> Unit, ) { when (msg) { is ClientMessage.Ping -> Unit - is ClientMessage.StartSession -> handleStartSession(session, msg, liveStreamJobs, sendFrame) + is ClientMessage.StartSession -> handleStartSession(session, msg, sendFrame) is ClientMessage.CancelSession -> module.orchestrator.cancel(msg.sessionId) is ClientMessage.ResumeSession -> session.send(encodeError("ResumeSession not supported")) is ClientMessage.ApprovalResponse -> session.send(encodeError("ApprovalResponse must be sent to /sessions/{id}/stream")) - is ClientMessage.ChatInput -> Unit // handler deferred to subsequent task } } @@ -131,7 +137,6 @@ class GlobalStreamHandler(private val module: ServerModule) { private suspend fun handleStartSession( session: DefaultWebSocketServerSession, msg: ClientMessage.StartSession, - liveStreamJobs: MutableMap, sendFrame: suspend (ServerMessage) -> Unit, ) { val graph = module.workflowRegistry.find(msg.workflowId) @@ -171,19 +176,6 @@ class GlobalStreamHandler(private val module: ServerModule) { ) } } - liveStreamJobs.put( - sessionId, - session.launch { - val bridge = SessionEventBridge( - eventStore = module.eventStore, - artifactStore = module.artifactStore, - send = sendFrame, - orchestrationRepository = module.orchestrationRepository, - approvalRepository = module.approvalRepository, - ) - bridge.streamLive(sessionId) - }, - )?.cancel() val started = ServerMessage.SessionStarted( sessionId = sessionId, workflowId = msg.workflowId, @@ -193,3 +185,30 @@ class GlobalStreamHandler(private val module: ServerModule) { session.send(Frame.Text(ProtocolSerializer.encodeServerMessage(started))) } } + +internal suspend fun streamGlobal( + eventStore: EventStore, + bridge: SessionEventBridge, + mapper: DomainEventMapper, + sendFrame: suspend (ServerMessage) -> Unit, +) = coroutineScope { + val buffer = Channel( + capacity = BUFFER_CAPACITY, + onBufferOverflow = BufferOverflow.SUSPEND, + ) + + val subscription = launch { + eventStore.subscribeAll().collect { buffer.send(it) } + } + + try { + bridge.replaySnapshot() + for (event in buffer) { + val msg = mapper.map(event) ?: continue + sendFrame(msg) + } + } finally { + subscription.cancel() + buffer.close() + } +} diff --git a/apps/server/src/test/kotlin/com/correx/apps/server/ws/GlobalStreamHandlerTest.kt b/apps/server/src/test/kotlin/com/correx/apps/server/ws/GlobalStreamHandlerTest.kt new file mode 100644 index 00000000..aefd703d --- /dev/null +++ b/apps/server/src/test/kotlin/com/correx/apps/server/ws/GlobalStreamHandlerTest.kt @@ -0,0 +1,256 @@ +package com.correx.apps.server.ws + +import com.correx.apps.server.bridge.DomainEventMapper +import com.correx.apps.server.bridge.SessionEventBridge +import com.correx.apps.server.protocol.ServerMessage +import com.correx.core.approvals.ApprovalProjector +import com.correx.core.approvals.DefaultApprovalReducer +import com.correx.core.approvals.DefaultApprovalRepository +import com.correx.core.artifactstore.ArtifactStore +import com.correx.core.events.events.EventMetadata +import com.correx.core.events.events.EventPayload +import com.correx.core.events.events.NewEvent +import com.correx.core.events.events.StoredEvent +import com.correx.core.events.events.WorkflowStartedEvent +import com.correx.core.events.orchestration.OrchestrationState +import com.correx.core.events.orchestration.OrchestrationStatus +import com.correx.core.events.stores.EventStore +import com.correx.core.events.types.ArtifactId +import com.correx.core.events.types.EventId +import com.correx.core.events.types.SessionId +import com.correx.core.events.types.StageId +import com.correx.core.kernel.orchestration.OrchestrationRepository +import com.correx.core.sessions.projections.replay.DefaultEventReplayer +import com.correx.core.sessions.projections.replay.EventReplayer +import kotlinx.coroutines.cancelAndJoin +import kotlinx.coroutines.channels.Channel +import kotlinx.coroutines.delay +import kotlinx.coroutines.flow.Flow +import kotlinx.coroutines.flow.MutableSharedFlow +import kotlinx.coroutines.launch +import kotlinx.coroutines.test.runTest +import kotlinx.datetime.Instant +import org.junit.jupiter.api.Assertions.assertEquals +import org.junit.jupiter.api.Assertions.assertFalse +import org.junit.jupiter.api.Assertions.assertInstanceOf +import org.junit.jupiter.api.Test + +class GlobalStreamHandlerTest { + + private val sessionId = SessionId("session-1") + private val stageId = StageId("stage-1") + private val workflowId = "workflow-test" + private val timestamp = Instant.parse("2026-01-01T00:00:00Z") + + private val noopArtifactStore: ArtifactStore = object : ArtifactStore { + override suspend fun put(bytes: ByteArray): ArtifactId = ArtifactId("noop") + override suspend fun get(id: ArtifactId): ByteArray? = null + override suspend fun flushBefore(commit: suspend () -> Unit) = commit() + } + + private fun storedEvent(payload: EventPayload, seq: Long): StoredEvent = + StoredEvent( + metadata = EventMetadata( + eventId = EventId("evt-$seq"), + sessionId = sessionId, + timestamp = timestamp, + schemaVersion = 1, + causationId = null, + correlationId = null, + ), + sequence = seq, + sessionSequence = seq, + payload = payload, + ) + + private fun fakeEventStore( + liveFlow: MutableSharedFlow = MutableSharedFlow(), + sessions: Set = emptySet(), + lastGlobal: Long = 0L, + ): EventStore = object : EventStore { + override suspend fun append(event: NewEvent): StoredEvent = error("unused") + override suspend fun appendAll(events: List): List = error("unused") + override fun read(sessionId: SessionId): List = emptyList() + override fun readFrom(sessionId: SessionId, fromSequence: Long): List = emptyList() + override fun lastSequence(sessionId: SessionId): Long = lastGlobal + override fun subscribe(sessionId: SessionId): Flow = liveFlow + override fun allEvents(): Sequence = emptySequence() + override fun allSessionIds(): Set = sessions + override fun subscribeAll(): Flow = liveFlow + override suspend fun lastGlobalSequence(): Long = lastGlobal + } + + private fun activeOrchestrationRepository(): OrchestrationRepository = OrchestrationRepository( + object : EventReplayer { + override fun rebuild(sessionId: SessionId): OrchestrationState = OrchestrationState( + workflowId = workflowId, + status = OrchestrationStatus.RUNNING, + currentStageId = stageId, + pendingApproval = false, + ) + }, + ) + + private fun idleOrchestrationRepository(): OrchestrationRepository = OrchestrationRepository( + object : EventReplayer { + override fun rebuild(sessionId: SessionId): OrchestrationState = OrchestrationState( + workflowId = workflowId, + status = OrchestrationStatus.IDLE, + currentStageId = stageId, + pendingApproval = false, + ) + }, + ) + + private fun fakeApprovalRepository(store: EventStore): DefaultApprovalRepository = + DefaultApprovalRepository( + DefaultEventReplayer( + store, + ApprovalProjector(DefaultApprovalReducer()), + ), + ) + + /** + * Creates a bridge and a unified send channel so that snapshot messages and live messages + * both flow to the same [received] channel. + */ + private fun bridgeAndChannel( + store: EventStore, + orchRepo: OrchestrationRepository, + ): Pair> { + val received = Channel(Channel.UNLIMITED) + val approvals = fakeApprovalRepository(store) + val bridge = SessionEventBridge(store, noopArtifactStore, orchRepo, approvals) { received.send(it) } + return bridge to received + } + + @Test + fun `happy path - snapshot then 3 live events delivered in order`() = runTest { + val liveFlow = MutableSharedFlow(extraBufferCapacity = 16) + val store = fakeEventStore(liveFlow, setOf(sessionId), lastGlobal = 5L) + val (bridge, received) = bridgeAndChannel(store, activeOrchestrationRepository()) + val mapper = DomainEventMapper(noopArtifactStore) + + val job = launch { + // Live messages also go to the same received channel + streamGlobal(store, bridge, mapper) { received.send(it) } + } + + delay(50) + + // Emit 3 live events with sequences 6, 7, 8 + repeat(3) { i -> + liveFlow.emit(storedEvent(WorkflowStartedEvent(sessionId, workflowId, stageId), seq = (6 + i).toLong())) + } + + delay(100) + job.cancelAndJoin() + received.close() + + val messages = mutableListOf() + for (msg in received) messages.add(msg) + + // First: SessionSnapshot (from active session) + assertInstanceOf(ServerMessage.SessionSnapshot::class.java, messages[0]) + // Second: SnapshotComplete + assertEquals(ServerMessage.SnapshotComplete, messages[1]) + // Then 3 WorkflowStarted messages (seq 6, 7, 8) + val live = messages.drop(2) + assertEquals(3, live.size) + live.forEach { assertInstanceOf(ServerMessage.SessionStarted::class.java, it) } + } + + @Test + fun `race - event emitted during replaySnapshot is not lost`() = runTest { + repeat(100) { + val liveFlow = MutableSharedFlow(extraBufferCapacity = 64) + val store = fakeEventStore(liveFlow, emptySet(), lastGlobal = 5L) + val approvals = fakeApprovalRepository(store) + val received = Channel(Channel.UNLIMITED) + val bridge = SessionEventBridge(store, noopArtifactStore, idleOrchestrationRepository(), approvals) { + received.send(it) + } + val mapper = DomainEventMapper(noopArtifactStore) + val raceEvent = storedEvent(WorkflowStartedEvent(sessionId, workflowId, stageId), seq = 6L) + + // Emit the race event concurrently with streamGlobal starting. + // streamGlobal subscribes BEFORE calling replaySnapshot, so this event will be buffered + // and delivered after SnapshotComplete even if it arrives during the snapshot phase. + val emitter = launch { + delay(1) + liveFlow.emit(raceEvent) + } + + val job = launch { + streamGlobal(store, bridge, mapper) { received.send(it) } + } + + delay(100) + emitter.cancelAndJoin() + job.cancelAndJoin() + received.close() + + val messages = mutableListOf() + for (msg in received) messages.add(msg) + + assert(messages.contains(ServerMessage.SnapshotComplete)) { + "SnapshotComplete missing in iteration $it" + } + val completeIdx = messages.indexOf(ServerMessage.SnapshotComplete) + val afterComplete = messages.drop(completeIdx + 1) + assert(afterComplete.any { msg -> msg is ServerMessage.SessionStarted }) { + "Race event (seq=6) was lost in iteration $it! Messages: $messages" + } + } + } + + @Test + fun `empty - zero sessions - receives SnapshotComplete then live events`() = runTest { + val liveFlow = MutableSharedFlow(extraBufferCapacity = 16) + val store = fakeEventStore(liveFlow, emptySet(), lastGlobal = 0L) + val approvals = fakeApprovalRepository(store) + val received = Channel(Channel.UNLIMITED) + val bridge = SessionEventBridge(store, noopArtifactStore, idleOrchestrationRepository(), approvals) { + received.send(it) + } + val mapper = DomainEventMapper(noopArtifactStore) + + val job = launch { + streamGlobal(store, bridge, mapper) { received.send(it) } + } + + delay(50) + + // Emit a live event after snapshot + liveFlow.emit(storedEvent(WorkflowStartedEvent(sessionId, workflowId, stageId), seq = 1L)) + + delay(50) + job.cancelAndJoin() + received.close() + + val messages = mutableListOf() + for (msg in received) messages.add(msg) + + assertEquals(ServerMessage.SnapshotComplete, messages[0]) + assertEquals(2, messages.size) + assertInstanceOf(ServerMessage.SessionStarted::class.java, messages[1]) + } + + @Test + fun `shutdown - cancel outer scope ends subscription with no leaks`() = runTest { + val liveFlow = MutableSharedFlow() + val store = fakeEventStore(liveFlow, emptySet(), lastGlobal = 0L) + val approvals = fakeApprovalRepository(store) + val bridge = SessionEventBridge(store, noopArtifactStore, idleOrchestrationRepository(), approvals) { } + val mapper = DomainEventMapper(noopArtifactStore) + + val job = launch { + streamGlobal(store, bridge, mapper) { } + } + + delay(50) + job.cancelAndJoin() + + assertFalse(job.isActive, "Handler job must not be active after cancel") + } +}