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, toolName = null,
preview = null, preview = null,
sequence = 0L,
sessionSequence = 0L,
) )
broadcast(event.sessionId, msg) broadcast(event.sessionId, msg)
scheduleTimeout(event.requestId, event.sessionId, event.stageId, event.tier) scheduleTimeout(event.requestId, event.sessionId, event.stageId, event.tier)
@@ -64,12 +66,25 @@ class ApprovalCoordinator(
fun handleResponse(msg: ClientMessage.ApprovalResponse, sessionId: SessionId): ServerMessage? { fun handleResponse(msg: ClientMessage.ApprovalResponse, sessionId: SessionId): ServerMessage? {
if (resolved.putIfAbsent(msg.requestId, true) != null) { 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() timeoutJobs.remove(msg.requestId)?.cancel()
val domain = msg.toDomainDecision(sessionId, null, Tier.T2) val domain = msg.toDomainDecision(sessionId, null, Tier.T2)
return runCatching { orchestrator.submitApprovalDecision(msg.requestId, domain) } 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) { 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.artifactstore.ArtifactStore
import com.correx.core.events.events.ApprovalRequestedEvent import com.correx.core.events.events.ApprovalRequestedEvent
import com.correx.core.events.events.InferenceCompletedEvent 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.InferenceStartedEvent
import com.correx.core.events.events.InferenceTimeoutEvent import com.correx.core.events.events.InferenceTimeoutEvent
import com.correx.core.events.events.OrchestrationPausedEvent import com.correx.core.events.events.OrchestrationPausedEvent
@@ -33,20 +34,47 @@ private object NoopArtifactStore : ArtifactStore {
} }
@Suppress("CyclomaticComplexMethod") @Suppress("CyclomaticComplexMethod")
suspend fun domainEventToServerMessage(event: StoredEvent, artifactStore: ArtifactStore): ServerMessage? = suspend fun domainEventToServerMessage(
when (val p = event.payload) { event: StoredEvent,
is WorkflowCompletedEvent -> ServerMessage.SessionCompleted(sessionId = p.sessionId) artifactStore: ArtifactStore,
is WorkflowFailedEvent -> ServerMessage.SessionFailed(sessionId = p.sessionId, reason = p.reason) 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( is TransitionExecutedEvent -> ServerMessage.StageStarted(
sessionId = p.sessionId, sessionId = p.sessionId,
stageId = p.to, stageId = p.to,
occurredAt = event.metadata.timestamp.toEpochMilliseconds(), occurredAt = event.metadata.timestamp.toEpochMilliseconds(),
sequence = seq,
sessionSequence = sessionSequence,
) )
is StageCompletedEvent -> ServerMessage.StageCompleted( is StageCompletedEvent -> ServerMessage.StageCompleted(
sessionId = p.sessionId, sessionId = p.sessionId,
stageId = p.stageId, stageId = p.stageId,
occurredAt = event.metadata.timestamp.toEpochMilliseconds(), occurredAt = event.metadata.timestamp.toEpochMilliseconds(),
sequence = seq,
sessionSequence = sessionSequence,
) )
is StageFailedEvent -> ServerMessage.StageFailed( is StageFailedEvent -> ServerMessage.StageFailed(
@@ -54,17 +82,33 @@ suspend fun domainEventToServerMessage(event: StoredEvent, artifactStore: Artifa
stageId = p.stageId, stageId = p.stageId,
reason = p.reason, reason = p.reason,
occurredAt = event.metadata.timestamp.toEpochMilliseconds(), occurredAt = event.metadata.timestamp.toEpochMilliseconds(),
sequence = seq,
sessionSequence = sessionSequence,
) )
is OrchestrationPausedEvent -> mapOrchestrationPaused(p) is OrchestrationPausedEvent -> mapOrchestrationPaused(p, seq, sessionSequence)
is InferenceStartedEvent -> ServerMessage.InferenceStarted(sessionId = p.sessionId, stageId = p.stageId) is InferenceStartedEvent -> ServerMessage.InferenceStarted(
is InferenceCompletedEvent -> mapInferenceCompleted(event, p, artifactStore) sessionId = p.sessionId,
stageId = p.stageId,
sequence = seq,
sessionSequence = sessionSequence,
)
is InferenceCompletedEvent -> mapInferenceCompleted(event, p, artifactStore, sessionSequence)
is InferenceTimeoutEvent -> ServerMessage.InferenceTimedOut( 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( 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( is ToolExecutionCompletedEvent -> ServerMessage.ToolCompleted(
@@ -72,6 +116,8 @@ suspend fun domainEventToServerMessage(event: StoredEvent, artifactStore: Artifa
toolName = p.toolName, toolName = p.toolName,
outputSummary = p.receipt.outputSummary, outputSummary = p.receipt.outputSummary,
occurredAt = event.metadata.timestamp.toEpochMilliseconds(), occurredAt = event.metadata.timestamp.toEpochMilliseconds(),
sequence = seq,
sessionSequence = sessionSequence,
) )
is ToolExecutionFailedEvent -> ServerMessage.ToolFailed( is ToolExecutionFailedEvent -> ServerMessage.ToolFailed(
@@ -79,25 +125,42 @@ suspend fun domainEventToServerMessage(event: StoredEvent, artifactStore: Artifa
toolName = p.toolName, toolName = p.toolName,
reason = p.reason, reason = p.reason,
occurredAt = event.metadata.timestamp.toEpochMilliseconds(), occurredAt = event.metadata.timestamp.toEpochMilliseconds(),
sequence = seq,
sessionSequence = sessionSequence,
) )
is ToolExecutionRejectedEvent -> ServerMessage.ToolRejected( 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 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 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( private suspend fun mapInferenceCompleted(
event: StoredEvent, event: StoredEvent,
p: InferenceCompletedEvent, p: InferenceCompletedEvent,
artifactStore: ArtifactStore, artifactStore: ArtifactStore,
sessionSequence: Long,
): ServerMessage = runCatching { ): ServerMessage = runCatching {
artifactStore.get(p.responseArtifactId)?.toString(Charsets.UTF_8) ?: "" artifactStore.get(p.responseArtifactId)?.toString(Charsets.UTF_8) ?: ""
}.getOrElse { "" }.run { }.getOrElse { "" }.run {
@@ -107,10 +170,16 @@ private suspend fun mapInferenceCompleted(
outputSummary = this, outputSummary = this,
responseText = this, responseText = this,
occurredAt = event.metadata.timestamp.toEpochMilliseconds(), 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( ServerMessage.ApprovalRequired(
sessionId = p.sessionId, sessionId = p.sessionId,
requestId = p.requestId, requestId = p.requestId,
@@ -118,4 +187,6 @@ private fun mapApprovalRequested(p: ApprovalRequestedEvent): ServerMessage =
riskSummary = RiskSummaryDto("unknown", emptyList(), ""), riskSummary = RiskSummaryDto("unknown", emptyList(), ""),
toolName = p.toolName, toolName = p.toolName,
preview = p.preview, preview = p.preview,
sequence = seq,
sessionSequence = sessionSequence,
) )
@@ -1,6 +1,7 @@
package com.correx.apps.server.bridge package com.correx.apps.server.bridge
import com.correx.apps.server.protocol.ServerMessage import com.correx.apps.server.protocol.ServerMessage
import com.correx.apps.server.protocol.SessionStateDto
import com.correx.core.approvals.DefaultApprovalRepository import com.correx.core.approvals.DefaultApprovalRepository
import com.correx.core.artifactstore.ArtifactStore import com.correx.core.artifactstore.ArtifactStore
import com.correx.core.events.orchestration.OrchestrationStatus import com.correx.core.events.orchestration.OrchestrationStatus
@@ -29,6 +30,7 @@ class SessionEventBridge(
approvalState.decisions.values.none { it.requestId == req.id } approvalState.decisions.values.none { it.requestId == req.id }
} }
val lastSeq = eventStore.lastSequence(sessionId) ?: 0L
send(ServerMessage.SessionSnapshot( send(ServerMessage.SessionSnapshot(
sessionId = sessionId, sessionId = sessionId,
workflowId = orchState.workflowId, workflowId = orchState.workflowId,
@@ -40,13 +42,25 @@ class SessionEventBridge(
approvalTier = pendingRequest?.tier?.name, approvalTier = pendingRequest?.tier?.name,
approvalToolName = pendingRequest?.toolName, approvalToolName = pendingRequest?.toolName,
approvalPreview = pendingRequest?.preview, 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) { suspend fun streamLive(sessionId: SessionId) {
var sessionSequence = 0L
eventStore.subscribe(sessionId).collect { event -> 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?, 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 @Serializable
enum class PauseReason { enum class PauseReason {
APPROVAL_PENDING, APPROVAL_PENDING,
@@ -8,18 +8,69 @@ import kotlinx.serialization.SerialName
import kotlinx.serialization.Serializable import kotlinx.serialization.Serializable
@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 @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 @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 @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 @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 @Serializable
@SerialName("session_snapshot") @SerialName("session_snapshot")
@@ -34,55 +85,128 @@ sealed class ServerMessage {
val approvalTier: String?, val approvalTier: String?,
val approvalToolName: String?, val approvalToolName: String?,
val approvalPreview: 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 @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 @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 @Serializable
data class StageFailed(val sessionId: SessionId, val stageId: StageId, val reason: String, val occurredAt: Long) : @SerialName("stage.failed")
ServerMessage() 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 @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 @Serializable
@SerialName("inference.completed")
data class InferenceCompleted( data class InferenceCompleted(
val sessionId: SessionId, val sessionId: SessionId,
val stageId: StageId, val stageId: StageId,
val outputSummary: String, val outputSummary: String,
val responseText: String = "", val responseText: String = "",
val occurredAt: Long, val occurredAt: Long,
) : ServerMessage() override val sequence: Long,
override val sessionSequence: Long,
) : ServerMessage, SessionMessage
@Serializable @Serializable
data class InferenceTimedOut(val sessionId: SessionId, val stageId: StageId, val elapsedMs: Long) : @SerialName("inference.timed_out")
ServerMessage() 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 @Serializable
data class ToolStarted(val sessionId: SessionId, val toolName: String, val tier: Tier) : @SerialName("tool.started")
ServerMessage() data class ToolStarted(
val sessionId: SessionId,
val toolName: String,
val tier: Tier,
override val sequence: Long,
override val sessionSequence: Long,
) : ServerMessage, SessionMessage
@Serializable @Serializable
@SerialName("tool.completed")
data class ToolCompleted( data class ToolCompleted(
val sessionId: SessionId, val sessionId: SessionId,
val toolName: String, val toolName: String,
val outputSummary: String, val outputSummary: String,
val occurredAt: Long, val occurredAt: Long,
) : ServerMessage() override val sequence: Long,
override val sessionSequence: Long,
) : ServerMessage, SessionMessage
@Serializable @Serializable
data class ToolFailed(val sessionId: SessionId, val toolName: String, val reason: String, val occurredAt: Long) : @SerialName("tool.failed")
ServerMessage() 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 @Serializable
data class ToolRejected(val sessionId: SessionId, val toolName: String, val reason: String) : @SerialName("tool.rejected")
ServerMessage() data class ToolRejected(
val sessionId: SessionId,
val toolName: String,
val reason: String,
override val sequence: Long,
override val sessionSequence: Long,
) : ServerMessage, SessionMessage
// -- Approval --
@Serializable @Serializable
@SerialName("approval.required")
data class ApprovalRequired( data class ApprovalRequired(
val sessionId: SessionId, val sessionId: SessionId,
val requestId: ApprovalRequestId, val requestId: ApprovalRequestId,
@@ -90,19 +214,48 @@ sealed class ServerMessage {
val riskSummary: RiskSummaryDto, val riskSummary: RiskSummaryDto,
val toolName: String?, val toolName: String?,
val preview: String?, val preview: String?,
) : ServerMessage() override val sequence: Long,
override val sessionSequence: Long,
) : ServerMessage, SessionMessage
// -- System / infra --
@Serializable @Serializable
data class ProviderStatusChanged(val providerId: String, val status: ProviderHealthDto) : @SerialName("provider.status_changed")
ServerMessage() data class ProviderStatusChanged(
val providerId: String,
val status: ProviderHealthDto,
override val sequence: Long? = null,
override val sessionSequence: Long? = null,
) : ServerMessage, NonEventMessage
@Serializable @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 @Serializable
@SerialName("router.response")
data class RouterResponseMessage( data class RouterResponseMessage(
val sessionId: SessionId, val sessionId: SessionId,
val content: String, val content: String,
val steeringEmitted: Boolean, 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 { .onFailure {
log.warn("decode error: {}", it.message) 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))) session.send(Frame.Text(ProtocolSerializer.encodeServerMessage(error)))
} }
} }
@@ -88,6 +92,8 @@ class GlobalStreamHandler(private val module: ServerModule) {
val msg = ServerMessage.ProviderStatusChanged( val msg = ServerMessage.ProviderStatusChanged(
providerId = providerId.value, providerId = providerId.value,
status = ProviderHealthDto(providerId.value, status, null), status = ProviderHealthDto(providerId.value, status, null),
sequence = null,
sessionSequence = null,
) )
session.send(Frame.Text(ProtocolSerializer.encodeServerMessage(msg))) session.send(Frame.Text(ProtocolSerializer.encodeServerMessage(msg)))
} }
@@ -112,7 +118,15 @@ class GlobalStreamHandler(private val module: ServerModule) {
} }
private fun encodeError(message: String): Frame.Text = 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( private suspend fun handleStartSession(
session: DefaultWebSocketServerSession, session: DefaultWebSocketServerSession,
@@ -123,7 +137,11 @@ class GlobalStreamHandler(private val module: ServerModule) {
val graph = module.workflowRegistry.find(msg.workflowId) val graph = module.workflowRegistry.find(msg.workflowId)
log.info("find returned: {}", graph) log.info("find returned: {}", graph)
if (graph == null) { 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))) session.send(Frame.Text(ProtocolSerializer.encodeServerMessage(error)))
return return
} }
@@ -166,7 +184,12 @@ class GlobalStreamHandler(private val module: ServerModule) {
bridge.streamLive(sessionId) bridge.streamLive(sessionId)
}, },
)?.cancel() )?.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))) session.send(Frame.Text(ProtocolSerializer.encodeServerMessage(started)))
} }
} }
@@ -65,6 +65,7 @@ class DomainEventMapperTest {
correlationId = null, correlationId = null,
), ),
sequence = 1L, sequence = 1L,
sessionSequence = 1L,
payload = payload, payload = payload,
) )
@@ -72,8 +73,11 @@ class DomainEventMapperTest {
fun `WorkflowStartedEvent maps to SessionStarted`(): Unit = runTest { fun `WorkflowStartedEvent maps to SessionStarted`(): Unit = runTest {
val event = val event =
storedEvent(WorkflowStartedEvent(sessionId = sessionId, startStageId = stageId, workflowId = workflowId)) storedEvent(WorkflowStartedEvent(sessionId = sessionId, startStageId = stageId, workflowId = workflowId))
val result = domainEventToServerMessage(event, noopStore) val result = domainEventToServerMessage(event, noopStore, sessionSequence = 0L)
assertEquals(ServerMessage.SessionStarted(sessionId, workflowId), result) assertEquals(
ServerMessage.SessionStarted(sessionId, workflowId, event.sequence, 0L),
result,
)
} }
@Test @Test
@@ -85,8 +89,11 @@ class DomainEventMapperTest {
totalStages = 1, totalStages = 1,
), ),
) )
val result = domainEventToServerMessage(event, noopStore) val result = domainEventToServerMessage(event, noopStore, sessionSequence = 0L)
assertEquals(ServerMessage.SessionCompleted(sessionId), result) assertEquals(
ServerMessage.SessionCompleted(sessionId, event.sequence, 0L),
result,
)
} }
@Test @Test
@@ -94,8 +101,11 @@ class DomainEventMapperTest {
val event = storedEvent( val event = storedEvent(
WorkflowFailedEvent(sessionId = sessionId, stageId = stageId, reason = "boom", retryExhausted = false), WorkflowFailedEvent(sessionId = sessionId, stageId = stageId, reason = "boom", retryExhausted = false),
) )
val result = domainEventToServerMessage(event, noopStore) val result = domainEventToServerMessage(event, noopStore, sessionSequence = 0L)
assertEquals(ServerMessage.SessionFailed(sessionId, "boom"), result) assertEquals(
ServerMessage.SessionFailed(sessionId, "boom", event.sequence, 0L),
result,
)
} }
@Test @Test
@@ -105,8 +115,11 @@ class DomainEventMapperTest {
val event = storedEvent( val event = storedEvent(
TransitionExecutedEvent(sessionId = sessionId, from = from, to = to, transitionId = TransitionId("t1")), TransitionExecutedEvent(sessionId = sessionId, from = from, to = to, transitionId = TransitionId("t1")),
) )
val result = domainEventToServerMessage(event, noopStore) val result = domainEventToServerMessage(event, noopStore, sessionSequence = 0L)
assertEquals(ServerMessage.StageStarted(sessionId, to, occurredAt), result) assertEquals(
ServerMessage.StageStarted(sessionId, to, occurredAt, event.sequence, 0L),
result,
)
} }
@Test @Test
@@ -114,19 +127,25 @@ class DomainEventMapperTest {
val event = storedEvent( val event = storedEvent(
StageCompletedEvent(sessionId = sessionId, stageId = stageId, transitionId = TransitionId("t1")), StageCompletedEvent(sessionId = sessionId, stageId = stageId, transitionId = TransitionId("t1")),
) )
val result = domainEventToServerMessage(event, noopStore) val result = domainEventToServerMessage(event, noopStore, sessionSequence = 0L)
assertEquals(ServerMessage.StageCompleted(sessionId, stageId, occurredAt), result) assertEquals(
ServerMessage.StageCompleted(sessionId, stageId, occurredAt, event.sequence, 0L),
result,
)
} }
@Test @Test
fun `StageFailedEvent maps to StageFailed`(): Unit = runTest { fun `StageFailedEvent maps to StageFailed`(): Unit = runTest {
val event = storedEvent( val event = storedEvent(
StageFailedEvent( StageFailedEvent(
sessionId = sessionId, stageId = stageId, transitionId = TransitionId("t1"), reason = "err", sessionId = sessionId, stageId = stageId, transitionId = TransitionId("t1"), reason = "err",
), ),
) )
val result = domainEventToServerMessage(event, noopStore) val result = domainEventToServerMessage(event, noopStore, sessionSequence = 0L)
assertEquals(ServerMessage.StageFailed(sessionId, stageId, "err", occurredAt), result) assertEquals(
ServerMessage.StageFailed(sessionId, stageId, "err", occurredAt, event.sequence, 0L),
result,
)
} }
@Test @Test
@@ -134,8 +153,11 @@ class DomainEventMapperTest {
val event = storedEvent( val event = storedEvent(
OrchestrationPausedEvent(sessionId = sessionId, stageId = stageId, reason = "APPROVAL_PENDING"), OrchestrationPausedEvent(sessionId = sessionId, stageId = stageId, reason = "APPROVAL_PENDING"),
) )
val result = domainEventToServerMessage(event, noopStore) val result = domainEventToServerMessage(event, noopStore, sessionSequence = 0L)
assertEquals(ServerMessage.SessionPaused(sessionId, PauseReason.APPROVAL_PENDING), result) assertEquals(
ServerMessage.SessionPaused(sessionId, PauseReason.APPROVAL_PENDING, event.sequence, 0L),
result,
)
} }
@Test @Test
@@ -143,8 +165,11 @@ class DomainEventMapperTest {
val event = storedEvent( val event = storedEvent(
OrchestrationPausedEvent(sessionId = sessionId, stageId = stageId, reason = "USER_REQUESTED"), OrchestrationPausedEvent(sessionId = sessionId, stageId = stageId, reason = "USER_REQUESTED"),
) )
val result = domainEventToServerMessage(event, noopStore) val result = domainEventToServerMessage(event, noopStore, sessionSequence = 0L)
assertEquals(ServerMessage.SessionPaused(sessionId, PauseReason.USER_REQUESTED), result) assertEquals(
ServerMessage.SessionPaused(sessionId, PauseReason.USER_REQUESTED, event.sequence, 0L),
result,
)
} }
@Test @Test
@@ -158,8 +183,11 @@ class DomainEventMapperTest {
promptArtifactId = ArtifactId("art-prompt"), promptArtifactId = ArtifactId("art-prompt"),
), ),
) )
val result = domainEventToServerMessage(event, noopStore) val result = domainEventToServerMessage(event, noopStore, sessionSequence = 0L)
assertEquals(ServerMessage.InferenceStarted(sessionId, stageId), result) assertEquals(
ServerMessage.InferenceStarted(sessionId, stageId, event.sequence, 0L),
result,
)
} }
@Test @Test
@@ -184,9 +212,11 @@ class DomainEventMapperTest {
responseArtifactId = artifactId, responseArtifactId = artifactId,
), ),
) )
val result = domainEventToServerMessage(event, store) val result = domainEventToServerMessage(event, store, sessionSequence = 0L)
assertEquals( assertEquals(
ServerMessage.InferenceCompleted(sessionId, stageId, responseText, responseText, occurredAt), ServerMessage.InferenceCompleted(
sessionId, stageId, responseText, responseText, occurredAt, event.sequence, 0L,
),
result, result,
) )
} }
@@ -204,8 +234,11 @@ class DomainEventMapperTest {
responseArtifactId = ArtifactId("missing"), responseArtifactId = ArtifactId("missing"),
), ),
) )
val result = domainEventToServerMessage(event, noopStore) val result = domainEventToServerMessage(event, noopStore, sessionSequence = 0L)
assertEquals(ServerMessage.InferenceCompleted(sessionId, stageId, "", "", occurredAt), result) assertEquals(
ServerMessage.InferenceCompleted(sessionId, stageId, "", "", occurredAt, event.sequence, 0L),
result,
)
} }
@Test @Test
@@ -219,8 +252,11 @@ class DomainEventMapperTest {
timeoutMs = 5000L, timeoutMs = 5000L,
), ),
) )
val result = domainEventToServerMessage(event, noopStore) val result = domainEventToServerMessage(event, noopStore, sessionSequence = 0L)
assertEquals(ServerMessage.InferenceTimedOut(sessionId, stageId, 5000L), result) assertEquals(
ServerMessage.InferenceTimedOut(sessionId, stageId, 5000L, event.sequence, 0L),
result,
)
} }
@Test @Test
@@ -238,8 +274,11 @@ class DomainEventMapperTest {
), ),
), ),
) )
val result = domainEventToServerMessage(event, noopStore) val result = domainEventToServerMessage(event, noopStore, sessionSequence = 0L)
assertEquals(ServerMessage.ToolStarted(sessionId, "file_write", Tier.T2), result) assertEquals(
ServerMessage.ToolStarted(sessionId, "file_write", Tier.T2, event.sequence, 0L),
result,
)
} }
@Test @Test
@@ -259,8 +298,11 @@ class DomainEventMapperTest {
invocationId = invId, sessionId = sessionId, toolName = "file_write", receipt = receipt, invocationId = invId, sessionId = sessionId, toolName = "file_write", receipt = receipt,
), ),
) )
val result = domainEventToServerMessage(event, noopStore) val result = domainEventToServerMessage(event, noopStore, sessionSequence = 0L)
assertEquals(ServerMessage.ToolCompleted(sessionId, "file_write", "wrote 3 lines", occurredAt), result) assertEquals(
ServerMessage.ToolCompleted(sessionId, "file_write", "wrote 3 lines", occurredAt, event.sequence, 0L),
result,
)
} }
@Test @Test
@@ -273,8 +315,11 @@ class DomainEventMapperTest {
reason = "disk full", reason = "disk full",
), ),
) )
val result = domainEventToServerMessage(event, noopStore) val result = domainEventToServerMessage(event, noopStore, sessionSequence = 0L)
assertEquals(ServerMessage.ToolFailed(sessionId, "file_write", "disk full", occurredAt), result) assertEquals(
ServerMessage.ToolFailed(sessionId, "file_write", "disk full", occurredAt, event.sequence, 0L),
result,
)
} }
@Test @Test
@@ -288,8 +333,11 @@ class DomainEventMapperTest {
reason = "policy denied", reason = "policy denied",
), ),
) )
val result = domainEventToServerMessage(event, noopStore) val result = domainEventToServerMessage(event, noopStore, sessionSequence = 0L)
assertEquals(ServerMessage.ToolRejected(sessionId, "file_write", "policy denied"), result) assertEquals(
ServerMessage.ToolRejected(sessionId, "file_write", "policy denied", event.sequence, 0L),
result,
)
} }
@Test @Test
@@ -306,17 +354,18 @@ class DomainEventMapperTest {
projectId = null, 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(requestId, result.requestId)
assertEquals(Tier.T3, result.tier) assertEquals(Tier.T3, result.tier)
assertNull(result.toolName) assertNull(result.toolName)
assertNull(result.preview) assertNull(result.preview)
assertEquals(event.sequence, result.sequence)
} }
@Test @Test
fun `unmapped event returns null`(): Unit = runTest { fun `unmapped event returns null`(): Unit = runTest {
val event = storedEvent(SessionStartedEvent(sessionId = sessionId)) val event = storedEvent(SessionStartedEvent(sessionId = sessionId))
val result = domainEventToServerMessage(event, noopStore) val result = domainEventToServerMessage(event, noopStore, sessionSequence = 0L)
assertNull(result) 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.StoredEvent
import com.correx.core.events.events.WorkflowCompletedEvent import com.correx.core.events.events.WorkflowCompletedEvent
import com.correx.core.events.events.WorkflowStartedEvent 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.stores.EventStore
import com.correx.core.events.types.ArtifactId import com.correx.core.events.types.ArtifactId
import com.correx.core.events.types.EventId import com.correx.core.events.types.EventId
import com.correx.core.events.types.SessionId import com.correx.core.events.types.SessionId
import com.correx.core.events.types.StageId 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.kernel.orchestration.OrchestrationRepository
import com.correx.core.sessions.projections.replay.DefaultEventReplayer import com.correx.core.sessions.projections.replay.DefaultEventReplayer
import com.correx.core.sessions.projections.replay.EventReplayer
import kotlinx.coroutines.channels.Channel import kotlinx.coroutines.channels.Channel
import kotlinx.coroutines.flow.Flow import kotlinx.coroutines.flow.Flow
import kotlinx.coroutines.flow.flow import kotlinx.coroutines.flow.flow
@@ -54,6 +55,7 @@ class SessionEventBridgeTest {
correlationId = null, correlationId = null,
), ),
sequence = seq, sequence = seq,
sessionSequence = seq,
payload = payload, payload = payload,
) )
@@ -65,17 +67,26 @@ class SessionEventBridgeTest {
override suspend fun appendAll(events: List<NewEvent>): List<StoredEvent> = error("unused") override suspend fun appendAll(events: List<NewEvent>): List<StoredEvent> = error("unused")
override fun read(sessionId: SessionId): List<StoredEvent> = allEventsList override fun read(sessionId: SessionId): List<StoredEvent> = allEventsList
override fun readFrom(sessionId: SessionId, fromSequence: Long): 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 subscribe(sessionId: SessionId): Flow<StoredEvent> = liveFlow
override fun allEvents(): Sequence<StoredEvent> = allEventsList.asSequence() override fun allEvents(): Sequence<StoredEvent> = allEventsList.asSequence()
override fun allSessionIds(): Set<SessionId> = allEventsList.map { it.metadata.sessionId }.toSet() 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( private fun activeOrchestrationRepository(
DefaultEventReplayer( status: OrchestrationStatus = OrchestrationStatus.RUNNING,
fakeEventStore(), ): OrchestrationRepository = OrchestrationRepository(
OrchestrationProjector(DefaultOrchestrationReducer()), object : EventReplayer<OrchestrationState> {
), override fun rebuild(sessionId: SessionId): OrchestrationState = OrchestrationState(
workflowId = workflowId,
status = status,
currentStageId = stageId,
pendingApproval = false,
)
},
) )
private val approvalRepository = DefaultApprovalRepository( private val approvalRepository = DefaultApprovalRepository(
@@ -92,19 +103,23 @@ class SessionEventBridgeTest {
storedEvent(WorkflowCompletedEvent(sessionId, stageId, 1), seq = 2L), storedEvent(WorkflowCompletedEvent(sessionId, stageId, 1), seq = 2L),
) )
val store = fakeEventStore(allEventsList = events) val store = fakeEventStore(allEventsList = events)
val orchRepo = activeOrchestrationRepository()
val sent = mutableListOf<ServerMessage>() val sent = mutableListOf<ServerMessage>()
val bridge = SessionEventBridge( val bridge = SessionEventBridge(
store, store,
noopArtifactStore, noopArtifactStore,
orchestrationRepository, orchRepo,
approvalRepository, approvalRepository,
) { sent.add(it) } ) { sent.add(it) }
bridge.replaySnapshot() bridge.replaySnapshot()
assertEquals(2, sent.size) assertEquals(2, sent.size)
assertEquals(ServerMessage.SessionStarted(sessionId, workflowId), sent[0]) val snapshot = sent[0] as ServerMessage.SessionSnapshot
assertEquals(ServerMessage.SessionCompleted(sessionId), sent[1]) assertEquals(sessionId, snapshot.sessionId)
assertEquals(2L, snapshot.lastSequence)
assertEquals(0L, snapshot.lastSessionSequence)
assertEquals(ServerMessage.SnapshotComplete, sent[1])
} }
@Test @Test
@@ -114,18 +129,21 @@ class SessionEventBridgeTest {
storedEvent(WorkflowCompletedEvent(sessionId, stageId, 1), seq = 2L), storedEvent(WorkflowCompletedEvent(sessionId, stageId, 1), seq = 2L),
) )
val store = fakeEventStore(allEventsList = events) val store = fakeEventStore(allEventsList = events)
val orchRepo = activeOrchestrationRepository()
val sent = mutableListOf<ServerMessage>() val sent = mutableListOf<ServerMessage>()
val bridge = SessionEventBridge( val bridge = SessionEventBridge(
store, store,
noopArtifactStore, noopArtifactStore,
orchestrationRepository, orchRepo,
approvalRepository, approvalRepository,
) { sent.add(it) } ) { sent.add(it) }
bridge.replaySnapshot() bridge.replaySnapshot()
assertEquals(1, sent.size) assertEquals(2, sent.size)
assertEquals(ServerMessage.SessionCompleted(sessionId), sent[0]) val snapshot = sent[0] as ServerMessage.SessionSnapshot
assertEquals(sessionId, snapshot.sessionId)
assertEquals(ServerMessage.SnapshotComplete, sent[1])
} }
@Test @Test
@@ -140,15 +158,21 @@ class SessionEventBridgeTest {
val bridge = SessionEventBridge( val bridge = SessionEventBridge(
store, store,
noopArtifactStore, noopArtifactStore,
orchestrationRepository, activeOrchestrationRepository(),
approvalRepository, approvalRepository,
) { sent.add(it) } ) { sent.add(it) }
bridge.streamLive(sessionId) bridge.streamLive(sessionId)
assertEquals(2, sent.size) assertEquals(2, sent.size)
assertEquals(ServerMessage.SessionStarted(sessionId, workflowId), sent[0]) assertEquals(
assertEquals(ServerMessage.SessionCompleted(sessionId), sent[1]) ServerMessage.SessionStarted(sessionId, workflowId, 1L, 1L),
sent[0],
)
assertEquals(
ServerMessage.SessionCompleted(sessionId, 2L, 2L),
sent[1],
)
} }
@Test @Test
@@ -162,7 +186,7 @@ class SessionEventBridgeTest {
val bridge = SessionEventBridge( val bridge = SessionEventBridge(
store, store,
noopArtifactStore, noopArtifactStore,
orchestrationRepository, activeOrchestrationRepository(),
approvalRepository, approvalRepository,
) { sent.add(it) } ) { 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"), riskSummary = RiskSummaryDto(level = "HIGH", factors = emptyList(), recommendedAction = "review"),
toolName = "bash", toolName = "bash",
preview = "rm -rf /", preview = "rm -rf /",
sequence = 1L,
sessionSequence = 1L,
) )
val (state, effects) = reduce(action = Action.ServerEventReceived(msg)) val (state, effects) = reduce(action = Action.ServerEventReceived(msg))
assertEquals("req-2", state.active?.requestId) assertEquals("req-2", state.active?.requestId)
@@ -43,11 +43,15 @@ class RootReducerTest {
val startMsg = com.correx.apps.server.protocol.ServerMessage.SessionStarted( val startMsg = com.correx.apps.server.protocol.ServerMessage.SessionStarted(
sessionId = com.correx.core.events.types.SessionId("s1"), sessionId = com.correx.core.events.types.SessionId("s1"),
workflowId = "wf", workflowId = "wf",
sequence = 1L,
sessionSequence = 1L,
) )
val (state1, _) = RootReducer.reduce(state0, Action.ServerEventReceived(startMsg), fixedClock) val (state1, _) = RootReducer.reduce(state0, Action.ServerEventReceived(startMsg), fixedClock)
val startMsg2 = com.correx.apps.server.protocol.ServerMessage.SessionStarted( val startMsg2 = com.correx.apps.server.protocol.ServerMessage.SessionStarted(
sessionId = com.correx.core.events.types.SessionId("s2"), sessionId = com.correx.core.events.types.SessionId("s2"),
workflowId = "wf", workflowId = "wf",
sequence = 2L,
sessionSequence = 1L,
) )
val (state2, _) = RootReducer.reduce(state1, Action.ServerEventReceived(startMsg2), fixedClock) val (state2, _) = RootReducer.reduce(state1, Action.ServerEventReceived(startMsg2), fixedClock)
// selected is s1 (first); navigating down should move to s2 // selected is s1 (first); navigating down should move to s2
@@ -91,7 +91,7 @@ class SessionsReducerTest {
@Test @Test
fun `ServerEventReceived SessionStarted appends session and sets selection`() { 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)) val (state, effects) = reduce(action = Action.ServerEventReceived(msg))
assertEquals(1, state.sessions.size) assertEquals(1, state.sessions.size)
assertEquals("s1", state.sessions[0].id) assertEquals("s1", state.sessions[0].id)
@@ -103,7 +103,7 @@ class SessionsReducerTest {
@Test @Test
fun `ServerEventReceived SessionStarted preserves existing selection`() { fun `ServerEventReceived SessionStarted preserves existing selection`() {
val existing = SessionsState(sessions = listOf(session("existing")), selectedId = "existing") 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)) val (state, _) = reduce(sessions = existing, action = Action.ServerEventReceived(msg))
assertEquals("existing", state.selectedId) assertEquals("existing", state.selectedId)
assertEquals(2, state.sessions.size) assertEquals(2, state.sessions.size)
@@ -112,7 +112,7 @@ class SessionsReducerTest {
@Test @Test
fun `ServerEventReceived SessionCompleted updates status to COMPLETED`() { fun `ServerEventReceived SessionCompleted updates status to COMPLETED`() {
val s = SessionsState(sessions = listOf(session("s1")), selectedId = "s1") 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)) val (state, _) = reduce(sessions = s, action = Action.ServerEventReceived(msg))
assertEquals("COMPLETED", state.sessions[0].status) assertEquals("COMPLETED", state.sessions[0].status)
assertEquals(1000L, state.sessions[0].lastEventAt) assertEquals(1000L, state.sessions[0].lastEventAt)
@@ -121,7 +121,7 @@ class SessionsReducerTest {
@Test @Test
fun `ServerEventReceived SessionFailed updates status to FAILED`() { fun `ServerEventReceived SessionFailed updates status to FAILED`() {
val s = SessionsState(sessions = listOf(session("s1")), selectedId = "s1") 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)) val (state, _) = reduce(sessions = s, action = Action.ServerEventReceived(msg))
assertEquals("FAILED", state.sessions[0].status) assertEquals("FAILED", state.sessions[0].status)
} }
@@ -136,6 +136,8 @@ class SessionsReducerTest {
sessionId = SessionId("s1"), sessionId = SessionId("s1"),
stageId = StageId("stage-1"), stageId = StageId("stage-1"),
occurredAt = fixedClock(), occurredAt = fixedClock(),
sequence = 1L,
sessionSequence = 1L,
) )
val (state, _) = reduce(sessions = s, action = Action.ServerEventReceived(msg)) val (state, _) = reduce(sessions = s, action = Action.ServerEventReceived(msg))
assertEquals(null, state.sessions[0].currentStage) assertEquals(null, state.sessions[0].currentStage)
@@ -153,6 +155,8 @@ class SessionsReducerTest {
stageId = StageId("stage-1"), stageId = StageId("stage-1"),
reason = "error", reason = "error",
occurredAt = fixedClock(), occurredAt = fixedClock(),
sequence = 1L,
sessionSequence = 1L,
) )
val (state, _) = reduce(sessions = s, action = Action.ServerEventReceived(msg)) val (state, _) = reduce(sessions = s, action = Action.ServerEventReceived(msg))
assertEquals(null, state.sessions[0].currentStage) assertEquals(null, state.sessions[0].currentStage)
@@ -183,6 +187,8 @@ class SessionsReducerTest {
sessionId = SessionId("s1"), stageId = StageId("stage-1"), sessionId = SessionId("s1"), stageId = StageId("stage-1"),
outputSummary = "summary", responseText = "full response", outputSummary = "summary", responseText = "full response",
occurredAt = fixedClock(), occurredAt = fixedClock(),
sequence = 1L,
sessionSequence = 1L,
) )
val (state, _) = reduce(sessions = s, action = Action.ServerEventReceived(msg)) val (state, _) = reduce(sessions = s, action = Action.ServerEventReceived(msg))
assertEquals("summary", state.sessions[0].lastOutput) assertEquals("summary", state.sessions[0].lastOutput)
@@ -197,6 +203,8 @@ class SessionsReducerTest {
sessionId = SessionId("s1"), stageId = StageId("stage-1"), sessionId = SessionId("s1"), stageId = StageId("stage-1"),
outputSummary = "", responseText = "", outputSummary = "", responseText = "",
occurredAt = fixedClock(), occurredAt = fixedClock(),
sequence = 1L,
sessionSequence = 1L,
) )
val (state, _) = reduce(sessions = s, action = Action.ServerEventReceived(msg)) val (state, _) = reduce(sessions = s, action = Action.ServerEventReceived(msg))
assertEquals(null, state.sessions[0].lastResponseText) assertEquals(null, state.sessions[0].lastResponseText)
@@ -209,6 +217,8 @@ class SessionsReducerTest {
sessionId = SessionId("s1"), sessionId = SessionId("s1"),
stageId = StageId("stage-1"), stageId = StageId("stage-1"),
occurredAt = fixedClock(), occurredAt = fixedClock(),
sequence = 1L,
sessionSequence = 1L,
) )
val (state, _) = reduce(sessions = s, action = Action.ServerEventReceived(msg)) val (state, _) = reduce(sessions = s, action = Action.ServerEventReceived(msg))
assertEquals("stage-1", state.sessions[0].currentStage) assertEquals("stage-1", state.sessions[0].currentStage)
@@ -223,6 +233,8 @@ class SessionsReducerTest {
toolName = "file_write", toolName = "file_write",
outputSummary = "wrote 3 lines", outputSummary = "wrote 3 lines",
occurredAt = fixedClock(), occurredAt = fixedClock(),
sequence = 1L,
sessionSequence = 1L,
) )
val (state, _) = reduce(sessions = s, action = Action.ServerEventReceived(msg)) val (state, _) = reduce(sessions = s, action = Action.ServerEventReceived(msg))
assertEquals("file_write: wrote 3 lines", state.sessions[0].lastOutput) assertEquals("file_write: wrote 3 lines", state.sessions[0].lastOutput)
@@ -232,7 +244,7 @@ class SessionsReducerTest {
@Test @Test
fun `ServerEventReceived SessionPaused with APPROVAL_PENDING sets status label`() { fun `ServerEventReceived SessionPaused with APPROVAL_PENDING sets status label`() {
val s = SessionsState(sessions = listOf(session("s1")), selectedId = "s1") 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)) val (state, _) = reduce(sessions = s, action = Action.ServerEventReceived(msg))
assertEquals("PAUSED awaiting approval", state.sessions[0].status) assertEquals("PAUSED awaiting approval", state.sessions[0].status)
} }
@@ -38,6 +38,7 @@ class DefaultApprovalReducerTest {
correlationId = null correlationId = null
), ),
sequence = sequence, sequence = sequence,
sessionSequence = sequence,
payload = payload payload = payload
) )
} }
@@ -59,6 +59,7 @@ class DefaultToolReducerTest {
correlationId = null correlationId = null
), ),
sequence = 1L, sequence = 1L,
sessionSequence = 1L,
payload = payload payload = payload
) )