diff --git a/apps/server/src/main/kotlin/com/correx/apps/server/ServerModule.kt b/apps/server/src/main/kotlin/com/correx/apps/server/ServerModule.kt index cdef5095..eecc4ad3 100644 --- a/apps/server/src/main/kotlin/com/correx/apps/server/ServerModule.kt +++ b/apps/server/src/main/kotlin/com/correx/apps/server/ServerModule.kt @@ -107,47 +107,47 @@ class ServerModule( * duplicating inference calls (the root cause of the healthcheck script corruption). */ fun launchSessionRun(sessionId: SessionId, graph: com.correx.core.transitions.graph.WorkflowGraph) { - if (activeSessionJobs.containsKey(sessionId)) return - val job = moduleScope.launch { - runCatching { - orchestrator.run(sessionId, graph, defaultOrchestrationConfig) - }.onFailure { ex -> - log.error("runSession: session={} failed: {}", sessionId.value, ex.message, ex) - eventStore.append( - NewEvent( - metadata = EventMetadata( - eventId = EventId(java.util.UUID.randomUUID().toString()), - sessionId = sessionId, - timestamp = Clock.System.now(), - schemaVersion = 1, - causationId = null, - correlationId = null, + activeSessionJobs.computeIfAbsent(sessionId) { + moduleScope.launch { + runCatching { + orchestrator.run(sessionId, graph, defaultOrchestrationConfig) + }.onFailure { ex -> + log.error("runSession: session={} failed: {}", sessionId.value, ex.message, ex) + eventStore.append( + NewEvent( + metadata = EventMetadata( + eventId = EventId(java.util.UUID.randomUUID().toString()), + sessionId = sessionId, + timestamp = Clock.System.now(), + schemaVersion = 1, + causationId = null, + correlationId = null, + ), + payload = WorkflowFailedEvent( + sessionId = sessionId, + stageId = graph.start, + reason = ex.message ?: "unexpected orchestrator failure", + retryExhausted = false, + ), ), - payload = WorkflowFailedEvent( - sessionId = sessionId, - stageId = graph.start, - reason = ex.message ?: "unexpected orchestrator failure", - retryExhausted = false, - ), - ), - ) + ) + } + activeSessionJobs.remove(sessionId) } - activeSessionJobs.remove(sessionId) } - activeSessionJobs[sessionId] = job } private fun launchSessionResume(sessionId: SessionId, graph: com.correx.core.transitions.graph.WorkflowGraph) { - if (activeSessionJobs.containsKey(sessionId)) return - val job = moduleScope.launch { - runCatching { - orchestrator.resume(sessionId, graph, defaultOrchestrationConfig) - }.onFailure { e -> - log.error("resumeSession: session={} failed", sessionId.value, e) + activeSessionJobs.computeIfAbsent(sessionId) { + moduleScope.launch { + runCatching { + orchestrator.resume(sessionId, graph, defaultOrchestrationConfig) + }.onFailure { e -> + log.error("resumeSession: session={} failed", sessionId.value, e) + } + activeSessionJobs.remove(sessionId) } - activeSessionJobs.remove(sessionId) } - activeSessionJobs[sessionId] = job } private fun resumeAbandonedSessions() { diff --git a/apps/tui/src/main/kotlin/com/correx/apps/tui/components/RouterPanel.kt b/apps/tui/src/main/kotlin/com/correx/apps/tui/components/RouterPanel.kt index 95298f84..9065a325 100644 --- a/apps/tui/src/main/kotlin/com/correx/apps/tui/components/RouterPanel.kt +++ b/apps/tui/src/main/kotlin/com/correx/apps/tui/components/RouterPanel.kt @@ -17,7 +17,8 @@ fun routerPanelWidget(state: TuiState): Paragraph { if (!state.routerConnected) { add(Line.from(Span.styled("not connected — epic 14", dimStyle))) } else { - for (msg in state.routerMessages) { + val msgs = state.routerMessages[state.sessions.selectedId].orEmpty() + for (msg in msgs) { add(Line.from(Span.raw(msg))) } } diff --git a/apps/tui/src/main/kotlin/com/correx/apps/tui/reducer/RootReducer.kt b/apps/tui/src/main/kotlin/com/correx/apps/tui/reducer/RootReducer.kt index 115a0498..4ea6436b 100644 --- a/apps/tui/src/main/kotlin/com/correx/apps/tui/reducer/RootReducer.kt +++ b/apps/tui/src/main/kotlin/com/correx/apps/tui/reducer/RootReducer.kt @@ -134,6 +134,11 @@ object RootReducer { ) } + // Optimistic session created by SessionsReducer for IDLE chat submit. + sessions.selectedId != null && prevInputMode == InputMode.ROUTER && + withBgReset.displayState == DisplayState.IDLE -> + withBgReset.copy(sessionEntered = true) + else -> withBgReset } } @@ -159,13 +164,18 @@ object RootReducer { } msg is ServerMessage.RouterResponseMessage -> { withBgReset.copy( - routerMessages = withBgReset.routerMessages + msg.content, + routerMessages = withBgReset.routerMessages.run { + val key = msg.sessionId.value + this + (key to (this[key].orEmpty() + msg.content)) + }, ) } msg is ServerMessage.ToolCompleted && msg.diff != null -> { withBgReset.copy( - routerMessages = withBgReset.routerMessages + - "--- diff from ${msg.toolName} ---", + routerMessages = withBgReset.routerMessages.run { + val key = msg.sessionId.value + this + (key to (this[key].orEmpty() + msg.diff!!)) + }, ) } else -> withBgReset diff --git a/apps/tui/src/main/kotlin/com/correx/apps/tui/reducer/SessionsReducer.kt b/apps/tui/src/main/kotlin/com/correx/apps/tui/reducer/SessionsReducer.kt index c35257c6..8870f676 100644 --- a/apps/tui/src/main/kotlin/com/correx/apps/tui/reducer/SessionsReducer.kt +++ b/apps/tui/src/main/kotlin/com/correx/apps/tui/reducer/SessionsReducer.kt @@ -62,7 +62,17 @@ object SessionsReducer { text.isNotBlank() -> { val sessionId = SessionId(java.util.UUID.randomUUID().toString()) - sessions to listOf( + val optimistic = SessionSummary( + id = sessionId.value, + status = "STARTING", + workflowId = "chat", + name = "chat", + lastEventAt = clock(), + ) + sessions.copy( + sessions = sessions.sessions + optimistic, + selectedId = sessionId.value, + ) to listOf( Effect.SendWs(ClientMessage.StartChatSession(sessionId = sessionId, text = text)), ) } @@ -559,8 +569,11 @@ object SessionsReducer { clock: () -> Long, sessions: SessionsState, ): Pair> { - // TODO(RF-3): If StageToolManifest created a placeholder, dedup/replace instead of - // unconditional append to avoid duplicate session entries. + // Remove any optimistic (STARTING) session created by IDLE submit, then insert the + // real session from the server. This prevents duplicate session entries when the + // server-generated session ID differs from the client-generated optimistic ID. + val hadOptimistic = sessions.sessions.any { it.status == "STARTING" } + val cleaned = sessions.sessions.filter { it.status != "STARTING" } val summary = SessionSummary( id = msg.sessionId.value, status = "ACTIVE", @@ -570,8 +583,13 @@ object SessionsReducer { currentStage = null, lastOutput = null, ) - val selected = sessions.selectedId ?: msg.sessionId.value - return sessions.copy(sessions = sessions.sessions + summary, selectedId = selected) to emptyList() + // When replacing an optimistic session, switch selectedId to the real session ID. + // Otherwise preserve the existing selection (normal SessionStarted for a new session). + val selected = if (hadOptimistic) msg.sessionId.value else (sessions.selectedId ?: msg.sessionId.value) + return sessions.copy( + sessions = cleaned + summary, + selectedId = selected, + ) to emptyList() } private fun processInferenceCompletedMessage( diff --git a/apps/tui/src/main/kotlin/com/correx/apps/tui/state/TuiState.kt b/apps/tui/src/main/kotlin/com/correx/apps/tui/state/TuiState.kt index 0144f2c0..c3b0a093 100644 --- a/apps/tui/src/main/kotlin/com/correx/apps/tui/state/TuiState.kt +++ b/apps/tui/src/main/kotlin/com/correx/apps/tui/state/TuiState.kt @@ -46,7 +46,7 @@ data class TuiState( val currentModel: String? = null, val providerType: ProviderType = ProviderType.LOCAL, val routerConnected: Boolean = false, - val routerMessages: List = emptyList(), + val routerMessages: Map> = emptyMap(), /** * True from connect until SnapshotComplete observed. Reset to true on Disconnected. * When true, event-bearing messages are buffered in [pendingEvents]; snapshots are applied immediately. diff --git a/apps/tui/src/test/kotlin/com/correx/apps/tui/reducer/RootReducerTest.kt b/apps/tui/src/test/kotlin/com/correx/apps/tui/reducer/RootReducerTest.kt index fdadce31..6bc1dc83 100644 --- a/apps/tui/src/test/kotlin/com/correx/apps/tui/reducer/RootReducerTest.kt +++ b/apps/tui/src/test/kotlin/com/correx/apps/tui/reducer/RootReducerTest.kt @@ -17,6 +17,7 @@ import com.correx.core.approvals.Tier import com.correx.core.events.types.ApprovalRequestId import com.correx.core.events.types.SessionId import org.junit.jupiter.api.Assertions.assertEquals +import org.junit.jupiter.api.Assertions.assertNotNull import org.junit.jupiter.api.Assertions.assertTrue import org.junit.jupiter.api.Test @@ -388,4 +389,72 @@ class RootReducerTest { assertEquals("s2", state.sessions.selectedId) assertTrue(state.sessionEntered) } + + @Test + fun `ToolCompleted with diff appends diff content to per-session routerMessages`() { + val initial = TuiState(snapshotPhase = false) + val toolCompleted = ServerMessage.ToolCompleted( + sessionId = SessionId("s1"), + toolName = "file_write", + outputSummary = "wrote 3 lines", + occurredAt = fixedClock(), + diff = "--- a/file.txt\n+++ b/file.txt\n@@ -1 +1 @@\n-old content\n+new content", + sequence = 1L, + sessionSequence = 1L, + ) + val (state, _) = RootReducer.reduce(initial, Action.ServerEventReceived(toolCompleted), fixedClock) + val msgs = state.routerMessages["s1"] + assertEquals(1, msgs?.size) + assertTrue(msgs?.first()?.contains("--- a/file.txt") == true) + assertTrue(msgs?.first()?.contains("+new content") == true) + } + + @Test + fun `ToolCompleted without diff does not append to routerMessages`() { + val initial = TuiState(snapshotPhase = false, routerMessages = mapOf("s1" to listOf("existing msg"))) + val toolCompleted = ServerMessage.ToolCompleted( + sessionId = SessionId("s1"), + toolName = "file_write", + outputSummary = "wrote 3 lines", + occurredAt = fixedClock(), + diff = null, + sequence = 1L, + sessionSequence = 1L, + ) + val (state, _) = RootReducer.reduce(initial, Action.ServerEventReceived(toolCompleted), fixedClock) + val msgs = state.routerMessages["s1"] + assertEquals(1, msgs?.size) + assertEquals("existing msg", msgs?.first()) + } + + @Test + fun `routerMessages are isolated per session`() { + val initial = TuiState(snapshotPhase = false) + val msg1 = ServerMessage.RouterResponseMessage( + sessionId = SessionId("s1"), content = "hello from s1", + steeringEmitted = false, sequence = null, sessionSequence = null, + ) + val msg2 = ServerMessage.RouterResponseMessage( + sessionId = SessionId("s2"), content = "hello from s2", + steeringEmitted = false, sequence = null, sessionSequence = null, + ) + val after1 = RootReducer.reduce(initial, Action.ServerEventReceived(msg1), fixedClock).first + val after2 = RootReducer.reduce(after1, Action.ServerEventReceived(msg2), fixedClock).first + assertEquals(listOf("hello from s1"), after2.routerMessages["s1"]) + assertEquals(listOf("hello from s2"), after2.routerMessages["s2"]) + } + + @Test + fun `SubmitInput in IDLE with text creates optimistic session and enters session`() { + val initial = TuiState(snapshotPhase = false, inputMode = InputMode.ROUTER, inputBuffer = "my chat text") + val (state, effects) = RootReducer.reduce(initial, Action.SubmitInput, fixedClock) + assertTrue(state.sessionEntered) + assertNotNull(state.sessions.selectedId) + assertEquals(1, state.sessions.sessions.size) + assertEquals("STARTING", state.sessions.sessions[0].status) + assertEquals(state.sessions.selectedId, state.sessions.sessions[0].id) + val wsEffects = effects.filterIsInstance() + assertEquals(1, wsEffects.size) + assertTrue(wsEffects[0].message is ClientMessage.StartChatSession) + } } diff --git a/apps/tui/src/test/kotlin/com/correx/apps/tui/reducer/SessionsReducerTest.kt b/apps/tui/src/test/kotlin/com/correx/apps/tui/reducer/SessionsReducerTest.kt index 15a6322d..b14e5c26 100644 --- a/apps/tui/src/test/kotlin/com/correx/apps/tui/reducer/SessionsReducerTest.kt +++ b/apps/tui/src/test/kotlin/com/correx/apps/tui/reducer/SessionsReducerTest.kt @@ -389,4 +389,32 @@ class SessionsReducerTest { assertEquals("SessionCompleted", events[2].type) assertEquals("", events[2].detail) } + + @Test + fun `SubmitInput in IDLE creates optimistic STARTING session`() { + val (state, effects) = reduce(action = Action.SubmitInput, inputText = "hello world") + assertEquals(1, state.sessions.size) + assertEquals("STARTING", state.sessions[0].status) + assertEquals(state.sessions[0].id, state.selectedId) + val wsEffects = effects.filterIsInstance() + assertEquals(1, wsEffects.size) + assertTrue(wsEffects[0].message is ClientMessage.StartChatSession) + val chatSession = wsEffects[0].message as ClientMessage.StartChatSession + assertEquals("hello world", chatSession.text) + } + + @Test + fun `SessionStarted reconciles by removing optimistic STARTING session`() { + val optimisticId = SessionId("opt-1") + val optimistic = SessionSummary(id = "opt-1", status = "STARTING", workflowId = "chat", name = "chat", lastEventAt = 0L) + val s = SessionsState(sessions = listOf(optimistic), selectedId = "opt-1") + val msg = ServerMessage.SessionStarted( + sessionId = SessionId("real-1"), workflowId = "wf", sequence = 1L, sessionSequence = 1L, + ) + val (state, _) = reduce(sessions = s, action = Action.ServerEventReceived(msg)) + assertEquals(1, state.sessions.size) + assertEquals("real-1", state.sessions[0].id) + assertEquals("ACTIVE", state.sessions[0].status) + assertEquals("real-1", state.selectedId) + } } diff --git a/core/kernel/src/main/kotlin/com/correx/core/kernel/orchestration/SessionOrchestrator.kt b/core/kernel/src/main/kotlin/com/correx/core/kernel/orchestration/SessionOrchestrator.kt index fb572e4c..0d2783d2 100644 --- a/core/kernel/src/main/kotlin/com/correx/core/kernel/orchestration/SessionOrchestrator.kt +++ b/core/kernel/src/main/kotlin/com/correx/core/kernel/orchestration/SessionOrchestrator.kt @@ -407,6 +407,8 @@ abstract class SessionOrchestrator( } // grant auto-approved — fall through to execute } else { + val deferred = CompletableDeferred() + pendingApprovals[requestId] = deferred emit(sessionId, OrchestrationPausedEvent(sessionId, stageId, "APPROVAL_PENDING")) emit( sessionId, @@ -422,8 +424,6 @@ abstract class SessionOrchestrator( preview = toolPreview ?: toolCall.function.arguments.take(200), ), ) - val deferred = CompletableDeferred() - pendingApprovals[requestId] = deferred val userDecision = try { deferred.await() } finally { @@ -830,9 +830,11 @@ abstract class SessionOrchestrator( outcome: ValidationOutcome.NeedsApproval, ): StageExecutionResult { log.debug("[Orchestrator] approval required session={} stage={}", sessionId.value, stageId.value) - emit(sessionId, OrchestrationPausedEvent(sessionId, stageId, "APPROVAL_PENDING")) val domainRequest = buildApprovalRequest(sessionId, stageId, outcome) + val deferred = CompletableDeferred() + pendingApprovals[domainRequest.id] = deferred + emit(sessionId, OrchestrationPausedEvent(sessionId, stageId, "APPROVAL_PENDING")) emit( sessionId, ApprovalRequestedEvent( @@ -846,9 +848,6 @@ abstract class SessionOrchestrator( ), ) - val deferred = CompletableDeferred() - pendingApprovals[domainRequest.id] = deferred - return try { log.debug( "[Orchestrator] awaiting approval session={} requestId={}", diff --git a/testing/kernel/src/test/kotlin/ApprovalRegisterBeforeEmitTest.kt b/testing/kernel/src/test/kotlin/ApprovalRegisterBeforeEmitTest.kt new file mode 100644 index 00000000..366289f3 --- /dev/null +++ b/testing/kernel/src/test/kotlin/ApprovalRegisterBeforeEmitTest.kt @@ -0,0 +1,95 @@ +import kotlinx.coroutines.CompletableDeferred +import kotlinx.coroutines.test.runTest +import org.junit.jupiter.api.Assertions.assertNotNull +import org.junit.jupiter.api.Assertions.assertTrue +import org.junit.jupiter.api.Test +import java.util.concurrent.ConcurrentHashMap + +/** + * Validates the register-before-emit invariant: a [CompletableDeferred] registered + * in a [ConcurrentHashMap] before a subsequent operation (e.g. event emission) is + * guaranteed to be visible both to the current thread and to any thread observing + * the subsequent operation's side effects. + * + * This mirrors the fix in [SessionOrchestrator.dispatchToolCalls] and + * [SessionOrchestrator.handleApproval], where [pendingApprovals] registration was + * moved BEFORE the [ApprovalRequestedEvent] emission. The old code (register AFTER + * emit) left a window where the event store subscriber could observe the event and + * call [submitApprovalDecision] before the [CompletableDeferred] was in the map, + * creating an orphan deferred that permanently blocked the session. + */ +class ApprovalRegisterBeforeEmitTest { + + @Test + fun `deferred registered before emit is findable after emit`(): Unit = runTest { + val registry = ConcurrentHashMap>() + val requestId = "req-1" + + // --- Fixed pattern: register BEFORE emit --- + val deferred = CompletableDeferred() + registry[requestId] = deferred + + // Simulate the "emit" step — any concurrent observer sees this AFTER the put + val emitPayload: String = "ApprovalRequested($requestId)" + + // Now simulate a concurrent submitApprovalDecision arriving after the emit + val found = registry[requestId] + + assertNotNull(found, "deferred must be in map when submitApprovalDecision observes the event") + assertTrue(found === deferred, "found deferred must be the same instance that was registered") + } + + @Test + fun `deferred registered before emit can be completed by submitApprovalDecision`(): Unit = runTest { + val registry = ConcurrentHashMap>() + val requestId = "req-2" + + // Fixed pattern + val deferred = CompletableDeferred() + registry[requestId] = deferred + + // Simulate emit (event store subscribers will see this) + @Suppress("UNUSED_VARIABLE") + val emitted = true + + // Simulate submitApprovalDecision looking up and completing the deferred + val lookup = registry[requestId] + assertNotNull(lookup) + lookup!!.complete(true) + + // The original awaiter gets the value + val result = deferred.await() + assertTrue(result, "deferred must complete with the decision value") + } + + @Test + fun `old pattern race window eliminated by register-before-emit ordering`(): Unit = runTest { + // In the OLD code, registration happened AFTER emit: + // emit(ApprovalRequestedEvent(requestId)) // step 1 + // pendingApprovals[requestId] = deferred // step 2 + // + // Between step 1 and step 2, a concurrently executing submitApprovalDecision + // (triggered by the event store subscription processing the emitted event) + // would find no deferred in pendingApprovals — creating an orphan that + // permanently blocked the session coroutine on deferred.await(). + // + // This race is inherently non-deterministic and cannot be reproduced + // reliably in a unit test. The fix is STRUCTURAL: reorder step 2 before + // step 1 so the deferred is guaranteed to be in the map when any observer + // processes the event. + // + // This test confirms the fixed ordering compiles and the ConcurrentHashMap + // contract guarantees visibility: after put(K, V), get(K) returns V. + + val registry = ConcurrentHashMap>() + val requestId = "req-3" + + // Fixed pattern: register before the point that triggers concurrent observers + val deferred = CompletableDeferred() + registry[requestId] = deferred + + // "Emit" happens after registration. Any observer of this point will find the deferred. + assertNotNull(registry[requestId], "deferred must be findable after put") + assertTrue(registry[requestId] === deferred, "map entry must reference the same deferred") + } +} diff --git a/testing/kernel/src/test/kotlin/LaunchRegistrationRaceTest.kt b/testing/kernel/src/test/kotlin/LaunchRegistrationRaceTest.kt new file mode 100644 index 00000000..caa24bb7 --- /dev/null +++ b/testing/kernel/src/test/kotlin/LaunchRegistrationRaceTest.kt @@ -0,0 +1,76 @@ +import kotlinx.coroutines.CoroutineScope +import kotlinx.coroutines.Dispatchers +import kotlinx.coroutines.Job +import kotlinx.coroutines.SupervisorJob +import kotlinx.coroutines.async +import kotlinx.coroutines.awaitAll +import kotlinx.coroutines.cancel +import kotlinx.coroutines.coroutineScope +import kotlinx.coroutines.delay +import kotlinx.coroutines.launch +import kotlinx.coroutines.test.runTest +import org.junit.jupiter.api.Assertions.assertEquals +import org.junit.jupiter.api.Test +import java.util.concurrent.ConcurrentHashMap +import java.util.concurrent.atomic.AtomicInteger + +/** + * Validates that [ConcurrentHashMap.computeIfAbsent] combined with [CoroutineScope.launch] + * prevents double execution of the launch body under concurrent access. + * + * This is the exact idiom used in [ServerModule.launchSessionRun] and + * [ServerModule.launchSessionResume] to atomically register the running job + * before launching it, replacing the non-atomic check-then-act pattern + * (containsKey → launch → put) which allowed duplicate launches. + */ +class LaunchRegistrationRaceTest { + + @Test + fun `concurrent computeIfAbsent only invokes mapping lambda once per key`(): Unit = runTest { + val registry = ConcurrentHashMap() + val launchCount = AtomicInteger(0) + val scope = CoroutineScope(SupervisorJob() + Dispatchers.Default) + val key = "session-1" + + coroutineScope { + (1..50).map { + async(Dispatchers.Default) { + registry.computeIfAbsent(key) { _: String -> + scope.launch { + launchCount.incrementAndGet() + delay(500) + } + } + } + }.awaitAll() + } + + assertEquals(1, launchCount.get(), "mapping lambda must execute exactly once") + assertEquals(1, registry.size, "registry must contain exactly one entry") + scope.cancel() + } + + @Test + fun `concurrent computeIfAbsent on different keys invokes each lambda independently`(): Unit = runTest { + val registry = ConcurrentHashMap() + val launchCount = AtomicInteger(0) + val scope = CoroutineScope(SupervisorJob() + Dispatchers.Default) + + coroutineScope { + (1..50).map { i -> + async(Dispatchers.Default) { + registry.computeIfAbsent("session-$i") { _: String -> + scope.launch { + launchCount.incrementAndGet() + delay(100) + } + } + } + }.awaitAll() + } + + assertEquals(50, launchCount.get(), "each distinct key must execute its lambda") + assertEquals(50, registry.size, "registry must contain 50 entries") + scope.cancel() + } +}