feat(server-01): extend wire protocol with sequence cursors and SnapshotComplete

Every event-derived ServerMessage gains sequence (global) and sessionSequence
(per-session) cursor fields. SessionSnapshot adds lastSequence and
lastSessionSequence at projection-read time. SnapshotComplete data object
terminates the snapshot phase and is registered in the polymorphic module.
replaySnapshot() now emits SnapshotComplete after all SessionSnapshot messages.
SessionStarted audit deferred with TODO comment. DomainEventMapper propagates
both cursors to all mapped messages including InferenceCompleted.
This commit is contained in:
2026-05-24 21:47:30 +04:00
parent fc7b879891
commit 0cfb784187
14 changed files with 569 additions and 105 deletions
@@ -57,6 +57,8 @@ class ApprovalCoordinator(
),
toolName = null,
preview = null,
sequence = 0L,
sessionSequence = 0L,
)
broadcast(event.sessionId, msg)
scheduleTimeout(event.requestId, event.sessionId, event.stageId, event.tier)
@@ -64,12 +66,25 @@ class ApprovalCoordinator(
fun handleResponse(msg: ClientMessage.ApprovalResponse, sessionId: SessionId): ServerMessage? {
if (resolved.putIfAbsent(msg.requestId, true) != null) {
return ServerMessage.ProtocolError("Approval request ${msg.requestId.value} already resolved")
return ServerMessage.ProtocolError(
message = "Approval request ${msg.requestId.value} already resolved",
sequence = null,
sessionSequence = null,
)
}
timeoutJobs.remove(msg.requestId)?.cancel()
val domain = msg.toDomainDecision(sessionId, null, Tier.T2)
return runCatching { orchestrator.submitApprovalDecision(msg.requestId, domain) }
.fold(onSuccess = { null }, onFailure = { ServerMessage.ProtocolError(it.message ?: "Unknown error") })
.fold(
onSuccess = { null },
onFailure = {
ServerMessage.ProtocolError(
message = it.message ?: "Unknown error",
sequence = null,
sessionSequence = null,
)
},
)
}
private suspend fun broadcast(sessionId: SessionId, msg: ServerMessage) {
@@ -6,6 +6,7 @@ import com.correx.apps.server.protocol.ServerMessage
import com.correx.core.artifactstore.ArtifactStore
import com.correx.core.events.events.ApprovalRequestedEvent
import com.correx.core.events.events.InferenceCompletedEvent
import com.correx.core.events.events.WorkflowStartedEvent
import com.correx.core.events.events.InferenceStartedEvent
import com.correx.core.events.events.InferenceTimeoutEvent
import com.correx.core.events.events.OrchestrationPausedEvent
@@ -33,20 +34,47 @@ private object NoopArtifactStore : ArtifactStore {
}
@Suppress("CyclomaticComplexMethod")
suspend fun domainEventToServerMessage(event: StoredEvent, artifactStore: ArtifactStore): ServerMessage? =
when (val p = event.payload) {
is WorkflowCompletedEvent -> ServerMessage.SessionCompleted(sessionId = p.sessionId)
is WorkflowFailedEvent -> ServerMessage.SessionFailed(sessionId = p.sessionId, reason = p.reason)
suspend fun domainEventToServerMessage(
event: StoredEvent,
artifactStore: ArtifactStore,
sessionSequence: Long = 0L,
): ServerMessage? {
val seq = event.sequence
return when (val p = event.payload) {
is WorkflowStartedEvent -> ServerMessage.SessionStarted(
sessionId = p.sessionId,
workflowId = p.workflowId,
sequence = seq,
sessionSequence = sessionSequence,
)
is WorkflowCompletedEvent -> ServerMessage.SessionCompleted(
sessionId = p.sessionId,
sequence = seq,
sessionSequence = sessionSequence,
)
is WorkflowFailedEvent -> ServerMessage.SessionFailed(
sessionId = p.sessionId,
reason = p.reason,
sequence = seq,
sessionSequence = sessionSequence,
)
is TransitionExecutedEvent -> ServerMessage.StageStarted(
sessionId = p.sessionId,
stageId = p.to,
occurredAt = event.metadata.timestamp.toEpochMilliseconds(),
sequence = seq,
sessionSequence = sessionSequence,
)
is StageCompletedEvent -> ServerMessage.StageCompleted(
sessionId = p.sessionId,
stageId = p.stageId,
occurredAt = event.metadata.timestamp.toEpochMilliseconds(),
sequence = seq,
sessionSequence = sessionSequence,
)
is StageFailedEvent -> ServerMessage.StageFailed(
@@ -54,17 +82,33 @@ suspend fun domainEventToServerMessage(event: StoredEvent, artifactStore: Artifa
stageId = p.stageId,
reason = p.reason,
occurredAt = event.metadata.timestamp.toEpochMilliseconds(),
sequence = seq,
sessionSequence = sessionSequence,
)
is OrchestrationPausedEvent -> mapOrchestrationPaused(p)
is InferenceStartedEvent -> ServerMessage.InferenceStarted(sessionId = p.sessionId, stageId = p.stageId)
is InferenceCompletedEvent -> mapInferenceCompleted(event, p, artifactStore)
is OrchestrationPausedEvent -> mapOrchestrationPaused(p, seq, sessionSequence)
is InferenceStartedEvent -> ServerMessage.InferenceStarted(
sessionId = p.sessionId,
stageId = p.stageId,
sequence = seq,
sessionSequence = sessionSequence,
)
is InferenceCompletedEvent -> mapInferenceCompleted(event, p, artifactStore, sessionSequence)
is InferenceTimeoutEvent -> ServerMessage.InferenceTimedOut(
sessionId = p.sessionId, stageId = p.stageId, elapsedMs = p.timeoutMs,
sessionId = p.sessionId,
stageId = p.stageId,
elapsedMs = p.timeoutMs,
sequence = seq,
sessionSequence = sessionSequence,
)
is ToolInvocationRequestedEvent -> ServerMessage.ToolStarted(
sessionId = p.sessionId, toolName = p.toolName, tier = p.tier,
sessionId = p.sessionId,
toolName = p.toolName,
tier = p.tier,
sequence = seq,
sessionSequence = sessionSequence,
)
is ToolExecutionCompletedEvent -> ServerMessage.ToolCompleted(
@@ -72,6 +116,8 @@ suspend fun domainEventToServerMessage(event: StoredEvent, artifactStore: Artifa
toolName = p.toolName,
outputSummary = p.receipt.outputSummary,
occurredAt = event.metadata.timestamp.toEpochMilliseconds(),
sequence = seq,
sessionSequence = sessionSequence,
)
is ToolExecutionFailedEvent -> ServerMessage.ToolFailed(
@@ -79,25 +125,42 @@ suspend fun domainEventToServerMessage(event: StoredEvent, artifactStore: Artifa
toolName = p.toolName,
reason = p.reason,
occurredAt = event.metadata.timestamp.toEpochMilliseconds(),
sequence = seq,
sessionSequence = sessionSequence,
)
is ToolExecutionRejectedEvent -> ServerMessage.ToolRejected(
sessionId = p.sessionId, toolName = p.toolName, reason = p.reason,
sessionId = p.sessionId,
toolName = p.toolName,
reason = p.reason,
sequence = seq,
sessionSequence = sessionSequence,
)
is ApprovalRequestedEvent -> mapApprovalRequested(p)
is ApprovalRequestedEvent -> mapApprovalRequested(p, seq, sessionSequence)
else -> null
}
}
private fun mapOrchestrationPaused(p: OrchestrationPausedEvent): ServerMessage {
private fun mapOrchestrationPaused(
p: OrchestrationPausedEvent,
seq: Long,
sessionSequence: Long,
): ServerMessage {
val reason = if (p.reason == "APPROVAL_PENDING") PauseReason.APPROVAL_PENDING else PauseReason.USER_REQUESTED
return ServerMessage.SessionPaused(sessionId = p.sessionId, reason = reason)
return ServerMessage.SessionPaused(
sessionId = p.sessionId,
reason = reason,
sequence = seq,
sessionSequence = sessionSequence,
)
}
private suspend fun mapInferenceCompleted(
event: StoredEvent,
p: InferenceCompletedEvent,
artifactStore: ArtifactStore,
sessionSequence: Long,
): ServerMessage = runCatching {
artifactStore.get(p.responseArtifactId)?.toString(Charsets.UTF_8) ?: ""
}.getOrElse { "" }.run {
@@ -107,10 +170,16 @@ private suspend fun mapInferenceCompleted(
outputSummary = this,
responseText = this,
occurredAt = event.metadata.timestamp.toEpochMilliseconds(),
sequence = event.sequence,
sessionSequence = sessionSequence,
)
}
private fun mapApprovalRequested(p: ApprovalRequestedEvent): ServerMessage =
private fun mapApprovalRequested(
p: ApprovalRequestedEvent,
seq: Long,
sessionSequence: Long,
): ServerMessage =
ServerMessage.ApprovalRequired(
sessionId = p.sessionId,
requestId = p.requestId,
@@ -118,4 +187,6 @@ private fun mapApprovalRequested(p: ApprovalRequestedEvent): ServerMessage =
riskSummary = RiskSummaryDto("unknown", emptyList(), ""),
toolName = p.toolName,
preview = p.preview,
sequence = seq,
sessionSequence = sessionSequence,
)
@@ -1,6 +1,7 @@
package com.correx.apps.server.bridge
import com.correx.apps.server.protocol.ServerMessage
import com.correx.apps.server.protocol.SessionStateDto
import com.correx.core.approvals.DefaultApprovalRepository
import com.correx.core.artifactstore.ArtifactStore
import com.correx.core.events.orchestration.OrchestrationStatus
@@ -29,6 +30,7 @@ class SessionEventBridge(
approvalState.decisions.values.none { it.requestId == req.id }
}
val lastSeq = eventStore.lastSequence(sessionId) ?: 0L
send(ServerMessage.SessionSnapshot(
sessionId = sessionId,
workflowId = orchState.workflowId,
@@ -40,13 +42,25 @@ class SessionEventBridge(
approvalTier = pendingRequest?.tier?.name,
approvalToolName = pendingRequest?.toolName,
approvalPreview = pendingRequest?.preview,
state = SessionStateDto(
status = orchState.status.name,
currentStageId = orchState.currentStageId?.value,
pauseReason = orchState.pauseReason,
),
pendingApprovals = emptyList(),
lastSequence = lastSeq,
lastSessionSequence = 0L,
))
}
send(ServerMessage.SnapshotComplete)
}
suspend fun streamLive(sessionId: SessionId) {
var sessionSequence = 0L
eventStore.subscribe(sessionId).collect { event ->
domainEventToServerMessage(event, artifactStore)?.let { send(it) }
sessionSequence++
domainEventToServerMessage(event, artifactStore, sessionSequence = sessionSequence)
?.let { send(it) }
}
}
}
@@ -22,6 +22,19 @@ data class SessionConfigDto(
val retryPolicy: String?,
)
@Serializable
data class SessionStateDto(
val status: String,
val currentStageId: String?,
val pauseReason: String?,
)
@Serializable
data class ApprovalDto(
val requestId: String,
val tier: String,
)
@Serializable
enum class PauseReason {
APPROVAL_PENDING,
@@ -8,18 +8,69 @@ import kotlinx.serialization.SerialName
import kotlinx.serialization.Serializable
@Serializable
sealed class ServerMessage {
sealed interface ServerMessage {
/** Global monotonic sequence cursor. Non-null for event-derived messages. */
val sequence: Long?
/** Per-session monotonic sequence cursor. Non-null for event-derived messages. */
val sessionSequence: Long?
/** Marker for event-derived messages — cursors must be non-null. */
@Serializable
data class SessionStarted(val sessionId: SessionId, val workflowId: String) : ServerMessage()
sealed interface SessionMessage : ServerMessage {
override val sequence: Long
override val sessionSequence: Long
}
/** Marker for control / infra messages — cursors are always null. */
@Serializable
sealed interface NonEventMessage : ServerMessage {
override val sequence: Long? get() = null
override val sessionSequence: Long? get() = null
}
// -- Session lifecycle --
/**
* @see SessionStartedEvent — emitted when a session begins.
* TODO(cleanup-NN): retire SessionStarted once TUI migrates to event-derived ordering.
*/
@Serializable
@SerialName("session.started")
data class SessionStarted(
val sessionId: SessionId,
val workflowId: String,
override val sequence: Long,
override val sessionSequence: Long,
) : ServerMessage, SessionMessage
@Serializable
data class SessionPaused(val sessionId: SessionId, val reason: PauseReason) : ServerMessage()
@SerialName("session.paused")
data class SessionPaused(
val sessionId: SessionId,
val reason: PauseReason,
override val sequence: Long,
override val sessionSequence: Long,
) : ServerMessage, SessionMessage
@Serializable
data class SessionCompleted(val sessionId: SessionId) : ServerMessage()
@SerialName("session.completed")
data class SessionCompleted(
val sessionId: SessionId,
override val sequence: Long,
override val sessionSequence: Long,
) : ServerMessage, SessionMessage
@Serializable
data class SessionFailed(val sessionId: SessionId, val reason: String) : ServerMessage()
@SerialName("session.failed")
data class SessionFailed(
val sessionId: SessionId,
val reason: String,
override val sequence: Long,
override val sessionSequence: Long,
) : ServerMessage, SessionMessage
// -- State snapshot --
@Serializable
@SerialName("session_snapshot")
@@ -34,55 +85,128 @@ sealed class ServerMessage {
val approvalTier: String?,
val approvalToolName: String?,
val approvalPreview: String?,
) : ServerMessage()
val state: SessionStateDto,
val pendingApprovals: List<ApprovalDto>,
val lastSequence: Long,
val lastSessionSequence: Long,
override val sequence: Long? = null,
override val sessionSequence: Long? = null,
) : ServerMessage, NonEventMessage
// -- Stage lifecycle --
@Serializable
data class StageStarted(val sessionId: SessionId, val stageId: StageId, val occurredAt: Long) : ServerMessage()
@SerialName("stage.started")
data class StageStarted(
val sessionId: SessionId,
val stageId: StageId,
val occurredAt: Long,
override val sequence: Long,
override val sessionSequence: Long,
) : ServerMessage, SessionMessage
@Serializable
data class StageCompleted(val sessionId: SessionId, val stageId: StageId, val occurredAt: Long) : ServerMessage()
@SerialName("stage.completed")
data class StageCompleted(
val sessionId: SessionId,
val stageId: StageId,
val occurredAt: Long,
override val sequence: Long,
override val sessionSequence: Long,
) : ServerMessage, SessionMessage
@Serializable
data class StageFailed(val sessionId: SessionId, val stageId: StageId, val reason: String, val occurredAt: Long) :
ServerMessage()
@SerialName("stage.failed")
data class StageFailed(
val sessionId: SessionId,
val stageId: StageId,
val reason: String,
val occurredAt: Long,
override val sequence: Long,
override val sessionSequence: Long,
) : ServerMessage, SessionMessage
// -- Inference --
@Serializable
data class InferenceStarted(val sessionId: SessionId, val stageId: StageId) : ServerMessage()
@SerialName("inference.started")
data class InferenceStarted(
val sessionId: SessionId,
val stageId: StageId,
override val sequence: Long,
override val sessionSequence: Long,
) : ServerMessage, SessionMessage
@Serializable
@SerialName("inference.completed")
data class InferenceCompleted(
val sessionId: SessionId,
val stageId: StageId,
val outputSummary: String,
val responseText: String = "",
val occurredAt: Long,
) : ServerMessage()
override val sequence: Long,
override val sessionSequence: Long,
) : ServerMessage, SessionMessage
@Serializable
data class InferenceTimedOut(val sessionId: SessionId, val stageId: StageId, val elapsedMs: Long) :
ServerMessage()
@SerialName("inference.timed_out")
data class InferenceTimedOut(
val sessionId: SessionId,
val stageId: StageId,
val elapsedMs: Long,
override val sequence: Long,
override val sessionSequence: Long,
) : ServerMessage, SessionMessage
// -- Tool execution --
@Serializable
data class ToolStarted(val sessionId: SessionId, val toolName: String, val tier: Tier) :
ServerMessage()
@SerialName("tool.started")
data class ToolStarted(
val sessionId: SessionId,
val toolName: String,
val tier: Tier,
override val sequence: Long,
override val sessionSequence: Long,
) : ServerMessage, SessionMessage
@Serializable
@SerialName("tool.completed")
data class ToolCompleted(
val sessionId: SessionId,
val toolName: String,
val outputSummary: String,
val occurredAt: Long,
) : ServerMessage()
override val sequence: Long,
override val sessionSequence: Long,
) : ServerMessage, SessionMessage
@Serializable
data class ToolFailed(val sessionId: SessionId, val toolName: String, val reason: String, val occurredAt: Long) :
ServerMessage()
@SerialName("tool.failed")
data class ToolFailed(
val sessionId: SessionId,
val toolName: String,
val reason: String,
val occurredAt: Long,
override val sequence: Long,
override val sessionSequence: Long,
) : ServerMessage, SessionMessage
@Serializable
data class ToolRejected(val sessionId: SessionId, val toolName: String, val reason: String) :
ServerMessage()
@SerialName("tool.rejected")
data class ToolRejected(
val sessionId: SessionId,
val toolName: String,
val reason: String,
override val sequence: Long,
override val sessionSequence: Long,
) : ServerMessage, SessionMessage
// -- Approval --
@Serializable
@SerialName("approval.required")
data class ApprovalRequired(
val sessionId: SessionId,
val requestId: ApprovalRequestId,
@@ -90,19 +214,48 @@ sealed class ServerMessage {
val riskSummary: RiskSummaryDto,
val toolName: String?,
val preview: String?,
) : ServerMessage()
override val sequence: Long,
override val sessionSequence: Long,
) : ServerMessage, SessionMessage
// -- System / infra --
@Serializable
data class ProviderStatusChanged(val providerId: String, val status: ProviderHealthDto) :
ServerMessage()
@SerialName("provider.status_changed")
data class ProviderStatusChanged(
val providerId: String,
val status: ProviderHealthDto,
override val sequence: Long? = null,
override val sessionSequence: Long? = null,
) : ServerMessage, NonEventMessage
@Serializable
data class ProtocolError(val message: String) : ServerMessage()
@SerialName("protocol_error")
data class ProtocolError(
val message: String,
override val sequence: Long? = null,
override val sessionSequence: Long? = null,
) : ServerMessage, NonEventMessage
// -- Router --
@Serializable
@SerialName("router.response")
data class RouterResponseMessage(
val sessionId: SessionId,
val content: String,
val steeringEmitted: Boolean,
) : ServerMessage()
override val sequence: Long? = null,
override val sessionSequence: Long? = null,
) : ServerMessage, NonEventMessage
// -- Snapshot phase markers --
/** Marker indicating the snapshot phase has completed. No cursor fields. */
@Serializable
@SerialName("snapshot_complete")
data object SnapshotComplete : ServerMessage, NonEventMessage {
override val sequence: Long? get() = null
override val sessionSequence: Long? get() = null
}
}
@@ -61,7 +61,11 @@ class GlobalStreamHandler(private val module: ServerModule) {
}
.onFailure {
log.warn("decode error: {}", it.message)
val error = ServerMessage.ProtocolError("Unknown message: ${it.message}")
val error = ServerMessage.ProtocolError(
message = "Unknown message: ${it.message}",
sequence = null,
sessionSequence = null,
)
session.send(Frame.Text(ProtocolSerializer.encodeServerMessage(error)))
}
}
@@ -88,6 +92,8 @@ class GlobalStreamHandler(private val module: ServerModule) {
val msg = ServerMessage.ProviderStatusChanged(
providerId = providerId.value,
status = ProviderHealthDto(providerId.value, status, null),
sequence = null,
sessionSequence = null,
)
session.send(Frame.Text(ProtocolSerializer.encodeServerMessage(msg)))
}
@@ -112,7 +118,15 @@ class GlobalStreamHandler(private val module: ServerModule) {
}
private fun encodeError(message: String): Frame.Text =
Frame.Text(ProtocolSerializer.encodeServerMessage(ServerMessage.ProtocolError(message)))
Frame.Text(
ProtocolSerializer.encodeServerMessage(
ServerMessage.ProtocolError(
message = message,
sequence = null,
sessionSequence = null,
),
),
)
private suspend fun handleStartSession(
session: DefaultWebSocketServerSession,
@@ -123,7 +137,11 @@ class GlobalStreamHandler(private val module: ServerModule) {
val graph = module.workflowRegistry.find(msg.workflowId)
log.info("find returned: {}", graph)
if (graph == null) {
val error = ServerMessage.ProtocolError("Unknown workflowId: ${msg.workflowId}")
val error = ServerMessage.ProtocolError(
message = "Unknown workflowId: ${msg.workflowId}",
sequence = null,
sessionSequence = null,
)
session.send(Frame.Text(ProtocolSerializer.encodeServerMessage(error)))
return
}
@@ -166,7 +184,12 @@ class GlobalStreamHandler(private val module: ServerModule) {
bridge.streamLive(sessionId)
},
)?.cancel()
val started = ServerMessage.SessionStarted(sessionId = sessionId, workflowId = msg.workflowId)
val started = ServerMessage.SessionStarted(
sessionId = sessionId,
workflowId = msg.workflowId,
sequence = 0L,
sessionSequence = 0L,
)
session.send(Frame.Text(ProtocolSerializer.encodeServerMessage(started)))
}
}
@@ -65,6 +65,7 @@ class DomainEventMapperTest {
correlationId = null,
),
sequence = 1L,
sessionSequence = 1L,
payload = payload,
)
@@ -72,8 +73,11 @@ class DomainEventMapperTest {
fun `WorkflowStartedEvent maps to SessionStarted`(): Unit = runTest {
val event =
storedEvent(WorkflowStartedEvent(sessionId = sessionId, startStageId = stageId, workflowId = workflowId))
val result = domainEventToServerMessage(event, noopStore)
assertEquals(ServerMessage.SessionStarted(sessionId, workflowId), result)
val result = domainEventToServerMessage(event, noopStore, sessionSequence = 0L)
assertEquals(
ServerMessage.SessionStarted(sessionId, workflowId, event.sequence, 0L),
result,
)
}
@Test
@@ -85,8 +89,11 @@ class DomainEventMapperTest {
totalStages = 1,
),
)
val result = domainEventToServerMessage(event, noopStore)
assertEquals(ServerMessage.SessionCompleted(sessionId), result)
val result = domainEventToServerMessage(event, noopStore, sessionSequence = 0L)
assertEquals(
ServerMessage.SessionCompleted(sessionId, event.sequence, 0L),
result,
)
}
@Test
@@ -94,8 +101,11 @@ class DomainEventMapperTest {
val event = storedEvent(
WorkflowFailedEvent(sessionId = sessionId, stageId = stageId, reason = "boom", retryExhausted = false),
)
val result = domainEventToServerMessage(event, noopStore)
assertEquals(ServerMessage.SessionFailed(sessionId, "boom"), result)
val result = domainEventToServerMessage(event, noopStore, sessionSequence = 0L)
assertEquals(
ServerMessage.SessionFailed(sessionId, "boom", event.sequence, 0L),
result,
)
}
@Test
@@ -105,8 +115,11 @@ class DomainEventMapperTest {
val event = storedEvent(
TransitionExecutedEvent(sessionId = sessionId, from = from, to = to, transitionId = TransitionId("t1")),
)
val result = domainEventToServerMessage(event, noopStore)
assertEquals(ServerMessage.StageStarted(sessionId, to, occurredAt), result)
val result = domainEventToServerMessage(event, noopStore, sessionSequence = 0L)
assertEquals(
ServerMessage.StageStarted(sessionId, to, occurredAt, event.sequence, 0L),
result,
)
}
@Test
@@ -114,8 +127,11 @@ class DomainEventMapperTest {
val event = storedEvent(
StageCompletedEvent(sessionId = sessionId, stageId = stageId, transitionId = TransitionId("t1")),
)
val result = domainEventToServerMessage(event, noopStore)
assertEquals(ServerMessage.StageCompleted(sessionId, stageId, occurredAt), result)
val result = domainEventToServerMessage(event, noopStore, sessionSequence = 0L)
assertEquals(
ServerMessage.StageCompleted(sessionId, stageId, occurredAt, event.sequence, 0L),
result,
)
}
@Test
@@ -125,8 +141,11 @@ class DomainEventMapperTest {
sessionId = sessionId, stageId = stageId, transitionId = TransitionId("t1"), reason = "err",
),
)
val result = domainEventToServerMessage(event, noopStore)
assertEquals(ServerMessage.StageFailed(sessionId, stageId, "err", occurredAt), result)
val result = domainEventToServerMessage(event, noopStore, sessionSequence = 0L)
assertEquals(
ServerMessage.StageFailed(sessionId, stageId, "err", occurredAt, event.sequence, 0L),
result,
)
}
@Test
@@ -134,8 +153,11 @@ class DomainEventMapperTest {
val event = storedEvent(
OrchestrationPausedEvent(sessionId = sessionId, stageId = stageId, reason = "APPROVAL_PENDING"),
)
val result = domainEventToServerMessage(event, noopStore)
assertEquals(ServerMessage.SessionPaused(sessionId, PauseReason.APPROVAL_PENDING), result)
val result = domainEventToServerMessage(event, noopStore, sessionSequence = 0L)
assertEquals(
ServerMessage.SessionPaused(sessionId, PauseReason.APPROVAL_PENDING, event.sequence, 0L),
result,
)
}
@Test
@@ -143,8 +165,11 @@ class DomainEventMapperTest {
val event = storedEvent(
OrchestrationPausedEvent(sessionId = sessionId, stageId = stageId, reason = "USER_REQUESTED"),
)
val result = domainEventToServerMessage(event, noopStore)
assertEquals(ServerMessage.SessionPaused(sessionId, PauseReason.USER_REQUESTED), result)
val result = domainEventToServerMessage(event, noopStore, sessionSequence = 0L)
assertEquals(
ServerMessage.SessionPaused(sessionId, PauseReason.USER_REQUESTED, event.sequence, 0L),
result,
)
}
@Test
@@ -158,8 +183,11 @@ class DomainEventMapperTest {
promptArtifactId = ArtifactId("art-prompt"),
),
)
val result = domainEventToServerMessage(event, noopStore)
assertEquals(ServerMessage.InferenceStarted(sessionId, stageId), result)
val result = domainEventToServerMessage(event, noopStore, sessionSequence = 0L)
assertEquals(
ServerMessage.InferenceStarted(sessionId, stageId, event.sequence, 0L),
result,
)
}
@Test
@@ -184,9 +212,11 @@ class DomainEventMapperTest {
responseArtifactId = artifactId,
),
)
val result = domainEventToServerMessage(event, store)
val result = domainEventToServerMessage(event, store, sessionSequence = 0L)
assertEquals(
ServerMessage.InferenceCompleted(sessionId, stageId, responseText, responseText, occurredAt),
ServerMessage.InferenceCompleted(
sessionId, stageId, responseText, responseText, occurredAt, event.sequence, 0L,
),
result,
)
}
@@ -204,8 +234,11 @@ class DomainEventMapperTest {
responseArtifactId = ArtifactId("missing"),
),
)
val result = domainEventToServerMessage(event, noopStore)
assertEquals(ServerMessage.InferenceCompleted(sessionId, stageId, "", "", occurredAt), result)
val result = domainEventToServerMessage(event, noopStore, sessionSequence = 0L)
assertEquals(
ServerMessage.InferenceCompleted(sessionId, stageId, "", "", occurredAt, event.sequence, 0L),
result,
)
}
@Test
@@ -219,8 +252,11 @@ class DomainEventMapperTest {
timeoutMs = 5000L,
),
)
val result = domainEventToServerMessage(event, noopStore)
assertEquals(ServerMessage.InferenceTimedOut(sessionId, stageId, 5000L), result)
val result = domainEventToServerMessage(event, noopStore, sessionSequence = 0L)
assertEquals(
ServerMessage.InferenceTimedOut(sessionId, stageId, 5000L, event.sequence, 0L),
result,
)
}
@Test
@@ -238,8 +274,11 @@ class DomainEventMapperTest {
),
),
)
val result = domainEventToServerMessage(event, noopStore)
assertEquals(ServerMessage.ToolStarted(sessionId, "file_write", Tier.T2), result)
val result = domainEventToServerMessage(event, noopStore, sessionSequence = 0L)
assertEquals(
ServerMessage.ToolStarted(sessionId, "file_write", Tier.T2, event.sequence, 0L),
result,
)
}
@Test
@@ -259,8 +298,11 @@ class DomainEventMapperTest {
invocationId = invId, sessionId = sessionId, toolName = "file_write", receipt = receipt,
),
)
val result = domainEventToServerMessage(event, noopStore)
assertEquals(ServerMessage.ToolCompleted(sessionId, "file_write", "wrote 3 lines", occurredAt), result)
val result = domainEventToServerMessage(event, noopStore, sessionSequence = 0L)
assertEquals(
ServerMessage.ToolCompleted(sessionId, "file_write", "wrote 3 lines", occurredAt, event.sequence, 0L),
result,
)
}
@Test
@@ -273,8 +315,11 @@ class DomainEventMapperTest {
reason = "disk full",
),
)
val result = domainEventToServerMessage(event, noopStore)
assertEquals(ServerMessage.ToolFailed(sessionId, "file_write", "disk full", occurredAt), result)
val result = domainEventToServerMessage(event, noopStore, sessionSequence = 0L)
assertEquals(
ServerMessage.ToolFailed(sessionId, "file_write", "disk full", occurredAt, event.sequence, 0L),
result,
)
}
@Test
@@ -288,8 +333,11 @@ class DomainEventMapperTest {
reason = "policy denied",
),
)
val result = domainEventToServerMessage(event, noopStore)
assertEquals(ServerMessage.ToolRejected(sessionId, "file_write", "policy denied"), result)
val result = domainEventToServerMessage(event, noopStore, sessionSequence = 0L)
assertEquals(
ServerMessage.ToolRejected(sessionId, "file_write", "policy denied", event.sequence, 0L),
result,
)
}
@Test
@@ -306,17 +354,18 @@ class DomainEventMapperTest {
projectId = null,
),
)
val result = domainEventToServerMessage(event, noopStore) as ServerMessage.ApprovalRequired
val result = domainEventToServerMessage(event, noopStore, sessionSequence = 0L) as ServerMessage.ApprovalRequired
assertEquals(requestId, result.requestId)
assertEquals(Tier.T3, result.tier)
assertNull(result.toolName)
assertNull(result.preview)
assertEquals(event.sequence, result.sequence)
}
@Test
fun `unmapped event returns null`(): Unit = runTest {
val event = storedEvent(SessionStartedEvent(sessionId = sessionId))
val result = domainEventToServerMessage(event, noopStore)
val result = domainEventToServerMessage(event, noopStore, sessionSequence = 0L)
assertNull(result)
}
}
@@ -12,15 +12,16 @@ import com.correx.core.events.events.OrchestrationResumedEvent
import com.correx.core.events.events.StoredEvent
import com.correx.core.events.events.WorkflowCompletedEvent
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.DefaultOrchestrationReducer
import com.correx.core.kernel.orchestration.OrchestrationProjector
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.channels.Channel
import kotlinx.coroutines.flow.Flow
import kotlinx.coroutines.flow.flow
@@ -54,6 +55,7 @@ class SessionEventBridgeTest {
correlationId = null,
),
sequence = seq,
sessionSequence = seq,
payload = payload,
)
@@ -65,17 +67,26 @@ class SessionEventBridgeTest {
override suspend fun appendAll(events: List<NewEvent>): List<StoredEvent> = error("unused")
override fun read(sessionId: SessionId): List<StoredEvent> = allEventsList
override fun readFrom(sessionId: SessionId, fromSequence: Long): List<StoredEvent> = allEventsList
override fun lastSequence(sessionId: SessionId): Long? = null
override fun lastSequence(sessionId: SessionId): Long? =
allEventsList.maxOfOrNull { it.sequence } ?: 0L
override fun subscribe(sessionId: SessionId): Flow<StoredEvent> = liveFlow
override fun allEvents(): Sequence<StoredEvent> = allEventsList.asSequence()
override fun allSessionIds(): Set<SessionId> = allEventsList.map { it.metadata.sessionId }.toSet()
override fun subscribeAll(): Flow<StoredEvent> = TODO("Not needed in this test context")
override suspend fun lastGlobalSequence(): Long = TODO("Not needed in this test context")
}
private val orchestrationRepository = OrchestrationRepository(
DefaultEventReplayer(
fakeEventStore(),
OrchestrationProjector(DefaultOrchestrationReducer()),
),
private fun activeOrchestrationRepository(
status: OrchestrationStatus = OrchestrationStatus.RUNNING,
): OrchestrationRepository = OrchestrationRepository(
object : EventReplayer<OrchestrationState> {
override fun rebuild(sessionId: SessionId): OrchestrationState = OrchestrationState(
workflowId = workflowId,
status = status,
currentStageId = stageId,
pendingApproval = false,
)
},
)
private val approvalRepository = DefaultApprovalRepository(
@@ -92,19 +103,23 @@ class SessionEventBridgeTest {
storedEvent(WorkflowCompletedEvent(sessionId, stageId, 1), seq = 2L),
)
val store = fakeEventStore(allEventsList = events)
val orchRepo = activeOrchestrationRepository()
val sent = mutableListOf<ServerMessage>()
val bridge = SessionEventBridge(
store,
noopArtifactStore,
orchestrationRepository,
orchRepo,
approvalRepository,
) { sent.add(it) }
bridge.replaySnapshot()
assertEquals(2, sent.size)
assertEquals(ServerMessage.SessionStarted(sessionId, workflowId), sent[0])
assertEquals(ServerMessage.SessionCompleted(sessionId), sent[1])
val snapshot = sent[0] as ServerMessage.SessionSnapshot
assertEquals(sessionId, snapshot.sessionId)
assertEquals(2L, snapshot.lastSequence)
assertEquals(0L, snapshot.lastSessionSequence)
assertEquals(ServerMessage.SnapshotComplete, sent[1])
}
@Test
@@ -114,18 +129,21 @@ class SessionEventBridgeTest {
storedEvent(WorkflowCompletedEvent(sessionId, stageId, 1), seq = 2L),
)
val store = fakeEventStore(allEventsList = events)
val orchRepo = activeOrchestrationRepository()
val sent = mutableListOf<ServerMessage>()
val bridge = SessionEventBridge(
store,
noopArtifactStore,
orchestrationRepository,
orchRepo,
approvalRepository,
) { sent.add(it) }
bridge.replaySnapshot()
assertEquals(1, sent.size)
assertEquals(ServerMessage.SessionCompleted(sessionId), sent[0])
assertEquals(2, sent.size)
val snapshot = sent[0] as ServerMessage.SessionSnapshot
assertEquals(sessionId, snapshot.sessionId)
assertEquals(ServerMessage.SnapshotComplete, sent[1])
}
@Test
@@ -140,15 +158,21 @@ class SessionEventBridgeTest {
val bridge = SessionEventBridge(
store,
noopArtifactStore,
orchestrationRepository,
activeOrchestrationRepository(),
approvalRepository,
) { sent.add(it) }
bridge.streamLive(sessionId)
assertEquals(2, sent.size)
assertEquals(ServerMessage.SessionStarted(sessionId, workflowId), sent[0])
assertEquals(ServerMessage.SessionCompleted(sessionId), sent[1])
assertEquals(
ServerMessage.SessionStarted(sessionId, workflowId, 1L, 1L),
sent[0],
)
assertEquals(
ServerMessage.SessionCompleted(sessionId, 2L, 2L),
sent[1],
)
}
@Test
@@ -162,7 +186,7 @@ class SessionEventBridgeTest {
val bridge = SessionEventBridge(
store,
noopArtifactStore,
orchestrationRepository,
activeOrchestrationRepository(),
approvalRepository,
) { sent.add(it) }
@@ -0,0 +1,82 @@
package com.correx.apps.server.protocol
import com.correx.core.events.types.SessionId
import kotlinx.serialization.json.Json
import org.junit.jupiter.api.Assertions.assertEquals
import org.junit.jupiter.api.Test
class ServerMessageSerializationTest {
private val json = Json {
classDiscriminator = "type"
ignoreUnknownKeys = true
}
@Test
fun `SessionStarted encodes sequence and sessionSequence`() {
val msg = ServerMessage.SessionStarted(
sessionId = SessionId("sess-1"),
workflowId = "wf-1",
sequence = 42,
sessionSequence = 7,
)
val jsonStr = ProtocolSerializer.encodeServerMessage(msg)
assert(jsonStr.contains("\"type\":\"session.started\"")) { "expected type=session.started" }
assert(jsonStr.contains("\"sessionId\":\"sess-1\"")) { "expected sessionId" }
assert(jsonStr.contains("\"workflowId\":\"wf-1\"")) { "expected workflowId" }
assert(jsonStr.contains("\"sequence\":42")) { "expected sequence=42" }
assert(jsonStr.contains("\"sessionSequence\":7")) { "expected sessionSequence=7" }
}
@Test
fun `SnapshotComplete encodes with only type field`() {
val msg = ServerMessage.SnapshotComplete
val jsonStr = ProtocolSerializer.encodeServerMessage(msg)
assert(jsonStr.contains("\"type\":\"snapshot_complete\"")) { "expected type=snapshot_complete" }
}
@Test
fun `SessionSnapshot encodes lastSequence and lastSessionSequence`() {
val msg = ServerMessage.SessionSnapshot(
sessionId = SessionId("sess-1"),
workflowId = "wf-1",
status = "running",
currentStageId = "stage-1",
pauseReason = null,
pendingApproval = false,
approvalRequestId = null,
approvalTier = null,
approvalToolName = null,
approvalPreview = null,
state = SessionStateDto(
status = "running",
currentStageId = "stage-1",
pauseReason = null,
),
pendingApprovals = emptyList(),
lastSequence = 100,
lastSessionSequence = 15,
)
val jsonStr = ProtocolSerializer.encodeServerMessage(msg)
assert(jsonStr.contains("\"type\":\"session_snapshot\"")) { "expected type=session_snapshot" }
assert(jsonStr.contains("\"sessionId\":\"sess-1\"")) { "expected sessionId" }
assert(jsonStr.contains("\"lastSequence\":100")) { "expected lastSequence=100" }
assert(jsonStr.contains("\"lastSessionSequence\":15")) { "expected lastSessionSequence=15" }
}
@Test
fun `SessionStarted round-trips through JSON`() {
val original = ServerMessage.SessionStarted(
sessionId = SessionId("sess-42"),
workflowId = "wf-abc",
sequence = 200,
sessionSequence = 33,
)
val jsonStr = ProtocolSerializer.encodeServerMessage(original)
val decoded = json.decodeFromString<ServerMessage.SessionStarted>(jsonStr)
assertEquals(original.sessionId, decoded.sessionId)
assertEquals(original.workflowId, decoded.workflowId)
assertEquals(200L, decoded.sequence)
assertEquals(33L, decoded.sessionSequence)
}
}
@@ -98,6 +98,8 @@ class ApprovalReducerTest {
riskSummary = RiskSummaryDto(level = "HIGH", factors = emptyList(), recommendedAction = "review"),
toolName = "bash",
preview = "rm -rf /",
sequence = 1L,
sessionSequence = 1L,
)
val (state, effects) = reduce(action = Action.ServerEventReceived(msg))
assertEquals("req-2", state.active?.requestId)
@@ -43,11 +43,15 @@ class RootReducerTest {
val startMsg = com.correx.apps.server.protocol.ServerMessage.SessionStarted(
sessionId = com.correx.core.events.types.SessionId("s1"),
workflowId = "wf",
sequence = 1L,
sessionSequence = 1L,
)
val (state1, _) = RootReducer.reduce(state0, Action.ServerEventReceived(startMsg), fixedClock)
val startMsg2 = com.correx.apps.server.protocol.ServerMessage.SessionStarted(
sessionId = com.correx.core.events.types.SessionId("s2"),
workflowId = "wf",
sequence = 2L,
sessionSequence = 1L,
)
val (state2, _) = RootReducer.reduce(state1, Action.ServerEventReceived(startMsg2), fixedClock)
// selected is s1 (first); navigating down should move to s2
@@ -91,7 +91,7 @@ class SessionsReducerTest {
@Test
fun `ServerEventReceived SessionStarted appends session and sets selection`() {
val msg = ServerMessage.SessionStarted(sessionId = SessionId("s1"), workflowId = "wf1")
val msg = ServerMessage.SessionStarted(sessionId = SessionId("s1"), workflowId = "wf1", sequence = 1L, sessionSequence = 1L)
val (state, effects) = reduce(action = Action.ServerEventReceived(msg))
assertEquals(1, state.sessions.size)
assertEquals("s1", state.sessions[0].id)
@@ -103,7 +103,7 @@ class SessionsReducerTest {
@Test
fun `ServerEventReceived SessionStarted preserves existing selection`() {
val existing = SessionsState(sessions = listOf(session("existing")), selectedId = "existing")
val msg = ServerMessage.SessionStarted(sessionId = SessionId("new"), workflowId = "wf")
val msg = ServerMessage.SessionStarted(sessionId = SessionId("new"), workflowId = "wf", sequence = 1L, sessionSequence = 1L)
val (state, _) = reduce(sessions = existing, action = Action.ServerEventReceived(msg))
assertEquals("existing", state.selectedId)
assertEquals(2, state.sessions.size)
@@ -112,7 +112,7 @@ class SessionsReducerTest {
@Test
fun `ServerEventReceived SessionCompleted updates status to COMPLETED`() {
val s = SessionsState(sessions = listOf(session("s1")), selectedId = "s1")
val msg = ServerMessage.SessionCompleted(sessionId = SessionId("s1"))
val msg = ServerMessage.SessionCompleted(sessionId = SessionId("s1"), sequence = 1L, sessionSequence = 1L)
val (state, _) = reduce(sessions = s, action = Action.ServerEventReceived(msg))
assertEquals("COMPLETED", state.sessions[0].status)
assertEquals(1000L, state.sessions[0].lastEventAt)
@@ -121,7 +121,7 @@ class SessionsReducerTest {
@Test
fun `ServerEventReceived SessionFailed updates status to FAILED`() {
val s = SessionsState(sessions = listOf(session("s1")), selectedId = "s1")
val msg = ServerMessage.SessionFailed(sessionId = SessionId("s1"), reason = "oops")
val msg = ServerMessage.SessionFailed(sessionId = SessionId("s1"), reason = "oops", sequence = 1L, sessionSequence = 1L)
val (state, _) = reduce(sessions = s, action = Action.ServerEventReceived(msg))
assertEquals("FAILED", state.sessions[0].status)
}
@@ -136,6 +136,8 @@ class SessionsReducerTest {
sessionId = SessionId("s1"),
stageId = StageId("stage-1"),
occurredAt = fixedClock(),
sequence = 1L,
sessionSequence = 1L,
)
val (state, _) = reduce(sessions = s, action = Action.ServerEventReceived(msg))
assertEquals(null, state.sessions[0].currentStage)
@@ -153,6 +155,8 @@ class SessionsReducerTest {
stageId = StageId("stage-1"),
reason = "error",
occurredAt = fixedClock(),
sequence = 1L,
sessionSequence = 1L,
)
val (state, _) = reduce(sessions = s, action = Action.ServerEventReceived(msg))
assertEquals(null, state.sessions[0].currentStage)
@@ -183,6 +187,8 @@ class SessionsReducerTest {
sessionId = SessionId("s1"), stageId = StageId("stage-1"),
outputSummary = "summary", responseText = "full response",
occurredAt = fixedClock(),
sequence = 1L,
sessionSequence = 1L,
)
val (state, _) = reduce(sessions = s, action = Action.ServerEventReceived(msg))
assertEquals("summary", state.sessions[0].lastOutput)
@@ -197,6 +203,8 @@ class SessionsReducerTest {
sessionId = SessionId("s1"), stageId = StageId("stage-1"),
outputSummary = "", responseText = "",
occurredAt = fixedClock(),
sequence = 1L,
sessionSequence = 1L,
)
val (state, _) = reduce(sessions = s, action = Action.ServerEventReceived(msg))
assertEquals(null, state.sessions[0].lastResponseText)
@@ -209,6 +217,8 @@ class SessionsReducerTest {
sessionId = SessionId("s1"),
stageId = StageId("stage-1"),
occurredAt = fixedClock(),
sequence = 1L,
sessionSequence = 1L,
)
val (state, _) = reduce(sessions = s, action = Action.ServerEventReceived(msg))
assertEquals("stage-1", state.sessions[0].currentStage)
@@ -223,6 +233,8 @@ class SessionsReducerTest {
toolName = "file_write",
outputSummary = "wrote 3 lines",
occurredAt = fixedClock(),
sequence = 1L,
sessionSequence = 1L,
)
val (state, _) = reduce(sessions = s, action = Action.ServerEventReceived(msg))
assertEquals("file_write: wrote 3 lines", state.sessions[0].lastOutput)
@@ -232,7 +244,7 @@ class SessionsReducerTest {
@Test
fun `ServerEventReceived SessionPaused with APPROVAL_PENDING sets status label`() {
val s = SessionsState(sessions = listOf(session("s1")), selectedId = "s1")
val msg = ServerMessage.SessionPaused(sessionId = SessionId("s1"), reason = PauseReason.APPROVAL_PENDING)
val msg = ServerMessage.SessionPaused(sessionId = SessionId("s1"), reason = PauseReason.APPROVAL_PENDING, sequence = 1L, sessionSequence = 1L)
val (state, _) = reduce(sessions = s, action = Action.ServerEventReceived(msg))
assertEquals("PAUSED awaiting approval", state.sessions[0].status)
}
@@ -38,6 +38,7 @@ class DefaultApprovalReducerTest {
correlationId = null
),
sequence = sequence,
sessionSequence = sequence,
payload = payload
)
}
@@ -59,6 +59,7 @@ class DefaultToolReducerTest {
correlationId = null
),
sequence = 1L,
sessionSequence = 1L,
payload = payload
)