import com.correx.core.events.events.InferenceCompletedEvent import com.correx.core.events.events.InferenceFailedEvent import com.correx.core.events.events.InferenceStartedEvent import com.correx.core.events.events.InferenceTimeoutEvent import com.correx.core.events.events.ModelLoadedEvent 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.DefaultInferenceReducer import com.correx.core.inference.InferenceRecord import com.correx.core.inference.InferenceReducer import com.correx.core.inference.InferenceState import com.correx.core.inference.InferenceStatus import com.correx.core.inference.TokenUsage import com.correx.core.utils.TypeId import com.correx.testing.fixtures.EventFixtures.stored import org.junit.jupiter.api.Assertions.assertEquals import org.junit.jupiter.api.BeforeEach import org.junit.jupiter.api.Test class InferenceReducerTest { private lateinit var reducer: InferenceReducer @BeforeEach fun setup() { reducer = DefaultInferenceReducer() } @Test fun `should start a new record on InferenceStartedEvent`() { val initialState = InferenceState(emptyList()) val event = stored( payload = InferenceStartedEvent( requestId = InferenceRequestId("req1"), sessionId = SessionId("sess1"), stageId = StageId("stageA"), providerId = ProviderId("providerX"), promptArtifactId = TypeId("00".repeat(32)), ), ) val newState = reducer.reduce(initialState, event) assertEquals(1, newState.records.size) val record = newState.records.first() assertEquals(InferenceStatus.STARTED, record.status) assertEquals("req1", record.requestId.value) } @Test fun `should complete a record on InferenceCompletedEvent`() { val startedRecord = InferenceRecord( requestId = InferenceRequestId("req1"), sessionId = SessionId("sess1"), stageId = StageId("stageA"), providerId = ProviderId("providerX"), status = InferenceStatus.STARTED, ) val initialState = InferenceState(listOf(startedRecord)) val event = stored( payload = InferenceCompletedEvent( requestId = InferenceRequestId("req1"), sessionId = SessionId("sess1"), stageId = StageId("stageA"), providerId = ProviderId("providerX"), tokensUsed = TokenUsage(50, 50), latencyMs = 500L, responseArtifactId = TypeId("00".repeat(32)), ), ) val newState = reducer.reduce(initialState, event) assertEquals(1, newState.records.size) val record = newState.records.first() assertEquals(InferenceStatus.COMPLETED, record.status) assertEquals(100, record.tokensUsed?.totalTokens) assertEquals(500L, record.latencyMs) } @Test fun `should fail a record on InferenceFailedEvent`() { val startedRecord = InferenceRecord( requestId = InferenceRequestId("req1"), sessionId = SessionId("sess1"), stageId = StageId("stageA"), providerId = ProviderId("providerX"), status = InferenceStatus.STARTED, ) val initialState = InferenceState(listOf(startedRecord)) val event = stored( payload = InferenceFailedEvent( requestId = InferenceRequestId("req1"), sessionId = SessionId("sess1"), stageId = StageId("stageA"), providerId = ProviderId("providerX"), reason = "Network failure", ), ) val newState = reducer.reduce(initialState, event) assertEquals(1, newState.records.size) val record = newState.records.first() assertEquals(InferenceStatus.FAILED, record.status) assertEquals("Network failure", record.failureReason) } @Test fun `should timeout a record on InferenceTimeoutEvent`() { val startedRecord = InferenceRecord( requestId = InferenceRequestId("req1"), sessionId = SessionId("sess1"), stageId = StageId("stageA"), providerId = ProviderId("providerX"), status = InferenceStatus.STARTED, ) val initialState = InferenceState(listOf(startedRecord)) val event = stored( payload = InferenceTimeoutEvent( requestId = InferenceRequestId("req1"), sessionId = SessionId("sess1"), stageId = StageId("stageA"), providerId = ProviderId("providerX"), timeoutMs = 10000L, ), ) val newState = reducer.reduce(initialState, event) assertEquals(1, newState.records.size) val record = newState.records.first() assertEquals(InferenceStatus.TIMED_OUT, record.status) assertEquals(10000L, record.latencyMs) } @Test fun `should pass state unchanged on unrelated events`() { val startedRecord = InferenceRecord( requestId = InferenceRequestId("req1"), sessionId = SessionId("sess1"), stageId = StageId("stageA"), providerId = ProviderId("providerX"), status = InferenceStatus.STARTED, ) val initialState = InferenceState(listOf(startedRecord)) val event = stored( payload = ModelLoadedEvent( sessionId = SessionId("sess1"), providerId = ProviderId("providerX"), modelId = "model", ), ) val newState = reducer.reduce(initialState, event) assertEquals(1, newState.records.size) assertEquals(InferenceStatus.STARTED, newState.records.first().status) } @Test fun `should not modify record if requestId does not match`() { val startedRecord = InferenceRecord( requestId = InferenceRequestId("req1"), sessionId = SessionId("sess1"), stageId = StageId("stageA"), providerId = ProviderId("providerX"), status = InferenceStatus.STARTED, ) val initialState = InferenceState(listOf(startedRecord)) // Event for a different request val event = stored( payload = InferenceCompletedEvent( requestId = InferenceRequestId("req2"), sessionId = SessionId("sess1"), stageId = StageId("stageA"), providerId = ProviderId("providerX"), tokensUsed = TokenUsage(50, 50), latencyMs = 500L, responseArtifactId = TypeId("00".repeat(32)), ), ) val newState = reducer.reduce(initialState, event) // The record should remain unchanged assertEquals(1, newState.records.size) assertEquals(InferenceStatus.STARTED, newState.records.first().status) } }