feat: L3 memory retrieval on the router read path (record-only)
When a user sends input, embed it, query L3 across all sessions, dedup against in-session turns, and record the retrieval as an event so replay is deterministic (invariant #9). Hits land in RouterState; context injection follows in a later slice. - Add L3MemoryRetrievedEvent + L3RetrievedHit (registered in eventModule) - RouterState.lastRetrievedMemory + reducer case; RouterTurn carries turnId - RouterConfig.retrievalK (default 5) - Harden L3 write path: runCatching + visible logging, cancellation re-thrown; event append stays ahead of the L3 write - Warn prominently when the non-durable in_memory L3 backend is selected
This commit is contained in:
@@ -1,5 +1,10 @@
|
||||
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.EventMetadata
|
||||
import com.correx.core.events.events.L3MemoryRetrievedEvent
|
||||
import com.correx.core.events.events.L3RetrievedHit
|
||||
import com.correx.core.events.events.NewEvent
|
||||
import com.correx.core.events.events.StoredEvent
|
||||
import com.correx.core.events.stores.EventStore
|
||||
@@ -29,6 +34,7 @@ import com.correx.core.router.RouterFacade
|
||||
import com.correx.core.router.RouterProjector
|
||||
import com.correx.core.router.RouterRepository
|
||||
import com.correx.core.router.l3.InMemoryL3MemoryStore
|
||||
import com.correx.core.router.l3.L3MemoryEntry
|
||||
import com.correx.core.router.l3.L3MemoryStore
|
||||
import com.correx.core.router.l3.L3Query
|
||||
import com.correx.core.router.model.RouterConfig
|
||||
@@ -37,6 +43,7 @@ import com.correx.core.router.model.RouterState
|
||||
import com.correx.core.router.model.TurnRole
|
||||
import com.correx.core.router.model.WorkflowStatus
|
||||
import com.correx.core.sessions.projections.replay.DefaultEventReplayer
|
||||
import com.correx.testing.fixtures.EventFixtures
|
||||
import com.correx.testing.fixtures.inference.MockTokenizer
|
||||
import kotlinx.coroutines.flow.Flow
|
||||
import kotlinx.coroutines.runBlocking
|
||||
@@ -73,14 +80,15 @@ class RouterFacadeTest {
|
||||
val mockStore = mockEventStore()
|
||||
val facade = facadeWithMocks(eventStore = mockStore, chatMode = ChatMode.CHAT)
|
||||
facade.onUserInput(sessionId = SessionId("test-session"), input = "Hello!")
|
||||
assertEquals(2, mockStore.appendedEvents.size)
|
||||
val userEvent = mockStore.appendedEvents[0].payload
|
||||
val routerEvent = mockStore.appendedEvents[1].payload
|
||||
assertTrue(userEvent is com.correx.core.events.events.ChatTurnEvent)
|
||||
assertTrue(routerEvent is com.correx.core.events.events.ChatTurnEvent)
|
||||
assertEquals("Hello!", (userEvent as com.correx.core.events.events.ChatTurnEvent).content)
|
||||
val chatEvents = mockStore.appendedEvents
|
||||
.map { it.payload }
|
||||
.filterIsInstance<com.correx.core.events.events.ChatTurnEvent>()
|
||||
assertEquals(2, chatEvents.size)
|
||||
val userEvent = chatEvents[0]
|
||||
val routerEvent = chatEvents[1]
|
||||
assertEquals("Hello!", userEvent.content)
|
||||
assertEquals(com.correx.core.events.events.ChatTurnRole.USER, userEvent.role)
|
||||
assertEquals("inference response", (routerEvent as com.correx.core.events.events.ChatTurnEvent).content)
|
||||
assertEquals("inference response", routerEvent.content)
|
||||
assertEquals(com.correx.core.events.events.ChatTurnRole.ROUTER, routerEvent.role)
|
||||
}
|
||||
|
||||
@@ -101,16 +109,16 @@ class RouterFacadeTest {
|
||||
val mockStore = mockEventStore()
|
||||
val facade = facadeWithMocks(eventStore = mockStore, chatMode = ChatMode.STEERING)
|
||||
facade.onUserInput(sessionId = SessionId("session-xyz"), input = "steer this way")
|
||||
assertEquals(3, mockStore.appendedEvents.size)
|
||||
val userEvent = mockStore.appendedEvents[0].payload
|
||||
val routerEvent = mockStore.appendedEvents[1].payload
|
||||
val steeringEvent = mockStore.appendedEvents[2].payload
|
||||
assertTrue(userEvent is com.correx.core.events.events.ChatTurnEvent)
|
||||
assertTrue(routerEvent is com.correx.core.events.events.ChatTurnEvent)
|
||||
assertTrue(steeringEvent is com.correx.core.events.events.SteeringNoteAddedEvent)
|
||||
assertEquals("steer this way", (userEvent as com.correx.core.events.events.ChatTurnEvent).content)
|
||||
assertEquals("inference response", (routerEvent as com.correx.core.events.events.ChatTurnEvent).content)
|
||||
assertEquals("inference response", (steeringEvent as com.correx.core.events.events.SteeringNoteAddedEvent).content)
|
||||
val payloads = mockStore.appendedEvents.map { it.payload }
|
||||
val chatEvents = payloads.filterIsInstance<com.correx.core.events.events.ChatTurnEvent>()
|
||||
val steeringEvents = payloads.filterIsInstance<com.correx.core.events.events.SteeringNoteAddedEvent>()
|
||||
assertEquals(2, chatEvents.size)
|
||||
assertEquals(1, steeringEvents.size)
|
||||
assertEquals("steer this way", chatEvents[0].content)
|
||||
assertEquals(com.correx.core.events.events.ChatTurnRole.USER, chatEvents[0].role)
|
||||
assertEquals("inference response", chatEvents[1].content)
|
||||
assertEquals(com.correx.core.events.events.ChatTurnRole.ROUTER, chatEvents[1].role)
|
||||
assertEquals("inference response", steeringEvents[0].content)
|
||||
}
|
||||
|
||||
@Test
|
||||
@@ -136,10 +144,10 @@ class RouterFacadeTest {
|
||||
l3MemoryStore = InMemoryL3MemoryStore(),
|
||||
)
|
||||
facade.onUserInput(sessionId = SessionId("test-session"), input = "Hello!", mode = ChatMode.STEERING)
|
||||
assertEquals(3, mockStore.appendedEvents.size)
|
||||
val steeringEvent = mockStore.appendedEvents[2].payload
|
||||
assertTrue(steeringEvent is com.correx.core.events.events.SteeringNoteAddedEvent)
|
||||
assertEquals("steering response", (steeringEvent as com.correx.core.events.events.SteeringNoteAddedEvent).content)
|
||||
val payloads = mockStore.appendedEvents.map { it.payload }
|
||||
val steeringEvents = payloads.filterIsInstance<com.correx.core.events.events.SteeringNoteAddedEvent>()
|
||||
assertEquals(1, steeringEvents.size)
|
||||
assertEquals("steering response", steeringEvents[0].content)
|
||||
}
|
||||
|
||||
// --------------------------------------------------------------------------
|
||||
@@ -503,7 +511,9 @@ class RouterFacadeTest {
|
||||
assertNotNull(response)
|
||||
assertEquals("inference response", response.content)
|
||||
assertTrue(response.steeringEmitted)
|
||||
assertEquals(3, mockStore.appendedEvents.size)
|
||||
val payloads = mockStore.appendedEvents.map { it.payload }
|
||||
assertEquals(2, payloads.filterIsInstance<com.correx.core.events.events.ChatTurnEvent>().size)
|
||||
assertEquals(1, payloads.filterIsInstance<com.correx.core.events.events.SteeringNoteAddedEvent>().size)
|
||||
}
|
||||
|
||||
@Test
|
||||
@@ -643,6 +653,200 @@ class RouterFacadeTest {
|
||||
override fun capabilities(): Set<CapabilityScore> = emptySet()
|
||||
}
|
||||
|
||||
// --------------------------------------------------------------------------
|
||||
// L3 retrieval slice 2a
|
||||
// --------------------------------------------------------------------------
|
||||
|
||||
@Test
|
||||
fun `retrieval emits L3MemoryRetrievedEvent with cross-session hit`(): 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("other-session")
|
||||
val crossSessionEntryId = "cross-entry-id"
|
||||
val crossSessionTurnId = "cross-turn-id"
|
||||
l3Store.store(
|
||||
L3MemoryEntry(
|
||||
id = crossSessionEntryId,
|
||||
sessionId = otherSessionId,
|
||||
turnId = crossSessionTurnId,
|
||||
text = "memory from another session",
|
||||
vector = knownVector.copyOf(),
|
||||
timestampMs = 1000L,
|
||||
)
|
||||
)
|
||||
|
||||
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 {
|
||||
override suspend fun build(state: RouterState, budget: TokenBudget): ContextPack = emptyContextPack()
|
||||
},
|
||||
inferenceRouter = mockInferenceRouter("router reply"),
|
||||
eventStore = eventStore,
|
||||
config = RouterConfig(tokenBudget = TokenBudget(limit = 5000), retrievalK = 5),
|
||||
embedder = stubEmbedder,
|
||||
l3MemoryStore = l3Store,
|
||||
)
|
||||
|
||||
facade.onUserInput(sessionId = SessionId("current-session"), input = "hello")
|
||||
|
||||
val l3Events = eventStore.appendedEvents.filter { it.payload is L3MemoryRetrievedEvent }
|
||||
assertEquals(1, l3Events.size)
|
||||
val retrieved = l3Events[0].payload as L3MemoryRetrievedEvent
|
||||
assertEquals(1, retrieved.hits.size)
|
||||
assertEquals(crossSessionEntryId, retrieved.hits[0].entryId)
|
||||
assertEquals(otherSessionId, retrieved.hits[0].sourceSessionId)
|
||||
assertEquals(crossSessionTurnId, retrieved.hits[0].sourceTurnId)
|
||||
assertEquals("memory from another session", retrieved.hits[0].text)
|
||||
}
|
||||
|
||||
@Test
|
||||
fun `dedup filters hit whose turnId is already in current session history`(): 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 currentSessionId = SessionId("dedup-session")
|
||||
val inSessionTurnId = "in-session-turn-id"
|
||||
|
||||
// Pre-populate L3 with an entry whose turnId matches an in-session turn
|
||||
l3Store.store(
|
||||
L3MemoryEntry(
|
||||
id = inSessionTurnId,
|
||||
sessionId = currentSessionId,
|
||||
turnId = inSessionTurnId,
|
||||
text = "already in session history",
|
||||
vector = knownVector.copyOf(),
|
||||
timestampMs = 500L,
|
||||
)
|
||||
)
|
||||
|
||||
val eventStore = mockEventStore()
|
||||
val replayer = DefaultEventReplayer<RouterState>(
|
||||
store = eventStore,
|
||||
projection = RouterProjector(DefaultRouterReducer()),
|
||||
)
|
||||
|
||||
// Bootstrap session with a ChatTurnEvent whose turnId matches the L3 entry
|
||||
eventStore.append(
|
||||
NewEvent(
|
||||
metadata = EventMetadata(
|
||||
eventId = EventId("seed-event-id"),
|
||||
sessionId = currentSessionId,
|
||||
timestamp = kotlinx.datetime.Clock.System.now(),
|
||||
schemaVersion = 1,
|
||||
causationId = null,
|
||||
correlationId = null,
|
||||
),
|
||||
payload = ChatTurnEvent(
|
||||
sessionId = currentSessionId,
|
||||
turnId = inSessionTurnId,
|
||||
role = ChatTurnRole.USER,
|
||||
content = "already in session history",
|
||||
timestampMs = 500L,
|
||||
),
|
||||
)
|
||||
)
|
||||
|
||||
val facade = DefaultRouterFacade(
|
||||
routerRepository = object : RouterRepository {
|
||||
override suspend fun getRouterState(sessionId: SessionId): RouterState =
|
||||
replayer.rebuild(sessionId)
|
||||
},
|
||||
routerContextBuilder = object : RouterContextBuilder {
|
||||
override suspend fun build(state: RouterState, budget: TokenBudget): ContextPack = emptyContextPack()
|
||||
},
|
||||
inferenceRouter = mockInferenceRouter("router reply"),
|
||||
eventStore = eventStore,
|
||||
config = RouterConfig(tokenBudget = TokenBudget(limit = 5000), retrievalK = 5),
|
||||
embedder = stubEmbedder,
|
||||
l3MemoryStore = l3Store,
|
||||
)
|
||||
|
||||
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")
|
||||
}
|
||||
|
||||
@Test
|
||||
fun `retrieval failure is non-fatal - CHAT turn still completes`(): Unit = runBlocking {
|
||||
val throwingEmbedder = object : Embedder {
|
||||
override val dimension: Int = 8
|
||||
override suspend fun embed(text: String): FloatArray =
|
||||
throw RuntimeException("embedding service unavailable")
|
||||
}
|
||||
val eventStore = mockEventStore()
|
||||
val facade = DefaultRouterFacade(
|
||||
routerRepository = object : RouterRepository {
|
||||
override suspend fun getRouterState(sessionId: SessionId): RouterState = RouterState()
|
||||
},
|
||||
routerContextBuilder = object : RouterContextBuilder {
|
||||
override suspend fun build(state: RouterState, budget: TokenBudget): ContextPack = emptyContextPack()
|
||||
},
|
||||
inferenceRouter = mockInferenceRouter("response"),
|
||||
eventStore = eventStore,
|
||||
config = RouterConfig(tokenBudget = TokenBudget(limit = 5000)),
|
||||
embedder = throwingEmbedder,
|
||||
l3MemoryStore = InMemoryL3MemoryStore(),
|
||||
)
|
||||
|
||||
// Should not throw despite embedder failure
|
||||
val response = facade.onUserInput(sessionId = SessionId("fail-session"), input = "test")
|
||||
assertEquals("response", response.content)
|
||||
|
||||
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")
|
||||
val l3Events = payloads.filterIsInstance<L3MemoryRetrievedEvent>()
|
||||
assertTrue(l3Events.isEmpty(), "No L3RetrievedEvent expected on retrieval failure")
|
||||
}
|
||||
|
||||
@Test
|
||||
fun `replay of L3MemoryRetrievedEvent populates lastRetrievedMemory in RouterState`(): Unit = runBlocking {
|
||||
val reducer = DefaultRouterReducer()
|
||||
val projector = RouterProjector(reducer)
|
||||
val sessionId = SessionId("replay-session")
|
||||
val hits = listOf(
|
||||
L3RetrievedHit(
|
||||
entryId = "entry-1",
|
||||
sourceSessionId = SessionId("source-session"),
|
||||
sourceTurnId = "source-turn-1",
|
||||
text = "remembered text",
|
||||
score = 0.95f,
|
||||
)
|
||||
)
|
||||
val event = EventFixtures.stored(
|
||||
payload = L3MemoryRetrievedEvent(
|
||||
sessionId = sessionId,
|
||||
queryTurnId = "query-turn-id",
|
||||
hits = hits,
|
||||
timestampMs = 9999L,
|
||||
)
|
||||
)
|
||||
val state = projector.apply(projector.initial(), event)
|
||||
assertEquals(1, state.lastRetrievedMemory.size)
|
||||
assertEquals("entry-1", state.lastRetrievedMemory[0].entryId)
|
||||
assertEquals("remembered text", state.lastRetrievedMemory[0].text)
|
||||
assertEquals(0.95f, state.lastRetrievedMemory[0].score)
|
||||
}
|
||||
|
||||
private fun emptyContextPack(): ContextPack = ContextPack(
|
||||
id = ContextPackId("empty"),
|
||||
sessionId = SessionId("unknown"),
|
||||
|
||||
Reference in New Issue
Block a user