epic-12: after epic audit and init commit
This commit is contained in:
@@ -0,0 +1,187 @@
|
||||
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.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"),
|
||||
),
|
||||
)
|
||||
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,
|
||||
),
|
||||
)
|
||||
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,
|
||||
),
|
||||
)
|
||||
val newState = reducer.reduce(initialState, event)
|
||||
|
||||
// The record should remain unchanged
|
||||
assertEquals(1, newState.records.size)
|
||||
assertEquals(InferenceStatus.STARTED, newState.records.first().status)
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user