From 8e6a3e1470041c3a449565aff31accc1f3e9e4bf Mon Sep 17 00:00:00 2001 From: kami Date: Sat, 30 May 2026 18:40:08 +0400 Subject: [PATCH] feat: inject recalled L3 memory into router context with budget Record L3 retrieval as an event carrying hit text (invariant #9), then rebuild router state and inject recalled memories as a SYSTEM L3 layer in the context pack. Apply token budget: protected frames (L0 immutable + current user turn) are never dropped; honest budgetUsed is reported and 'BudgetExceeded' is flagged in appliedStrategies when they overflow. L1/L2 fit newest-first so oldest entries drop first. Emit ContextTruncatedEvent when entries are dropped. L3MemoryRetrievedEvent is emitted on every CHAT turn (empty hits reset recalled memory). --- .../correx/core/events/events/RouterEvents.kt | 10 + .../events/serialization/Serialization.kt | 2 + .../core/router/RouterContextBuilder.kt | 137 +++++- .../com/correx/core/router/RouterFacade.kt | 81 +++- .../test/kotlin/RouterContextBuilderTest.kt | 427 ++++++++++++++++-- .../src/test/kotlin/RouterFacadeTest.kt | 272 ++++++++++- 6 files changed, 852 insertions(+), 77 deletions(-) diff --git a/core/events/src/main/kotlin/com/correx/core/events/events/RouterEvents.kt b/core/events/src/main/kotlin/com/correx/core/events/events/RouterEvents.kt index 671b9eaa..8f51a5b0 100644 --- a/core/events/src/main/kotlin/com/correx/core/events/events/RouterEvents.kt +++ b/core/events/src/main/kotlin/com/correx/core/events/events/RouterEvents.kt @@ -37,3 +37,13 @@ data class L3RetrievedHit( val text: String, val score: Float, ) + +@Serializable +@SerialName("ContextTruncated") +data class ContextTruncatedEvent( + val sessionId: SessionId, + val turnId: String, + val entriesDropped: Int, + val truncatedLayers: List, + val timestampMs: Long, +) : EventPayload diff --git a/core/events/src/main/kotlin/com/correx/core/events/serialization/Serialization.kt b/core/events/src/main/kotlin/com/correx/core/events/serialization/Serialization.kt index 2e0fb8c9..62d82958 100644 --- a/core/events/src/main/kotlin/com/correx/core/events/serialization/Serialization.kt +++ b/core/events/src/main/kotlin/com/correx/core/events/serialization/Serialization.kt @@ -9,6 +9,7 @@ import com.correx.core.events.events.ArtifactValidatedEvent import com.correx.core.events.events.ArtifactValidatingEvent import com.correx.core.events.events.ChatSessionStartedEvent import com.correx.core.events.events.ChatTurnEvent +import com.correx.core.events.events.ContextTruncatedEvent import com.correx.core.events.events.EventPayload import com.correx.core.events.events.L3MemoryRetrievedEvent import com.correx.core.events.events.InferenceCompletedEvent @@ -75,6 +76,7 @@ val eventModule = SerializersModule { subclass(ChatSessionStartedEvent::class) subclass(ChatTurnEvent::class) subclass(L3MemoryRetrievedEvent::class) + subclass(ContextTruncatedEvent::class) } } diff --git a/core/router/src/main/kotlin/com/correx/core/router/RouterContextBuilder.kt b/core/router/src/main/kotlin/com/correx/core/router/RouterContextBuilder.kt index 2c70f455..4c483720 100644 --- a/core/router/src/main/kotlin/com/correx/core/router/RouterContextBuilder.kt +++ b/core/router/src/main/kotlin/com/correx/core/router/RouterContextBuilder.kt @@ -17,6 +17,7 @@ import com.correx.core.router.model.RouterState import com.correx.core.router.model.RouterTurn import com.correx.core.router.model.TurnRole import java.util.* +import kotlinx.coroutines.CancellationException interface RouterContextBuilder { suspend fun build(state: RouterState, budget: TokenBudget): ContextPack @@ -30,56 +31,128 @@ class DefaultRouterContextBuilder( companion object { private const val SYSTEM_PROMPT = "You are a routing assistant. Provide guidance based on workflow state and conversation context." + private const val RECALLED_MEMORY_PREFIX = "[recalled memory]" } override suspend fun build(state: RouterState, budget: TokenBudget): ContextPack { - var remainingBudget = budget.limit - val allEntries = mutableListOf() - var droppedCount = 0 - + // Protected frame: system prompt and workflow status are ALWAYS included regardless of budget. val systemPrompt = buildContextEntry( sourceType = "systemPrompt", sourceId = "router-system", content = SYSTEM_PROMPT, + layer = ContextLayer.L0, role = EntryRole.SYSTEM, ) - remainingBudget -= systemPrompt.tokenEstimate - if (remainingBudget < 0) remainingBudget = 0 - allEntries += systemPrompt - val workflowStatusEntry = buildContextEntry( sourceType = "workflowStatus", sourceId = state.currentStageId?.value ?: "none", content = buildWorkflowStatusContent(state), + layer = ContextLayer.L0, role = EntryRole.SYSTEM, ) - remainingBudget -= workflowStatusEntry.tokenEstimate - if (remainingBudget < 0) remainingBudget = 0 + + // Compute remaining budget after protected frame. + val protectedTokens = systemPrompt.tokenEstimate + workflowStatusEntry.tokenEstimate + var remainingBudget = (budget.limit - protectedTokens).coerceAtLeast(0) + + val allEntries = mutableListOf() + allEntries += systemPrompt allEntries += workflowStatusEntry - val recentTurns = state.conversationHistory.takeLast(config.conversationKeepLast) - for (turn in recentTurns) { + var droppedCount = 0 + val truncatedLayers = mutableSetOf() + + // --- Conversation turns (L1) --- + // The current (protected) user turn is the last element in conversationHistory, + // regardless of conversationKeepLast — it must always be present. + val protectedUserTurn: RouterTurn? = state.conversationHistory.lastOrNull() + + // Apply conversationKeepLast cap. The protected turn may or may not fall within this window. + val cappedTurns = state.conversationHistory.takeLast(config.conversationKeepLast) + + // Build entries for the capped window, excluding the protected turn (to avoid double-counting). + val nonProtectedCapped = if (protectedUserTurn != null && cappedTurns.lastOrNull() == protectedUserTurn) { + cappedTurns.dropLast(1) + } else { + cappedTurns + } + + // Build the protected user turn entry. + val protectedTurnEntry: ContextEntry? = protectedUserTurn?.let { turn -> val role = when (turn.role) { TurnRole.USER -> EntryRole.USER TurnRole.ROUTER -> EntryRole.ASSISTANT } - val entry = buildContextEntry( + buildContextEntry( sourceType = "conversation", sourceId = "${turn.role.name}-${turn.hashCode()}", content = turn.content, layer = ContextLayer.L1, role = role, ) + } + + // Reserve budget for the protected user turn first. + if (protectedTurnEntry != null) { + remainingBudget = (remainingBudget - protectedTurnEntry.tokenEstimate).coerceAtLeast(0) + } + + // Fit remaining capped turns newest-to-oldest (excluding protected), so oldest get dropped when tight. + val nonProtectedTurnEntries = nonProtectedCapped.map { turn -> + val role = when (turn.role) { + TurnRole.USER -> EntryRole.USER + TurnRole.ROUTER -> EntryRole.ASSISTANT + } + buildContextEntry( + sourceType = "conversation", + sourceId = "${turn.role.name}-${turn.hashCode()}", + content = turn.content, + layer = ContextLayer.L1, + role = role, + ) + } + + val fittedTurns = mutableListOf() + for (entry in nonProtectedTurnEntries.asReversed()) { if (remainingBudget >= entry.tokenEstimate) { remainingBudget -= entry.tokenEstimate - if (remainingBudget < 0) remainingBudget = 0 - allEntries += entry + fittedTurns.add(entry) } else { droppedCount++ + truncatedLayers.add(ContextLayer.L1) } } - for (l2Entry in state.l2Memory) { + // Re-order fitted turns oldest-to-newest (reverse back) and append protected turn at end. + fittedTurns.reverse() + allEntries.addAll(fittedTurns) + protectedTurnEntry?.let { allEntries.add(it) } + + // --- L3 memory (recalled cross-session) --- + // Sort by score descending (highest relevance first), add what fits. + val sortedHits = state.lastRetrievedMemory.sortedByDescending { it.score } + for (hit in sortedHits) { + val content = "$RECALLED_MEMORY_PREFIX ${hit.text}" + val entry = buildContextEntry( + sourceType = "recalledMemory", + sourceId = hit.entryId, + content = content, + layer = ContextLayer.L3, + role = EntryRole.SYSTEM, + ) + if (remainingBudget >= entry.tokenEstimate) { + remainingBudget -= entry.tokenEstimate + allEntries += entry + } else { + droppedCount++ + truncatedLayers.add(ContextLayer.L3) + } + } + + // --- L2 memory (compressed session summaries) --- + // Iterate newest-to-oldest so that when budget is tight the OLDEST entries are dropped. + val fittedL2 = mutableListOf() + for (l2Entry in state.l2Memory.asReversed()) { val content = buildL2Content(l2Entry) val entry = buildContextEntry( sourceType = "stageSummary", @@ -90,16 +163,31 @@ class DefaultRouterContextBuilder( ) if (remainingBudget >= entry.tokenEstimate) { remainingBudget -= entry.tokenEstimate - if (remainingBudget < 0) remainingBudget = 0 - allEntries += entry + fittedL2.add(entry) } else { droppedCount++ + truncatedLayers.add(ContextLayer.L2) } } + // Restore chronological (oldest-to-newest) order. + fittedL2.reverse() + allEntries.addAll(fittedL2) + + // Build layers map sorted by ContextLayer ordinal for deterministic ordering. + val layers = allEntries + .groupBy { it.layer } + .toSortedMap(compareBy { it.ordinal }) - val layers = allEntries.groupBy { it.layer } val budgetUsed = allEntries.sumOf { it.tokenEstimate } + // Build appliedStrategies list — deterministic, stable order. + val appliedStrategies = buildList { + add("L0Immutable") + add("Conversation") + if (state.lastRetrievedMemory.isNotEmpty()) add("L3Recall") + if (budgetUsed > budget.limit) add("BudgetExceeded") + } + return ContextPack( id = ContextPackId("${state.sessionId?.value ?: "unknown"}-router-pack"), sessionId = state.sessionId ?: SessionId("unknown"), @@ -108,8 +196,8 @@ class DefaultRouterContextBuilder( budgetUsed = budgetUsed, budgetLimit = budget.limit, compressionMetadata = CompressionMetadata( - appliedStrategies = listOf("L0Immutable", "Conversation"), - truncatedLayers = emptyList(), + appliedStrategies = appliedStrategies, + truncatedLayers = truncatedLayers.toList(), entriesDropped = droppedCount, ), ) @@ -150,7 +238,10 @@ class DefaultRouterContextBuilder( private suspend fun estimateTokens(content: String): Int { val t = tokenizer if (t != null) { - return runCatching { t.countTokens(content) }.getOrElse { fallbackEstimate(content) } + return runCatching { t.countTokens(content) }.getOrElse { e -> + if (e is CancellationException) throw e + fallbackEstimate(content) + } } return fallbackEstimate(content) } @@ -158,4 +249,4 @@ class DefaultRouterContextBuilder( private fun fallbackEstimate(content: String): Int { return (content.length / 4).coerceAtLeast(1) } -} \ No newline at end of file +} diff --git a/core/router/src/main/kotlin/com/correx/core/router/RouterFacade.kt b/core/router/src/main/kotlin/com/correx/core/router/RouterFacade.kt index c6083c6e..d7f2f1a0 100644 --- a/core/router/src/main/kotlin/com/correx/core/router/RouterFacade.kt +++ b/core/router/src/main/kotlin/com/correx/core/router/RouterFacade.kt @@ -2,6 +2,7 @@ package com.correx.core.router import com.correx.core.events.events.ChatTurnEvent import com.correx.core.events.events.ChatTurnRole +import com.correx.core.events.events.ContextTruncatedEvent import com.correx.core.events.events.EventMetadata import com.correx.core.events.events.L3MemoryRetrievedEvent import com.correx.core.events.events.L3RetrievedHit @@ -57,10 +58,12 @@ class DefaultRouterFacade( val userTurnId = emitChatTurn(sessionId, input, ChatTurnRole.USER) // Rebuild state with user turn appended - val stateWithUserTurn = routerRepository.getRouterState(sessionId) - val effectiveStageId = stateWithUserTurn.currentStageId ?: StageId.NONE + var state = routerRepository.getRouterState(sessionId) + val effectiveStageId = state.currentStageId ?: StageId.NONE - // L3 retrieval — non-fatal; does NOT alter what contextBuilder receives in this slice + // L3 retrieval — non-fatal; results fed back into state via event. + // Always emit L3MemoryRetrievedEvent for every CHAT turn where retrieval was attempted, + // including with empty hits — so lastRetrievedMemory is always current-turn-scoped. val retrieved = runCatching { val queryVector = embedder.embed(input) l3MemoryStore.query(L3Query(vector = queryVector, k = config.retrievalK)) @@ -71,40 +74,67 @@ class DefaultRouterFacade( } }.getOrElse { emptyList() } - val inSessionTurnIds = stateWithUserTurn.conversationHistory.map { it.turnId }.toSet() + val inSessionTurnIds = state.conversationHistory.map { it.turnId }.toSet() val deduped = retrieved.filter { it.entry.turnId !in inSessionTurnIds } - if (deduped.isNotEmpty()) { - val nowMs = Clock.System.now().toEpochMilliseconds() + // Always emit L3MemoryRetrievedEvent (even with empty hits) so that + // lastRetrievedMemory always reflects this turn and a prior turn's hits + // are never injected into the next turn's context. + val retrievalNow = Clock.System.now() + eventStore.append( + NewEvent( + metadata = EventMetadata( + eventId = EventId(UUID.randomUUID().toString()), + sessionId = sessionId, + timestamp = retrievalNow, + schemaVersion = 1, + causationId = null, + correlationId = null, + ), + payload = L3MemoryRetrievedEvent( + sessionId = sessionId, + queryTurnId = userTurnId, + hits = deduped.map { hit -> + L3RetrievedHit( + entryId = hit.entry.id, + sourceSessionId = hit.entry.sessionId, + sourceTurnId = hit.entry.turnId, + text = hit.entry.text, + score = hit.score, + ) + }, + timestampMs = retrievalNow.toEpochMilliseconds(), + ), + ), + ) + // Rebuild state so lastRetrievedMemory reflects this turn + state = routerRepository.getRouterState(sessionId) + + val contextPack = routerContextBuilder.build(state, config.tokenBudget) + + if (contextPack.compressionMetadata.entriesDropped > 0) { + val truncNow = Clock.System.now() eventStore.append( NewEvent( metadata = EventMetadata( eventId = EventId(UUID.randomUUID().toString()), sessionId = sessionId, - timestamp = Clock.System.now(), + timestamp = truncNow, schemaVersion = 1, causationId = null, correlationId = null, ), - payload = L3MemoryRetrievedEvent( + payload = ContextTruncatedEvent( sessionId = sessionId, - queryTurnId = userTurnId, - hits = deduped.map { hit -> - L3RetrievedHit( - entryId = hit.entry.id, - sourceSessionId = hit.entry.sessionId, - sourceTurnId = hit.entry.turnId, - text = hit.entry.text, - score = hit.score, - ) - }, - timestampMs = nowMs, + turnId = userTurnId, + entriesDropped = contextPack.compressionMetadata.entriesDropped, + truncatedLayers = contextPack.compressionMetadata.truncatedLayers.map { it.name }, + timestampMs = truncNow.toEpochMilliseconds(), ), ), ) } - val contextPack = routerContextBuilder.build(stateWithUserTurn, config.tokenBudget) val provider = inferenceRouter.route(effectiveStageId, setOf(ModelCapability.General)) val inferenceRequest = InferenceRequest( requestId = InferenceRequestId(UUID.randomUUID().toString()), @@ -132,13 +162,13 @@ class DefaultRouterFacade( private suspend fun emitChatTurn(sessionId: SessionId, content: String, role: ChatTurnRole): String { val turnId = UUID.randomUUID().toString() - val nowMs = Clock.System.now().toEpochMilliseconds() + val now = Clock.System.now() eventStore.append( NewEvent( metadata = EventMetadata( eventId = EventId(UUID.randomUUID().toString()), sessionId = sessionId, - timestamp = Clock.System.now(), + timestamp = now, schemaVersion = 1, causationId = null, correlationId = null, @@ -148,7 +178,7 @@ class DefaultRouterFacade( turnId = turnId, role = role, content = content, - timestampMs = nowMs, + timestampMs = now.toEpochMilliseconds(), ), ), ) @@ -162,7 +192,7 @@ class DefaultRouterFacade( turnId = turnId, text = content, vector = vector, - timestampMs = nowMs, + timestampMs = now.toEpochMilliseconds(), ) ) }.also { result -> @@ -176,12 +206,13 @@ class DefaultRouterFacade( } private suspend fun emitSteeringNote(sessionId: SessionId, content: String, effectiveStageId: StageId) { + val now = Clock.System.now() eventStore.append( NewEvent( metadata = EventMetadata( eventId = EventId(UUID.randomUUID().toString()), sessionId = sessionId, - timestamp = Clock.System.now(), + timestamp = now, schemaVersion = 1, causationId = null, correlationId = null, diff --git a/testing/deterministic/src/test/kotlin/RouterContextBuilderTest.kt b/testing/deterministic/src/test/kotlin/RouterContextBuilderTest.kt index 3817cab8..7918859b 100644 --- a/testing/deterministic/src/test/kotlin/RouterContextBuilderTest.kt +++ b/testing/deterministic/src/test/kotlin/RouterContextBuilderTest.kt @@ -2,7 +2,7 @@ import com.correx.core.context.model.ContextLayer import com.correx.core.context.model.ContextPack import com.correx.core.context.model.EntryRole import com.correx.core.context.model.TokenBudget -import kotlinx.coroutines.runBlocking +import com.correx.core.events.events.L3RetrievedHit import com.correx.core.events.types.SessionId import com.correx.core.events.types.StageId import com.correx.core.router.DefaultRouterContextBuilder @@ -13,6 +13,7 @@ import com.correx.core.router.model.RouterTurn import com.correx.core.router.model.StageOutcomeKind import com.correx.core.router.model.TurnRole import com.correx.core.router.model.WorkflowStatus +import kotlinx.coroutines.runBlocking import kotlinx.datetime.Clock import kotlinx.datetime.Instant import org.junit.jupiter.api.Assertions.assertEquals @@ -76,27 +77,32 @@ class RouterContextBuilderTest { } @Test - fun `build drops entries oldest-first for L2 memory`() { - val short = "ok" + fun `build drops entries oldest-first for L2 memory when budget forces partial drop`() { + // Each L2 entry: "Stage sN (SUCCESS): " + 200-char summary ≈ 55 tokens each. + // Use a config with conversationKeepLast=0 so only L0 + L2 compete for budget. + val configNoConv = RouterConfig(conversationKeepLast = 0, tokenBudget = TokenBudget(limit = 10000)) + val builderNoConv = DefaultRouterContextBuilder(configNoConv) + val summary = "x".repeat(200) // ~50 tokens; with prefix "Stage sN (SUCCESS): " ~55 tokens total val state = RouterState( sessionId = sessionId, workflowStatus = WorkflowStatus.RUNNING, currentStageId = stageId, l2Memory = listOf( - RouterL2Entry(StageId("s1"), short, StageOutcomeKind.SUCCESS, Instant.parse("2026-01-01T00:00:00Z")), - RouterL2Entry(StageId("s2"), short, StageOutcomeKind.SUCCESS, Instant.parse("2026-01-02T00:00:00Z")), - RouterL2Entry(StageId("s3"), short, StageOutcomeKind.SUCCESS, Instant.parse("2026-01-03T00:00:00Z")), + RouterL2Entry(StageId("s1"), summary, StageOutcomeKind.SUCCESS, Instant.parse("2026-01-01T00:00:00Z")), + RouterL2Entry(StageId("s2"), summary, StageOutcomeKind.SUCCESS, Instant.parse("2026-01-02T00:00:00Z")), + RouterL2Entry(StageId("s3"), summary, StageOutcomeKind.SUCCESS, Instant.parse("2026-01-03T00:00:00Z")), ), ) - // L0 consumes ~33 tokens; budget 93 leaves ~60 for L2. - // Each L2 entry is ~1 token; all 3 fit - val pack = buildPack(state, TokenBudget(limit = 93)) + // L0 ~30 tokens; budget 145 leaves ~115. Two entries (~55 each = ~110) fit; the third (oldest) is dropped. + val pack = runBlocking { builderNoConv.build(state, TokenBudget(limit = 145)) } val l2Entries = pack.layers[ContextLayer.L2] ?: emptyList() - val remainingStageIds = l2Entries.map { it.sourceId }.toSet() - assertEquals(3, l2Entries.size) - assertTrue(remainingStageIds.contains("s1")) - assertTrue(remainingStageIds.contains("s2")) - assertTrue(remainingStageIds.contains("s3")) + val retainedIds = l2Entries.map { it.sourceId }.toSet() + // Newest two survive; oldest is dropped + assertEquals(2, l2Entries.size, "Only 2 of 3 L2 entries should fit") + assertFalse(retainedIds.contains("s1"), "Oldest L2 entry (s1) should be dropped first") + assertTrue(retainedIds.contains("s2"), "Second-newest L2 entry (s2) should be retained") + assertTrue(retainedIds.contains("s3"), "Newest L2 entry (s3) should be retained") + assertEquals(1, pack.compressionMetadata.entriesDropped) } @Test @@ -182,10 +188,10 @@ class RouterContextBuilderTest { ) val pack = buildPack(state, TokenBudget(limit = 50)) val l0Entries = pack.layers[ContextLayer.L0] ?: emptyList() - val l1Entries = pack.layers[ContextLayer.L1] ?: emptyList() val l2Entries = pack.layers[ContextLayer.L2] ?: emptyList() assertEquals(2, l0Entries.size) - assertEquals(0, l1Entries.size) + // The current (last) user turn is always protected and present even at budget=50 + assertEquals(1, (pack.layers[ContextLayer.L1] ?: emptyList()).size) assertEquals(0, l2Entries.size) assertTrue(pack.compressionMetadata.entriesDropped > 0) } @@ -195,8 +201,11 @@ class RouterContextBuilderTest { // -------------------------------------------------------------------------- @Test - fun `L2 entries are evicted in insertion order (oldest first)`() { - val configTight = RouterConfig(conversationKeepLast = 0, tokenBudget = TokenBudget(limit = 10)) + fun `L2 entries are evicted oldest-first and retained in chronological order`() { + // Each L2 entry: "Stage sN (SUCCESS): " + 100-char summary ≈ 30 tokens each. + // L0 ~30 tokens; budget 90 leaves ~60. Two entries (~30 each = ~60) fit; oldest dropped. + val summary = "x".repeat(100) // ~25 tokens; with prefix ~30 tokens total + val configTight = RouterConfig(conversationKeepLast = 0, tokenBudget = TokenBudget(limit = 10000)) val builderTight = DefaultRouterContextBuilder(configTight) val baseTime = Instant.parse("2026-01-01T00:00:00Z") val state = RouterState( @@ -204,16 +213,25 @@ class RouterContextBuilderTest { workflowStatus = WorkflowStatus.RUNNING, currentStageId = stageId, l2Memory = listOf( - RouterL2Entry(StageId("s1"), "old", StageOutcomeKind.SUCCESS, baseTime), - RouterL2Entry(StageId("s2"), "mid", StageOutcomeKind.SUCCESS, Instant.parse("2026-01-01T01:00:00Z")), - RouterL2Entry(StageId("s3"), "new", StageOutcomeKind.SUCCESS, Instant.parse("2026-01-01T02:00:00Z")), + RouterL2Entry(StageId("s1"), summary, StageOutcomeKind.SUCCESS, baseTime), + RouterL2Entry(StageId("s2"), summary, StageOutcomeKind.SUCCESS, Instant.parse("2026-01-01T01:00:00Z")), + RouterL2Entry(StageId("s3"), summary, StageOutcomeKind.SUCCESS, Instant.parse("2026-01-01T02:00:00Z")), ), ) - // L0 consumes ~33 tokens; budget 76 leaves ~43 for L2. - // Each L2 entry is ~1 token; all 3 fit - val pack = runBlocking { builderTight.build(state, TokenBudget(limit = 76)) } + // "Stage sN (SUCCESS): " + 100x = ~30 tokens each; L0 ~30 tokens. + // Budget = L0 (30) + two entries (30+30=60) = 90. Third entry doesn't fit → oldest (s1) dropped. + val pack = runBlocking { builderTight.build(state, TokenBudget(limit = 90)) } val l2Entries = pack.layers[ContextLayer.L2] ?: emptyList() - assertEquals(3, l2Entries.size) + val retainedIds = l2Entries.map { it.sourceId } + // Oldest entry (s1) should be dropped; s2 and s3 retained in chronological order + assertFalse(retainedIds.contains("s1"), "Oldest L2 entry (s1) should be dropped first") + assertTrue(retainedIds.contains("s2"), "Second L2 entry (s2) should be retained") + assertTrue(retainedIds.contains("s3"), "Newest L2 entry (s3) should be retained") + // Chronological ordering preserved + if (retainedIds.size == 2) { + assertEquals("s2", retainedIds[0], "s2 should appear before s3 in chronological order") + assertEquals("s3", retainedIds[1], "s3 should appear after s2 in chronological order") + } } @Test @@ -516,7 +534,7 @@ class RouterContextBuilderTest { ) val pack = runBlocking { builder.build(state, TokenBudget(limit = 65)) } // L0 consumes ~30 tokens, leaving 35 — both L2 entries (each ~130 tokens) dropped - // conversationKeepLast=0 means conversation entry is not included + // conversationKeepLast=0 but the single user turn is the protected current turn; it fits and is not dropped assertEquals(2, pack.compressionMetadata.entriesDropped) } @@ -585,4 +603,361 @@ class RouterContextBuilderTest { val pack = buildPack(state, TokenBudget(limit = 1000)) assertTrue(pack.id.value.contains("unknown")) } + + // -------------------------------------------------------------------------- + // L3 memory injection (slice 2b) + // -------------------------------------------------------------------------- + + @Test + fun `L3 hits injected into pack ordered by score descending`() { + val state = RouterState( + sessionId = sessionId, + workflowStatus = WorkflowStatus.RUNNING, + currentStageId = stageId, + lastRetrievedMemory = listOf( + L3RetrievedHit( + entryId = "entry-low", + sourceSessionId = SessionId("other"), + sourceTurnId = "turn-low", + text = "low relevance memory", + score = 0.3f, + ), + L3RetrievedHit( + entryId = "entry-high", + sourceSessionId = SessionId("other"), + sourceTurnId = "turn-high", + text = "high relevance memory", + score = 0.9f, + ), + L3RetrievedHit( + entryId = "entry-mid", + sourceSessionId = SessionId("other"), + sourceTurnId = "turn-mid", + text = "mid relevance memory", + score = 0.6f, + ), + ), + ) + val pack = buildPack(state, TokenBudget(limit = 10000)) + val l3Entries = pack.layers[ContextLayer.L3] + assertNotNull(l3Entries) + assertEquals(3, l3Entries!!.size) + // Verify ordering: highest score first + assertEquals("entry-high", l3Entries[0].sourceId) + assertEquals("entry-mid", l3Entries[1].sourceId) + assertEquals("entry-low", l3Entries[2].sourceId) + } + + @Test + fun `L3 entries carry recalled memory label in content`() { + val state = RouterState( + sessionId = sessionId, + workflowStatus = WorkflowStatus.RUNNING, + lastRetrievedMemory = listOf( + L3RetrievedHit( + entryId = "e1", + sourceSessionId = SessionId("other"), + sourceTurnId = "t1", + text = "some past knowledge", + score = 0.8f, + ), + ), + ) + val pack = buildPack(state, TokenBudget(limit = 10000)) + val l3Entry = pack.layers[ContextLayer.L3]?.first() + assertNotNull(l3Entry) + assertTrue(l3Entry!!.content.contains("[recalled memory]")) + assertTrue(l3Entry.content.contains("some past knowledge")) + assertEquals(EntryRole.SYSTEM, l3Entry.role) + assertEquals(ContextLayer.L3, l3Entry.layer) + assertEquals("recalledMemory", l3Entry.sourceType) + } + + @Test + fun `L3 hits dropped when budget exhausted and counted in entriesDropped`() { + // Budget tight enough for L0 but not L3 entries with long text + val longText = "x".repeat(800) // ~200 tokens + val state = RouterState( + sessionId = sessionId, + workflowStatus = WorkflowStatus.RUNNING, + lastRetrievedMemory = listOf( + L3RetrievedHit( + entryId = "big-entry", + sourceSessionId = SessionId("other"), + sourceTurnId = "t1", + text = longText, + score = 0.9f, + ), + ), + ) + // Budget 50: L0 frame ~30 tokens leaves ~20; L3 entry ~200 tokens -> dropped + val pack = buildPack(state, TokenBudget(limit = 50)) + assertTrue(pack.layers[ContextLayer.L3].isNullOrEmpty()) + assertEquals(1, pack.compressionMetadata.entriesDropped) + assertTrue(pack.compressionMetadata.truncatedLayers.contains(ContextLayer.L3)) + } + + // -------------------------------------------------------------------------- + // Budget priority: protect system + current user turn, drop oldest first + // -------------------------------------------------------------------------- + + @Test + fun `system prompt and current user turn survive tight budget`() { + val longContent = "x".repeat(800) // ~200 tokens each + val state = RouterState( + sessionId = sessionId, + workflowStatus = WorkflowStatus.RUNNING, + currentStageId = stageId, + conversationHistory = listOf( + RouterTurn(TurnRole.USER, longContent, Instant.parse("2026-01-01T00:00:00Z"), turnId = "old-turn"), + RouterTurn(TurnRole.ROUTER, longContent, Instant.parse("2026-01-02T00:00:00Z"), turnId = "mid-turn"), + RouterTurn(TurnRole.USER, "current user input", Instant.parse("2026-01-03T00:00:00Z"), turnId = "current-turn"), + ), + ) + // Budget 60: L0 frame ~30 tokens -> remaining ~30; current user turn ~4 tokens -> protected. + // Older turns (each ~200 tokens) cannot fit -> dropped. + val pack = buildPack(state, TokenBudget(limit = 60)) + val l0 = pack.layers[ContextLayer.L0] + assertNotNull(l0) + assertEquals(2, l0!!.size) // system prompt + workflow status always present + + val l1 = pack.layers[ContextLayer.L1] ?: emptyList() + // Current user turn ("current user input") must be present + assertTrue(l1.any { it.content == "current user input" }) + // Older turns dropped due to budget + assertFalse(l1.any { it.content == longContent }) + assertTrue(pack.compressionMetadata.entriesDropped > 0) + } + + @Test + fun `oldest conversation turn is dropped before newest when budget is tight`() { + val config3 = RouterConfig(conversationKeepLast = 3, tokenBudget = TokenBudget(limit = 10000)) + val builder3 = DefaultRouterContextBuilder(config3) + val longContent = "x".repeat(800) // ~200 tokens each + val state = RouterState( + sessionId = sessionId, + workflowStatus = WorkflowStatus.RUNNING, + currentStageId = stageId, + conversationHistory = listOf( + RouterTurn(TurnRole.USER, "oldest turn $longContent", Instant.parse("2026-01-01T00:00:00Z"), turnId = "t1"), + RouterTurn(TurnRole.ROUTER, "middle turn $longContent", Instant.parse("2026-01-02T00:00:00Z"), turnId = "t2"), + RouterTurn(TurnRole.USER, "newest turn", Instant.parse("2026-01-03T00:00:00Z"), turnId = "t3"), + ), + ) + // Budget 100: L0 frame ~30 tokens -> remaining ~70; newest turn ~3 tokens -> protected. + // Middle turn (~200+) doesn't fit -> dropped. Oldest turn also dropped. + val pack = runBlocking { builder3.build(state, TokenBudget(limit = 100)) } + val l1 = pack.layers[ContextLayer.L1] ?: emptyList() + // Newest turn must survive + assertTrue(l1.any { it.content == "newest turn" }) + // Oldest turn should be dropped (too large) + assertFalse(l1.any { it.content.contains("oldest turn") && it.content.contains(longContent) }) + assertTrue(pack.compressionMetadata.entriesDropped > 0) + assertTrue(pack.compressionMetadata.truncatedLayers.contains(ContextLayer.L1)) + } + + @Test + fun `L3 and L2 dropped before recent conversation when budget is tight`() { + val longContent = "x".repeat(200) // ~50 tokens each + val state = RouterState( + sessionId = sessionId, + workflowStatus = WorkflowStatus.RUNNING, + currentStageId = stageId, + conversationHistory = listOf( + RouterTurn(TurnRole.USER, "current user message", clock.now(), turnId = "curr"), + ), + lastRetrievedMemory = listOf( + L3RetrievedHit( + entryId = "l3-entry", + sourceSessionId = SessionId("other"), + sourceTurnId = "l3-turn", + text = longContent, + score = 0.9f, + ), + ), + l2Memory = listOf( + RouterL2Entry(StageId("s1"), longContent, StageOutcomeKind.SUCCESS, clock.now()), + ), + ) + // Budget 80: L0 frame ~30 -> remaining ~50; current user turn ~5 -> protected, remaining ~45 + // L3 entry ~50 tokens -> just barely doesn't fit (50 > 45) -> dropped + // L2 entry ~50 tokens -> dropped + val pack = buildPack(state, TokenBudget(limit = 80)) + val l1 = pack.layers[ContextLayer.L1] ?: emptyList() + // Current user message must survive + assertTrue(l1.any { it.content == "current user message" }) + // L3 and L2 must be absent + assertTrue(pack.layers[ContextLayer.L3].isNullOrEmpty()) + assertTrue(pack.layers[ContextLayer.L2].isNullOrEmpty()) + assertTrue(pack.compressionMetadata.entriesDropped >= 2) + } + + @Test + fun `truncatedLayers populated with affected layers when drops occur`() { + val longText = "x".repeat(800) + val state = RouterState( + sessionId = sessionId, + workflowStatus = WorkflowStatus.RUNNING, + currentStageId = stageId, + conversationHistory = listOf( + RouterTurn(TurnRole.USER, longText, clock.now(), turnId = "old"), + RouterTurn(TurnRole.USER, "current", clock.now(), turnId = "curr"), + ), + lastRetrievedMemory = listOf( + L3RetrievedHit( + entryId = "l3e", + sourceSessionId = SessionId("other"), + sourceTurnId = "lt", + text = longText, + score = 0.8f, + ), + ), + ) + val pack = buildPack(state, TokenBudget(limit = 60)) + val truncated = pack.compressionMetadata.truncatedLayers + // At least L1 and L3 should be truncated + assertTrue(truncated.isNotEmpty()) + } + + // -------------------------------------------------------------------------- + // N2/B2: protected-frame-overflow — budget smaller than system + current user turn + // -------------------------------------------------------------------------- + + @Test + fun `budget overflow - system prompt and current user turn present even when they exceed budgetLimit`() { + // Budget so tight that the protected frame alone (system prompt + workflow status) exceeds it. + // The current user turn must still be included; budgetUsed must report the true total; + // appliedStrategies must contain "BudgetExceeded". + val longUserTurn = "x".repeat(200) // ~50 tokens + val state = RouterState( + sessionId = sessionId, + workflowStatus = WorkflowStatus.RUNNING, + currentStageId = stageId, + conversationHistory = listOf( + RouterTurn(TurnRole.USER, longUserTurn, clock.now(), turnId = "current"), + ), + ) + // Budget 1: system prompt + workflow status alone exceed this. + val pack = buildPack(state, TokenBudget(limit = 1)) + + // System prompt must be present + val l0 = pack.layers[ContextLayer.L0] ?: emptyList() + assertTrue(l0.any { it.sourceType == "systemPrompt" }, "System prompt must always be present") + assertTrue(l0.any { it.sourceType == "workflowStatus" }, "Workflow status must always be present") + + // Current user turn must be present + val l1 = pack.layers[ContextLayer.L1] ?: emptyList() + assertTrue(l1.any { it.content == longUserTurn }, "Current user turn must be present even when over budget") + + // budgetUsed must be the true total (exceeds budgetLimit) + val allEntries = pack.layers.values.flatten() + val trueTotal = allEntries.sumOf { it.tokenEstimate } + assertEquals(trueTotal, pack.budgetUsed, "budgetUsed must report the true token total, not clamped") + assertTrue(pack.budgetUsed > pack.budgetLimit, "budgetUsed must exceed budgetLimit in overflow scenario") + + // BudgetExceeded must be signalled in appliedStrategies + assertTrue( + pack.compressionMetadata.appliedStrategies.contains("BudgetExceeded"), + "appliedStrategies must contain 'BudgetExceeded' when frame exceeds budget", + ) + } + + // -------------------------------------------------------------------------- + // S1: conversationKeepLast = 0 must still include the current user turn + // -------------------------------------------------------------------------- + + @Test + fun `conversationKeepLast zero still includes the current user turn`() { + val configZero = RouterConfig(conversationKeepLast = 0, tokenBudget = TokenBudget(limit = 10000)) + val builderZero = DefaultRouterContextBuilder(configZero) + val state = RouterState( + sessionId = sessionId, + workflowStatus = WorkflowStatus.RUNNING, + currentStageId = stageId, + conversationHistory = listOf( + RouterTurn(TurnRole.USER, "old turn 1", Instant.parse("2026-01-01T00:00:00Z"), turnId = "t1"), + RouterTurn(TurnRole.ROUTER, "old reply", Instant.parse("2026-01-02T00:00:00Z"), turnId = "t2"), + RouterTurn(TurnRole.USER, "current user input", Instant.parse("2026-01-03T00:00:00Z"), turnId = "t3"), + ), + ) + val pack = runBlocking { builderZero.build(state, TokenBudget(limit = 10000)) } + val l1 = pack.layers[ContextLayer.L1] ?: emptyList() + // conversationKeepLast=0 caps the window to 0 turns, but the protected current turn must appear + assertTrue(l1.any { it.content == "current user input" }, "Current user turn must be present even with conversationKeepLast=0") + // The older turns (capped out) must NOT be present + assertFalse(l1.any { it.content == "old turn 1" }, "Older turns must not appear with conversationKeepLast=0") + assertFalse(l1.any { it.content == "old reply" }, "Older turns must not appear with conversationKeepLast=0") + } + + // -------------------------------------------------------------------------- + // S4: appliedStrategies includes L3Recall when L3 memory is injected + // -------------------------------------------------------------------------- + + @Test + fun `appliedStrategies includes L3Recall when lastRetrievedMemory is non-empty`() { + val state = RouterState( + sessionId = sessionId, + workflowStatus = WorkflowStatus.RUNNING, + currentStageId = stageId, + lastRetrievedMemory = listOf( + L3RetrievedHit( + entryId = "e1", + sourceSessionId = SessionId("other"), + sourceTurnId = "t1", + text = "recalled fact", + score = 0.7f, + ), + ), + ) + val pack = buildPack(state, TokenBudget(limit = 10000)) + assertTrue( + pack.compressionMetadata.appliedStrategies.contains("L3Recall"), + "appliedStrategies must include 'L3Recall' when L3 memory is injected", + ) + } + + @Test + fun `appliedStrategies does not include L3Recall when lastRetrievedMemory is empty`() { + val state = RouterState( + sessionId = sessionId, + workflowStatus = WorkflowStatus.RUNNING, + ) + val pack = buildPack(state, TokenBudget(limit = 10000)) + assertFalse( + pack.compressionMetadata.appliedStrategies.contains("L3Recall"), + "appliedStrategies must not include 'L3Recall' when no L3 memory is present", + ) + } + + // -------------------------------------------------------------------------- + // N3: ContextPack.layers is sorted by ContextLayer ordinal + // -------------------------------------------------------------------------- + + @Test + fun `layers map is sorted by ContextLayer ordinal`() { + val state = RouterState( + sessionId = sessionId, + workflowStatus = WorkflowStatus.RUNNING, + currentStageId = stageId, + conversationHistory = listOf( + RouterTurn(TurnRole.USER, "hello", clock.now()), + ), + l2Memory = listOf( + RouterL2Entry(StageId("s1"), "summary", StageOutcomeKind.SUCCESS, clock.now()), + ), + lastRetrievedMemory = listOf( + L3RetrievedHit( + entryId = "e1", + sourceSessionId = SessionId("other"), + sourceTurnId = "t1", + text = "recalled", + score = 0.5f, + ), + ), + ) + val pack = buildPack(state, TokenBudget(limit = 10000)) + val layerKeys = pack.layers.keys.toList() + val ordinals = layerKeys.map { it.ordinal } + assertEquals(ordinals.sorted(), ordinals, "ContextPack.layers must be sorted by ContextLayer ordinal") + } } diff --git a/testing/deterministic/src/test/kotlin/RouterFacadeTest.kt b/testing/deterministic/src/test/kotlin/RouterFacadeTest.kt index d78c6354..c3603aef 100644 --- a/testing/deterministic/src/test/kotlin/RouterFacadeTest.kt +++ b/testing/deterministic/src/test/kotlin/RouterFacadeTest.kt @@ -1,7 +1,9 @@ +import com.correx.core.context.model.ContextLayer import com.correx.core.context.model.ContextPack import com.correx.core.context.model.TokenBudget import com.correx.core.events.events.ChatTurnEvent import com.correx.core.events.events.ChatTurnRole +import com.correx.core.events.events.ContextTruncatedEvent import com.correx.core.events.events.EventMetadata import com.correx.core.events.events.L3MemoryRetrievedEvent import com.correx.core.events.events.L3RetrievedHit @@ -27,6 +29,7 @@ import com.correx.core.inference.ProviderHealth import com.correx.core.inference.ResponseFormat import com.correx.core.inference.TokenUsage import com.correx.core.router.ChatMode +import com.correx.core.router.DefaultRouterContextBuilder import com.correx.core.router.DefaultRouterFacade import com.correx.core.router.DefaultRouterReducer import com.correx.core.router.RouterContextBuilder @@ -40,6 +43,7 @@ import com.correx.core.router.l3.L3Query import com.correx.core.router.model.RouterConfig import com.correx.core.router.model.RouterResponse import com.correx.core.router.model.RouterState +import com.correx.core.router.model.RouterTurn import com.correx.core.router.model.TurnRole import com.correx.core.router.model.WorkflowStatus import com.correx.core.sessions.projections.replay.DefaultEventReplayer @@ -47,6 +51,7 @@ import com.correx.testing.fixtures.EventFixtures import com.correx.testing.fixtures.inference.MockTokenizer import kotlinx.coroutines.flow.Flow import kotlinx.coroutines.runBlocking +import kotlinx.datetime.Instant import org.junit.jupiter.api.Assertions.assertEquals import org.junit.jupiter.api.Assertions.assertFalse import org.junit.jupiter.api.Assertions.assertNotNull @@ -781,8 +786,10 @@ class RouterFacadeTest { facade.onUserInput(sessionId = currentSessionId, input = "a new message") val l3Events = eventStore.appendedEvents.filter { it.payload is L3MemoryRetrievedEvent } - // The only hit was in-session, so deduped to empty → no L3MemoryRetrievedEvent - assertTrue(l3Events.isEmpty(), "Expected no L3RetrievedEvent when all hits are deduplicated") + // The only hit was in-session, so deduped to empty → L3MemoryRetrievedEvent emitted with empty hits + assertEquals(1, l3Events.size, "Expected one L3MemoryRetrievedEvent even when all hits are deduplicated") + val deduplicatedRetrieved = l3Events[0].payload as L3MemoryRetrievedEvent + assertTrue(deduplicatedRetrieved.hits.isEmpty(), "Expected empty hits list when all hits are deduplicated") } @Test @@ -814,8 +821,10 @@ class RouterFacadeTest { val payloads = eventStore.appendedEvents.map { it.payload } val chatEvents = payloads.filterIsInstance() assertEquals(2, chatEvents.size, "USER and ROUTER ChatTurnEvents must be present") + // Retrieval failed → deduped is empty; L3MemoryRetrievedEvent still emitted with empty hits val l3Events = payloads.filterIsInstance() - assertTrue(l3Events.isEmpty(), "No L3RetrievedEvent expected on retrieval failure") + assertEquals(1, l3Events.size, "L3MemoryRetrievedEvent must be emitted even on retrieval failure") + assertTrue(l3Events[0].hits.isEmpty(), "Expected empty hits when retrieval failed") } @Test @@ -847,6 +856,263 @@ class RouterFacadeTest { assertEquals(0.95f, state.lastRetrievedMemory[0].score) } + // -------------------------------------------------------------------------- + // Slice 2b: ContextTruncatedEvent + end-to-end L3 recall injection + // -------------------------------------------------------------------------- + + @Test + fun `ContextTruncatedEvent emitted when context builder drops entries`(): Unit = runBlocking { + // Use the real DefaultRouterContextBuilder with a tight budget and large conversation + val longContent = "x".repeat(800) // ~200 tokens each + val sessionId = SessionId("truncation-session") + val eventStore = mockEventStore() + val replayer = DefaultEventReplayer( + store = eventStore, + projection = RouterProjector(DefaultRouterReducer()), + ) + val facade = DefaultRouterFacade( + routerRepository = object : RouterRepository { + override suspend fun getRouterState(sessionId: SessionId): RouterState = + replayer.rebuild(sessionId) + }, + routerContextBuilder = DefaultRouterContextBuilder( + config = RouterConfig( + conversationKeepLast = 4, + tokenBudget = TokenBudget(limit = 80), + ), + ), + inferenceRouter = mockInferenceRouter("response"), + eventStore = eventStore, + config = RouterConfig( + conversationKeepLast = 4, + tokenBudget = TokenBudget(limit = 80), + ), + embedder = NoopEmbedder(dimension = 8), + l3MemoryStore = InMemoryL3MemoryStore(), + ) + + // Pre-populate with large conversation turns so budget overflows + eventStore.append( + NewEvent( + metadata = EventMetadata( + eventId = EventId("seed-1"), + sessionId = sessionId, + timestamp = kotlinx.datetime.Clock.System.now(), + schemaVersion = 1, + causationId = null, + correlationId = null, + ), + payload = ChatTurnEvent( + sessionId = sessionId, + turnId = "old-turn", + role = ChatTurnRole.USER, + content = longContent, + timestampMs = 1000L, + ), + ) + ) + + facade.onUserInput(sessionId = sessionId, input = "current input") + + val payloads = eventStore.appendedEvents.map { it.payload } + val truncEvents = payloads.filterIsInstance() + assertEquals(1, truncEvents.size) + val truncEvent = truncEvents[0] + assertEquals(sessionId, truncEvent.sessionId) + assertTrue(truncEvent.entriesDropped > 0) + assertTrue(truncEvent.truncatedLayers.isNotEmpty()) + } + + @Test + fun `no ContextTruncatedEvent when budget is generous`(): Unit = runBlocking { + val eventStore = mockEventStore() + val facade = DefaultRouterFacade( + routerRepository = object : RouterRepository { + override suspend fun getRouterState(sessionId: SessionId): RouterState = RouterState() + }, + routerContextBuilder = DefaultRouterContextBuilder( + config = RouterConfig(tokenBudget = TokenBudget(limit = 10000)), + ), + inferenceRouter = mockInferenceRouter("response"), + eventStore = eventStore, + config = RouterConfig(tokenBudget = TokenBudget(limit = 10000)), + embedder = NoopEmbedder(dimension = 8), + l3MemoryStore = InMemoryL3MemoryStore(), + ) + facade.onUserInput(sessionId = SessionId("generous-session"), input = "hello") + val payloads = eventStore.appendedEvents.map { it.payload } + val truncEvents = payloads.filterIsInstance() + assertTrue(truncEvents.isEmpty(), "No ContextTruncatedEvent expected when budget is generous") + } + + @Test + fun `ContextTruncatedEvent carries correct turnId matching userTurnId`(): Unit = runBlocking { + val longContent = "x".repeat(800) + val sessionId = SessionId("turnid-session") + val eventStore = mockEventStore() + val replayer = DefaultEventReplayer( + store = eventStore, + projection = RouterProjector(DefaultRouterReducer()), + ) + + // Pre-seed with a large turn so the budget overflows on the next call + eventStore.append( + NewEvent( + metadata = EventMetadata( + eventId = EventId("seed-x"), + sessionId = sessionId, + timestamp = kotlinx.datetime.Clock.System.now(), + schemaVersion = 1, + causationId = null, + correlationId = null, + ), + payload = ChatTurnEvent( + sessionId = sessionId, + turnId = "seeded-turn", + role = ChatTurnRole.USER, + content = longContent, + timestampMs = 500L, + ), + ) + ) + + val facade = DefaultRouterFacade( + routerRepository = object : RouterRepository { + override suspend fun getRouterState(sessionId: SessionId): RouterState = + replayer.rebuild(sessionId) + }, + routerContextBuilder = DefaultRouterContextBuilder( + config = RouterConfig(conversationKeepLast = 4, tokenBudget = TokenBudget(limit = 80)), + ), + inferenceRouter = mockInferenceRouter("response"), + eventStore = eventStore, + config = RouterConfig(conversationKeepLast = 4, tokenBudget = TokenBudget(limit = 80)), + embedder = NoopEmbedder(dimension = 8), + l3MemoryStore = InMemoryL3MemoryStore(), + ) + + facade.onUserInput(sessionId = sessionId, input = "new input") + + val payloads = eventStore.appendedEvents.map { it.payload } + val userChatEvent = payloads.filterIsInstance().first { it.role == ChatTurnRole.USER && it.content == "new input" } + val truncEvent = payloads.filterIsInstance().firstOrNull() + assertNotNull(truncEvent) + assertEquals(userChatEvent.turnId, truncEvent!!.turnId) + } + + @Test + fun `empty retrieval always emits L3MemoryRetrievedEvent so next turn does not see stale memory`(): Unit = runBlocking { + // Scenario: two consecutive CHAT turns with no cross-session L3 entries. + // Each turn must emit L3MemoryRetrievedEvent (with empty hits), so that + // lastRetrievedMemory is always current-turn-scoped and never stale. + val eventStore = mockEventStore() + val replayer = DefaultEventReplayer( + store = eventStore, + projection = RouterProjector(DefaultRouterReducer()), + ) + val capturedStates = mutableListOf() + val facade = DefaultRouterFacade( + routerRepository = object : RouterRepository { + override suspend fun getRouterState(sessionId: SessionId): RouterState = + replayer.rebuild(sessionId) + }, + routerContextBuilder = object : RouterContextBuilder { + private val real = DefaultRouterContextBuilder( + config = RouterConfig(tokenBudget = TokenBudget(limit = 10000)), + ) + override suspend fun build(state: RouterState, budget: TokenBudget): ContextPack { + capturedStates.add(state) + return real.build(state, budget) + } + }, + inferenceRouter = mockInferenceRouter("reply"), + eventStore = eventStore, + config = RouterConfig(tokenBudget = TokenBudget(limit = 10000), retrievalK = 5), + embedder = NoopEmbedder(dimension = 8), + l3MemoryStore = InMemoryL3MemoryStore(), + ) + + val sid = SessionId("stale-test-session") + facade.onUserInput(sessionId = sid, input = "first input") + facade.onUserInput(sessionId = sid, input = "second input") + + // Exactly two L3MemoryRetrievedEvents must have been emitted — one per CHAT turn + val l3Events = eventStore.appendedEvents.filter { it.payload is L3MemoryRetrievedEvent } + assertEquals(2, l3Events.size, "One L3MemoryRetrievedEvent must be emitted per CHAT turn") + assertTrue((l3Events[0].payload as L3MemoryRetrievedEvent).hits.isEmpty(), "First turn: no cross-session entries → empty hits") + assertTrue((l3Events[1].payload as L3MemoryRetrievedEvent).hits.isEmpty(), "Second turn: no cross-session entries → empty hits") + + // State passed to builder on both calls must have empty lastRetrievedMemory + // (because the L3MemoryRetrievedEvent with empty hits was emitted and replayed before each build) + assertEquals(2, capturedStates.size) + assertTrue( + capturedStates[0].lastRetrievedMemory.isEmpty(), + "First build: lastRetrievedMemory must be empty (no prior hits)", + ) + assertTrue( + capturedStates[1].lastRetrievedMemory.isEmpty(), + "Second build: lastRetrievedMemory must be empty (stale memory guard)", + ) + } + + @Test + fun `end-to-end L3 hit from cross-session injected into context pack reaching inference`(): Unit = runBlocking { + val dimension = 8 + val knownVector = FloatArray(dimension) { if (it == 0) 1f else 0f } + val stubEmbedder = object : Embedder { + override val dimension: Int = dimension + override suspend fun embed(text: String): FloatArray = knownVector.copyOf() + } + val l3Store = InMemoryL3MemoryStore() + val otherSessionId = SessionId("cross-session") + val crossSessionText = "important cross-session memory" + l3Store.store( + L3MemoryEntry( + id = "cross-entry", + sessionId = otherSessionId, + turnId = "cross-turn", + text = crossSessionText, + vector = knownVector.copyOf(), + timestampMs = 1000L, + ) + ) + + val capturedPacks = mutableListOf() + val eventStore = mockEventStore() + val replayer = DefaultEventReplayer( + store = eventStore, + projection = RouterProjector(DefaultRouterReducer()), + ) + val facade = DefaultRouterFacade( + routerRepository = object : RouterRepository { + override suspend fun getRouterState(sessionId: SessionId): RouterState = + replayer.rebuild(sessionId) + }, + routerContextBuilder = object : RouterContextBuilder { + private val real = DefaultRouterContextBuilder( + config = RouterConfig(tokenBudget = TokenBudget(limit = 10000)), + ) + override suspend fun build(state: RouterState, budget: TokenBudget): ContextPack = + real.build(state, budget).also { capturedPacks.add(it) } + }, + inferenceRouter = mockInferenceRouter("router reply"), + eventStore = eventStore, + config = RouterConfig(tokenBudget = TokenBudget(limit = 10000), retrievalK = 5), + embedder = stubEmbedder, + l3MemoryStore = l3Store, + ) + + facade.onUserInput(sessionId = SessionId("current-session"), input = "tell me something") + + assertEquals(1, capturedPacks.size) + val pack = capturedPacks[0] + val l3Entries = pack.layers[ContextLayer.L3] + assertNotNull(l3Entries) + assertTrue(l3Entries!!.isNotEmpty()) + assertTrue(l3Entries.any { it.content.contains(crossSessionText) }) + assertTrue(l3Entries.any { it.content.contains("[recalled memory]") }) + } + private fun emptyContextPack(): ContextPack = ContextPack( id = ContextPackId("empty"), sessionId = SessionId("unknown"),