feat: correx-managed model lifecycle slice 4 — resource telemetry

Vendor-agnostic ResourceProbe (commons): NvidiaResourceProbe reads nvidia-smi
(injectable runner) for whole-device VRAM/util and /proc/<pid>/status for the
managed llama-server RSS, both fail-soft to null; UnavailableProbe for non-GPU
hosts / the static path. DefaultModelManager.currentPid() + LlamaProcess.pid
expose the managed pid. ServerMessage.ResourceStatus (resource.status,
NonEventMessage, all-nullable) pushed every 2.5s on the global stream plus one
in the initial snapshot. Live gauge only — never event-sourced (feeds no core
decision), so invariants #8/#9 hold. Main wires the NVIDIA probe on the managed
path when nvidia-smi is present, else UnavailableProbe.

Tests: csv parsing (single/multi/malformed), gpu+rss combine, fallbacks.

Plan: docs/plans/2026-05-31-model-lifecycle-management.md (slice 4 of 5).
This commit is contained in:
2026-06-01 12:36:38 +04:00
parent 7341ab578e
commit 7b1df95627
10 changed files with 273 additions and 0 deletions
@@ -49,6 +49,9 @@ import com.correx.infrastructure.InfrastructureModule
import com.correx.infrastructure.inference.DefaultProviderRegistry import com.correx.infrastructure.inference.DefaultProviderRegistry
import com.correx.infrastructure.inference.FirstAvailableRoutingStrategy import com.correx.infrastructure.inference.FirstAvailableRoutingStrategy
import com.correx.infrastructure.inference.commons.ManagedInferenceRouter import com.correx.infrastructure.inference.commons.ManagedInferenceRouter
import com.correx.infrastructure.inference.commons.NvidiaResourceProbe
import com.correx.infrastructure.inference.commons.ResourceProbe
import com.correx.infrastructure.inference.commons.UnavailableProbe
import com.correx.core.inference.InferenceProvider import com.correx.core.inference.InferenceProvider
import com.correx.infrastructure.inference.llama.cpp.LlamaCppInferenceProvider import com.correx.infrastructure.inference.llama.cpp.LlamaCppInferenceProvider
import com.correx.infrastructure.tools.FileEditConfig import com.correx.infrastructure.tools.FileEditConfig
@@ -83,6 +86,8 @@ fun main() {
val inferenceRouter: InferenceRouter val inferenceRouter: InferenceRouter
// Non-null only on the managed path; backs manual model swap/pin from clients. // Non-null only on the managed path; backs manual model swap/pin from clients.
var modelSwapper: ManagedInferenceRouter? = null var modelSwapper: ManagedInferenceRouter? = null
// Live GPU/RAM gauge; NVIDIA-backed on the managed path when nvidia-smi is present, else unavailable.
var resourceProbe: ResourceProbe = UnavailableProbe
if (correxConfig.models.isNotEmpty()) { if (correxConfig.models.isNotEmpty()) {
val settings = correxConfig.modelsSettings val settings = correxConfig.modelsSettings
@@ -108,6 +113,11 @@ fun main() {
val managedRouter = ManagedInferenceRouter(modelManager, descriptors, targetModelConfig.id) val managedRouter = ManagedInferenceRouter(modelManager, descriptors, targetModelConfig.id)
inferenceRouter = managedRouter inferenceRouter = managedRouter
modelSwapper = managedRouter modelSwapper = managedRouter
resourceProbe = if (NvidiaResourceProbe.isAvailable()) {
NvidiaResourceProbe(pidSupplier = { modelManager.currentPid() })
} else {
UnavailableProbe
}
// Shutdown hook to kill the managed llama-server // Shutdown hook to kill the managed llama-server
Runtime.getRuntime().addShutdownHook(Thread { Runtime.getRuntime().addShutdownHook(Thread {
log.info("Shutdown: unloading managed model '{}'", targetModelConfig.id) log.info("Shutdown: unloading managed model '{}'", targetModelConfig.id)
@@ -243,6 +253,7 @@ fun main() {
toolRegistry = toolRegistry, toolRegistry = toolRegistry,
sessionUndoService = sessionUndoService, sessionUndoService = sessionUndoService,
modelSwapper = modelSwapper, modelSwapper = modelSwapper,
resourceProbe = resourceProbe,
) )
module.start() module.start()
log.info("==============================") log.info("==============================")
@@ -52,6 +52,10 @@ class ServerModule(
// The managed-model router when correx owns the local model process ([[models]] configured); // The managed-model router when correx owns the local model process ([[models]] configured);
// null on the static-provider path. Backs the manual SwapModel / ClearModelPin operations. // null on the static-provider path. Backs the manual SwapModel / ClearModelPin operations.
val modelSwapper: com.correx.infrastructure.inference.commons.ManagedInferenceRouter? = null, val modelSwapper: com.correx.infrastructure.inference.commons.ManagedInferenceRouter? = null,
// Live GPU/RAM gauge pushed to clients on the global stream. Defaults to the always-unavailable
// probe (static path / non-GPU host); the managed path wires an NVIDIA-backed probe.
val resourceProbe: com.correx.infrastructure.inference.commons.ResourceProbe =
com.correx.infrastructure.inference.commons.UnavailableProbe,
// Long-lived scope owned by the module — backs the event-store subscription. // Long-lived scope owned by the module — backs the event-store subscription.
// SupervisorJob so one failure doesn't kill the whole module; // SupervisorJob so one failure doesn't kill the whole module;
// Dispatchers.Default since the work is non-blocking and CPU-light. // Dispatchers.Default since the work is non-blocking and CPU-light.
@@ -283,6 +283,22 @@ sealed interface ServerMessage {
override val sessionSequence: Long? = null, override val sessionSequence: Long? = null,
) : ServerMessage, NonEventMessage ) : ServerMessage, NonEventMessage
/**
* Live resource gauge pushed periodically on the global stream (not event-derived). All fields
* are nullable: GPU fields are null on a non-NVIDIA host, `processRssMb` is null when no model
* is resident.
*/
@Serializable
@SerialName("resource.status")
data class ResourceStatus(
val gpuMemoryUsedMb: Long?,
val gpuMemoryTotalMb: Long?,
val gpuUtilizationPct: Int?,
val processRssMb: Long?,
override val sequence: Long? = null,
override val sessionSequence: Long? = null,
) : ServerMessage, NonEventMessage
@Serializable @Serializable
@SerialName("protocol_error") @SerialName("protocol_error")
data class ProtocolError( data class ProtocolError(
@@ -25,24 +25,39 @@ import com.correx.core.events.types.SessionId
import com.correx.core.events.types.StageId import com.correx.core.events.types.StageId
import com.correx.core.inference.ProviderHealth import com.correx.core.inference.ProviderHealth
import com.correx.core.utils.TypeId import com.correx.core.utils.TypeId
import com.correx.infrastructure.inference.commons.ResourceProbe
import com.correx.infrastructure.inference.commons.ResourceSnapshot
import io.ktor.server.websocket.DefaultWebSocketServerSession import io.ktor.server.websocket.DefaultWebSocketServerSession
import io.ktor.websocket.Frame import io.ktor.websocket.Frame
import io.ktor.websocket.readText import io.ktor.websocket.readText
import kotlinx.coroutines.CompletableDeferred import kotlinx.coroutines.CompletableDeferred
import kotlinx.coroutines.Dispatchers
import kotlinx.coroutines.channels.BufferOverflow import kotlinx.coroutines.channels.BufferOverflow
import kotlinx.coroutines.channels.Channel import kotlinx.coroutines.channels.Channel
import kotlinx.coroutines.channels.ClosedReceiveChannelException import kotlinx.coroutines.channels.ClosedReceiveChannelException
import kotlinx.coroutines.coroutineScope import kotlinx.coroutines.coroutineScope
import kotlinx.coroutines.delay
import kotlinx.coroutines.flow.SharedFlow import kotlinx.coroutines.flow.SharedFlow
import kotlinx.coroutines.flow.onStart import kotlinx.coroutines.flow.onStart
import kotlinx.coroutines.flow.onSubscription import kotlinx.coroutines.flow.onSubscription
import kotlinx.coroutines.launch import kotlinx.coroutines.launch
import kotlinx.coroutines.withContext
import kotlinx.datetime.Clock import kotlinx.datetime.Clock
import org.slf4j.LoggerFactory import org.slf4j.LoggerFactory
import java.util.UUID import java.util.UUID
private val log = LoggerFactory.getLogger(GlobalStreamHandler::class.java) private val log = LoggerFactory.getLogger(GlobalStreamHandler::class.java)
private const val BUFFER_CAPACITY = 1024 private const val BUFFER_CAPACITY = 1024
private const val RESOURCE_PUSH_INTERVAL_MS = 2500L
private const val BYTES_PER_MB = 1024L * 1024L
private fun ResourceSnapshot.toResourceStatus(): ServerMessage.ResourceStatus =
ServerMessage.ResourceStatus(
gpuMemoryUsedMb = gpu?.memoryUsedMb,
gpuMemoryTotalMb = gpu?.memoryTotalMb,
gpuUtilizationPct = gpu?.utilizationPct,
processRssMb = processRssBytes?.let { it / BYTES_PER_MB },
)
class GlobalStreamHandler(private val module: ServerModule) { class GlobalStreamHandler(private val module: ServerModule) {
@@ -72,6 +87,10 @@ class GlobalStreamHandler(private val module: ServerModule) {
launch { launch {
streamGlobal(module.eventStore, bridge, mapper, sendFrame) streamGlobal(module.eventStore, bridge, mapper, sendFrame)
} }
// Live resource gauge (GPU/RAM) pushed periodically. Cancelled with the session scope.
launch {
streamResources(module.resourceProbe, sendFrame)
}
try { try {
for (frame in session.incoming) { for (frame in session.incoming) {
@@ -129,6 +148,22 @@ class GlobalStreamHandler(private val module: ServerModule) {
session.send(Frame.Text(ProtocolSerializer.encodeServerMessage( session.send(Frame.Text(ProtocolSerializer.encodeServerMessage(
ServerMessage.WorkflowList(workflows = workflows), ServerMessage.WorkflowList(workflows = workflows),
))) )))
// Initial resource gauge so a freshly-connected client renders VRAM/RAM immediately.
val snapshot = withContext(Dispatchers.IO) { module.resourceProbe.probe() }
session.send(Frame.Text(ProtocolSerializer.encodeServerMessage(snapshot.toResourceStatus())))
}
/** Pushes a [ServerMessage.ResourceStatus] every [RESOURCE_PUSH_INTERVAL_MS]. */
private suspend fun streamResources(
probe: ResourceProbe,
sendFrame: suspend (ServerMessage) -> Unit,
) {
while (true) {
delay(RESOURCE_PUSH_INTERVAL_MS)
val snapshot = withContext(Dispatchers.IO) { probe.probe() }
sendFrame(snapshot.toResourceStatus())
}
} }
private suspend fun handleClientMessage( private suspend fun handleClientMessage(
@@ -16,6 +16,7 @@ dependencies {
implementation "io.ktor:ktor-client-cio:$ktor_version" implementation "io.ktor:ktor-client-cio:$ktor_version"
implementation "io.ktor:ktor-client-content-negotiation:$ktor_version" implementation "io.ktor:ktor-client-content-negotiation:$ktor_version"
implementation "io.ktor:ktor-serialization-kotlinx-json:$ktor_version" implementation "io.ktor:ktor-serialization-kotlinx-json:$ktor_version"
implementation "org.slf4j:slf4j-api:2.0.16"
} }
tasks.named("koverVerify").configure { enabled = false } tasks.named("koverVerify").configure { enabled = false }
@@ -0,0 +1,80 @@
package com.correx.infrastructure.inference.commons
import org.slf4j.LoggerFactory
import java.io.File
private val log = LoggerFactory.getLogger("NvidiaResourceProbe")
/**
* [ResourceProbe] backed by `nvidia-smi` for GPU stats and `/proc/<pid>/status` for the managed
* model process RSS. Both reads are isolated and fail soft to null, so a missing `nvidia-smi`,
* a non-NVIDIA host, or a dead pid never throws.
*
* The `nvidia-smi` invocation and RSS read are injectable so the parsing is unit-testable without
* a GPU or a live process.
*
* @param pidSupplier returns the managed `llama-server` pid, or null when no model is resident.
* @param nvidiaSmiRunner returns raw `nvidia-smi` CSV stdout, or null when unavailable.
* @param rssReader returns the RSS in bytes for a pid, or null when unreadable.
*/
class NvidiaResourceProbe(
private val pidSupplier: () -> Long?,
private val nvidiaSmiRunner: () -> String? = ::runNvidiaSmi,
private val rssReader: (Long) -> Long? = ::readRssBytes,
) : ResourceProbe {
override fun probe(): ResourceSnapshot {
val gpu = nvidiaSmiRunner()?.let { parseGpu(it) }
val rss = pidSupplier()?.let { rssReader(it) }
return ResourceSnapshot(gpu = gpu, processRssBytes = rss)
}
companion object {
private const val EXPECTED_FIELDS = 3
private val QUERY = listOf(
"nvidia-smi",
"--query-gpu=memory.used,memory.total,utilization.gpu",
"--format=csv,noheader,nounits",
)
/**
* Parses the first GPU row of `nvidia-smi --format=csv,noheader,nounits` output
* (`memory.used, memory.total, utilization.gpu`, all in MiB / percent). Returns null on any
* shape mismatch so the caller surfaces "unavailable" instead of partial data.
*/
fun parseGpu(csv: String): GpuStatus? {
val row = csv.lineSequence().map { it.trim() }.firstOrNull { it.isNotEmpty() } ?: return null
val fields = row.split(',').map { it.trim() }
if (fields.size < EXPECTED_FIELDS) return null
val used = fields[0].toLongOrNull() ?: return null
val total = fields[1].toLongOrNull() ?: return null
val util = fields[2].toIntOrNull() ?: return null
return GpuStatus(memoryUsedMb = used, memoryTotalMb = total, utilizationPct = util)
}
/** Returns true if `nvidia-smi -L` runs successfully — used to pick the probe at boot. */
fun isAvailable(): Boolean = runCatching {
val proc = ProcessBuilder("nvidia-smi", "-L")
.redirectErrorStream(true)
.start()
proc.waitFor() == 0
}.getOrElse { false }
private fun runNvidiaSmi(): String? = runCatching {
val proc = ProcessBuilder(QUERY).redirectErrorStream(true).start()
val out = proc.inputStream.bufferedReader().readText()
if (proc.waitFor() == 0) out else null
}.onFailure { log.debug("nvidia-smi query failed: {}", it.message) }.getOrNull()
private fun readRssBytes(pid: Long): Long? = runCatching {
val statusFile = File("/proc/$pid/status")
if (!statusFile.exists()) return null
val vmRss = statusFile.readLines().firstOrNull { it.startsWith("VmRSS:") } ?: return null
// Format: "VmRSS:\t 12345 kB"
val kb = vmRss.removePrefix("VmRSS:").trim().substringBefore(' ').toLongOrNull() ?: return null
kb * BYTES_PER_KB
}.onFailure { log.debug("RSS read failed for pid {}: {}", pid, it.message) }.getOrNull()
private const val BYTES_PER_KB = 1024L
}
}
@@ -0,0 +1,40 @@
package com.correx.infrastructure.inference.commons
/**
* Whole-device GPU usage at an instant. All fields are device-wide (not per-process) —
* per-model VRAM attribution is an explicit non-goal.
*/
data class GpuStatus(
val memoryUsedMb: Long,
val memoryTotalMb: Long,
val utilizationPct: Int,
)
/**
* A point-in-time reading of local-model resource usage. Both fields are nullable so that
* "unavailable" is representable without exceptions: [gpu] is null when no GPU probe is present
* (non-NVIDIA box / `nvidia-smi` absent), [processRssBytes] is null when no model is resident.
*/
data class ResourceSnapshot(
val gpu: GpuStatus?,
val processRssBytes: Long?,
) {
companion object {
val UNAVAILABLE = ResourceSnapshot(gpu = null, processRssBytes = null)
}
}
/**
* Vendor-agnostic live gauge of GPU/process resource usage. This is a **live gauge, never recorded
* as events** — it feeds no core decision, so invariant #9 does not require event-sourcing it, and
* replay (invariant #8) never reads it. Implementations must degrade gracefully (return
* [ResourceSnapshot.UNAVAILABLE]-shaped data) rather than throw.
*/
fun interface ResourceProbe {
fun probe(): ResourceSnapshot
}
/** Always-unavailable probe — used on non-GPU boxes and the static (unmanaged) provider path. */
object UnavailableProbe : ResourceProbe {
override fun probe(): ResourceSnapshot = ResourceSnapshot.UNAVAILABLE
}
@@ -129,6 +129,13 @@ class DefaultModelManager(
override fun currentModel(): ModelDescriptor? = currentDescriptor override fun currentModel(): ModelDescriptor? = currentDescriptor
/**
* The managed `llama-server` OS pid, or null when no model is resident. Read without the mutex
* for the live resource gauge — a stale/null read for one probe tick during a swap is acceptable
* (telemetry only, feeds no core decision).
*/
fun currentPid(): Long? = currentProcess?.pid
override fun healthCheck(): ProviderHealth { override fun healthCheck(): ProviderHealth {
val process = currentProcess ?: return ProviderHealth.Unavailable("No model loaded") val process = currentProcess ?: return ProviderHealth.Unavailable("No model loaded")
return if (process.isAlive) { return if (process.isAlive) {
@@ -16,6 +16,10 @@ class LlamaProcess(
val isAlive: Boolean val isAlive: Boolean
get() = process?.isAlive ?: false get() = process?.isAlive ?: false
/** The OS pid of the running process, or null when not started / not alive. */
val pid: Long?
get() = process?.takeIf { it.isAlive }?.pid()
suspend fun start() = withContext(Dispatchers.IO) { suspend fun start() = withContext(Dispatchers.IO) {
if (isAlive) return@withContext if (isAlive) return@withContext
@@ -0,0 +1,75 @@
package com.correx.testing.contracts.model.inference
import com.correx.infrastructure.inference.commons.NvidiaResourceProbe
import com.correx.infrastructure.inference.commons.UnavailableProbe
import org.junit.jupiter.api.Assertions.assertEquals
import org.junit.jupiter.api.Assertions.assertNull
import org.junit.jupiter.api.Test
class NvidiaResourceProbeTest {
@Test
fun `parses a single GPU csv row`() {
val gpu = NvidiaResourceProbe.parseGpu("8192, 16384, 37")
assertEquals(8192L, gpu?.memoryUsedMb)
assertEquals(16384L, gpu?.memoryTotalMb)
assertEquals(37, gpu?.utilizationPct)
}
@Test
fun `parses the first row when multiple GPUs are present`() {
val gpu = NvidiaResourceProbe.parseGpu("8192, 16384, 37\n1024, 16384, 5")
assertEquals(8192L, gpu?.memoryUsedMb)
assertEquals(37, gpu?.utilizationPct)
}
@Test
fun `returns null on malformed csv`() {
assertNull(NvidiaResourceProbe.parseGpu("not,enough"))
assertNull(NvidiaResourceProbe.parseGpu("abc, def, ghi"))
assertNull(NvidiaResourceProbe.parseGpu(""))
}
@Test
fun `probe combines gpu csv and rss reader`() {
val probe = NvidiaResourceProbe(
pidSupplier = { 4242L },
nvidiaSmiRunner = { "8192, 16384, 37" },
rssReader = { pid -> if (pid == 4242L) 1_048_576L else null },
)
val snapshot = probe.probe()
assertEquals(8192L, snapshot.gpu?.memoryUsedMb)
assertEquals(1_048_576L, snapshot.processRssBytes)
}
@Test
fun `probe reports gpu unavailable when nvidia-smi yields null`() {
val probe = NvidiaResourceProbe(
pidSupplier = { 4242L },
nvidiaSmiRunner = { null },
rssReader = { 1_048_576L },
)
val snapshot = probe.probe()
assertNull(snapshot.gpu)
assertEquals(1_048_576L, snapshot.processRssBytes)
}
@Test
fun `probe reports rss null when no pid is resident`() {
val probe = NvidiaResourceProbe(
pidSupplier = { null },
nvidiaSmiRunner = { "8192, 16384, 37" },
rssReader = { error("must not be called when pid is null") },
)
val snapshot = probe.probe()
assertEquals(8192L, snapshot.gpu?.memoryUsedMb)
assertNull(snapshot.processRssBytes)
}
@Test
fun `unavailable probe is fully unavailable`() {
val snapshot = UnavailableProbe.probe()
assertNull(snapshot.gpu)
assertNull(snapshot.processRssBytes)
}
}