Files
correx/testing/deterministic/src/test/kotlin/RouterFacadeTest.kt
T

1402 lines
66 KiB
Kotlin

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
import com.correx.core.events.events.NewEvent
import com.correx.core.events.events.StoredEvent
import com.correx.core.events.stores.EventStore
import com.correx.core.events.types.ContextPackId
import com.correx.core.events.types.EventId
import com.correx.core.events.types.InferenceRequestId
import com.correx.core.events.types.ProviderId
import com.correx.core.events.types.SessionId
import com.correx.core.events.types.StageId
import com.correx.core.inference.CapabilityScore
import com.correx.core.inference.Embedder
import com.correx.core.inference.FinishReason
import com.correx.core.inference.InferenceProvider
import com.correx.core.inference.InferenceRequest
import com.correx.core.inference.InferenceResponse
import com.correx.core.inference.InferenceRouter
import com.correx.core.inference.ModelCapability
import com.correx.core.inference.NoopEmbedder
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
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
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.WorkflowSummary
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
import kotlinx.datetime.Instant
import kotlinx.serialization.decodeFromString
import org.junit.jupiter.api.Assertions.assertEquals
import org.junit.jupiter.api.Assertions.assertFalse
import org.junit.jupiter.api.Assertions.assertNotNull
import org.junit.jupiter.api.Assertions.assertNull
import org.junit.jupiter.api.Assertions.assertTrue
import org.junit.jupiter.api.Test
class RouterFacadeTest {
// --------------------------------------------------------------------------
// CHAT mode
// --------------------------------------------------------------------------
@Test
fun `CHAT mode returns inference response content`(): Unit = runBlocking {
val mockStore = mockEventStore()
val facade = facadeWithMocks(eventStore = mockStore, chatMode = ChatMode.CHAT)
val response = facade.onUserInput(sessionId = SessionId("test-session"), input = "Hello, world!")
assertEquals("inference response", response.content)
}
@Test
fun `CHAT mode sets steeringEmitted to false`(): Unit = runBlocking {
val mockStore = mockEventStore()
val facade = facadeWithMocks(eventStore = mockStore, chatMode = ChatMode.CHAT)
val response = facade.onUserInput(sessionId = SessionId("test-session"), input = "Hello!")
assertFalse(response.steeringEmitted)
}
@Test
fun `CHAT mode appends USER and ROUTER ChatTurnEvent to EventStore`(): Unit = runBlocking {
val mockStore = mockEventStore()
val facade = facadeWithMocks(eventStore = mockStore, chatMode = ChatMode.CHAT)
facade.onUserInput(sessionId = SessionId("test-session"), input = "Hello!")
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.content)
assertEquals(com.correx.core.events.events.ChatTurnRole.ROUTER, routerEvent.role)
}
// --------------------------------------------------------------------------
// STEERING mode
// --------------------------------------------------------------------------
@Test
fun `STEERING mode sets steeringEmitted to true`(): Unit = runBlocking {
val mockStore = mockEventStore()
val facade = facadeWithMocks(eventStore = mockStore, chatMode = ChatMode.STEERING)
val response = facade.onUserInput(sessionId = SessionId("test-session"), input = "Hello!")
assertTrue(response.steeringEmitted)
}
@Test
fun `STEERING mode appends USER ChatTurnEvent, ROUTER ChatTurnEvent, and SteeringNoteAddedEvent to store`(): Unit = runBlocking {
val mockStore = mockEventStore()
val facade = facadeWithMocks(eventStore = mockStore, chatMode = ChatMode.STEERING)
facade.onUserInput(sessionId = SessionId("session-xyz"), input = "steer this way")
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
fun `STEERING mode appends SteeringNoteAddedEvent with current stageId`(): Unit = runBlocking {
val mockStore = mockEventStore()
val stageId = StageId("stage-A")
val facade = DefaultRouterFacade(
routerRepository = object : RouterRepository {
override suspend fun getRouterState(sessionId: SessionId): RouterState =
RouterState(
sessionId = sessionId,
workflowStatus = WorkflowStatus.RUNNING,
currentStageId = stageId,
)
},
routerContextBuilder = object : RouterContextBuilder {
override suspend fun build(state: RouterState, budget: TokenBudget, availableWorkflows: List<WorkflowSummary>, projectProfileText: String?): ContextPack = emptyContextPack()
},
inferenceRouter = mockInferenceRouter("steering response"),
eventStore = mockStore,
config = RouterConfig(tokenBudget = TokenBudget(limit = 5000)),
embedder = NoopEmbedder(dimension = 8),
l3MemoryStore = InMemoryL3MemoryStore(),
)
facade.onUserInput(sessionId = SessionId("test-session"), input = "Hello!", mode = ChatMode.STEERING)
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)
}
// --------------------------------------------------------------------------
// In-memory conversation history
// --------------------------------------------------------------------------
@Test
fun `conversation history grows per call - user and router turns appended`(): Unit = runBlocking {
val capturedStates = mutableListOf<RouterState>()
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, availableWorkflows: List<WorkflowSummary>, projectProfileText: String?): ContextPack {
capturedStates.add(state)
return emptyContextPack()
}
},
inferenceRouter = mockInferenceRouter("router reply"),
eventStore = eventStore,
config = RouterConfig(tokenBudget = TokenBudget(limit = 5000)),
embedder = NoopEmbedder(dimension = 8),
l3MemoryStore = InMemoryL3MemoryStore(),
)
val sessionId = SessionId("history-session")
facade.onUserInput(sessionId = sessionId, input = "first message")
facade.onUserInput(sessionId = sessionId, input = "second message")
// Second call's context builder sees two user turns + one router turn from the first call
val stateOnSecondCall = capturedStates[1]
assertEquals(3, stateOnSecondCall.conversationHistory.size)
assertEquals(TurnRole.USER, stateOnSecondCall.conversationHistory[0].role)
assertEquals("first message", stateOnSecondCall.conversationHistory[0].content)
assertEquals(TurnRole.ROUTER, stateOnSecondCall.conversationHistory[1].role)
assertEquals("router reply", stateOnSecondCall.conversationHistory[1].content)
assertEquals(TurnRole.USER, stateOnSecondCall.conversationHistory[2].role)
assertEquals("second message", stateOnSecondCall.conversationHistory[2].content)
}
@Test
fun `conversation history is session-scoped - different sessions do not share history`(): Unit = runBlocking {
val capturedStates = mutableListOf<RouterState>()
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, availableWorkflows: List<WorkflowSummary>, projectProfileText: String?): ContextPack {
capturedStates.add(state)
return emptyContextPack()
}
},
inferenceRouter = mockInferenceRouter("response"),
eventStore = eventStore,
config = RouterConfig(tokenBudget = TokenBudget(limit = 5000)),
embedder = NoopEmbedder(dimension = 8),
l3MemoryStore = InMemoryL3MemoryStore(),
)
facade.onUserInput(sessionId = SessionId("session-A"), input = "message A")
facade.onUserInput(sessionId = SessionId("session-B"), input = "message B")
val stateA = capturedStates[0]
val stateB = capturedStates[1]
assertEquals(1, stateA.conversationHistory.size)
assertEquals(1, stateB.conversationHistory.size)
assertEquals("message A", stateA.conversationHistory[0].content)
assertEquals("message B", stateB.conversationHistory[0].content)
}
// --------------------------------------------------------------------------
// Orchestration
// --------------------------------------------------------------------------
@Test
fun `state is passed through to context builder`(): Unit = runBlocking {
val capturedState = mutableListOf<RouterState>()
val mockContextBuilder = object : RouterContextBuilder {
override suspend fun build(state: RouterState, budget: TokenBudget, availableWorkflows: List<WorkflowSummary>, projectProfileText: String?): ContextPack {
capturedState.add(state)
return emptyContextPack()
}
}
val facade = DefaultRouterFacade(
routerRepository = object : RouterRepository {
override suspend fun getRouterState(sessionId: SessionId): RouterState =
RouterState(
sessionId = sessionId,
workflowStatus = WorkflowStatus.RUNNING,
currentStageId = StageId("s1"),
)
},
routerContextBuilder = mockContextBuilder,
inferenceRouter = mockInferenceRouter("inference response"),
eventStore = mockEventStore(),
config = RouterConfig(tokenBudget = TokenBudget(limit = 5000)),
embedder = NoopEmbedder(dimension = 8),
l3MemoryStore = InMemoryL3MemoryStore(),
)
facade.onUserInput(sessionId = SessionId("test-session"), input = "Hello!")
assertEquals(1, capturedState.size)
assertEquals(SessionId("test-session"), capturedState[0].sessionId)
assertEquals(WorkflowStatus.RUNNING, capturedState[0].workflowStatus)
}
@Test
fun `budget is passed through to context builder`(): Unit = runBlocking {
val capturedBudget = mutableListOf<TokenBudget>()
val mockContextBuilder = object : RouterContextBuilder {
override suspend fun build(state: RouterState, budget: TokenBudget, availableWorkflows: List<WorkflowSummary>, projectProfileText: String?): ContextPack {
capturedBudget.add(budget)
return emptyContextPack()
}
}
val facade = DefaultRouterFacade(
routerRepository = object : RouterRepository {
override suspend fun getRouterState(sessionId: SessionId): RouterState = RouterState()
},
routerContextBuilder = mockContextBuilder,
inferenceRouter = mockInferenceRouter("response"),
eventStore = mockEventStore(),
config = RouterConfig(tokenBudget = TokenBudget(limit = 4200)),
embedder = NoopEmbedder(dimension = 8),
l3MemoryStore = InMemoryL3MemoryStore(),
)
facade.onUserInput(sessionId = SessionId("test-session"), input = "Hello!")
assertEquals(1, capturedBudget.size)
assertEquals(4200, capturedBudget[0].limit)
}
@Test
fun `stage ID uses state currentStageId`(): Unit = runBlocking {
val capturedStageId = mutableListOf<StageId>()
val mockInferenceRouter = object : InferenceRouter {
override suspend fun route(
stageId: StageId,
requiredCapabilities: Set<ModelCapability>,
): InferenceProvider {
capturedStageId.add(stageId)
return mockProvider("response")
}
}
val stateStageId = StageId("state-stage")
val facade = DefaultRouterFacade(
routerRepository = object : RouterRepository {
override suspend fun getRouterState(sessionId: SessionId): RouterState = RouterState(
sessionId = sessionId,
workflowStatus = WorkflowStatus.RUNNING,
currentStageId = stateStageId,
)
},
routerContextBuilder = object : RouterContextBuilder {
override suspend fun build(state: RouterState, budget: TokenBudget, availableWorkflows: List<WorkflowSummary>, projectProfileText: String?): ContextPack = emptyContextPack()
},
inferenceRouter = mockInferenceRouter,
eventStore = mockEventStore(),
config = RouterConfig(tokenBudget = TokenBudget(limit = 5000)),
embedder = NoopEmbedder(dimension = 8),
l3MemoryStore = InMemoryL3MemoryStore(),
)
facade.onUserInput(sessionId = SessionId("test-session"), input = "Hello!")
assertEquals(stateStageId, capturedStageId[0])
}
@Test
fun `stage ID falls back to StageId none when state has no currentStageId`(): Unit = runBlocking {
val capturedStageId = mutableListOf<StageId>()
val mockInferenceRouter = object : InferenceRouter {
override suspend fun route(
stageId: StageId,
requiredCapabilities: Set<ModelCapability>,
): InferenceProvider {
capturedStageId.add(stageId)
return mockProvider("response")
}
}
val facade = DefaultRouterFacade(
routerRepository = object : RouterRepository {
override suspend fun getRouterState(sessionId: SessionId): RouterState = RouterState(
sessionId = sessionId,
workflowStatus = WorkflowStatus.IDLE,
currentStageId = null,
)
},
routerContextBuilder = object : RouterContextBuilder {
override suspend fun build(state: RouterState, budget: TokenBudget, availableWorkflows: List<WorkflowSummary>, projectProfileText: String?): ContextPack = emptyContextPack()
},
inferenceRouter = mockInferenceRouter,
eventStore = mockEventStore(),
config = RouterConfig(tokenBudget = TokenBudget(limit = 5000)),
embedder = NoopEmbedder(dimension = 8),
l3MemoryStore = InMemoryL3MemoryStore(),
)
facade.onUserInput(sessionId = SessionId("test-session"), input = "Hello!")
assertEquals(StageId("none"), capturedStageId[0])
}
@Test
fun `new InferenceRequestId per call`(): Unit = runBlocking {
val capturedRequestIds = mutableListOf<InferenceRequestId>()
val mockInferenceRouter = object : InferenceRouter {
override suspend fun route(
stageId: StageId,
requiredCapabilities: Set<ModelCapability>,
): InferenceProvider {
return mockProviderWithCapture("response", capturedRequestIds)
}
}
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, availableWorkflows: List<WorkflowSummary>, projectProfileText: String?): ContextPack = emptyContextPack()
},
inferenceRouter = mockInferenceRouter,
eventStore = mockEventStore(),
config = RouterConfig(tokenBudget = TokenBudget(limit = 5000)),
embedder = NoopEmbedder(dimension = 8),
l3MemoryStore = InMemoryL3MemoryStore(),
)
facade.onUserInput(sessionId = SessionId("s1"), input = "first")
facade.onUserInput(sessionId = SessionId("s1"), input = "second")
assertEquals(2, capturedRequestIds.size)
assertFalse(capturedRequestIds[0] == capturedRequestIds[1])
}
@Test
fun `GenerationConfig defaults temperature 0 7 topP 0 9 maxTokens 512`(): Unit = runBlocking {
val capturedRequests = mutableListOf<InferenceRequest>()
val mockInferenceRouter = object : InferenceRouter {
override suspend fun route(
stageId: StageId,
requiredCapabilities: Set<ModelCapability>,
): InferenceProvider {
return mockProviderWithRequestCapture("response", capturedRequests)
}
}
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, availableWorkflows: List<WorkflowSummary>, projectProfileText: String?): ContextPack = emptyContextPack()
},
inferenceRouter = mockInferenceRouter,
eventStore = mockEventStore(),
config = RouterConfig(tokenBudget = TokenBudget(limit = 5000)),
embedder = NoopEmbedder(dimension = 8),
l3MemoryStore = InMemoryL3MemoryStore(),
)
facade.onUserInput(sessionId = SessionId("test-session"), input = "Hello!")
val req = capturedRequests[0]
assertEquals(0.7, req.generationConfig.temperature)
assertEquals(0.9, req.generationConfig.topP)
assertEquals(512, req.generationConfig.maxTokens)
}
@Test
fun `context pack is passed to inference provider`(): Unit = runBlocking {
val capturedContextPacks = mutableListOf<ContextPack>()
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, availableWorkflows: List<WorkflowSummary>, projectProfileText: String?): ContextPack {
val pack = ContextPack(
id = ContextPackId("test-pack"),
sessionId = state.sessionId ?: SessionId("unknown"),
stageId = StageId("test"),
layers = emptyMap(),
budgetUsed = 0,
budgetLimit = budget.limit,
)
capturedContextPacks.add(pack)
return pack
}
},
inferenceRouter = object : InferenceRouter {
override suspend fun route(
stageId: StageId,
requiredCapabilities: Set<ModelCapability>,
): InferenceProvider {
return object : InferenceProvider {
override val id = ProviderId("mock")
override val name = "Mock"
override val tokenizer = MockTokenizer()
override suspend fun infer(request: InferenceRequest): InferenceResponse {
assertEquals(1, capturedContextPacks.size)
assertEquals(capturedContextPacks[0].id, request.contextPack.id)
return InferenceResponse(
requestId = request.requestId,
text = "response",
finishReason = FinishReason.Stop,
tokensUsed = TokenUsage(promptTokens = 10, completionTokens = 5),
latencyMs = 0,
)
}
override suspend fun healthCheck(): ProviderHealth =
ProviderHealth.Healthy
override fun capabilities(): Set<CapabilityScore> = emptySet()
}
}
},
eventStore = mockEventStore(),
config = RouterConfig(tokenBudget = TokenBudget(limit = 5000)),
embedder = NoopEmbedder(dimension = 8),
l3MemoryStore = InMemoryL3MemoryStore(),
)
facade.onUserInput(sessionId = SessionId("test-session"), input = "Hello!")
}
@Test
fun `responseFormat defaults to Text`(): Unit = runBlocking {
val capturedRequests = mutableListOf<InferenceRequest>()
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, availableWorkflows: List<WorkflowSummary>, projectProfileText: String?): ContextPack = emptyContextPack()
},
inferenceRouter = object : InferenceRouter {
override suspend fun route(
stageId: StageId,
requiredCapabilities: Set<ModelCapability>,
): InferenceProvider {
return mockProviderWithRequestCapture("response", capturedRequests)
}
},
eventStore = mockEventStore(),
config = RouterConfig(tokenBudget = TokenBudget(limit = 5000)),
embedder = NoopEmbedder(dimension = 8),
l3MemoryStore = InMemoryL3MemoryStore(),
)
facade.onUserInput(sessionId = SessionId("test-session"), input = "Hello!")
val req = capturedRequests[0]
assertTrue(req.responseFormat is ResponseFormat.Text)
}
@Test
fun `onUserInput returns RouterResponse with content and mode flag`(): Unit = runBlocking {
val mockStore = mockEventStore()
val facade = facadeWithMocks(eventStore = mockStore, chatMode = ChatMode.STEERING)
val response = facade.onUserInput(sessionId = SessionId("test-session"), input = "Hello!")
assertNotNull(response)
assertEquals("inference response", response.content)
assertTrue(response.steeringEmitted)
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
fun `CHAT round-trip stores user and router turns in L3 memory`(): Unit = runBlocking {
val mockStore = mockEventStore()
val l3Store = InMemoryL3MemoryStore()
val embedder = NoopEmbedder(dimension = 8)
val facade = facadeWithMocks(
eventStore = mockStore,
chatMode = ChatMode.CHAT,
embedder = embedder,
l3MemoryStore = l3Store,
)
val sessionId = SessionId("test-session")
facade.onUserInput(sessionId = sessionId, input = "Hello!")
val hits = l3Store.query(
L3Query(
vector = FloatArray(8),
k = 10,
sessionIdFilter = sessionId,
)
)
assertEquals(2, hits.size)
}
// --------------------------------------------------------------------------
// Task 2.1 — latency + token metrics on ROUTER ChatTurnEvent
// --------------------------------------------------------------------------
@Test
fun `ROUTER ChatTurnEvent carries latencyMs and tokensUsed from inference and USER turn has both null`(): Unit = runBlocking {
val capturedRequests = mutableListOf<InferenceRequest>()
val knownLatencyMs = 42L
val knownTokenUsage = TokenUsage(promptTokens = 10, completionTokens = 20)
val metricsProvider = object : InferenceProvider {
override val id = ProviderId("mock")
override val name = "Mock"
override val tokenizer = MockTokenizer()
override suspend fun infer(request: InferenceRequest): InferenceResponse {
capturedRequests.add(request)
return InferenceResponse(
requestId = request.requestId,
text = "metrics response",
finishReason = FinishReason.Stop,
tokensUsed = knownTokenUsage,
latencyMs = knownLatencyMs,
)
}
override suspend fun healthCheck(): ProviderHealth = ProviderHealth.Healthy
override fun capabilities(): Set<CapabilityScore> = setOf(CapabilityScore(ModelCapability.General, 1.0))
}
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, availableWorkflows: List<WorkflowSummary>, projectProfileText: String?): ContextPack = emptyContextPack()
},
inferenceRouter = object : InferenceRouter {
override suspend fun route(stageId: StageId, requiredCapabilities: Set<ModelCapability>): InferenceProvider =
metricsProvider
},
eventStore = eventStore,
config = RouterConfig(tokenBudget = TokenBudget(limit = 5000)),
embedder = NoopEmbedder(dimension = 8),
l3MemoryStore = InMemoryL3MemoryStore(),
)
facade.onUserInput(sessionId = SessionId("metrics-session"), input = "hello metrics")
val chatEvents = eventStore.appendedEvents.map { it.payload }.filterIsInstance<ChatTurnEvent>()
assertEquals(2, chatEvents.size)
val userEvent = chatEvents.first { it.role == ChatTurnRole.USER }
val routerEvent = chatEvents.first { it.role == ChatTurnRole.ROUTER }
assertNull(userEvent.latencyMs, "USER ChatTurnEvent must have null latencyMs")
assertNull(userEvent.tokensUsed, "USER ChatTurnEvent must have null tokensUsed")
assertEquals(knownLatencyMs, routerEvent.latencyMs)
assertNotNull(routerEvent.tokensUsed)
assertEquals(30, routerEvent.tokensUsed!!.totalTokens)
}
@Test
fun `ChatTurnEvent backward-compat - legacy JSON without latencyMs and tokensUsed deserializes with both null`() {
val legacyJson = """{"sessionId":"s1","turnId":"t1","role":"USER","content":"hello","timestampMs":1000}"""
val event = com.correx.core.events.serialization.eventJson.decodeFromString<ChatTurnEvent>(legacyJson)
assertEquals(SessionId("s1"), event.sessionId)
assertEquals("hello", event.content)
assertNull(event.latencyMs)
assertNull(event.tokensUsed)
}
// --------------------------------------------------------------------------
// Helpers
// --------------------------------------------------------------------------
private fun mockEventStore(): MapBackedEventStore = MapBackedEventStore()
private fun facadeWithMocks(
eventStore: EventStore,
chatMode: ChatMode = ChatMode.CHAT,
embedder: Embedder = NoopEmbedder(dimension = 8),
l3MemoryStore: L3MemoryStore = InMemoryL3MemoryStore(),
): RouterFacade = DefaultRouterFacade(
routerRepository = object : RouterRepository {
override suspend fun getRouterState(sessionId: SessionId): RouterState =
RouterState(
sessionId = sessionId,
workflowStatus = WorkflowStatus.RUNNING,
currentStageId = StageId("s1"),
)
},
routerContextBuilder = object : RouterContextBuilder {
override suspend fun build(state: RouterState, budget: TokenBudget, availableWorkflows: List<WorkflowSummary>, projectProfileText: String?): ContextPack = emptyContextPack()
},
inferenceRouter = mockInferenceRouter("inference response"),
eventStore = eventStore,
config = RouterConfig(tokenBudget = TokenBudget(limit = 5000)),
embedder = embedder,
l3MemoryStore = l3MemoryStore,
).let { impl ->
object : RouterFacade {
override suspend fun onUserInput(sessionId: SessionId, input: String, mode: ChatMode): RouterResponse =
impl.onUserInput(sessionId = sessionId, input = input, mode = chatMode)
override suspend fun narrate(sessionId: SessionId, trigger: com.correx.core.router.model.NarrationTrigger) =
impl.narrate(sessionId = sessionId, trigger = trigger)
}
}
private fun mockInferenceRouter(responseText: String): InferenceRouter =
object : InferenceRouter {
override suspend fun route(
stageId: StageId,
requiredCapabilities: Set<ModelCapability>,
): InferenceProvider =
mockProvider(responseText)
}
private fun mockProvider(responseText: String): InferenceProvider =
object : InferenceProvider {
override val id = ProviderId("mock")
override val name = "Mock"
override val tokenizer = MockTokenizer()
override suspend fun infer(request: InferenceRequest): InferenceResponse = InferenceResponse(
requestId = request.requestId,
text = responseText,
finishReason = FinishReason.Stop,
tokensUsed = TokenUsage(promptTokens = 10, completionTokens = 5),
latencyMs = 0,
)
override suspend fun healthCheck(): ProviderHealth =
ProviderHealth.Healthy
override fun capabilities(): Set<CapabilityScore> =
setOf(CapabilityScore(ModelCapability.General, 1.0))
}
private fun mockProviderWithCapture(
responseText: String,
requestIds: MutableList<InferenceRequestId>,
): InferenceProvider =
object : InferenceProvider {
override val id = ProviderId("mock")
override val name = "Mock"
override val tokenizer = MockTokenizer()
override suspend fun infer(request: InferenceRequest): InferenceResponse {
requestIds.add(request.requestId)
return InferenceResponse(
requestId = request.requestId,
text = responseText,
finishReason = FinishReason.Stop,
tokensUsed = TokenUsage(promptTokens = 10, completionTokens = 5),
latencyMs = 0,
)
}
override suspend fun healthCheck(): ProviderHealth =
ProviderHealth.Healthy
override fun capabilities(): Set<CapabilityScore> = emptySet()
}
private fun mockProviderWithRequestCapture(
responseText: String,
requests: MutableList<InferenceRequest>,
): InferenceProvider =
object : InferenceProvider {
override val id = ProviderId("mock")
override val name = "Mock"
override val tokenizer = MockTokenizer()
override suspend fun infer(request: InferenceRequest): InferenceResponse {
requests.add(request)
return InferenceResponse(
requestId = request.requestId,
text = responseText,
finishReason = FinishReason.Stop,
tokensUsed = TokenUsage(promptTokens = 10, completionTokens = 5),
latencyMs = 0,
)
}
override suspend fun healthCheck(): ProviderHealth =
ProviderHealth.Healthy
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, availableWorkflows: List<WorkflowSummary>, projectProfileText: String?): 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, availableWorkflows: List<WorkflowSummary>, projectProfileText: String?): 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 → 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
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, availableWorkflows: List<WorkflowSummary>, projectProfileText: String?): 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")
// Retrieval failed → deduped is empty; L3MemoryRetrievedEvent still emitted with empty hits
val l3Events = payloads.filterIsInstance<L3MemoryRetrievedEvent>()
assertEquals(1, l3Events.size, "L3MemoryRetrievedEvent must be emitted even on retrieval failure")
assertTrue(l3Events[0].hits.isEmpty(), "Expected empty hits when retrieval failed")
}
@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)
}
// --------------------------------------------------------------------------
// 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, availableWorkflows: List<WorkflowSummary>, projectProfileText: String?): 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, availableWorkflows: List<WorkflowSummary>, projectProfileText: String?): 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]") })
}
// --------------------------------------------------------------------------
// Narration model ID routing
// --------------------------------------------------------------------------
@Test
fun `narrate passes narrationModelId to inferenceRouter 3-arg route`(): Unit = runBlocking {
val capturedModelId = mutableListOf<String?>()
val mockInferenceRouter = object : InferenceRouter {
override suspend fun route(
stageId: StageId,
requiredCapabilities: Set<ModelCapability>,
): InferenceProvider = mockProvider("narration text")
override suspend fun route(
stageId: StageId,
requiredCapabilities: Set<ModelCapability>,
modelId: String?,
): InferenceProvider {
capturedModelId.add(modelId)
return mockProvider("narration text")
}
}
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, availableWorkflows: List<WorkflowSummary>, projectProfileText: String?): ContextPack = emptyContextPack()
},
inferenceRouter = mockInferenceRouter,
eventStore = mockEventStore(),
config = RouterConfig(
tokenBudget = TokenBudget(limit = 5000),
narrationModelId = "llama-cpp:phi-3-mini",
),
embedder = NoopEmbedder(dimension = 8),
l3MemoryStore = InMemoryL3MemoryStore(),
)
facade.narrate(
sessionId = SessionId("test-session"),
trigger = com.correx.core.router.model.NarrationTrigger(
kind = "test",
instruction = "describe what happened",
),
)
assertEquals(1, capturedModelId.size)
assertEquals("llama-cpp:phi-3-mini", capturedModelId[0])
}
@Test
fun `onUserInput does NOT pass modelId — uses 2-arg route`(): Unit = runBlocking {
var threeArgCallCount = 0
val mockInferenceRouter = object : InferenceRouter {
override suspend fun route(
stageId: StageId,
requiredCapabilities: Set<ModelCapability>,
): InferenceProvider = mockProvider("response")
override suspend fun route(
stageId: StageId,
requiredCapabilities: Set<ModelCapability>,
modelId: String?,
): InferenceProvider {
threeArgCallCount++
return mockProvider("response")
}
}
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, availableWorkflows: List<WorkflowSummary>, projectProfileText: String?): ContextPack = emptyContextPack()
},
inferenceRouter = mockInferenceRouter,
eventStore = mockEventStore(),
config = RouterConfig(
tokenBudget = TokenBudget(limit = 5000),
narrationModelId = "llama-cpp:phi-3-mini",
),
embedder = NoopEmbedder(dimension = 8),
l3MemoryStore = InMemoryL3MemoryStore(),
)
facade.onUserInput(sessionId = SessionId("test-session"), input = "Hello!")
assertEquals(0, threeArgCallCount, "onUserInput must not call 3-arg route()")
}
// --------------------------------------------------------------------------
// A4 — workflowSummaryProvider and sessionProfileProvider forwarding
// --------------------------------------------------------------------------
@Test
fun `workflowSummaryProvider is called and forwarded to context builder on each turn`(): Unit = runBlocking {
val captured = mutableListOf<List<WorkflowSummary>>()
val capturingBuilder = object : RouterContextBuilder {
override suspend fun build(
state: RouterState,
budget: TokenBudget,
availableWorkflows: List<WorkflowSummary>,
projectProfileText: String?,
): ContextPack {
captured.add(availableWorkflows)
return emptyContextPack()
}
}
val workflow = WorkflowSummary("wf", "desc", listOf("s1"))
val facade = facadeWith(contextBuilder = capturingBuilder, workflowSummaryProvider = { listOf(workflow) })
facade.onUserInput(SessionId("s"), "hello")
assertEquals(listOf(listOf(workflow)), captured)
}
@Test
fun `sessionProfileProvider is called and forwarded to context builder`(): Unit = runBlocking {
val capturedProfiles = mutableListOf<String?>()
val capturingBuilder = object : RouterContextBuilder {
override suspend fun build(
state: RouterState,
budget: TokenBudget,
availableWorkflows: List<WorkflowSummary>,
projectProfileText: String?,
): ContextPack {
capturedProfiles.add(projectProfileText)
return emptyContextPack()
}
}
val facade = facadeWith(
contextBuilder = capturingBuilder,
sessionProfileProvider = { "profile text" },
)
facade.onUserInput(SessionId("s"), "hello")
assertEquals(listOf("profile text"), capturedProfiles)
}
private fun facadeWith(
contextBuilder: RouterContextBuilder = object : RouterContextBuilder {
override suspend fun build(
state: RouterState,
budget: TokenBudget,
availableWorkflows: List<WorkflowSummary>,
projectProfileText: String?,
): ContextPack = emptyContextPack()
},
workflowSummaryProvider: () -> List<WorkflowSummary> = { emptyList() },
sessionProfileProvider: suspend (SessionId) -> String? = { null },
): RouterFacade = DefaultRouterFacade(
routerRepository = object : RouterRepository {
override suspend fun getRouterState(sessionId: SessionId): RouterState = RouterState(sessionId = sessionId)
},
routerContextBuilder = contextBuilder,
inferenceRouter = mockInferenceRouter("inference response"),
eventStore = mockEventStore(),
config = RouterConfig(tokenBudget = TokenBudget(limit = 5000)),
embedder = NoopEmbedder(dimension = 8),
l3MemoryStore = InMemoryL3MemoryStore(),
workflowSummaryProvider = workflowSummaryProvider,
sessionProfileProvider = sessionProfileProvider,
)
private fun emptyContextPack(): ContextPack = ContextPack(
id = ContextPackId("empty"),
sessionId = SessionId("unknown"),
stageId = StageId("none"),
layers = emptyMap(),
budgetUsed = 0,
budgetLimit = 5000,
)
private class MapBackedEventStore : EventStore {
val appendedEvents = mutableListOf<NewEvent>()
private val storedEvents: MutableMap<EventId, StoredEvent> = mutableMapOf()
private var nextSequence = 1L
override suspend fun append(event: NewEvent): StoredEvent {
appendedEvents.add(event)
val stored = StoredEvent(
metadata = event.metadata,
sequence = nextSequence,
sessionSequence = nextSequence++,
payload = event.payload,
)
storedEvents[event.metadata.eventId] = stored
return stored
}
override suspend fun appendAll(events: List<NewEvent>): List<StoredEvent> =
events.map { append(it) }
override fun read(sessionId: SessionId): List<StoredEvent> =
storedEvents.values.filter { it.metadata.sessionId == sessionId }.toList()
override fun readFrom(
sessionId: SessionId,
fromSequence: Long,
): List<StoredEvent> =
read(sessionId).filter { it.sequence >= fromSequence }
override fun lastSequence(sessionId: SessionId): Long? =
read(sessionId).maxOfOrNull { it.sequence }
override fun subscribe(sessionId: SessionId): Flow<StoredEvent> =
throw UnsupportedOperationException("subscribe not implemented for mock")
override fun allEvents(): Sequence<StoredEvent> =
storedEvents.values.asSequence()
override fun allSessionIds(): Set<SessionId> = storedEvents.values.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")
}
}