diff --git a/apps/server/build.gradle b/apps/server/build.gradle index bfac1f8e..cce43481 100644 --- a/apps/server/build.gradle +++ b/apps/server/build.gradle @@ -54,6 +54,9 @@ dependencies { implementation "io.ktor:ktor-server-call-logging:$ktor_version" testImplementation "org.jetbrains.kotlinx:kotlinx-coroutines-test:$kotlinx_coroutines_version" + testImplementation "io.ktor:ktor-server-test-host:$ktor_version" + testImplementation "io.ktor:ktor-client-websockets:$ktor_version" + testImplementation "io.ktor:ktor-client-mock:$ktor_version" implementation "org.apache.logging.log4j:log4j-core:2.24.1" implementation "org.apache.logging.log4j:log4j-slf4j2-impl:2.24.1" diff --git a/apps/server/src/test/kotlin/com/correx/apps/server/lifecycle/LifecycleTestSupport.kt b/apps/server/src/test/kotlin/com/correx/apps/server/lifecycle/LifecycleTestSupport.kt new file mode 100644 index 00000000..8f9886e6 --- /dev/null +++ b/apps/server/src/test/kotlin/com/correx/apps/server/lifecycle/LifecycleTestSupport.kt @@ -0,0 +1,211 @@ +package com.correx.apps.server.lifecycle + +import com.correx.apps.server.ServerModule +import com.correx.apps.server.registry.ProviderRegistry +import com.correx.apps.server.registry.WorkflowRegistry +import com.correx.apps.server.registry.WorkflowSummary +import com.correx.apps.server.undo.SessionUndoService +import com.correx.core.approvals.domain.DefaultApprovalEngine +import com.correx.core.artifactstore.ArtifactStore +import com.correx.core.context.builder.DefaultContextPackBuilder +import com.correx.core.context.compression.DefaultContextCompressor +import com.correx.core.events.EventDispatcher +import com.correx.core.events.stores.EventStore +import com.correx.core.events.types.ArtifactId +import com.correx.core.inference.InferenceProjector +import com.correx.core.inference.InferenceProvider +import com.correx.core.inference.InferenceRepository +import com.correx.core.kernel.orchestration.DefaultOrchestrationReducer +import com.correx.core.kernel.orchestration.DefaultSessionOrchestrator +import com.correx.core.kernel.orchestration.OrchestrationConfig +import com.correx.core.kernel.orchestration.OrchestrationProjector +import com.correx.core.kernel.orchestration.OrchestrationRepository +import com.correx.core.kernel.orchestration.OrchestratorEngines +import com.correx.core.kernel.orchestration.OrchestratorRepositories +import com.correx.core.kernel.retry.DefaultRetryCoordinator +import com.correx.core.risk.DefaultRiskAssessor +import com.correx.core.router.model.RouterConfig +import com.correx.core.sessions.DefaultSessionReducer +import com.correx.core.sessions.DefaultSessionRepository +import com.correx.core.sessions.SessionProjector +import com.correx.core.sessions.projections.replay.DefaultEventReplayer +import com.correx.core.transitions.graph.WorkflowGraph +import com.correx.core.transitions.resolution.DefaultTransitionResolver +import com.correx.core.validation.pipeline.ValidationPipeline +import com.correx.infrastructure.InfrastructureModule +import com.correx.infrastructure.inference.DefaultProviderRegistry +import com.correx.infrastructure.inference.commons.ManagedInferenceRouter +import com.correx.infrastructure.inference.commons.ResourceProbe +import com.correx.infrastructure.inference.commons.UnavailableProbe +import com.correx.infrastructure.inference.llama.cpp.DefaultModelManager +import com.correx.infrastructure.persistence.InMemoryEventStore +import com.correx.infrastructure.tools.FileEditConfig +import com.correx.infrastructure.tools.FileReadConfig +import com.correx.infrastructure.tools.FileWriteConfig +import com.correx.infrastructure.tools.ShellConfig +import com.correx.infrastructure.tools.ToolConfig +import com.correx.infrastructure.tools.filesystem.FileMutationReverser +import io.ktor.client.HttpClient +import io.ktor.client.engine.mock.MockEngine +import io.ktor.client.engine.mock.respond +import io.ktor.http.ContentType +import io.ktor.http.HttpHeaders +import io.ktor.http.headersOf +import java.nio.file.Path + +/** + * Builds a fully-wired [ServerModule] on the managed-model path, using an in-memory event store + * and a no-op artifact store. The [modelSwapper] and [resourceProbe] are injected externally so + * tests can supply stubs or real implementations. + * + * The heavy infrastructure (SQLite, CAS, real tool executors) is deliberately avoided: the + * lifecycle tests only drive model-swap WS messages and never start sessions or invoke tools. + */ +fun buildTestServerModule( + eventStore: EventStore, + modelSwapper: ManagedInferenceRouter, + resourceProbe: ResourceProbe = UnavailableProbe, + tempDir: Path, +): ServerModule { + val artifactStore = noopArtifactStore() + val eventDispatcher = EventDispatcher(eventStore) + + val firstProvider: InferenceProvider = InfrastructureModule.createLlamaCppProvider( + modelId = modelSwapper.availableModelIds().first(), + modelPath = "/dev/null", + baseUrl = "http://127.0.0.1:1", + ) + + val infraRegistry: DefaultProviderRegistry = InfrastructureModule.createProviderRegistry(listOf(firstProvider)) + + val toolConfig = ToolConfig( + shell = ShellConfig(enabled = false, allowedExecutables = emptySet(), workingDir = tempDir), + fileRead = FileReadConfig(enabled = false, allowedPaths = setOf(tempDir)), + fileWrite = FileWriteConfig(enabled = false, allowedPaths = setOf(tempDir), workingDir = tempDir), + fileEdit = FileEditConfig(enabled = false, allowedPaths = setOf(tempDir), workingDir = tempDir), + ) + val toolRegistry = InfrastructureModule.createToolRegistry(toolConfig) + val toolExecutor = InfrastructureModule.createToolExecutor( + registry = toolRegistry, + eventDispatcher = eventDispatcher, + workDir = tempDir, + artifactStore = null, + ) + + val approvalEngine = DefaultApprovalEngine() + + val engines = OrchestratorEngines( + transitionResolver = DefaultTransitionResolver { condition, ctx -> condition.evaluate(ctx) }, + contextPackBuilder = DefaultContextPackBuilder(DefaultContextCompressor()), + inferenceRouter = modelSwapper, + validationPipeline = ValidationPipeline(validators = emptyList()), + approvalEngine = approvalEngine, + riskAssessor = DefaultRiskAssessor(), + toolRegistry = toolRegistry, + toolExecutor = toolExecutor, + ) + + val repositories = buildRepositories(eventStore) + + val orchestrator = DefaultSessionOrchestrator( + repositories = repositories, + engines = engines, + retryCoordinator = DefaultRetryCoordinator(eventStore), + artifactStore = artifactStore, + tokenizer = firstProvider.tokenizer, + ) + + val routerFacade = InfrastructureModule.createRouterFacade( + eventStore = eventStore, + inferenceRouter = modelSwapper, + config = RouterConfig(), + tokenizer = firstProvider.tokenizer, + ) + + val sessionUndoService = SessionUndoService( + eventStore = eventStore, + reverser = FileMutationReverser( + artifactStore = artifactStore, + allowedRoots = setOf(tempDir), + ), + ) + + val providerRegistry: ProviderRegistry = object : ProviderRegistry { + override fun listAll() = infraRegistry.listAll() + override suspend fun healthCheckAll() = infraRegistry.healthCheckAll() + } + + val noopWorkflowRegistry = object : WorkflowRegistry { + override fun listAll(): List = emptyList() + override fun find(workflowId: String): WorkflowGraph? = null + } + + return ServerModule( + orchestrator = orchestrator, + eventStore = eventStore, + artifactStore = artifactStore, + sessionRepository = repositories.sessionRepository, + workflowRegistry = noopWorkflowRegistry, + providerRegistry = providerRegistry, + defaultOrchestrationConfig = OrchestrationConfig(sandboxRoot = tempDir), + routerFacade = routerFacade, + orchestrationRepository = repositories.orchestrationRepository, + approvalRepository = repositories.approvalRepository, + toolRegistry = toolRegistry, + sessionUndoService = sessionUndoService, + modelSwapper = modelSwapper, + resourceProbe = resourceProbe, + ) +} + +private fun buildRepositories(eventStore: EventStore): OrchestratorRepositories = + OrchestratorRepositories( + eventStore = eventStore, + inferenceRepository = InferenceRepository(DefaultEventReplayer(eventStore, InferenceProjector())), + orchestrationRepository = OrchestrationRepository( + DefaultEventReplayer(eventStore, OrchestrationProjector(DefaultOrchestrationReducer())), + ), + sessionRepository = DefaultSessionRepository( + DefaultEventReplayer(eventStore, SessionProjector(DefaultSessionReducer())), + ), + artifactRepository = InfrastructureModule.createArtifactRepository(eventStore), + approvalRepository = InfrastructureModule.createApprovalRepository(eventStore), + ) + +private fun noopArtifactStore(): ArtifactStore = object : ArtifactStore { + override suspend fun put(bytes: ByteArray): ArtifactId = ArtifactId("noop") + override suspend fun get(id: ArtifactId): ByteArray? = null + override suspend fun flushBefore(commit: suspend () -> Unit) = commit() +} + +fun buildMockHttpClient(): HttpClient { + val engine = MockEngine { + respond( + content = """{"status":"healthy"}""", + headers = headersOf(HttpHeaders.ContentType, ContentType.Application.Json.toString()), + ) + } + return HttpClient(engine) +} + +fun buildTestModelManager( + llamaBin: String, + host: String, + port: Int, + healthTimeoutMs: Long, + eventStore: EventStore, + httpClient: HttpClient, +): DefaultModelManager { + val dispatcher = EventDispatcher(eventStore) + return DefaultModelManager( + llamaServerBin = llamaBin, + host = host, + port = port, + healthTimeoutMs = healthTimeoutMs, + eventStore = eventStore, + httpClient = httpClient, + eventDispatcher = dispatcher, + ) +} + +fun buildInMemoryEventStore(): InMemoryEventStore = InMemoryEventStore() diff --git a/apps/server/src/test/kotlin/com/correx/apps/server/lifecycle/ModelLifecycleLiveTest.kt b/apps/server/src/test/kotlin/com/correx/apps/server/lifecycle/ModelLifecycleLiveTest.kt new file mode 100644 index 00000000..61f177eb --- /dev/null +++ b/apps/server/src/test/kotlin/com/correx/apps/server/lifecycle/ModelLifecycleLiveTest.kt @@ -0,0 +1,193 @@ +package com.correx.apps.server.lifecycle + +import com.correx.apps.server.configureServer +import com.correx.apps.server.protocol.ClientMessage +import com.correx.apps.server.protocol.ServerMessage +import com.correx.infrastructure.inference.commons.ManagedInferenceRouter +import com.correx.infrastructure.inference.commons.ModelDescriptor +import com.correx.infrastructure.inference.commons.NvidiaResourceProbe +import com.correx.infrastructure.inference.commons.ResidencyMode +import com.correx.infrastructure.inference.commons.UnavailableProbe +import io.ktor.client.engine.cio.CIO +import io.ktor.client.plugins.websocket.WebSockets +import io.ktor.client.plugins.websocket.webSocket +import io.ktor.websocket.Frame +import io.ktor.websocket.readText +import kotlinx.coroutines.runBlocking +import kotlinx.coroutines.withTimeout +import kotlinx.serialization.encodeToString +import kotlinx.serialization.json.Json +import org.junit.jupiter.api.AfterEach +import org.junit.jupiter.api.Assumptions +import org.junit.jupiter.api.Test +import org.junit.jupiter.api.condition.EnabledIfEnvironmentVariable +import org.junit.jupiter.api.io.TempDir +import java.nio.file.Path +import kotlin.test.assertEquals +import kotlin.test.assertNotNull +import kotlin.test.assertTrue +import io.ktor.server.testing.testApplication + +/** + * Live integration test: spawns a real llama-server process and exercises the correx-managed + * model lifecycle WebSocket protocol end-to-end. + * + * Gated on the [CORREX_TEST_LLAMA_BIN] environment variable; skipped automatically in CI. + * + * **Run command:** + * ``` + * CORREX_TEST_LLAMA_BIN=/path/to/llama-server \ + * CORREX_TEST_MODEL=/path/to/model-a.gguf \ + * CORREX_TEST_MODEL_2=/path/to/model-b.gguf \ + * ./gradlew :apps:server:test --tests '*ModelLifecycleLiveTest*' --rerun-tasks + * ``` + * + * [CORREX_TEST_MODEL_2] is optional. When absent, the swap assertion is skipped via + * [Assumptions.assumeTrue] and only the load + telemetry assertions run. + */ +@EnabledIfEnvironmentVariable(named = "CORREX_TEST_LLAMA_BIN", matches = ".+") +class ModelLifecycleLiveTest { + + private val llamaBin: String = System.getenv("CORREX_TEST_LLAMA_BIN") + private val modelAPath: String = System.getenv("CORREX_TEST_MODEL") ?: "" + private val modelBPath: String? = System.getenv("CORREX_TEST_MODEL_2") + private val modelPort: Int = System.getenv("CORREX_TEST_MODEL_PORT")?.toIntOrNull() ?: 10000 + + private var modelManager: com.correx.infrastructure.inference.llama.cpp.DefaultModelManager? = null + private var currentModelId: String? = null + + private val protocolJson = Json { + classDiscriminator = "type" + ignoreUnknownKeys = true + } + + private fun decodeServerMessage(text: String): ServerMessage = + protocolJson.decodeFromString(text) + + private fun encodeClientMessage(msg: ClientMessage): String = + protocolJson.encodeToString(msg) + + @AfterEach + fun teardown() { + val mgr = modelManager ?: return + val id = currentModelId ?: return + runBlocking { + runCatching { mgr.unload(id) } + } + } + + @Test + fun `lifecycle protocol over websocket — live llama-server`(@TempDir tempDir: Path) { + val eventStore = buildInMemoryEventStore() + val httpClient = io.ktor.client.HttpClient(CIO) + + val descriptorA = ModelDescriptor( + modelId = "live-model-a", + modelPath = modelAPath, + residencyMode = ResidencyMode.PERSISTENT, + contextSize = 512, + ) + + val descriptors = if (modelBPath != null) { + listOf( + descriptorA, + ModelDescriptor( + modelId = "live-model-b", + modelPath = modelBPath, + residencyMode = ResidencyMode.PERSISTENT, + contextSize = 512, + ), + ) + } else { + listOf(descriptorA) + } + + val mgr = buildTestModelManager( + llamaBin = llamaBin, + host = "127.0.0.1", + port = modelPort, + healthTimeoutMs = 120_000L, + eventStore = eventStore, + httpClient = httpClient, + ) + modelManager = mgr + currentModelId = "live-model-a" + + val modelSwapper = ManagedInferenceRouter(mgr, descriptors, defaultModelId = "live-model-a") + + val resourceProbe = if (NvidiaResourceProbe.isAvailable()) { + NvidiaResourceProbe(pidSupplier = { mgr.currentPid() }) + } else { + UnavailableProbe + } + + val module = buildTestServerModule( + eventStore = eventStore, + modelSwapper = modelSwapper, + resourceProbe = resourceProbe, + tempDir = tempDir, + ) + + runBlocking { mgr.load(descriptorA) } + + testApplication { + application { configureServer(module) } + val client = createClient { install(WebSockets) } + + client.webSocket("/stream") { + var modelList: ServerMessage.ModelList? = null + var resourceStatus: ServerMessage.ResourceStatus? = null + + withTimeout(10_000L) { + while (modelList == null || resourceStatus == null) { + val frame = incoming.receive() + if (frame !is Frame.Text) continue + val raw = frame.readText() + when (val msg = decodeServerMessage(raw)) { + is ServerMessage.ModelList -> modelList = msg + is ServerMessage.ResourceStatus -> resourceStatus = msg + else -> Unit + } + } + } + + val ml = assertNotNull(modelList, "model.list must arrive in initial snapshot") + assertTrue(ml.models.contains("live-model-a"), "model-a must be listed") + assertEquals("live-model-a", ml.current) + + assertNotNull(resourceStatus, "resource.status must arrive in initial snapshot") + + // Swap assertion: only when a second model is configured + Assumptions.assumeTrue( + modelBPath != null, + "CORREX_TEST_MODEL_2 not set — skipping swap assertion", + ) + + currentModelId = "live-model-b" + val swapMsg = encodeClientMessage(ClientMessage.SwapModel("live-model-b")) + send(Frame.Text(swapMsg)) + + var modelChanged: ServerMessage.ModelChanged? = null + withTimeout(120_000L) { + while (modelChanged == null) { + val frame = incoming.receive() + if (frame !is Frame.Text) continue + val raw = frame.readText() + when (val msg = decodeServerMessage(raw)) { + is ServerMessage.ModelChanged -> modelChanged = msg + is ServerMessage.ProtocolError -> error("Unexpected protocol error during swap: ${msg.message}") + else -> Unit + } + } + } + + val mc = assertNotNull(modelChanged, "model.changed must arrive after SwapModel") + assertEquals("live-model-b", mc.modelId) + assertTrue(mc.loaded, "loaded must be true after swap") + + val clearMsg = encodeClientMessage(ClientMessage.ClearModelPin) + send(Frame.Text(clearMsg)) + } + } + } +} diff --git a/apps/server/src/test/kotlin/com/correx/apps/server/lifecycle/ModelLifecycleWiringTest.kt b/apps/server/src/test/kotlin/com/correx/apps/server/lifecycle/ModelLifecycleWiringTest.kt new file mode 100644 index 00000000..6581598e --- /dev/null +++ b/apps/server/src/test/kotlin/com/correx/apps/server/lifecycle/ModelLifecycleWiringTest.kt @@ -0,0 +1,184 @@ +package com.correx.apps.server.lifecycle + +import com.correx.apps.server.configureServer +import com.correx.apps.server.protocol.ClientMessage +import com.correx.apps.server.protocol.ServerMessage +import kotlinx.serialization.encodeToString +import kotlinx.serialization.json.Json +import com.correx.infrastructure.inference.commons.GpuStatus +import com.correx.infrastructure.inference.commons.ManagedInferenceRouter +import com.correx.infrastructure.inference.commons.ModelDescriptor +import com.correx.infrastructure.inference.commons.ResidencyMode +import com.correx.infrastructure.inference.commons.ResourceProbe +import com.correx.infrastructure.inference.commons.ResourceSnapshot +import io.ktor.client.plugins.websocket.WebSockets +import io.ktor.client.plugins.websocket.webSocket +import io.ktor.websocket.Frame +import io.ktor.websocket.readText +import kotlinx.coroutines.delay +import kotlinx.coroutines.withTimeout +import org.junit.jupiter.api.Test +import org.junit.jupiter.api.io.TempDir +import java.nio.file.Path +import kotlinx.coroutines.runBlocking +import kotlin.test.assertEquals +import kotlin.test.assertNotNull +import kotlin.test.assertTrue +import io.ktor.server.testing.testApplication + +/** + * CI test: exercises the correx-managed model lifecycle WebSocket protocol against a + * real Ktor [testApplication] server. No actual llama-server binary is spawned — the + * [DefaultModelManager] uses the system "sleep" binary + a MockEngine health client + * (mirrors the seam used by DefaultModelManagerTest). + * + * Assertions: + * 1. Initial snapshot contains model.list with [model-a, model-b], current=model-a. + * 2. Initial snapshot contains resource.status with the fake probe values. + * 3. SwapModel("model-b") → model.changed arrives with modelId=model-b, loaded=true. + * 4. ClearModelPin does not produce an error frame. + */ +class ModelLifecycleWiringTest { + + private val protocolJson = Json { + classDiscriminator = "type" + ignoreUnknownKeys = true + } + + private fun decodeServerMessage(text: String): ServerMessage = + protocolJson.decodeFromString(text) + + private fun encodeClientMessage(msg: ClientMessage): String = + protocolJson.encodeToString(msg) + + @Test + fun `lifecycle protocol over websocket`(@TempDir tempDir: Path): Unit { + val eventStore = buildInMemoryEventStore() + val httpClient = buildMockHttpClient() + + val descriptorA = ModelDescriptor( + modelId = "model-a", + modelPath = "/dev/null", + residencyMode = ResidencyMode.PERSISTENT, + contextSize = 512, + ) + val descriptorB = ModelDescriptor( + modelId = "model-b", + modelPath = "/dev/null", + residencyMode = ResidencyMode.PERSISTENT, + contextSize = 512, + ) + + val modelManager = buildTestModelManager( + llamaBin = "sleep", + host = "127.0.0.1", + port = 19888, + healthTimeoutMs = 5_000L, + eventStore = eventStore, + httpClient = httpClient, + ) + + // Pre-load model-a so that currentModelId() returns "model-a" in the initial snapshot, + // mirroring Main's runBlocking { modelManager.load(descriptor) } call. + runBlocking { modelManager.load(descriptorA) } + + val descriptors = listOf(descriptorA, descriptorB) + val modelSwapper = ManagedInferenceRouter(modelManager, descriptors, defaultModelId = "model-a") + + val fakeProbe = ResourceProbe { + ResourceSnapshot( + gpu = GpuStatus(memoryUsedMb = 2048L, memoryTotalMb = 8192L, utilizationPct = 42), + processRssBytes = 512L * 1024L * 1024L, + ) + } + + val module = buildTestServerModule( + eventStore = eventStore, + modelSwapper = modelSwapper, + resourceProbe = fakeProbe, + tempDir = tempDir, + ) + + try { + testApplication { + application { configureServer(module) } + val client = createClient { install(WebSockets) } + + client.webSocket("/stream") { + // Collect frames until we have both model.list and resource.status + var modelList: ServerMessage.ModelList? = null + var resourceStatus: ServerMessage.ResourceStatus? = null + + // Pass 1: collect initial snapshot messages + withTimeout(5_000L) { + while (modelList == null || resourceStatus == null) { + val frame = incoming.receive() + if (frame !is Frame.Text) continue + val raw = frame.readText() + when (val msg = decodeServerMessage(raw)) { + is ServerMessage.ModelList -> modelList = msg + is ServerMessage.ResourceStatus -> resourceStatus = msg + else -> Unit + } + } + } + + val ml = assertNotNull(modelList, "model.list must arrive in initial snapshot") + assertEquals(listOf("model-a", "model-b"), ml.models) + assertEquals("model-a", ml.current) + + val rs = assertNotNull(resourceStatus, "resource.status must arrive in initial snapshot") + assertEquals(2048L, rs.gpuMemoryUsedMb) + assertEquals(8192L, rs.gpuMemoryTotalMb) + assertEquals(42, rs.gpuUtilizationPct) + assertNotNull(rs.processRssMb) + + // Send SwapModel("model-b") + val swapMsg = encodeClientMessage(ClientMessage.SwapModel("model-b")) + send(Frame.Text(swapMsg)) + + // Wait for model.changed with modelId=model-b, loaded=true + var modelChanged: ServerMessage.ModelChanged? = null + withTimeout(8_000L) { + while (modelChanged == null) { + val frame = incoming.receive() + if (frame !is Frame.Text) continue + val raw = frame.readText() + when (val msg = decodeServerMessage(raw)) { + is ServerMessage.ModelChanged -> modelChanged = msg + is ServerMessage.ProtocolError -> error("Unexpected protocol error: ${msg.message}") + else -> Unit + } + } + } + + val mc = assertNotNull(modelChanged, "model.changed must arrive after SwapModel") + assertEquals("model-b", mc.modelId) + assertTrue(mc.loaded, "loaded must be true after swap") + + // Send ClearModelPin — the server responds with nothing (just clears the pin); + // verify no ProtocolError arrives within a brief window. + val clearMsg = encodeClientMessage(ClientMessage.ClearModelPin) + send(Frame.Text(clearMsg)) + + // Give the server a moment to process ClearModelPin; if it sends back a + // ProtocolError it would arrive almost immediately (well before 2500ms resource push). + delay(100) + val pending = incoming.tryReceive() + if (pending.isSuccess) { + val f = pending.getOrThrow() + if (f is Frame.Text) { + val msg = decodeServerMessage(f.readText()) + assertTrue(msg !is ServerMessage.ProtocolError, "ClearModelPin produced error: $msg") + } + } + // If no frame arrived — which is also fine (no error sent). + } + } + } finally { + runBlocking { + runCatching { modelManager.unload(modelSwapper.currentModelId() ?: "model-b") } + } + } + } +}