f827685ed0
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.
192 lines
7.1 KiB
Kotlin
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)
|
|
}
|
|
}
|