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).
This commit is contained in:
2026-05-30 18:40:08 +04:00
parent e6239515bd
commit 8e6a3e1470
6 changed files with 852 additions and 77 deletions
@@ -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<com.correx.core.events.events.ChatTurnEvent>()
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<L3MemoryRetrievedEvent>()
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<RouterState>(
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<ContextTruncatedEvent>()
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<ContextTruncatedEvent>()
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<RouterState>(
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<ChatTurnEvent>().first { it.role == ChatTurnRole.USER && it.content == "new input" }
val truncEvent = payloads.filterIsInstance<ContextTruncatedEvent>().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<RouterState>(
store = eventStore,
projection = RouterProjector(DefaultRouterReducer()),
)
val capturedStates = mutableListOf<RouterState>()
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<ContextPack>()
val eventStore = mockEventStore()
val replayer = DefaultEventReplayer<RouterState>(
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"),