Files
correx/testing/projections/src/test/kotlin/InferenceReducerTest.kt
T
kami f827685ed0 chore(damn): detekt, build, tests, formatting
fixed detekt issues where possible.
fixed disttar failing build because tools is added twice in the server module.
added workflowId where required.
fixed some tests not being recognized because of runBlocking without explicit return type.
formatting + imports.
2026-05-22 00:10:05 +04:00

192 lines
7.1 KiB
Kotlin

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)
}
}