From f26cf8986f2215a99c7531fcc913882e16a12056 Mon Sep 17 00:00:00 2001 From: Jussi Elo Date: Thu, 1 Oct 2026 12:07:55 +0000 Subject: [PATCH 1/3] ROCMAI-83: extract mcp.rs from apps/rocmd/src/lib.rs Sixth PR of Phase 5 (rocmd modularization, ROCMAI-27): pull the MCP stdio server, tool schema table, tool dispatch, and the rocm-subprocess capture/argv-building helpers behind the MCP tools into their own module. No behavior change. `CommandCapture` and `run_command_with_timeout` were independently claimed by this PR and by the already-landed common.rs extraction (both authored against the pre-stack monolith, unaware of each other); common.rs keeps ownership since it landed first in the merge order, so mcp.rs now imports both from `crate::common` instead of redefining them. The other crate-root reach-backs this module needed (`build_bridge_snapshot`, `gather_gpu_snapshot_for_config`, `bridge_engine_inventory`, `load_managed_services`) are now reached via `crate::common`/ `crate::persistence`, since those modules landed earlier in this stack; `stop_managed_service` stays reached via bare `crate::` since it has not moved out of lib.rs yet. `cli.rs`'s MCP dispatch arms and sandbox.rs's `run_rocm_capture_for_paths` import (both landed before this PR, so both still reached into lib.rs directly) are repointed to `crate::mcp::`. 15 tests that exercise this module's own logic (MCP tool-schema/ dispatch shape, read-only-verb classification, install_sdk/ install_engine/launch_server/watcher_enable argv building, and read_tail_lines) moved into mcp.rs's own #[cfg(test)] mod in this same PR, reusing `crate::test_support::workspace_test_artifact_dir` instead of a second local copy. Tests that share a name prefix or construct these types but actually exercise Cli parsing, run_daemon, or watcher-domain logic stayed in lib.rs for their own later PRs. Rebase note (main now at f9a80117): main gained a test since this commit was authored (install_sdk_rejects_system_prefix_reached_by_escaping_home, moved into mcp.rs's own test module alongside its sibling install_sdk_rejects_system_prefix_without_ack. Also: run_rocm_capture/run_rocm_capture_for_paths were independently duplicated into this commit's own mcp.rs (same pre-stack-aware-of-each- other story as the CommandCapture/run_command_with_timeout duplication already called out above) and into common.rs via #479. common.rs keeps ownership since it landed first; mcp.rs's ten call sites now go through common::run_rocm_capture instead of a local redefinition, and sandbox.rs's now-stale `use crate::mcp::run_rocm_capture_for_paths` (from when mcp.rs was this function's owner, pre-#479) is dropped in favor of its existing common::-qualified calls. Signed-off-by: Jussi Elo --- apps/rocmd/src/cli.rs | 14 +- apps/rocmd/src/lib.rs | 1537 +--------------------------------------- apps/rocmd/src/mcp.rs | 1549 +++++++++++++++++++++++++++++++++++++++++ docs/architecture.md | 2 +- 4 files changed, 1563 insertions(+), 1539 deletions(-) create mode 100644 apps/rocmd/src/mcp.rs diff --git a/apps/rocmd/src/cli.rs b/apps/rocmd/src/cli.rs index efa31e325..55071afa5 100644 --- a/apps/rocmd/src/cli.rs +++ b/apps/rocmd/src/cli.rs @@ -297,7 +297,7 @@ async fn run_cli(cli: Cli) -> Result<()> { allow_native_fallback, policy, )?; - crate::print_json(&value)?; + crate::mcp::print_json(&value)?; } Command::SandboxTool { tool, @@ -321,13 +321,13 @@ async fn run_cli(cli: Cli) -> Result<()> { message, policy, )?; - crate::print_json(&value)?; + crate::mcp::print_json(&value)?; } Command::McpServer => { - crate::run_mcp_server(&paths)?; + crate::mcp::run_mcp_server(&paths)?; } Command::McpToolsJson => { - crate::print_json(&json!({ "tools": crate::rocm_mcp_tools() }))?; + crate::mcp::print_json(&json!({ "tools": crate::mcp::rocm_mcp_tools() }))?; } Command::McpCall { name, @@ -340,15 +340,15 @@ async fn run_cli(cli: Cli) -> Result<()> { if !arguments.is_object() { bail!("--arguments-json for MCP tool `{name}` must be a JSON object"); } - crate::ensure_direct_mcp_call_allowed(&name, allow_mutation)?; - let result = crate::handle_mcp_tool_call( + crate::mcp::ensure_direct_mcp_call_allowed(&name, allow_mutation)?; + let result = crate::mcp::handle_mcp_tool_call( &paths, &json!({ "name": name, "arguments": arguments, }), )?; - crate::print_json(&result)?; + crate::mcp::print_json(&result)?; } } diff --git a/apps/rocmd/src/lib.rs b/apps/rocmd/src/lib.rs index fc3dbdf0c..efd6ad807 100644 --- a/apps/rocmd/src/lib.rs +++ b/apps/rocmd/src/lib.rs @@ -6,6 +6,7 @@ mod cli; mod common; +mod mcp; mod persistence; mod sandbox; #[cfg(test)] @@ -21,17 +22,15 @@ use rocm_core::AuditEventRecord; use rocm_core::AutomationEventRecord; use rocm_core::{ AppPaths, AutomationProposalRecord, AutomationRuntimeState, AutomationTriggerEvent, - CodexBridgeGpuSnapshot, DEFAULT_LOCAL_HOST, ExamineSummary, ManagedServiceRecord, - RocmCliConfig, WatcherMode, WatcherRuntimeSnapshot, append_automation_proposal, - builtin_watchers, daemon_binary_path, load_recent_automation_events, + CodexBridgeGpuSnapshot, ManagedServiceRecord, RocmCliConfig, WatcherMode, + WatcherRuntimeSnapshot, append_automation_proposal, builtin_watchers, daemon_binary_path, resolve_model_recipe_artifact, unix_time_millis, }; -use serde::Serialize; use serde_json::Value; use serde_json::json; -use std::collections::{HashSet, VecDeque}; +use std::collections::HashSet; use std::fs; -use std::io::{self, BufRead, Read, Seek, SeekFrom, Write}; +use std::io::{self, Read, Seek, SeekFrom, Write}; use std::path::Path; use std::process::{Command as ProcessCommand, Stdio}; use std::thread; @@ -49,1126 +48,6 @@ const GPU_THERMAL_MEMORY_PRESSURE_C: f64 = 95.0; const GPU_MEMORY_VRAM_PRESSURE_PERCENT: f64 = 95.0; const ARTIFACT_PREFETCH_TIMEOUT: Duration = Duration::from_mins(10); -fn run_mcp_server(paths: &AppPaths) -> Result<()> { - let stdin = io::stdin(); - let stdout = io::stdout(); - let mut reader = stdin.lock(); - let mut writer = stdout.lock(); - let mut line = String::new(); - - loop { - line.clear(); - let bytes_read = reader.read_line(&mut line)?; - if bytes_read == 0 { - break; - } - if line.trim().is_empty() { - continue; - } - - let message: Value = match serde_json::from_str(&line) { - Ok(value) => value, - Err(error) => { - write_json_line( - &mut writer, - &json!({ - "jsonrpc": "2.0", - "error": { - "code": -32700, - "message": format!("parse error: {error}"), - } - }), - )?; - continue; - } - }; - - let Some(method) = message.get("method").and_then(Value::as_str) else { - if message.get("id").is_some() { - write_json_line( - &mut writer, - &json!({ - "jsonrpc": "2.0", - "id": message.get("id").cloned().unwrap_or(Value::Null), - "error": { - "code": -32600, - "message": "invalid request: missing method", - } - }), - )?; - } - continue; - }; - - let id = message.get("id").cloned(); - let params = message.get("params").cloned().unwrap_or(Value::Null); - match method { - "initialize" => { - let protocol_version = params - .get("protocolVersion") - .and_then(Value::as_str) - .unwrap_or("2025-03-26"); - if let Some(id) = id { - write_json_line( - &mut writer, - &json!({ - "jsonrpc": "2.0", - "id": id, - "result": { - "protocolVersion": protocol_version, - "capabilities": { - "tools": { - "listChanged": true, - } - }, - "serverInfo": { - "name": "rocmd-mcp-server", - "title": "ROCm AI Command Center", - "version": env!("CARGO_PKG_VERSION"), - } - } - }), - )?; - } - } - "ping" => { - if let Some(id) = id { - write_json_line( - &mut writer, - &json!({ - "jsonrpc": "2.0", - "id": id, - "result": {} - }), - )?; - } - } - "notifications/initialized" => {} - "tools/list" => { - if let Some(id) = id { - write_json_line( - &mut writer, - &json!({ - "jsonrpc": "2.0", - "id": id, - "result": { - "tools": rocm_mcp_tools(), - "nextCursor": Value::Null, - } - }), - )?; - } - } - "tools/call" => { - if let Some(id) = id { - let result = match handle_mcp_tool_call(paths, ¶ms) { - Ok(result) => result, - Err(error) => tool_error( - format!("ROCm MCP tool call failed: {error:#}"), - json!({ - "tool": params.get("name").cloned().unwrap_or(Value::Null), - }), - ), - }; - write_json_line( - &mut writer, - &json!({ - "jsonrpc": "2.0", - "id": id, - "result": result, - }), - )?; - } - } - notification if notification.starts_with("notifications/") => {} - other => { - if let Some(id) = id { - write_json_line( - &mut writer, - &json!({ - "jsonrpc": "2.0", - "id": id, - "error": { - "code": -32601, - "message": format!("method not found: {other}"), - } - }), - )?; - } - } - } - } - - Ok(()) -} - -fn write_json_line(writer: &mut impl Write, value: &Value) -> Result<()> { - writer.write_all(serde_json::to_string(value)?.as_bytes())?; - writer.write_all(b"\n")?; - writer.flush()?; - Ok(()) -} - -fn print_json(value: &T) -> Result<()> { - println!( - "{}", - serde_json::to_string_pretty(value).context("failed to serialize json output")? - ); - Ok(()) -} - -fn rocm_mcp_tools() -> Vec { - vec![ - rocm_mcp_tool( - "examine", - "Read the current ROCm AI Command Center host summary.", - json!({ - "type": "object", - "properties": {}, - "additionalProperties": false - }), - true, - false, - ), - rocm_mcp_tool( - "bridge_snapshot", - "Read the full ROCm bridge snapshot including examine data, engines, services, automations, and gpu telemetry.", - json!({ - "type": "object", - "properties": {}, - "additionalProperties": false - }), - true, - false, - ), - rocm_mcp_tool( - "gpu_snapshot", - "Read the current amd-smi GPU telemetry snapshot if available.", - json!({ - "type": "object", - "properties": {}, - "additionalProperties": false - }), - true, - false, - ), - rocm_mcp_tool( - "engines", - "List available ROCm serving engines and whether each one is installed.", - json!({ - "type": "object", - "properties": {}, - "additionalProperties": false - }), - true, - false, - ), - rocm_mcp_tool( - "services", - "List managed model services and their current status.", - json!({ - "type": "object", - "properties": {}, - "additionalProperties": false - }), - true, - false, - ), - rocm_mcp_tool( - "service_logs", - "Read the tail of a managed service log file.", - json!({ - "type": "object", - "properties": { - "service_id": { - "type": "string" - }, - "lines": { - "type": "integer", - "minimum": 1, - "maximum": 500 - } - }, - "required": ["service_id"], - "additionalProperties": false - }), - true, - false, - ), - rocm_mcp_tool( - "automations", - "List automation runtime status, watcher events, and local webhook events.", - json!({ - "type": "object", - "properties": { - "event_limit": { - "type": "integer", - "minimum": 1, - "maximum": 64 - } - }, - "additionalProperties": false - }), - true, - false, - ), - rocm_mcp_tool( - "natural_language_plan", - "Ask `rocm` to translate a natural-language ROCm request into a visible plan without executing privileged work.", - json!({ - "type": "object", - "properties": { - "request": { - "type": "string" - } - }, - "required": ["request"], - "additionalProperties": false - }), - true, - false, - ), - rocm_mcp_tool( - "rocm_command", - "Run a supported read-only ROCm CLI command with argv-style arguments. Commands that change ROCm state are rejected here and must go through the ROCm CLI approval UI.", - json!({ - "type": "object", - "properties": { - "args": { - "type": "array", - "items": { - "type": "string" - }, - "minItems": 1, - "maxItems": 64 - }, - "reason": { - "type": "string" - } - }, - "required": ["args"], - "additionalProperties": false - }), - true, - false, - ), - rocm_mcp_tool( - "update_check", - "Run `rocm update` and return the current TheRock update status.", - json!({ - "type": "object", - "properties": {}, - "additionalProperties": false - }), - true, - false, - ), - rocm_mcp_tool( - "install_sdk_dry_run", - "Run a dry-run TheRock SDK install plan.", - json!({ - "type": "object", - "properties": { - "channel": { - "type": "string", - "enum": ["release", "nightly"] - }, - "format": { - "type": "string", - "enum": ["wheel", "tarball"] - }, - "prefix": { - "type": "string" - }, - "version": { - "type": "string" - }, - "build_date": { - "type": "string" - } - }, - "additionalProperties": false - }), - true, - false, - ), - rocm_mcp_tool( - "install_sdk", - "Install a TheRock SDK into the managed runtime area or an explicitly approved prefix.", - json!({ - "type": "object", - "properties": { - "channel": { - "type": "string", - "enum": ["release", "nightly"] - }, - "format": { - "type": "string", - "enum": ["wheel", "tarball"] - }, - "prefix": { - "type": "string" - }, - "version": { - "type": "string" - }, - "build_date": { - "type": "string" - }, - "allow_system_prefix": { - "type": "boolean" - } - }, - "additionalProperties": false - }), - false, - true, - ), - rocm_mcp_tool( - "install_engine", - "Install or refresh a managed serving engine environment.", - json!({ - "type": "object", - "properties": { - "engine": { - "type": "string" - }, - "runtime_id": { - "type": "string" - }, - "python_version": { - "type": "string" - }, - "reinstall": { - "type": "boolean" - } - }, - "required": ["engine"], - "additionalProperties": false - }), - false, - false, - ), - rocm_mcp_tool( - "launch_server", - "Launch a managed local model server through `rocm serve --managed`.", - json!({ - "type": "object", - "properties": { - "model": { - "type": "string" - }, - "engine": { - "type": "string" - }, - "device": { - "type": "string" - }, - "runtime_id": { - "type": "string" - }, - "env_id": { - "type": "string" - }, - "host": { - "type": "string" - }, - "port": { - "type": "integer", - "minimum": 1, - "maximum": 65535 - }, - "allow_public_bind": { - "type": "boolean" - } - }, - "required": ["model"], - "additionalProperties": false - }), - false, - true, - ), - rocm_mcp_tool( - "stop_server", - "Stop a managed service by service id and update its manifest status.", - json!({ - "type": "object", - "properties": { - "service_id": { - "type": "string" - } - }, - "required": ["service_id"], - "additionalProperties": false - }), - false, - true, - ), - rocm_mcp_tool( - "watcher_enable", - "Enable a watcher and optionally set its mode.", - json!({ - "type": "object", - "properties": { - "watcher": { - "type": "string" - }, - "mode": { - "type": "string", - "enum": ["observe", "propose", "contained"] - } - }, - "required": ["watcher"], - "additionalProperties": false - }), - false, - false, - ), - rocm_mcp_tool( - "watcher_disable", - "Disable a watcher.", - json!({ - "type": "object", - "properties": { - "watcher": { - "type": "string" - } - }, - "required": ["watcher"], - "additionalProperties": false - }), - false, - false, - ), - ] -} - -fn rocm_mcp_tool( - name: &str, - description: &str, - input_schema: Value, - read_only: bool, - destructive: bool, -) -> Value { - json!({ - "name": name, - "title": name.replace('_', " "), - "description": description, - "annotations": { - "readOnlyHint": read_only, - "destructiveHint": destructive, - "openWorldHint": false, - }, - "inputSchema": input_schema, - }) -} - -fn mcp_tool_requires_direct_approval(name: &str) -> bool { - matches!( - name, - "install_sdk" - | "install_engine" - | "launch_server" - | "stop_server" - | "watcher_enable" - | "watcher_disable" - ) -} - -fn ensure_direct_mcp_call_allowed(name: &str, allow_mutation: bool) -> Result<()> { - if mcp_tool_requires_direct_approval(name) && !allow_mutation { - bail!( - "MCP tool `{name}` changes local ROCm state; rerun `rocmd mcp-call {name}` with --allow-mutation only after an explicit user approval" - ); - } - Ok(()) -} - -fn handle_mcp_tool_call(paths: &AppPaths, params: &Value) -> Result { - let name = params - .get("name") - .and_then(Value::as_str) - .unwrap_or_default(); - let arguments = params - .get("arguments") - .and_then(Value::as_object) - .cloned() - .unwrap_or_default(); - - match name { - "examine" => { - let examine = ExamineSummary::gather()?; - let output = common::run_rocm_capture(&["examine"])?; - let text = command_capture_text(&output); - if output.exit_status == 0 { - Ok(tool_success(text, json!(examine))) - } else { - Ok(tool_error( - text, - json!({ - "examine": examine, - "argv": output.argv, - "exit_status": output.exit_status, - "stderr": output.stderr, - }), - )) - } - } - "bridge_snapshot" => { - let snapshot = common::build_bridge_snapshot(paths)?; - Ok(tool_success( - format!( - "Captured bridge snapshot for {} / {} with default engine `{}`.", - snapshot.examine.os, snapshot.examine.arch, snapshot.examine.default_engine - ), - json!(snapshot), - )) - } - "gpu_snapshot" => { - let config = RocmCliConfig::load(paths).unwrap_or_default(); - let gpu = common::gather_gpu_snapshot_for_config(&config); - let status = if !config.telemetry.local_inspection_enabled() { - "GPU telemetry is disabled by rocm-cli config." - } else if gpu.amd_smi_available { - "Captured amd-smi GPU snapshot." - } else { - "amd-smi is unavailable on this host." - }; - Ok(tool_success(status.to_owned(), json!(gpu))) - } - "engines" => { - let engines = common::bridge_engine_inventory(); - Ok(tool_success( - format!("Found {} engine entries.", engines.len()), - json!({ "engines": engines }), - )) - } - "services" => { - let services = persistence::load_managed_services(paths)?; - Ok(tool_success( - format!("Found {} managed services.", services.len()), - json!({ "services": services }), - )) - } - "service_logs" => { - let service_id = arguments - .get("service_id") - .and_then(Value::as_str) - .context("service_logs requires `service_id`")?; - let lines = arguments - .get("lines") - .and_then(Value::as_u64) - .unwrap_or(80) - .clamp(1, 500) as usize; - let record = persistence::load_managed_services(paths)? - .into_iter() - .find(|service| service.service_id == service_id) - .with_context(|| format!("managed service `{service_id}` not found"))?; - let tail = read_tail_lines(&record.log_path, lines)?; - Ok(tool_success( - format!( - "Read the last {} line(s) from service `{}`.", - lines, record.service_id - ), - json!({ - "service": record, - "lines": lines, - "tail": tail, - }), - )) - } - "automations" => { - let event_limit = arguments - .get("event_limit") - .and_then(Value::as_u64) - .unwrap_or(10) - .clamp(1, 64) as usize; - let runtime = AutomationRuntimeState::load(paths)?; - let events = load_recent_automation_events(paths, event_limit)?; - Ok(tool_success( - format!( - "Loaded automation runtime and {} recent events.", - events.len() - ), - json!({ - "runtime": runtime, - "recent_events": events, - }), - )) - } - "natural_language_plan" => { - let request = arguments - .get("request") - .and_then(Value::as_str) - .context("natural_language_plan requires `request`")?; - let output = common::run_rocm_capture(&[request])?; - Ok(tool_result_from_command( - "Ran natural-language planning through `rocm`.", - output, - false, - )) - } - "rocm_command" => { - let argv = normalized_rocm_command_args(&arguments)?; - ensure_rocm_command_is_read_only(&argv)?; - let refs = argv.iter().map(String::as_str).collect::>(); - let output = common::run_rocm_capture(&refs)?; - Ok(tool_result_from_command( - "Ran read-only `rocm` command.", - output, - false, - )) - } - "update_check" => { - let output = common::run_rocm_capture(&["update"])?; - Ok(tool_result_from_command( - "Ran `rocm update`.", - output, - false, - )) - } - "install_sdk_dry_run" => { - let argv = build_install_sdk_args(&arguments, true)?; - let refs = argv.iter().map(String::as_str).collect::>(); - let output = common::run_rocm_capture(&refs)?; - Ok(tool_result_from_command( - "Ran `rocm install sdk --dry-run`.", - output, - false, - )) - } - "install_sdk" => { - let argv = build_install_sdk_args(&arguments, false)?; - let refs = argv.iter().map(String::as_str).collect::>(); - let output = common::run_rocm_capture(&refs)?; - Ok(tool_result_from_command( - "Ran `rocm install sdk`.", - output, - false, - )) - } - "install_engine" => { - let argv = build_install_engine_args(&arguments)?; - let refs = argv.iter().map(String::as_str).collect::>(); - let output = common::run_rocm_capture(&refs)?; - Ok(tool_result_from_command( - "Ran `rocm engines install`.", - output, - false, - )) - } - "launch_server" => { - let argv = build_launch_server_args(&arguments)?; - let refs = argv.iter().map(String::as_str).collect::>(); - let output = common::run_rocm_capture(&refs)?; - Ok(tool_result_from_command( - "Ran `rocm serve --managed`.", - output, - false, - )) - } - "stop_server" => { - let service_id = arguments - .get("service_id") - .and_then(Value::as_str) - .context("stop_server requires `service_id`")?; - let stopped = stop_managed_service(paths, service_id)?; - Ok(tool_success( - format!("Stopped managed service `{service_id}`."), - stopped, - )) - } - "watcher_enable" => { - let argv = build_watcher_enable_args(&arguments)?; - let refs = argv.iter().map(String::as_str).collect::>(); - let output = common::run_rocm_capture(&refs)?; - Ok(tool_result_from_command( - "Ran `rocm automations enable`.", - output, - false, - )) - } - "watcher_disable" => { - let watcher = arguments - .get("watcher") - .and_then(Value::as_str) - .context("watcher_disable requires `watcher`")?; - let argv = [ - "automations".to_owned(), - "disable".to_owned(), - watcher.to_owned(), - ]; - let refs = argv.iter().map(String::as_str).collect::>(); - let output = common::run_rocm_capture(&refs)?; - Ok(tool_result_from_command( - "Ran `rocm automations disable`.", - output, - false, - )) - } - other => Ok(tool_error( - format!("Unknown ROCm MCP tool `{other}`."), - json!({ "tool": other }), - )), - } -} - -fn tool_success(text: String, structured: Value) -> Value { - json!({ - "content": [ - { - "type": "text", - "text": text, - } - ], - "structuredContent": structured, - "isError": false, - }) -} - -fn tool_error(text: String, structured: Value) -> Value { - json!({ - "content": [ - { - "type": "text", - "text": text, - } - ], - "structuredContent": structured, - "isError": true, - }) -} - -fn tool_result_from_command(prefix: &str, output: common::CommandCapture, is_error: bool) -> Value { - let text = format!("{prefix}\n\n{}", command_capture_text(&output)); - json!({ - "content": [ - { - "type": "text", - "text": text, - } - ], - "structuredContent": { - "argv": output.argv, - "exit_status": output.exit_status, - "stdout": output.stdout, - "stderr": output.stderr, - }, - "isError": is_error || output.exit_status != 0, - }) -} - -fn command_capture_text(output: &common::CommandCapture) -> String { - if output.stderr.trim().is_empty() { - output.stdout.trim().to_owned() - } else if output.stdout.trim().is_empty() { - format!("stderr:\n{}", output.stderr.trim()) - } else { - format!( - "stdout:\n{}\n\nstderr:\n{}", - output.stdout.trim(), - output.stderr.trim() - ) - } -} - -fn read_tail_lines(path: &std::path::Path, limit: usize) -> Result { - let content = - fs::read_to_string(path).with_context(|| format!("failed to read {}", path.display()))?; - let mut lines = VecDeque::with_capacity(limit); - for line in content.lines() { - if lines.len() == limit { - lines.pop_front(); - } - lines.push_back(line.to_owned()); - } - Ok(lines.into_iter().collect::>().join("\n")) -} - -fn normalized_rocm_command_args(arguments: &serde_json::Map) -> Result> { - let values = arguments - .get("args") - .and_then(Value::as_array) - .context("rocm_command requires `args`")?; - if values.is_empty() || values.len() > 64 { - bail!("rocm_command `args` must contain 1 to 64 strings"); - } - let mut args = Vec::with_capacity(values.len()); - for value in values { - let arg = value - .as_str() - .map(str::trim) - .filter(|value| !value.is_empty()) - .context("rocm_command `args` entries must be non-empty strings")?; - if arg.contains('\0') || arg.contains('\n') || arg.contains('\r') { - bail!("rocm_command arguments must not contain control characters"); - } - if arg.len() > 512 { - bail!("rocm_command argument is too long"); - } - args.push(arg.to_owned()); - } - if args - .first() - .is_some_and(|arg| arg.eq_ignore_ascii_case("rocm")) - { - args.remove(0); - } - if args - .first() - .is_some_and(|arg| arg.eq_ignore_ascii_case("comfy")) - { - args[0] = "comfyui".to_owned(); - } - if args.is_empty() { - bail!("rocm_command args should omit the leading `rocm` program name"); - } - Ok(args) -} - -fn ensure_rocm_command_is_read_only(args: &[String]) -> Result<()> { - let first = args.first().map(|value| value.to_ascii_lowercase()); - let second = args.get(1).map(|value| value.to_ascii_lowercase()); - let read_only = match first.as_deref() { - Some("examine" | "version" | "model" | "models" | "daemon" | "logs") => true, - Some("update") => !args.iter().any(|arg| arg == "--apply"), - Some("runtimes") => { - second.as_deref().is_none_or(|value| value == "list") - || (second - .as_deref() - .is_some_and(|value| value == "uninstall" || value == "remove") - && args.iter().any(|arg| arg == "--dry-run")) - } - Some("engines") => second.as_deref().is_some_and(|value| value == "list"), - Some("services") => second - .as_deref() - .is_none_or(|value| matches!(value, "list" | "logs")), - Some("automations") => second.as_deref().is_none_or(|value| value == "list"), - Some("config") => second.as_deref() == Some("show"), - Some("comfyui") => second - .as_deref() - .is_none_or(|value| matches!(value, "status" | "logs" | "log")), - Some("uninstall") => args.iter().any(|arg| arg == "--dry-run"), - // `storage report` (the default subcommand) only measures folders. The - // two `remove-*` verbs delete, so they stay off the read-only list. - Some("storage") => second.as_deref().is_none_or(|value| value == "report"), - // `setup status` reports first-time setup state (read-only); `setup reset` - // clears the completion/dismissal state and is mutating (it does not by - // itself reopen onboarding). Mirrors the bin's rocm_command classifier so - // the read-only allowlist is consistent across binaries. - Some("setup") => second.as_deref().is_none_or(|value| value == "status"), - // `remote targets` reads the local tailnet, `doctor` fetches another - // machine's state and scores it here, `status` probes sessions that - // already exist. None of them change anything on either machine. - // `serve`, `attach` and `stop` start, publish or tear down, so they stay - // off the list and go through the approval UI like any other mutation. - Some("remote") => second - .as_deref() - .is_some_and(|value| matches!(value, "targets" | "doctor" | "status")), - _ => false, - }; - if read_only { - return Ok(()); - } - bail!( - "rocm_command changes local ROCm state or is unsupported here; request it through the ROCm CLI approval UI instead" - ) -} - -fn build_install_sdk_args( - arguments: &serde_json::Map, - dry_run: bool, -) -> Result> { - let channel = arguments - .get("channel") - .and_then(Value::as_str) - .unwrap_or("release"); - let format = arguments - .get("format") - .and_then(Value::as_str) - .unwrap_or("wheel"); - let prefix = arguments.get("prefix").and_then(Value::as_str); - let version = arguments.get("version").and_then(Value::as_str); - let build_date = arguments.get("build_date").and_then(Value::as_str); - let allow_system_prefix = arguments - .get("allow_system_prefix") - .and_then(Value::as_bool) - .unwrap_or(false); - if version.is_some() && build_date.is_some() { - bail!("install_sdk accepts either `version` or `build_date`, not both"); - } - - let mut argv = vec![ - "install".to_owned(), - "sdk".to_owned(), - "--channel".to_owned(), - channel.to_owned(), - "--format".to_owned(), - format.to_owned(), - ]; - if let Some(prefix) = prefix { - let prefix_path = std::path::Path::new(prefix); - if system_prefix_requires_ack(prefix_path) && !allow_system_prefix { - bail!( - "install_sdk prefix `{}` is outside the user home; require `allow_system_prefix=true` before using system paths", - prefix_path.display() - ); - } - argv.push("--prefix".to_owned()); - argv.push(prefix.to_owned()); - } - if let Some(version) = version { - if version.trim().is_empty() { - bail!("install_sdk `version` cannot be empty"); - } - argv.push("--version".to_owned()); - argv.push(version.to_owned()); - } - if let Some(build_date) = build_date { - if build_date.trim().is_empty() { - bail!("install_sdk `build_date` cannot be empty"); - } - argv.push("--build-date".to_owned()); - argv.push(build_date.to_owned()); - } - if dry_run { - argv.push("--dry-run".to_owned()); - } else { - // `run_rocm_capture_for_paths` spawns `rocm` with null stdin, so - // `interactive_terminal()` is false in the child and an active default - // managed runtime would make the approval gate refuse with "re-run with - // `--approve-replacing-active-default`" — a flag no MCP caller of this - // tool can supply. - // - // Not `--yes` itself: that flag carries a second, unrelated consent — - // approving required system-package installs, which run `sudo`. This - // spawn has no terminal, so it could never answer a sudo password - // prompt; granting that consent would make the vLLM/OpenMPI step attempt - // an install it cannot complete and abort the engine auto-install that - // previously warned and continued. `--approve-replacing-active-default` - // grants only the runtime-displacement consent the gate asks for. - // - // Consent is not bypassed: `install_sdk` is in - // `mcp_tool_requires_direct_approval`, so a direct `rocmd mcp-call` - // needs `--allow-mutation` after an explicit user approval, and over the - // MCP protocol the tool is annotated `destructiveHint` for the client's - // approval UI. Mirrors the chat/MCP arm in `apps/rocm`. The dry-run - // branch never reaches the gate (it returns earlier), so it stays bare. - argv.push("--approve-replacing-active-default".to_owned()); - } - Ok(argv) -} - -fn build_install_engine_args(arguments: &serde_json::Map) -> Result> { - let engine = arguments - .get("engine") - .and_then(Value::as_str) - .context("install_engine requires `engine`")?; - let runtime_id = arguments - .get("runtime_id") - .and_then(Value::as_str) - .unwrap_or("therock-release"); - let python_version = arguments.get("python_version").and_then(Value::as_str); - let reinstall = arguments - .get("reinstall") - .and_then(Value::as_bool) - .unwrap_or(false); - - let mut argv = vec![ - "engines".to_owned(), - "install".to_owned(), - engine.to_owned(), - "--runtime-id".to_owned(), - runtime_id.to_owned(), - ]; - if let Some(python_version) = python_version { - argv.push("--python-version".to_owned()); - argv.push(python_version.to_owned()); - } - if reinstall { - argv.push("--reinstall".to_owned()); - } - Ok(argv) -} - -fn build_launch_server_args(arguments: &serde_json::Map) -> Result> { - let model = arguments - .get("model") - .and_then(Value::as_str) - .context("launch_server requires `model`")?; - let host = arguments - .get("host") - .and_then(Value::as_str) - .unwrap_or(DEFAULT_LOCAL_HOST); - let allow_public_bind = arguments - .get("allow_public_bind") - .and_then(Value::as_bool) - .unwrap_or(false); - if !is_loopback_host(host) && !allow_public_bind { - bail!( - "launch_server host `{host}` is not loopback; require `allow_public_bind=true` before binding a non-local interface" - ); - } - - let mut argv = vec!["serve".to_owned(), model.to_owned(), "--managed".to_owned()]; - if let Some(engine) = arguments.get("engine").and_then(Value::as_str) { - argv.push("--engine".to_owned()); - argv.push(engine.to_owned()); - } - if let Some(device) = arguments.get("device").and_then(Value::as_str) { - argv.push("--device".to_owned()); - argv.push(device.to_owned()); - } - if let Some(runtime_id) = arguments.get("runtime_id").and_then(Value::as_str) { - argv.push("--runtime-id".to_owned()); - argv.push(runtime_id.to_owned()); - } - if let Some(env_id) = arguments.get("env_id").and_then(Value::as_str) { - argv.push("--env-id".to_owned()); - argv.push(env_id.to_owned()); - } - argv.push("--host".to_owned()); - argv.push(host.to_owned()); - if allow_public_bind { - argv.push("--allow-public-bind".to_owned()); - } - if let Some(port) = arguments.get("port").and_then(Value::as_u64) { - argv.push("--port".to_owned()); - argv.push(port.to_string()); - } - Ok(argv) -} - -fn build_watcher_enable_args(arguments: &serde_json::Map) -> Result> { - let watcher = arguments - .get("watcher") - .and_then(Value::as_str) - .context("watcher_enable requires `watcher`")?; - let mut argv = vec![ - "automations".to_owned(), - "enable".to_owned(), - watcher.to_owned(), - ]; - if let Some(mode) = arguments.get("mode").and_then(Value::as_str) { - argv.push("--mode".to_owned()); - argv.push(mode.to_owned()); - } - Ok(argv) -} - -/// Inverse of [`rocm_engine_protocol::is_public_bind_host`], which owns the -/// policy so `rocm` and `rocmd` never classify the same host differently. -fn is_loopback_host(host: &str) -> bool { - !rocm_engine_protocol::is_public_bind_host(host) -} - -fn system_prefix_requires_ack(prefix: &std::path::Path) -> bool { - match rocm_core::runtime_home_dir() { - Some(home) => !rocm_core::runtime_path_is_same_or_inside(prefix, &home), - None => true, - } -} - fn stop_managed_service(paths: &AppPaths, service_id: &str) -> Result { let mut record = persistence::load_managed_services(paths)? .into_iter() @@ -3449,38 +2328,7 @@ fn detached_rocmd_command(rocmd_binary: &std::path::Path) -> ProcessCommand { #[cfg(test)] mod tests { use super::*; - use crate::test_support::{temp_app_paths, unique_test_root, workspace_test_artifact_dir}; - use std::path::PathBuf; - - #[test] - fn remote_read_only_verbs_are_allowed_and_mutating_ones_are_not() { - let allow = |args: &[&str]| { - let owned = args.iter().map(|a| (*a).to_owned()).collect::>(); - super::ensure_rocm_command_is_read_only(&owned) - }; - - // These read: the local tailnet, another machine's state, sessions that - // already exist. Rejecting them made the whole family unusable here even - // though none of them change anything. - for args in [ - &["remote", "targets"][..], - &["remote", "targets", "--tag", "gpu"][..], - &["remote", "doctor", "gpu-box"][..], - &["remote", "status"][..], - ] { - allow(args).unwrap_or_else(|error| panic!("{args:?} should be read-only: {error:#}")); - } - - // These start, publish or tear down, so they go through approval. - for args in [ - &["remote", "serve", "gpu-box", "a-model"][..], - &["remote", "attach", "sess"][..], - &["remote", "stop", "sess"][..], - &["remote"][..], - ] { - assert!(allow(args).is_err(), "{args:?} must not be read-only"); - } - } + use crate::test_support::{temp_app_paths, unique_test_root}; /// Drive `supervise_service` far enough to reach the key guard, and return /// what it did. @@ -3877,373 +2725,6 @@ mod tests { ); } - #[test] - fn rocm_mcp_tools_include_bridge_gaps() { - let tools = rocm_mcp_tools(); - let names = tools - .iter() - .filter_map(|tool| tool.get("name").and_then(Value::as_str).map(str::to_owned)) - .collect::>(); - assert!(names.contains(&"gpu_snapshot".to_owned())); - assert!(names.contains(&"service_logs".to_owned())); - assert!(names.contains(&"natural_language_plan".to_owned())); - assert!(names.contains(&"rocm_command".to_owned())); - assert!(names.contains(&"install_sdk".to_owned())); - assert!(names.contains(&"install_engine".to_owned())); - assert!(names.contains(&"launch_server".to_owned())); - assert!(names.contains(&"stop_server".to_owned())); - assert!(names.contains(&"watcher_enable".to_owned())); - assert!(names.contains(&"watcher_disable".to_owned())); - let automations = tools - .iter() - .find(|tool| tool.get("name").and_then(Value::as_str) == Some("automations")) - .expect("automations tool should be present"); - assert!( - automations - .get("description") - .and_then(Value::as_str) - .is_some_and(|description| description.contains("local webhook events")) - ); - } - - #[test] - fn direct_mcp_call_requires_approval_for_every_mutating_tool() { - for tool in rocm_mcp_tools() { - let name = tool - .get("name") - .and_then(Value::as_str) - .expect("tool should have a name"); - let read_only = tool - .get("annotations") - .and_then(|annotations| annotations.get("readOnlyHint")) - .and_then(Value::as_bool) - .unwrap_or(false); - assert_eq!( - mcp_tool_requires_direct_approval(name), - !read_only, - "hidden direct MCP helper approval classification drifted for `{name}`" - ); - } - } - - #[test] - fn direct_mcp_call_guard_blocks_mutation_without_explicit_ack() { - ensure_direct_mcp_call_allowed("examine", false) - .expect("read-only direct MCP helper calls should not need mutation approval"); - - let error = ensure_direct_mcp_call_allowed("install_sdk", false) - .expect_err("mutating direct MCP helper calls should require approval"); - assert!(error.to_string().contains("--allow-mutation"), "{error:#}"); - - ensure_direct_mcp_call_allowed("install_sdk", true) - .expect("explicitly approved direct MCP mutation should pass the helper guard"); - } - - #[test] - fn rocm_command_helper_allows_only_read_only_rocm_commands() -> Result<()> { - let status_args = normalized_rocm_command_args( - serde_json::json!({ - "args": ["rocm", "comfy", "status"] - }) - .as_object() - .expect("json object"), - )?; - assert_eq!(status_args, vec!["comfyui".to_owned(), "status".to_owned()]); - ensure_rocm_command_is_read_only(&status_args).expect("ComfyUI status should be read-only"); - - let log_args = normalized_rocm_command_args( - serde_json::json!({ - "args": ["comfyui", "logs"] - }) - .as_object() - .expect("json object"), - )?; - ensure_rocm_command_is_read_only(&log_args).expect("ComfyUI logs should be read-only"); - - let install_args = normalized_rocm_command_args( - serde_json::json!({ - "args": ["comfyui", "install"] - }) - .as_object() - .expect("json object"), - )?; - let error = ensure_rocm_command_is_read_only(&install_args) - .expect_err("ComfyUI install must go through approval"); - assert!(error.to_string().contains("approval UI")); - - let shell_args = normalized_rocm_command_args( - serde_json::json!({ - "args": ["powershell", "-Command", "whoami"] - }) - .as_object() - .expect("json object"), - )?; - let error = ensure_rocm_command_is_read_only(&shell_args) - .expect_err("non-rocm shell commands should be rejected"); - assert!(error.to_string().contains("approval UI")); - Ok(()) - } - - #[test] - fn rocm_command_helper_treats_setup_status_as_read_only_and_reset_as_mutating() -> Result<()> { - // Mirrors the bin's rocm_command classifier so `setup status` is read-only - // on every binary's tool surface while `setup reset` stays approval-gated. - let bare_args = normalized_rocm_command_args( - serde_json::json!({ "args": ["setup"] }) - .as_object() - .expect("json object"), - )?; - ensure_rocm_command_is_read_only(&bare_args).expect("bare setup should be read-only"); - - let status_args = normalized_rocm_command_args( - serde_json::json!({ "args": ["setup", "status"] }) - .as_object() - .expect("json object"), - )?; - ensure_rocm_command_is_read_only(&status_args).expect("setup status should be read-only"); - - let reset_args = normalized_rocm_command_args( - serde_json::json!({ "args": ["setup", "reset"] }) - .as_object() - .expect("json object"), - )?; - let error = ensure_rocm_command_is_read_only(&reset_args) - .expect_err("setup reset must go through approval"); - assert!(error.to_string().contains("approval UI")); - Ok(()) - } - - #[test] - fn rocm_command_helper_treats_runtimes_uninstall_dry_run_as_read_only() -> Result<()> { - // Mirrors the bin's chat_rocm_command_action_from_args classifier so a - // dry-run preview stays read-only on every binary's tool surface while - // an actual uninstall/remove still requires approval. - for verb in ["uninstall", "remove"] { - let dry_run_args = normalized_rocm_command_args( - serde_json::json!({ "args": ["runtimes", verb, "--dry-run"] }) - .as_object() - .expect("json object"), - )?; - ensure_rocm_command_is_read_only(&dry_run_args) - .unwrap_or_else(|_| panic!("runtimes {verb} --dry-run should be read-only")); - - let mutating_args = normalized_rocm_command_args( - serde_json::json!({ "args": ["runtimes", verb] }) - .as_object() - .expect("json object"), - )?; - let error = match ensure_rocm_command_is_read_only(&mutating_args) { - Ok(()) => panic!("runtimes {verb} without --dry-run must go through approval"), - Err(error) => error, - }; - assert!(error.to_string().contains("approval UI")); - } - Ok(()) - } - - #[test] - fn storage_report_is_read_only_but_removal_is_not() -> Result<()> { - for args in [vec!["storage"], vec!["storage", "report"]] { - let normalized = normalized_rocm_command_args( - serde_json::json!({ "args": args }) - .as_object() - .expect("json object"), - )?; - ensure_rocm_command_is_read_only(&normalized) - .unwrap_or_else(|_| panic!("storage {args:?} only measures folders")); - } - - for verb in ["remove-old-installs", "remove-downloads"] { - let normalized = normalized_rocm_command_args( - serde_json::json!({ "args": ["storage", verb] }) - .as_object() - .expect("json object"), - )?; - let error = ensure_rocm_command_is_read_only(&normalized) - .expect_err("storage removal must go through approval"); - assert!(error.to_string().contains("approval UI")); - } - Ok(()) - } - - #[test] - fn read_tail_lines_returns_last_lines_only() -> Result<()> { - let path = unique_test_path(&format!( - "rocmd-tail-test-{}-{}.log", - std::process::id(), - unix_time_millis() - )); - fs::write(&path, "line1\nline2\nline3\nline4\n")?; - let tail = read_tail_lines(&path, 2)?; - fs::remove_file(&path)?; - assert_eq!(tail, "line3\nline4"); - Ok(()) - } - - #[test] - fn launch_server_rejects_public_bind_without_ack() { - let arguments = serde_json::Map::from_iter([ - ("model".to_owned(), Value::String("tiny-gpt2".to_owned())), - ("host".to_owned(), Value::String("0.0.0.0".to_owned())), - ]); - let error = build_launch_server_args(&arguments).unwrap_err(); - assert!( - error.to_string().contains("allow_public_bind=true"), - "{error:#}" - ); - } - - #[test] - fn launch_server_forwards_public_bind_ack() -> Result<()> { - let arguments = serde_json::Map::from_iter([ - ("model".to_owned(), Value::String("tiny-gpt2".to_owned())), - ("host".to_owned(), Value::String("0.0.0.0".to_owned())), - ("allow_public_bind".to_owned(), Value::Bool(true)), - ]); - let args = build_launch_server_args(&arguments)?; - assert!(args.contains(&"--allow-public-bind".to_owned())); - Ok(()) - } - - #[test] - fn install_sdk_rejects_system_prefix_without_ack() { - let arguments = serde_json::Map::from_iter([( - "prefix".to_owned(), - Value::String("/opt/rocm".to_owned()), - )]); - let error = build_install_sdk_args(&arguments, false).unwrap_err(); - assert!( - error.to_string().contains("allow_system_prefix=true"), - "{error:#}" - ); - } - - /// The test above only ever hands `system_prefix_requires_ack` an - /// already-canonical path, so it cannot catch the bug this crate's fix - /// addresses: a `..`-respelled prefix that escapes `$HOME` used to compare - /// equal to a path still inside it (`Path::ancestors()` treats `..` as an - /// ordinary component), so acknowledgement was never required. Drive the - /// same check with a prefix built by walking `..` out of the real home - /// directory, which is exactly the shape the original bug let through. - #[test] - #[cfg(unix)] - fn install_sdk_rejects_system_prefix_reached_by_escaping_home() { - let home = rocm_core::runtime_home_dir().expect("a home directory"); - let escaped_prefix = format!("{}/../../usr", home.display()); - - let arguments = - serde_json::Map::from_iter([("prefix".to_owned(), Value::String(escaped_prefix))]); - let error = build_install_sdk_args(&arguments, false).unwrap_err(); - assert!( - error.to_string().contains("allow_system_prefix=true"), - "{error:#}" - ); - } - - /// The `install_sdk` MCP tool spawns `rocm` with null stdin, so a real - /// install over an active default managed runtime would hit the approval - /// gate's non-interactive refusal and bail asking for a flag no MCP caller - /// can pass. The real-install argv must therefore carry the consent flag; - /// the dry-run argv must not, because a dry run never reaches the gate and - /// the flag there would claim an approval the caller did not give. - /// - /// It must be `--approve-replacing-active-default` and never `--yes`: - /// `--yes` additionally approves running `sudo` for required system - /// packages, and a null-stdin spawn has no terminal on which that password - /// prompt could be answered. - #[test] - fn install_sdk_real_install_args_approve_only_the_runtime_replacement() -> Result<()> { - let arguments = serde_json::Map::new(); - - let real = build_install_sdk_args(&arguments, false)?; - assert!( - real.contains(&"--approve-replacing-active-default".to_owned()), - "real install argv must approve the replacement for the null-stdin spawn: {real:?}" - ); - assert!( - !real.contains(&"--yes".to_owned()), - "real install argv must not grant the system-package consent it cannot answer: {real:?}" - ); - assert!( - !real.contains(&"--dry-run".to_owned()), - "real install argv must not be a dry run: {real:?}" - ); - - let dry = build_install_sdk_args(&arguments, true)?; - assert!( - !dry.contains(&"--approve-replacing-active-default".to_owned()) - && !dry.contains(&"--yes".to_owned()), - "dry-run argv must not carry a consent flag: {dry:?}" - ); - assert!( - dry.contains(&"--dry-run".to_owned()), - "dry-run argv must carry --dry-run: {dry:?}" - ); - Ok(()) - } - - #[test] - fn install_sdk_forwards_requested_build_date_and_rejects_conflict() -> Result<()> { - let arguments = serde_json::Map::from_iter([( - "build_date".to_owned(), - Value::String("2026-06-05".to_owned()), - )]); - let argv = build_install_sdk_args(&arguments, true)?; - assert_eq!( - argv, - vec![ - "install".to_owned(), - "sdk".to_owned(), - "--channel".to_owned(), - "release".to_owned(), - "--format".to_owned(), - "wheel".to_owned(), - "--build-date".to_owned(), - "2026-06-05".to_owned(), - "--dry-run".to_owned(), - ] - ); - - let conflicting = serde_json::Map::from_iter([ - ( - "version".to_owned(), - Value::String("7.13.0a20260605".to_owned()), - ), - ( - "build_date".to_owned(), - Value::String("2026-06-05".to_owned()), - ), - ]); - let error = build_install_sdk_args(&conflicting, false) - .unwrap_err() - .to_string(); - assert!(error.contains("either `version` or `build_date`")); - Ok(()) - } - - #[test] - fn watcher_enable_builds_mode_args() -> Result<()> { - let arguments = serde_json::Map::from_iter([ - ( - "watcher".to_owned(), - Value::String("server-recover".to_owned()), - ), - ("mode".to_owned(), Value::String("contained".to_owned())), - ]); - let argv = build_watcher_enable_args(&arguments)?; - assert_eq!( - argv, - vec![ - "automations".to_owned(), - "enable".to_owned(), - "server-recover".to_owned(), - "--mode".to_owned(), - "contained".to_owned() - ] - ); - Ok(()) - } - #[test] fn watcher_policy_maps_modes_to_decisions() { assert_eq!( @@ -5882,12 +4363,6 @@ mod tests { ); } - fn unique_test_path(label: &str) -> PathBuf { - let root = workspace_test_artifact_dir(); - fs::create_dir_all(&root).expect("create workspace-local test dir"); - root.join(label) - } - fn test_watcher_snapshot( id: &str, mode: WatcherMode, diff --git a/apps/rocmd/src/mcp.rs b/apps/rocmd/src/mcp.rs new file mode 100644 index 000000000..e6995aca0 --- /dev/null +++ b/apps/rocmd/src/mcp.rs @@ -0,0 +1,1549 @@ +// Copyright © Advanced Micro Devices, Inc., or its affiliates. +// +// SPDX-License-Identifier: MIT + +use crate::common::{self, CommandCapture}; +use crate::persistence; +use anyhow::{Context, Result, bail}; +#[cfg(test)] +use rocm_core::unix_time_millis; +use rocm_core::{ + AppPaths, AutomationRuntimeState, DEFAULT_LOCAL_HOST, ExamineSummary, RocmCliConfig, + load_recent_automation_events, +}; +use serde::Serialize; +use serde_json::Value; +use serde_json::json; +use std::collections::VecDeque; +use std::fs; +use std::io::{self, BufRead, Write}; + +pub(crate) fn run_mcp_server(paths: &AppPaths) -> Result<()> { + let stdin = io::stdin(); + let stdout = io::stdout(); + let mut reader = stdin.lock(); + let mut writer = stdout.lock(); + let mut line = String::new(); + + loop { + line.clear(); + let bytes_read = reader.read_line(&mut line)?; + if bytes_read == 0 { + break; + } + if line.trim().is_empty() { + continue; + } + + let message: Value = match serde_json::from_str(&line) { + Ok(value) => value, + Err(error) => { + write_json_line( + &mut writer, + &json!({ + "jsonrpc": "2.0", + "error": { + "code": -32700, + "message": format!("parse error: {error}"), + } + }), + )?; + continue; + } + }; + + let Some(method) = message.get("method").and_then(Value::as_str) else { + if message.get("id").is_some() { + write_json_line( + &mut writer, + &json!({ + "jsonrpc": "2.0", + "id": message.get("id").cloned().unwrap_or(Value::Null), + "error": { + "code": -32600, + "message": "invalid request: missing method", + } + }), + )?; + } + continue; + }; + + let id = message.get("id").cloned(); + let params = message.get("params").cloned().unwrap_or(Value::Null); + match method { + "initialize" => { + let protocol_version = params + .get("protocolVersion") + .and_then(Value::as_str) + .unwrap_or("2025-03-26"); + if let Some(id) = id { + write_json_line( + &mut writer, + &json!({ + "jsonrpc": "2.0", + "id": id, + "result": { + "protocolVersion": protocol_version, + "capabilities": { + "tools": { + "listChanged": true, + } + }, + "serverInfo": { + "name": "rocmd-mcp-server", + "title": "ROCm AI Command Center", + "version": env!("CARGO_PKG_VERSION"), + } + } + }), + )?; + } + } + "ping" => { + if let Some(id) = id { + write_json_line( + &mut writer, + &json!({ + "jsonrpc": "2.0", + "id": id, + "result": {} + }), + )?; + } + } + "notifications/initialized" => {} + "tools/list" => { + if let Some(id) = id { + write_json_line( + &mut writer, + &json!({ + "jsonrpc": "2.0", + "id": id, + "result": { + "tools": rocm_mcp_tools(), + "nextCursor": Value::Null, + } + }), + )?; + } + } + "tools/call" => { + if let Some(id) = id { + let result = match handle_mcp_tool_call(paths, ¶ms) { + Ok(result) => result, + Err(error) => tool_error( + format!("ROCm MCP tool call failed: {error:#}"), + json!({ + "tool": params.get("name").cloned().unwrap_or(Value::Null), + }), + ), + }; + write_json_line( + &mut writer, + &json!({ + "jsonrpc": "2.0", + "id": id, + "result": result, + }), + )?; + } + } + notification if notification.starts_with("notifications/") => {} + other => { + if let Some(id) = id { + write_json_line( + &mut writer, + &json!({ + "jsonrpc": "2.0", + "id": id, + "error": { + "code": -32601, + "message": format!("method not found: {other}"), + } + }), + )?; + } + } + } + } + + Ok(()) +} + +fn write_json_line(writer: &mut impl Write, value: &Value) -> Result<()> { + writer.write_all(serde_json::to_string(value)?.as_bytes())?; + writer.write_all(b"\n")?; + writer.flush()?; + Ok(()) +} + +pub(crate) fn print_json(value: &T) -> Result<()> { + println!( + "{}", + serde_json::to_string_pretty(value).context("failed to serialize json output")? + ); + Ok(()) +} + +pub(crate) fn rocm_mcp_tools() -> Vec { + vec![ + rocm_mcp_tool( + "examine", + "Read the current ROCm AI Command Center host summary.", + json!({ + "type": "object", + "properties": {}, + "additionalProperties": false + }), + true, + false, + ), + rocm_mcp_tool( + "bridge_snapshot", + "Read the full ROCm bridge snapshot including examine data, engines, services, automations, and gpu telemetry.", + json!({ + "type": "object", + "properties": {}, + "additionalProperties": false + }), + true, + false, + ), + rocm_mcp_tool( + "gpu_snapshot", + "Read the current amd-smi GPU telemetry snapshot if available.", + json!({ + "type": "object", + "properties": {}, + "additionalProperties": false + }), + true, + false, + ), + rocm_mcp_tool( + "engines", + "List available ROCm serving engines and whether each one is installed.", + json!({ + "type": "object", + "properties": {}, + "additionalProperties": false + }), + true, + false, + ), + rocm_mcp_tool( + "services", + "List managed model services and their current status.", + json!({ + "type": "object", + "properties": {}, + "additionalProperties": false + }), + true, + false, + ), + rocm_mcp_tool( + "service_logs", + "Read the tail of a managed service log file.", + json!({ + "type": "object", + "properties": { + "service_id": { + "type": "string" + }, + "lines": { + "type": "integer", + "minimum": 1, + "maximum": 500 + } + }, + "required": ["service_id"], + "additionalProperties": false + }), + true, + false, + ), + rocm_mcp_tool( + "automations", + "List automation runtime status, watcher events, and local webhook events.", + json!({ + "type": "object", + "properties": { + "event_limit": { + "type": "integer", + "minimum": 1, + "maximum": 64 + } + }, + "additionalProperties": false + }), + true, + false, + ), + rocm_mcp_tool( + "natural_language_plan", + "Ask `rocm` to translate a natural-language ROCm request into a visible plan without executing privileged work.", + json!({ + "type": "object", + "properties": { + "request": { + "type": "string" + } + }, + "required": ["request"], + "additionalProperties": false + }), + true, + false, + ), + rocm_mcp_tool( + "rocm_command", + "Run a supported read-only ROCm CLI command with argv-style arguments. Commands that change ROCm state are rejected here and must go through the ROCm CLI approval UI.", + json!({ + "type": "object", + "properties": { + "args": { + "type": "array", + "items": { + "type": "string" + }, + "minItems": 1, + "maxItems": 64 + }, + "reason": { + "type": "string" + } + }, + "required": ["args"], + "additionalProperties": false + }), + true, + false, + ), + rocm_mcp_tool( + "update_check", + "Run `rocm update` and return the current TheRock update status.", + json!({ + "type": "object", + "properties": {}, + "additionalProperties": false + }), + true, + false, + ), + rocm_mcp_tool( + "install_sdk_dry_run", + "Run a dry-run TheRock SDK install plan.", + json!({ + "type": "object", + "properties": { + "channel": { + "type": "string", + "enum": ["release", "nightly"] + }, + "format": { + "type": "string", + "enum": ["wheel", "tarball"] + }, + "prefix": { + "type": "string" + }, + "version": { + "type": "string" + }, + "build_date": { + "type": "string" + } + }, + "additionalProperties": false + }), + true, + false, + ), + rocm_mcp_tool( + "install_sdk", + "Install a TheRock SDK into the managed runtime area or an explicitly approved prefix.", + json!({ + "type": "object", + "properties": { + "channel": { + "type": "string", + "enum": ["release", "nightly"] + }, + "format": { + "type": "string", + "enum": ["wheel", "tarball"] + }, + "prefix": { + "type": "string" + }, + "version": { + "type": "string" + }, + "build_date": { + "type": "string" + }, + "allow_system_prefix": { + "type": "boolean" + } + }, + "additionalProperties": false + }), + false, + true, + ), + rocm_mcp_tool( + "install_engine", + "Install or refresh a managed serving engine environment.", + json!({ + "type": "object", + "properties": { + "engine": { + "type": "string" + }, + "runtime_id": { + "type": "string" + }, + "python_version": { + "type": "string" + }, + "reinstall": { + "type": "boolean" + } + }, + "required": ["engine"], + "additionalProperties": false + }), + false, + false, + ), + rocm_mcp_tool( + "launch_server", + "Launch a managed local model server through `rocm serve --managed`.", + json!({ + "type": "object", + "properties": { + "model": { + "type": "string" + }, + "engine": { + "type": "string" + }, + "device": { + "type": "string" + }, + "runtime_id": { + "type": "string" + }, + "env_id": { + "type": "string" + }, + "host": { + "type": "string" + }, + "port": { + "type": "integer", + "minimum": 1, + "maximum": 65535 + }, + "allow_public_bind": { + "type": "boolean" + } + }, + "required": ["model"], + "additionalProperties": false + }), + false, + true, + ), + rocm_mcp_tool( + "stop_server", + "Stop a managed service by service id and update its manifest status.", + json!({ + "type": "object", + "properties": { + "service_id": { + "type": "string" + } + }, + "required": ["service_id"], + "additionalProperties": false + }), + false, + true, + ), + rocm_mcp_tool( + "watcher_enable", + "Enable a watcher and optionally set its mode.", + json!({ + "type": "object", + "properties": { + "watcher": { + "type": "string" + }, + "mode": { + "type": "string", + "enum": ["observe", "propose", "contained"] + } + }, + "required": ["watcher"], + "additionalProperties": false + }), + false, + false, + ), + rocm_mcp_tool( + "watcher_disable", + "Disable a watcher.", + json!({ + "type": "object", + "properties": { + "watcher": { + "type": "string" + } + }, + "required": ["watcher"], + "additionalProperties": false + }), + false, + false, + ), + ] +} + +fn rocm_mcp_tool( + name: &str, + description: &str, + input_schema: Value, + read_only: bool, + destructive: bool, +) -> Value { + json!({ + "name": name, + "title": name.replace('_', " "), + "description": description, + "annotations": { + "readOnlyHint": read_only, + "destructiveHint": destructive, + "openWorldHint": false, + }, + "inputSchema": input_schema, + }) +} + +fn mcp_tool_requires_direct_approval(name: &str) -> bool { + matches!( + name, + "install_sdk" + | "install_engine" + | "launch_server" + | "stop_server" + | "watcher_enable" + | "watcher_disable" + ) +} + +pub(crate) fn ensure_direct_mcp_call_allowed(name: &str, allow_mutation: bool) -> Result<()> { + if mcp_tool_requires_direct_approval(name) && !allow_mutation { + bail!( + "MCP tool `{name}` changes local ROCm state; rerun `rocmd mcp-call {name}` with --allow-mutation only after an explicit user approval" + ); + } + Ok(()) +} + +pub(crate) fn handle_mcp_tool_call(paths: &AppPaths, params: &Value) -> Result { + let name = params + .get("name") + .and_then(Value::as_str) + .unwrap_or_default(); + let arguments = params + .get("arguments") + .and_then(Value::as_object) + .cloned() + .unwrap_or_default(); + + match name { + "examine" => { + let examine = ExamineSummary::gather()?; + let output = common::run_rocm_capture(&["examine"])?; + let text = command_capture_text(&output); + if output.exit_status == 0 { + Ok(tool_success(text, json!(examine))) + } else { + Ok(tool_error( + text, + json!({ + "examine": examine, + "argv": output.argv, + "exit_status": output.exit_status, + "stderr": output.stderr, + }), + )) + } + } + "bridge_snapshot" => { + let snapshot = common::build_bridge_snapshot(paths)?; + Ok(tool_success( + format!( + "Captured bridge snapshot for {} / {} with default engine `{}`.", + snapshot.examine.os, snapshot.examine.arch, snapshot.examine.default_engine + ), + json!(snapshot), + )) + } + "gpu_snapshot" => { + let config = RocmCliConfig::load(paths).unwrap_or_default(); + let gpu = common::gather_gpu_snapshot_for_config(&config); + let status = if !config.telemetry.local_inspection_enabled() { + "GPU telemetry is disabled by rocm-cli config." + } else if gpu.amd_smi_available { + "Captured amd-smi GPU snapshot." + } else { + "amd-smi is unavailable on this host." + }; + Ok(tool_success(status.to_owned(), json!(gpu))) + } + "engines" => { + let engines = common::bridge_engine_inventory(); + Ok(tool_success( + format!("Found {} engine entries.", engines.len()), + json!({ "engines": engines }), + )) + } + "services" => { + let services = persistence::load_managed_services(paths)?; + Ok(tool_success( + format!("Found {} managed services.", services.len()), + json!({ "services": services }), + )) + } + "service_logs" => { + let service_id = arguments + .get("service_id") + .and_then(Value::as_str) + .context("service_logs requires `service_id`")?; + let lines = arguments + .get("lines") + .and_then(Value::as_u64) + .unwrap_or(80) + .clamp(1, 500) as usize; + let record = persistence::load_managed_services(paths)? + .into_iter() + .find(|service| service.service_id == service_id) + .with_context(|| format!("managed service `{service_id}` not found"))?; + let tail = read_tail_lines(&record.log_path, lines)?; + Ok(tool_success( + format!( + "Read the last {} line(s) from service `{}`.", + lines, record.service_id + ), + json!({ + "service": record, + "lines": lines, + "tail": tail, + }), + )) + } + "automations" => { + let event_limit = arguments + .get("event_limit") + .and_then(Value::as_u64) + .unwrap_or(10) + .clamp(1, 64) as usize; + let runtime = AutomationRuntimeState::load(paths)?; + let events = load_recent_automation_events(paths, event_limit)?; + Ok(tool_success( + format!( + "Loaded automation runtime and {} recent events.", + events.len() + ), + json!({ + "runtime": runtime, + "recent_events": events, + }), + )) + } + "natural_language_plan" => { + let request = arguments + .get("request") + .and_then(Value::as_str) + .context("natural_language_plan requires `request`")?; + let output = common::run_rocm_capture(&[request])?; + Ok(tool_result_from_command( + "Ran natural-language planning through `rocm`.", + output, + false, + )) + } + "rocm_command" => { + let argv = normalized_rocm_command_args(&arguments)?; + ensure_rocm_command_is_read_only(&argv)?; + let refs = argv.iter().map(String::as_str).collect::>(); + let output = common::run_rocm_capture(&refs)?; + Ok(tool_result_from_command( + "Ran read-only `rocm` command.", + output, + false, + )) + } + "update_check" => { + let output = common::run_rocm_capture(&["update"])?; + Ok(tool_result_from_command( + "Ran `rocm update`.", + output, + false, + )) + } + "install_sdk_dry_run" => { + let argv = build_install_sdk_args(&arguments, true)?; + let refs = argv.iter().map(String::as_str).collect::>(); + let output = common::run_rocm_capture(&refs)?; + Ok(tool_result_from_command( + "Ran `rocm install sdk --dry-run`.", + output, + false, + )) + } + "install_sdk" => { + let argv = build_install_sdk_args(&arguments, false)?; + let refs = argv.iter().map(String::as_str).collect::>(); + let output = common::run_rocm_capture(&refs)?; + Ok(tool_result_from_command( + "Ran `rocm install sdk`.", + output, + false, + )) + } + "install_engine" => { + let argv = build_install_engine_args(&arguments)?; + let refs = argv.iter().map(String::as_str).collect::>(); + let output = common::run_rocm_capture(&refs)?; + Ok(tool_result_from_command( + "Ran `rocm engines install`.", + output, + false, + )) + } + "launch_server" => { + let argv = build_launch_server_args(&arguments)?; + let refs = argv.iter().map(String::as_str).collect::>(); + let output = common::run_rocm_capture(&refs)?; + Ok(tool_result_from_command( + "Ran `rocm serve --managed`.", + output, + false, + )) + } + "stop_server" => { + let service_id = arguments + .get("service_id") + .and_then(Value::as_str) + .context("stop_server requires `service_id`")?; + let stopped = crate::stop_managed_service(paths, service_id)?; + Ok(tool_success( + format!("Stopped managed service `{service_id}`."), + stopped, + )) + } + "watcher_enable" => { + let argv = build_watcher_enable_args(&arguments)?; + let refs = argv.iter().map(String::as_str).collect::>(); + let output = common::run_rocm_capture(&refs)?; + Ok(tool_result_from_command( + "Ran `rocm automations enable`.", + output, + false, + )) + } + "watcher_disable" => { + let watcher = arguments + .get("watcher") + .and_then(Value::as_str) + .context("watcher_disable requires `watcher`")?; + let argv = [ + "automations".to_owned(), + "disable".to_owned(), + watcher.to_owned(), + ]; + let refs = argv.iter().map(String::as_str).collect::>(); + let output = common::run_rocm_capture(&refs)?; + Ok(tool_result_from_command( + "Ran `rocm automations disable`.", + output, + false, + )) + } + other => Ok(tool_error( + format!("Unknown ROCm MCP tool `{other}`."), + json!({ "tool": other }), + )), + } +} + +fn tool_success(text: String, structured: Value) -> Value { + json!({ + "content": [ + { + "type": "text", + "text": text, + } + ], + "structuredContent": structured, + "isError": false, + }) +} + +fn tool_error(text: String, structured: Value) -> Value { + json!({ + "content": [ + { + "type": "text", + "text": text, + } + ], + "structuredContent": structured, + "isError": true, + }) +} + +fn tool_result_from_command(prefix: &str, output: CommandCapture, is_error: bool) -> Value { + let text = format!("{prefix}\n\n{}", command_capture_text(&output)); + json!({ + "content": [ + { + "type": "text", + "text": text, + } + ], + "structuredContent": { + "argv": output.argv, + "exit_status": output.exit_status, + "stdout": output.stdout, + "stderr": output.stderr, + }, + "isError": is_error || output.exit_status != 0, + }) +} + +fn command_capture_text(output: &CommandCapture) -> String { + if output.stderr.trim().is_empty() { + output.stdout.trim().to_owned() + } else if output.stdout.trim().is_empty() { + format!("stderr:\n{}", output.stderr.trim()) + } else { + format!( + "stdout:\n{}\n\nstderr:\n{}", + output.stdout.trim(), + output.stderr.trim() + ) + } +} + +fn read_tail_lines(path: &std::path::Path, limit: usize) -> Result { + let content = + fs::read_to_string(path).with_context(|| format!("failed to read {}", path.display()))?; + let mut lines = VecDeque::with_capacity(limit); + for line in content.lines() { + if lines.len() == limit { + lines.pop_front(); + } + lines.push_back(line.to_owned()); + } + Ok(lines.into_iter().collect::>().join("\n")) +} + +fn normalized_rocm_command_args(arguments: &serde_json::Map) -> Result> { + let values = arguments + .get("args") + .and_then(Value::as_array) + .context("rocm_command requires `args`")?; + if values.is_empty() || values.len() > 64 { + bail!("rocm_command `args` must contain 1 to 64 strings"); + } + let mut args = Vec::with_capacity(values.len()); + for value in values { + let arg = value + .as_str() + .map(str::trim) + .filter(|value| !value.is_empty()) + .context("rocm_command `args` entries must be non-empty strings")?; + if arg.contains('\0') || arg.contains('\n') || arg.contains('\r') { + bail!("rocm_command arguments must not contain control characters"); + } + if arg.len() > 512 { + bail!("rocm_command argument is too long"); + } + args.push(arg.to_owned()); + } + if args + .first() + .is_some_and(|arg| arg.eq_ignore_ascii_case("rocm")) + { + args.remove(0); + } + if args + .first() + .is_some_and(|arg| arg.eq_ignore_ascii_case("comfy")) + { + args[0] = "comfyui".to_owned(); + } + if args.is_empty() { + bail!("rocm_command args should omit the leading `rocm` program name"); + } + Ok(args) +} + +fn ensure_rocm_command_is_read_only(args: &[String]) -> Result<()> { + let first = args.first().map(|value| value.to_ascii_lowercase()); + let second = args.get(1).map(|value| value.to_ascii_lowercase()); + let read_only = match first.as_deref() { + Some("examine" | "version" | "model" | "models" | "daemon" | "logs") => true, + Some("update") => !args.iter().any(|arg| arg == "--apply"), + Some("runtimes") => { + second.as_deref().is_none_or(|value| value == "list") + || (second + .as_deref() + .is_some_and(|value| value == "uninstall" || value == "remove") + && args.iter().any(|arg| arg == "--dry-run")) + } + Some("engines") => second.as_deref().is_some_and(|value| value == "list"), + Some("services") => second + .as_deref() + .is_none_or(|value| matches!(value, "list" | "logs")), + Some("automations") => second.as_deref().is_none_or(|value| value == "list"), + Some("config") => second.as_deref() == Some("show"), + Some("comfyui") => second + .as_deref() + .is_none_or(|value| matches!(value, "status" | "logs" | "log")), + Some("uninstall") => args.iter().any(|arg| arg == "--dry-run"), + // `storage report` (the default subcommand) only measures folders. The + // two `remove-*` verbs delete, so they stay off the read-only list. + Some("storage") => second.as_deref().is_none_or(|value| value == "report"), + // `setup status` reports first-time setup state (read-only); `setup reset` + // clears the completion/dismissal state and is mutating (it does not by + // itself reopen onboarding). Mirrors the bin's rocm_command classifier so + // the read-only allowlist is consistent across binaries. + Some("setup") => second.as_deref().is_none_or(|value| value == "status"), + // `remote targets` reads the local tailnet, `doctor` fetches another + // machine's state and scores it here, `status` probes sessions that + // already exist. None of them change anything on either machine. + // `serve`, `attach` and `stop` start, publish or tear down, so they stay + // off the list and go through the approval UI like any other mutation. + Some("remote") => second + .as_deref() + .is_some_and(|value| matches!(value, "targets" | "doctor" | "status")), + _ => false, + }; + if read_only { + return Ok(()); + } + bail!( + "rocm_command changes local ROCm state or is unsupported here; request it through the ROCm CLI approval UI instead" + ) +} + +fn build_install_sdk_args( + arguments: &serde_json::Map, + dry_run: bool, +) -> Result> { + let channel = arguments + .get("channel") + .and_then(Value::as_str) + .unwrap_or("release"); + let format = arguments + .get("format") + .and_then(Value::as_str) + .unwrap_or("wheel"); + let prefix = arguments.get("prefix").and_then(Value::as_str); + let version = arguments.get("version").and_then(Value::as_str); + let build_date = arguments.get("build_date").and_then(Value::as_str); + let allow_system_prefix = arguments + .get("allow_system_prefix") + .and_then(Value::as_bool) + .unwrap_or(false); + if version.is_some() && build_date.is_some() { + bail!("install_sdk accepts either `version` or `build_date`, not both"); + } + + let mut argv = vec![ + "install".to_owned(), + "sdk".to_owned(), + "--channel".to_owned(), + channel.to_owned(), + "--format".to_owned(), + format.to_owned(), + ]; + if let Some(prefix) = prefix { + let prefix_path = std::path::Path::new(prefix); + if system_prefix_requires_ack(prefix_path) && !allow_system_prefix { + bail!( + "install_sdk prefix `{}` is outside the user home; require `allow_system_prefix=true` before using system paths", + prefix_path.display() + ); + } + argv.push("--prefix".to_owned()); + argv.push(prefix.to_owned()); + } + if let Some(version) = version { + if version.trim().is_empty() { + bail!("install_sdk `version` cannot be empty"); + } + argv.push("--version".to_owned()); + argv.push(version.to_owned()); + } + if let Some(build_date) = build_date { + if build_date.trim().is_empty() { + bail!("install_sdk `build_date` cannot be empty"); + } + argv.push("--build-date".to_owned()); + argv.push(build_date.to_owned()); + } + if dry_run { + argv.push("--dry-run".to_owned()); + } else { + // `run_rocm_capture_for_paths` spawns `rocm` with null stdin, so + // `interactive_terminal()` is false in the child and an active default + // managed runtime would make the approval gate refuse with "re-run with + // `--approve-replacing-active-default`" — a flag no MCP caller of this + // tool can supply. + // + // Not `--yes` itself: that flag carries a second, unrelated consent — + // approving required system-package installs, which run `sudo`. This + // spawn has no terminal, so it could never answer a sudo password + // prompt; granting that consent would make the vLLM/OpenMPI step attempt + // an install it cannot complete and abort the engine auto-install that + // previously warned and continued. `--approve-replacing-active-default` + // grants only the runtime-displacement consent the gate asks for. + // + // Consent is not bypassed: `install_sdk` is in + // `mcp_tool_requires_direct_approval`, so a direct `rocmd mcp-call` + // needs `--allow-mutation` after an explicit user approval, and over the + // MCP protocol the tool is annotated `destructiveHint` for the client's + // approval UI. Mirrors the chat/MCP arm in `apps/rocm`. The dry-run + // branch never reaches the gate (it returns earlier), so it stays bare. + argv.push("--approve-replacing-active-default".to_owned()); + } + Ok(argv) +} + +fn build_install_engine_args(arguments: &serde_json::Map) -> Result> { + let engine = arguments + .get("engine") + .and_then(Value::as_str) + .context("install_engine requires `engine`")?; + let runtime_id = arguments + .get("runtime_id") + .and_then(Value::as_str) + .unwrap_or("therock-release"); + let python_version = arguments.get("python_version").and_then(Value::as_str); + let reinstall = arguments + .get("reinstall") + .and_then(Value::as_bool) + .unwrap_or(false); + + let mut argv = vec![ + "engines".to_owned(), + "install".to_owned(), + engine.to_owned(), + "--runtime-id".to_owned(), + runtime_id.to_owned(), + ]; + if let Some(python_version) = python_version { + argv.push("--python-version".to_owned()); + argv.push(python_version.to_owned()); + } + if reinstall { + argv.push("--reinstall".to_owned()); + } + Ok(argv) +} + +fn build_launch_server_args(arguments: &serde_json::Map) -> Result> { + let model = arguments + .get("model") + .and_then(Value::as_str) + .context("launch_server requires `model`")?; + let host = arguments + .get("host") + .and_then(Value::as_str) + .unwrap_or(DEFAULT_LOCAL_HOST); + let allow_public_bind = arguments + .get("allow_public_bind") + .and_then(Value::as_bool) + .unwrap_or(false); + if !is_loopback_host(host) && !allow_public_bind { + bail!( + "launch_server host `{host}` is not loopback; require `allow_public_bind=true` before binding a non-local interface" + ); + } + + let mut argv = vec!["serve".to_owned(), model.to_owned(), "--managed".to_owned()]; + if let Some(engine) = arguments.get("engine").and_then(Value::as_str) { + argv.push("--engine".to_owned()); + argv.push(engine.to_owned()); + } + if let Some(device) = arguments.get("device").and_then(Value::as_str) { + argv.push("--device".to_owned()); + argv.push(device.to_owned()); + } + if let Some(runtime_id) = arguments.get("runtime_id").and_then(Value::as_str) { + argv.push("--runtime-id".to_owned()); + argv.push(runtime_id.to_owned()); + } + if let Some(env_id) = arguments.get("env_id").and_then(Value::as_str) { + argv.push("--env-id".to_owned()); + argv.push(env_id.to_owned()); + } + argv.push("--host".to_owned()); + argv.push(host.to_owned()); + if allow_public_bind { + argv.push("--allow-public-bind".to_owned()); + } + if let Some(port) = arguments.get("port").and_then(Value::as_u64) { + argv.push("--port".to_owned()); + argv.push(port.to_string()); + } + Ok(argv) +} + +fn build_watcher_enable_args(arguments: &serde_json::Map) -> Result> { + let watcher = arguments + .get("watcher") + .and_then(Value::as_str) + .context("watcher_enable requires `watcher`")?; + let mut argv = vec![ + "automations".to_owned(), + "enable".to_owned(), + watcher.to_owned(), + ]; + if let Some(mode) = arguments.get("mode").and_then(Value::as_str) { + argv.push("--mode".to_owned()); + argv.push(mode.to_owned()); + } + Ok(argv) +} + +/// Inverse of [`rocm_engine_protocol::is_public_bind_host`], which owns the +/// policy so `rocm` and `rocmd` never classify the same host differently. +fn is_loopback_host(host: &str) -> bool { + !rocm_engine_protocol::is_public_bind_host(host) +} + +fn system_prefix_requires_ack(prefix: &std::path::Path) -> bool { + match rocm_core::runtime_home_dir() { + Some(home) => !rocm_core::runtime_path_is_same_or_inside(prefix, &home), + None => true, + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::test_support::workspace_test_artifact_dir; + use std::fs; + use std::path::PathBuf; + + fn unique_test_path(label: &str) -> PathBuf { + let root = workspace_test_artifact_dir(); + fs::create_dir_all(&root).expect("create workspace-local test dir"); + root.join(label) + } + + #[test] + fn remote_read_only_verbs_are_allowed_and_mutating_ones_are_not() { + let allow = |args: &[&str]| { + let owned = args.iter().map(|a| (*a).to_owned()).collect::>(); + super::ensure_rocm_command_is_read_only(&owned) + }; + + // These read: the local tailnet, another machine's state, sessions that + // already exist. Rejecting them made the whole family unusable here even + // though none of them change anything. + for args in [ + &["remote", "targets"][..], + &["remote", "targets", "--tag", "gpu"][..], + &["remote", "doctor", "gpu-box"][..], + &["remote", "status"][..], + ] { + allow(args).unwrap_or_else(|error| panic!("{args:?} should be read-only: {error:#}")); + } + + // These start, publish or tear down, so they go through approval. + for args in [ + &["remote", "serve", "gpu-box", "a-model"][..], + &["remote", "attach", "sess"][..], + &["remote", "stop", "sess"][..], + &["remote"][..], + ] { + assert!(allow(args).is_err(), "{args:?} must not be read-only"); + } + } + #[test] + fn rocm_mcp_tools_include_bridge_gaps() { + let tools = rocm_mcp_tools(); + let names = tools + .iter() + .filter_map(|tool| tool.get("name").and_then(Value::as_str).map(str::to_owned)) + .collect::>(); + assert!(names.contains(&"gpu_snapshot".to_owned())); + assert!(names.contains(&"service_logs".to_owned())); + assert!(names.contains(&"natural_language_plan".to_owned())); + assert!(names.contains(&"rocm_command".to_owned())); + assert!(names.contains(&"install_sdk".to_owned())); + assert!(names.contains(&"install_engine".to_owned())); + assert!(names.contains(&"launch_server".to_owned())); + assert!(names.contains(&"stop_server".to_owned())); + assert!(names.contains(&"watcher_enable".to_owned())); + assert!(names.contains(&"watcher_disable".to_owned())); + let automations = tools + .iter() + .find(|tool| tool.get("name").and_then(Value::as_str) == Some("automations")) + .expect("automations tool should be present"); + assert!( + automations + .get("description") + .and_then(Value::as_str) + .is_some_and(|description| description.contains("local webhook events")) + ); + } + + #[test] + fn direct_mcp_call_requires_approval_for_every_mutating_tool() { + for tool in rocm_mcp_tools() { + let name = tool + .get("name") + .and_then(Value::as_str) + .expect("tool should have a name"); + let read_only = tool + .get("annotations") + .and_then(|annotations| annotations.get("readOnlyHint")) + .and_then(Value::as_bool) + .unwrap_or(false); + assert_eq!( + mcp_tool_requires_direct_approval(name), + !read_only, + "hidden direct MCP helper approval classification drifted for `{name}`" + ); + } + } + + #[test] + fn direct_mcp_call_guard_blocks_mutation_without_explicit_ack() { + ensure_direct_mcp_call_allowed("examine", false) + .expect("read-only direct MCP helper calls should not need mutation approval"); + + let error = ensure_direct_mcp_call_allowed("install_sdk", false) + .expect_err("mutating direct MCP helper calls should require approval"); + assert!(error.to_string().contains("--allow-mutation"), "{error:#}"); + + ensure_direct_mcp_call_allowed("install_sdk", true) + .expect("explicitly approved direct MCP mutation should pass the helper guard"); + } + + #[test] + fn rocm_command_helper_allows_only_read_only_rocm_commands() -> Result<()> { + let status_args = normalized_rocm_command_args( + serde_json::json!({ + "args": ["rocm", "comfy", "status"] + }) + .as_object() + .expect("json object"), + )?; + assert_eq!(status_args, vec!["comfyui".to_owned(), "status".to_owned()]); + ensure_rocm_command_is_read_only(&status_args).expect("ComfyUI status should be read-only"); + + let log_args = normalized_rocm_command_args( + serde_json::json!({ + "args": ["comfyui", "logs"] + }) + .as_object() + .expect("json object"), + )?; + ensure_rocm_command_is_read_only(&log_args).expect("ComfyUI logs should be read-only"); + + let install_args = normalized_rocm_command_args( + serde_json::json!({ + "args": ["comfyui", "install"] + }) + .as_object() + .expect("json object"), + )?; + let error = ensure_rocm_command_is_read_only(&install_args) + .expect_err("ComfyUI install must go through approval"); + assert!(error.to_string().contains("approval UI")); + + let shell_args = normalized_rocm_command_args( + serde_json::json!({ + "args": ["powershell", "-Command", "whoami"] + }) + .as_object() + .expect("json object"), + )?; + let error = ensure_rocm_command_is_read_only(&shell_args) + .expect_err("non-rocm shell commands should be rejected"); + assert!(error.to_string().contains("approval UI")); + Ok(()) + } + + #[test] + fn rocm_command_helper_treats_setup_status_as_read_only_and_reset_as_mutating() -> Result<()> { + // Mirrors the bin's rocm_command classifier so `setup status` is read-only + // on every binary's tool surface while `setup reset` stays approval-gated. + let bare_args = normalized_rocm_command_args( + serde_json::json!({ "args": ["setup"] }) + .as_object() + .expect("json object"), + )?; + ensure_rocm_command_is_read_only(&bare_args).expect("bare setup should be read-only"); + + let status_args = normalized_rocm_command_args( + serde_json::json!({ "args": ["setup", "status"] }) + .as_object() + .expect("json object"), + )?; + ensure_rocm_command_is_read_only(&status_args).expect("setup status should be read-only"); + + let reset_args = normalized_rocm_command_args( + serde_json::json!({ "args": ["setup", "reset"] }) + .as_object() + .expect("json object"), + )?; + let error = ensure_rocm_command_is_read_only(&reset_args) + .expect_err("setup reset must go through approval"); + assert!(error.to_string().contains("approval UI")); + Ok(()) + } + + #[test] + fn rocm_command_helper_treats_runtimes_uninstall_dry_run_as_read_only() -> Result<()> { + // Mirrors the bin's chat_rocm_command_action_from_args classifier so a + // dry-run preview stays read-only on every binary's tool surface while + // an actual uninstall/remove still requires approval. + for verb in ["uninstall", "remove"] { + let dry_run_args = normalized_rocm_command_args( + serde_json::json!({ "args": ["runtimes", verb, "--dry-run"] }) + .as_object() + .expect("json object"), + )?; + ensure_rocm_command_is_read_only(&dry_run_args) + .unwrap_or_else(|_| panic!("runtimes {verb} --dry-run should be read-only")); + + let mutating_args = normalized_rocm_command_args( + serde_json::json!({ "args": ["runtimes", verb] }) + .as_object() + .expect("json object"), + )?; + let error = match ensure_rocm_command_is_read_only(&mutating_args) { + Ok(()) => panic!("runtimes {verb} without --dry-run must go through approval"), + Err(error) => error, + }; + assert!(error.to_string().contains("approval UI")); + } + Ok(()) + } + + #[test] + fn storage_report_is_read_only_but_removal_is_not() -> Result<()> { + for args in [vec!["storage"], vec!["storage", "report"]] { + let normalized = normalized_rocm_command_args( + serde_json::json!({ "args": args }) + .as_object() + .expect("json object"), + )?; + ensure_rocm_command_is_read_only(&normalized) + .unwrap_or_else(|_| panic!("storage {args:?} only measures folders")); + } + + for verb in ["remove-old-installs", "remove-downloads"] { + let normalized = normalized_rocm_command_args( + serde_json::json!({ "args": ["storage", verb] }) + .as_object() + .expect("json object"), + )?; + let error = ensure_rocm_command_is_read_only(&normalized) + .expect_err("storage removal must go through approval"); + assert!(error.to_string().contains("approval UI")); + } + Ok(()) + } + + #[test] + fn read_tail_lines_returns_last_lines_only() -> Result<()> { + let path = unique_test_path(&format!( + "rocmd-tail-test-{}-{}.log", + std::process::id(), + unix_time_millis() + )); + fs::write(&path, "line1\nline2\nline3\nline4\n")?; + let tail = read_tail_lines(&path, 2)?; + fs::remove_file(&path)?; + assert_eq!(tail, "line3\nline4"); + Ok(()) + } + + #[test] + fn launch_server_rejects_public_bind_without_ack() { + let arguments = serde_json::Map::from_iter([ + ("model".to_owned(), Value::String("tiny-gpt2".to_owned())), + ("host".to_owned(), Value::String("0.0.0.0".to_owned())), + ]); + let error = build_launch_server_args(&arguments).unwrap_err(); + assert!( + error.to_string().contains("allow_public_bind=true"), + "{error:#}" + ); + } + + #[test] + fn launch_server_forwards_public_bind_ack() -> Result<()> { + let arguments = serde_json::Map::from_iter([ + ("model".to_owned(), Value::String("tiny-gpt2".to_owned())), + ("host".to_owned(), Value::String("0.0.0.0".to_owned())), + ("allow_public_bind".to_owned(), Value::Bool(true)), + ]); + let args = build_launch_server_args(&arguments)?; + assert!(args.contains(&"--allow-public-bind".to_owned())); + Ok(()) + } + + #[test] + fn install_sdk_rejects_system_prefix_without_ack() { + let arguments = serde_json::Map::from_iter([( + "prefix".to_owned(), + Value::String("/opt/rocm".to_owned()), + )]); + let error = build_install_sdk_args(&arguments, false).unwrap_err(); + assert!( + error.to_string().contains("allow_system_prefix=true"), + "{error:#}" + ); + } + + /// The test above only ever hands `system_prefix_requires_ack` an + /// already-canonical path, so it cannot catch the bug this crate's fix + /// addresses: a `..`-respelled prefix that escapes `$HOME` used to compare + /// equal to a path still inside it (`Path::ancestors()` treats `..` as an + /// ordinary component), so acknowledgement was never required. Drive the + /// same check with a prefix built by walking `..` out of the real home + /// directory, which is exactly the shape the original bug let through. + #[test] + #[cfg(unix)] + fn install_sdk_rejects_system_prefix_reached_by_escaping_home() { + let home = rocm_core::runtime_home_dir().expect("a home directory"); + let escaped_prefix = format!("{}/../../usr", home.display()); + + let arguments = + serde_json::Map::from_iter([("prefix".to_owned(), Value::String(escaped_prefix))]); + let error = build_install_sdk_args(&arguments, false).unwrap_err(); + assert!( + error.to_string().contains("allow_system_prefix=true"), + "{error:#}" + ); + } + + /// The `install_sdk` MCP tool spawns `rocm` with null stdin, so a real + /// install over an active default managed runtime would hit the approval + /// gate's non-interactive refusal and bail asking for a flag no MCP caller + /// can pass. The real-install argv must therefore carry the consent flag; + /// the dry-run argv must not, because a dry run never reaches the gate and + /// the flag there would claim an approval the caller did not give. + /// + /// It must be `--approve-replacing-active-default` and never `--yes`: + /// `--yes` additionally approves running `sudo` for required system + /// packages, and a null-stdin spawn has no terminal on which that password + /// prompt could be answered. + #[test] + fn install_sdk_real_install_args_approve_only_the_runtime_replacement() -> Result<()> { + let arguments = serde_json::Map::new(); + + let real = build_install_sdk_args(&arguments, false)?; + assert!( + real.contains(&"--approve-replacing-active-default".to_owned()), + "real install argv must approve the replacement for the null-stdin spawn: {real:?}" + ); + assert!( + !real.contains(&"--yes".to_owned()), + "real install argv must not grant the system-package consent it cannot answer: {real:?}" + ); + assert!( + !real.contains(&"--dry-run".to_owned()), + "real install argv must not be a dry run: {real:?}" + ); + + let dry = build_install_sdk_args(&arguments, true)?; + assert!( + !dry.contains(&"--approve-replacing-active-default".to_owned()) + && !dry.contains(&"--yes".to_owned()), + "dry-run argv must not carry a consent flag: {dry:?}" + ); + assert!( + dry.contains(&"--dry-run".to_owned()), + "dry-run argv must carry --dry-run: {dry:?}" + ); + Ok(()) + } + + #[test] + fn install_sdk_forwards_requested_build_date_and_rejects_conflict() -> Result<()> { + let arguments = serde_json::Map::from_iter([( + "build_date".to_owned(), + Value::String("2026-06-05".to_owned()), + )]); + let argv = build_install_sdk_args(&arguments, true)?; + assert_eq!( + argv, + vec![ + "install".to_owned(), + "sdk".to_owned(), + "--channel".to_owned(), + "release".to_owned(), + "--format".to_owned(), + "wheel".to_owned(), + "--build-date".to_owned(), + "2026-06-05".to_owned(), + "--dry-run".to_owned(), + ] + ); + + let conflicting = serde_json::Map::from_iter([ + ( + "version".to_owned(), + Value::String("7.13.0a20260605".to_owned()), + ), + ( + "build_date".to_owned(), + Value::String("2026-06-05".to_owned()), + ), + ]); + let error = build_install_sdk_args(&conflicting, false) + .unwrap_err() + .to_string(); + assert!(error.contains("either `version` or `build_date`")); + Ok(()) + } + + #[test] + fn watcher_enable_builds_mode_args() -> Result<()> { + let arguments = serde_json::Map::from_iter([ + ( + "watcher".to_owned(), + Value::String("server-recover".to_owned()), + ), + ("mode".to_owned(), Value::String("contained".to_owned())), + ]); + let argv = build_watcher_enable_args(&arguments)?; + assert_eq!( + argv, + vec![ + "automations".to_owned(), + "enable".to_owned(), + "server-recover".to_owned(), + "--mode".to_owned(), + "contained".to_owned() + ] + ); + Ok(()) + } +} diff --git a/docs/architecture.md b/docs/architecture.md index f292ca17f..7117b9085 100644 --- a/docs/architecture.md +++ b/docs/architecture.md @@ -31,7 +31,7 @@ Subsystem modules already following full domain extraction (each owns its own ty ### `apps/rocmd` — background daemon -`lib.rs` modularization is in progress (ROCMAI-83, Phase 5 of EAI-7768's sequencing). Extracted so far: `persistence.rs` (`record_event`/`load_managed_services`, the automation-event/audit-log and managed-service-registry I/O shared across the daemon's sandbox, MCP, service-lifecycle, and watcher code), `common.rs` (helpers shared across ≥2 of those remaining clusters: GPU/amd-smi snapshotting, the bridge-snapshot diagnostic, `CommandCapture`/command-timeout plumbing including the shared `rocm`-subprocess capture helpers the sandbox and MCP clusters both call, and small arg/healthcheck/endpoint-key utilities), `webhook.rs` (the local webhook source: its axum routes, request validation, and watcher-kind allow-list), `cli.rs` (the `Cli`/`Command` clap definitions, `SandboxToolArg`/`SandboxToolPolicy`, and the top-level dispatch in `run_cli`/`run_bin_cli`/`run_from_args` — the crate's only two externally-consumed entry points are re-exported from here via `lib.rs`'s `pub use`), and `sandbox.rs` (bubblewrap/native sandbox execution, atomic-write helpers, artifact prefetch policy gating, and the sandbox-tool result shaping for `check_updates`/`driver_plan`). A helper earns a place in `common.rs` only once a second still-inline cluster calls it directly; a helper with exactly one caller stays in `lib.rs` next to that caller until its own cluster's extraction PR, even if it is conceptually similar to something that did move. Still pending: the MCP, service-lifecycle, and watcher clusters themselves — each landing as its own PR. +`lib.rs` modularization is in progress (ROCMAI-83, Phase 5 of EAI-7768's sequencing). Extracted so far: `persistence.rs` (`record_event`/`load_managed_services`, the automation-event/audit-log and managed-service-registry I/O shared across the daemon's sandbox, MCP, service-lifecycle, and watcher code), `common.rs` (helpers shared across ≥2 of those remaining clusters: GPU/amd-smi snapshotting, the bridge-snapshot diagnostic, `CommandCapture`/command-timeout plumbing including the shared `rocm`-subprocess capture helpers the sandbox and MCP clusters both call, and small arg/healthcheck/endpoint-key utilities), `webhook.rs` (the local webhook source: its axum routes, request validation, and watcher-kind allow-list), `cli.rs` (the `Cli`/`Command` clap definitions, `SandboxToolArg`/`SandboxToolPolicy`, and the top-level dispatch in `run_cli`/`run_bin_cli`/`run_from_args` — the crate's only two externally-consumed entry points are re-exported from here via `lib.rs`'s `pub use`), `sandbox.rs` (bubblewrap/native sandbox execution, atomic-write helpers, artifact prefetch policy gating, and the sandbox-tool result shaping for `check_updates`/`driver_plan`), and `mcp.rs` (the MCP stdio server, tool schema table, tool dispatch, and the `rocm`-subprocess capture/argv-building helpers behind the MCP tools). A helper earns a place in `common.rs` only once a second still-inline cluster calls it directly; a helper with exactly one caller stays in `lib.rs` next to that caller until its own cluster's extraction PR, even if it is conceptually similar to something that did move. Still pending: the service-lifecycle and watcher clusters themselves — each landing as its own PR. ### `crates/rocm-core` — core library From 41cef57079fdaf321c4708db9e52ecd27bb37327 Mon Sep 17 00:00:00 2001 From: Jussi Elo Date: Thu, 1 Oct 2026 13:15:34 +0000 Subject: [PATCH 2/3] ROCMAI-83: extract service.rs from apps/rocmd/src/lib.rs MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Seventh PR of Phase 5 (rocmd modularization, ROCMAI-27): pull the managed-service PID lifecycle (stop/terminate/descendant-pid discovery), run_daemon's foreground loop, supervise_service's spawn-and-recover path, and serve-log startup-phase polling into their own module. Non-contiguous in the source (the startup-phase/ shutdown-signal helpers sit far from the rest, separated by the whole test block) -- both halves moved in this PR. No behavior change. Reaches back into still-crate-root items (load_managed_services, record_event, start_local_webhook_source, receive_local_webhook_event, evaluate_watchers/evaluate_watchers_for_ events/reconcile_watcher_snapshots, WATCHER_TICK_INTERVAL, parse_gpu_indices_arg/optional_arg/ensure_public_service_has_endpoint_ key/apply_endpoint_key_env/engine_healthcheck_ready, load_service_record) via crate::, since the modules that will eventually own those (persistence.rs/webhook.rs/watchers.rs/ common.rs) are separate, independent PRs not present on this branch. 13 tests that exercise this module's own logic (the supervise key-guard regression tests, run_daemon's automation-loop gate, stop_managed_service, descendant-pid discovery, startup-phase log parsing, engine_serve_http_args) moved into service.rs's own sandbox/mcp-domain tests that construct ManagedServiceRecord or call stop_managed_service merely as setup stayed in lib.rs for watchers.rs/other PRs. Stacked on rocmai-83-mcp (ROCMAI-83 Phase 5 batch, AGENTS.md §11): rebased service.rs's crate:: reaches onto the sibling modules already landed by the predecessor PRs (persistence.rs, common.rs, webhook.rs), leaving only the references still owned by lib.rs (evaluate_watchers*, reconcile_watcher_snapshots, WATCHER_TICK_INTERVAL, load_service_record) until watchers.rs lands. Deduped the test-only temp_app_paths/ unique_test_root helpers in favor of the shared crate::test_support module, and repointed cli.rs/mcp.rs/sandbox.rs's stale crate::-root calls to service::stop_managed_service/run_daemon/supervise_service/ print_status now that this module owns them. Rebase note (main now at f9a80117): parse_gpu_indices_arg and engine_healthcheck_ready move here too, not to common.rs. Per this stack's admission rule (a helper earns common.rs only once a second still-inline cluster calls it directly), both have exactly one caller and it is this module (service.rs:... supervise_service, wait_for_service_ready), so common.rs is the wrong home for them even though an earlier draft of this stack put them there. Signed-off-by: Jussi Elo --- apps/rocmd/src/cli.rs | 6 +- apps/rocmd/src/lib.rs | 1335 +----------------------------------- apps/rocmd/src/mcp.rs | 2 +- apps/rocmd/src/sandbox.rs | 3 +- apps/rocmd/src/service.rs | 1352 +++++++++++++++++++++++++++++++++++++ docs/architecture.md | 2 +- 6 files changed, 1362 insertions(+), 1338 deletions(-) create mode 100644 apps/rocmd/src/service.rs diff --git a/apps/rocmd/src/cli.rs b/apps/rocmd/src/cli.rs index 55071afa5..6eecf526a 100644 --- a/apps/rocmd/src/cli.rs +++ b/apps/rocmd/src/cli.rs @@ -240,7 +240,7 @@ async fn run_cli(cli: Cli) -> Result<()> { Command::Run { automations_enabled, local_webhook_port, - } => crate::run_daemon(&paths, automations_enabled, local_webhook_port).await?, + } => crate::service::run_daemon(&paths, automations_enabled, local_webhook_port).await?, Command::Supervise { service_id, engine, @@ -253,7 +253,7 @@ async fn run_cli(cli: Cli) -> Result<()> { device_policy, gpu, engine_recipe_json, - } => crate::supervise_service( + } => crate::service::supervise_service( &paths, service_id, engine, @@ -268,7 +268,7 @@ async fn run_cli(cli: Cli) -> Result<()> { engine_recipe_json, )?, Command::Status => { - crate::print_status(&paths)?; + crate::service::print_status(&paths)?; } Command::BridgeSnapshot { pretty } => { print_bridge_snapshot(&paths, pretty)?; diff --git a/apps/rocmd/src/lib.rs b/apps/rocmd/src/lib.rs index efd6ad807..add9b110d 100644 --- a/apps/rocmd/src/lib.rs +++ b/apps/rocmd/src/lib.rs @@ -9,6 +9,7 @@ mod common; mod mcp; mod persistence; mod sandbox; +mod service; #[cfg(test)] mod test_support; mod webhook; @@ -23,19 +24,15 @@ use rocm_core::AutomationEventRecord; use rocm_core::{ AppPaths, AutomationProposalRecord, AutomationRuntimeState, AutomationTriggerEvent, CodexBridgeGpuSnapshot, ManagedServiceRecord, RocmCliConfig, WatcherMode, - WatcherRuntimeSnapshot, append_automation_proposal, builtin_watchers, daemon_binary_path, + WatcherRuntimeSnapshot, append_automation_proposal, builtin_watchers, resolve_model_recipe_artifact, unix_time_millis, }; use serde_json::Value; use serde_json::json; -use std::collections::HashSet; use std::fs; -use std::io::{self, Read, Seek, SeekFrom, Write}; -use std::path::Path; use std::process::{Command as ProcessCommand, Stdio}; use std::thread; use std::time::Duration; -use tokio::time::{self, MissedTickBehavior}; const WATCHER_TICK_INTERVAL: Duration = Duration::from_secs(5); const SERVER_RECOVER_BACKOFF_MS: u128 = 30_000; @@ -48,731 +45,6 @@ const GPU_THERMAL_MEMORY_PRESSURE_C: f64 = 95.0; const GPU_MEMORY_VRAM_PRESSURE_PERCENT: f64 = 95.0; const ARTIFACT_PREFETCH_TIMEOUT: Duration = Duration::from_mins(10); -fn stop_managed_service(paths: &AppPaths, service_id: &str) -> Result { - let mut record = persistence::load_managed_services(paths)? - .into_iter() - .find(|record| record.service_id == service_id) - .with_context(|| format!("managed service `{service_id}` not found"))?; - let mut signaled_pids = Vec::new(); - let mut skipped_pids = Vec::new(); - let mut root_pids = Vec::new(); - if let Some(engine_pid) = record.engine_pid - && engine_pid != 0 - { - if engine_pid == std::process::id() { - skipped_pids.push(engine_pid); - } else { - root_pids.push(engine_pid); - } - } - if record.supervisor_pid != 0 - && record.supervisor_pid != std::process::id() - && Some(record.supervisor_pid) != record.engine_pid - { - root_pids.push(record.supervisor_pid); - } - let mut pids_to_signal = descendant_pids_for_roots(&root_pids)?; - pids_to_signal.extend(root_pids); - let mut seen_pids = HashSet::new(); - for pid in pids_to_signal { - if !seen_pids.insert(pid) { - continue; - } - if terminate_process(pid)? { - signaled_pids.push(pid); - } else { - skipped_pids.push(pid); - } - } - let force_signaled_pids = force_terminate_remaining_processes(&signaled_pids)?; - record.status = "stopped".to_owned(); - record.write()?; - // Best-effort and idempotent: a missing key file is not an error, so this - // is safe to call unconditionally on every stop (including loopback - // services that never had a key, and repeated stops of an already-stopped - // service). Leaving the 0600 key file behind after stop would strand a - // plaintext secret on disk for a service that no longer exists. - let _ = std::fs::remove_file(rocm_engine_protocol::endpoint_key_file_path( - paths, service_id, - )); - Ok(json!({ - "service": record, - "signaled_pids": signaled_pids, - "force_signaled_pids": force_signaled_pids, - "skipped_pids": skipped_pids, - })) -} - -#[cfg(unix)] -fn descendant_pids_for_roots(root_pids: &[u32]) -> Result> { - if root_pids.is_empty() { - return Ok(Vec::new()); - } - let output = ProcessCommand::new("ps") - .args(["-eo", "pid=,ppid="]) - .stdin(Stdio::null()) - .stdout(Stdio::piped()) - .stderr(Stdio::piped()) - .output() - .context("failed to list process tree with ps")?; - if !output.status.success() { - let stderr = String::from_utf8_lossy(&output.stderr).trim().to_owned(); - let stdout = String::from_utf8_lossy(&output.stdout).trim().to_owned(); - bail!( - "failed to list process tree: {}", - if !stderr.is_empty() { - stderr - } else if !stdout.is_empty() { - stdout - } else { - format!("exit status {}", output.status) - } - ); - } - let output = String::from_utf8_lossy(&output.stdout); - Ok(descendant_pids_from_ps_output(&output, root_pids)) -} - -#[cfg(not(unix))] -fn descendant_pids_for_roots(_root_pids: &[u32]) -> Result> { - Ok(Vec::new()) -} - -#[cfg(any(unix, test))] -fn descendant_pids_from_ps_output(output: &str, root_pids: &[u32]) -> Vec { - fn append_descendants( - parent: u32, - processes: &[(u32, u32)], - seen: &mut HashSet, - output: &mut Vec, - ) { - for (pid, ppid) in processes { - if *ppid != parent || *pid == parent || !seen.insert(*pid) { - continue; - } - append_descendants(*pid, processes, seen, output); - output.push(*pid); - } - } - - let processes = output - .lines() - .filter_map(|line| { - let mut parts = line.split_whitespace(); - let pid = parts.next()?.parse::().ok()?; - let ppid = parts.next()?.parse::().ok()?; - Some((pid, ppid)) - }) - .collect::>(); - let mut seen = root_pids.iter().copied().collect::>(); - let mut descendants = Vec::new(); - for root in root_pids { - append_descendants(*root, &processes, &mut seen, &mut descendants); - } - descendants -} - -fn terminate_process(pid: u32) -> Result { - if pid == std::process::id() { - return Ok(false); - } - #[cfg(unix)] - { - let output = ProcessCommand::new("kill") - .arg("-TERM") - .arg(pid.to_string()) - .stdin(Stdio::null()) - .stdout(Stdio::piped()) - .stderr(Stdio::piped()) - .output() - .with_context(|| format!("failed to launch kill for pid {pid}"))?; - if output.status.success() { - Ok(true) - } else { - let stderr = String::from_utf8_lossy(&output.stderr).trim().to_owned(); - let stdout = String::from_utf8_lossy(&output.stdout).trim().to_owned(); - if stderr.contains("No such process") || stdout.contains("No such process") { - return Ok(false); - } - bail!( - "failed to signal pid {pid}: {}", - if !stderr.is_empty() { - stderr - } else if !stdout.is_empty() { - stdout - } else { - format!("exit status {}", output.status) - } - ) - } - } - #[cfg(windows)] - { - let output = ProcessCommand::new("taskkill") - .arg("/PID") - .arg(pid.to_string()) - .arg("/T") - .arg("/F") - .stdin(Stdio::null()) - .stdout(Stdio::piped()) - .stderr(Stdio::piped()) - .output() - .with_context(|| format!("failed to launch taskkill for pid {pid}"))?; - if output.status.success() { - Ok(true) - } else { - let stderr = String::from_utf8_lossy(&output.stderr).trim().to_owned(); - let stdout = String::from_utf8_lossy(&output.stdout).trim().to_owned(); - if stderr.contains("not found") || stdout.contains("not found") { - return Ok(false); - } - bail!( - "failed to stop pid {pid}: {}", - if !stderr.is_empty() { - stderr - } else if !stdout.is_empty() { - stdout - } else { - format!("exit status {}", output.status) - } - ) - } - } -} - -#[cfg(unix)] -fn force_terminate_remaining_processes(pids: &[u32]) -> Result> { - if pids.is_empty() { - return Ok(Vec::new()); - } - thread::sleep(Duration::from_millis(750)); - let mut force_signaled = Vec::new(); - for pid in pids { - if *pid == std::process::id() || !process_is_running(*pid)? { - continue; - } - let output = ProcessCommand::new("kill") - .arg("-KILL") - .arg(pid.to_string()) - .stdin(Stdio::null()) - .stdout(Stdio::piped()) - .stderr(Stdio::piped()) - .output() - .with_context(|| format!("failed to launch kill -KILL for pid {pid}"))?; - if output.status.success() { - force_signaled.push(*pid); - } else { - let stderr = String::from_utf8_lossy(&output.stderr).trim().to_owned(); - let stdout = String::from_utf8_lossy(&output.stdout).trim().to_owned(); - bail!( - "failed to force stop pid {pid}: {}", - if !stderr.is_empty() { - stderr - } else if !stdout.is_empty() { - stdout - } else { - format!("exit status {}", output.status) - } - ); - } - } - Ok(force_signaled) -} - -#[cfg(not(unix))] -fn force_terminate_remaining_processes(_pids: &[u32]) -> Result> { - Ok(Vec::new()) -} - -#[cfg(unix)] -fn process_is_running(pid: u32) -> Result { - let output = ProcessCommand::new("kill") - .arg("-0") - .arg(pid.to_string()) - .stdin(Stdio::null()) - .stdout(Stdio::piped()) - .stderr(Stdio::piped()) - .output() - .with_context(|| format!("failed to launch kill -0 for pid {pid}"))?; - Ok(output.status.success()) -} - -async fn run_daemon( - paths: &AppPaths, - automations_enabled: bool, - local_webhook_port: Option, -) -> Result<()> { - if local_webhook_port.is_some() && !automations_enabled { - bail!("local webhook source requires --automations-enabled"); - } - - let config = RocmCliConfig::load(paths)?; - let mut state = build_runtime_state(&config, automations_enabled); - let local_webhook = if let Some(port) = local_webhook_port { - Some(webhook::start_local_webhook_source(port).await?) - } else { - None - }; - let (local_webhook_endpoint, mut local_webhook_receiver, local_webhook_task) = - match local_webhook { - Some(source) => ( - Some(source.endpoint), - Some(source.receiver), - Some(source.task), - ), - None => (None, None, None), - }; - state.local_webhook_endpoint = local_webhook_endpoint.clone(); - - println!("rocmd run"); - println!(" automations enabled: {automations_enabled}"); - println!( - " lifecycle: {}", - if automations_enabled { - "persistent" - } else { - "on-demand" - } - ); - println!(" config: {}", paths.config_path().display()); - println!(" state: {}", paths.automation_state_path().display()); - println!( - " local_webhook_endpoint: {}", - local_webhook_endpoint.as_deref().unwrap_or("disabled") - ); - let enabled_count = state - .active_watchers - .iter() - .filter(|watcher| watcher.enabled) - .count(); - println!(" enabled watchers: {enabled_count}"); - // This banner is the foreground-loop readiness contract used by callers and - // integration tests. Flush it before any persistent work so piped stdout on - // Windows cannot retain the line in a userspace buffer indefinitely. - io::stdout() - .flush() - .context("failed to flush rocmd run banner")?; - - if !automations_enabled { - println!( - " note: rerun with --automations-enabled to keep rocmd alive for watcher execution" - ); - return Ok(()); - } - - paths.ensure()?; - state.write(paths)?; - persistence::record_event( - paths, - &mut state, - "rocmd", - "info", - "daemon_start", - "rocmd automation supervisor started", - None, - )?; - state.write(paths)?; - - evaluate_watchers(paths, &config, &mut state)?; - state.last_tick_unix_ms = unix_time_millis(); - state.write(paths)?; - - let shutdown = shutdown_signal(); - tokio::pin!(shutdown); - - let mut ticker = time::interval(WATCHER_TICK_INTERVAL); - ticker.set_missed_tick_behavior(MissedTickBehavior::Delay); - - loop { - tokio::select! { - _ = ticker.tick() => { - let config = RocmCliConfig::load(paths)?; - reconcile_watcher_snapshots(&config, &mut state); - evaluate_watchers(paths, &config, &mut state)?; - state.last_tick_unix_ms = unix_time_millis(); - state.write(paths)?; - } - event = webhook::receive_local_webhook_event(&mut local_webhook_receiver) => { - if let Some(event) = event { - let config = RocmCliConfig::load(paths)?; - reconcile_watcher_snapshots(&config, &mut state); - persistence::record_event( - paths, - &mut state, - "rocmd", - "info", - "local_webhook_event", - &format!( - "received local webhook event kind={} watcher_hint={}; dispatching through existing watcher policy; webhook payload grants no new action", - event.kind, - event.watcher_hint.as_deref().unwrap_or("") - ), - event.service_id.clone(), - )?; - if let Err(error) = - evaluate_watchers_for_events(paths, &config, &mut state, &[event]) - { - persistence::record_event( - paths, - &mut state, - "rocmd", - "error", - "local_webhook_dispatch_failed", - &format!( - "local webhook event could not be dispatched through watcher policy: {error}" - ), - None, - )?; - } - state.last_tick_unix_ms = unix_time_millis(); - state.write(paths)?; - } else { - local_webhook_receiver = None; - state.local_webhook_endpoint = None; - persistence::record_event( - paths, - &mut state, - "rocmd", - "warn", - "local_webhook_stopped", - "local webhook source stopped; automation daemon continues without webhook ingestion", - None, - )?; - state.write(paths)?; - } - } - () = &mut shutdown => { - break; - } - } - } - - state.running = false; - state.last_tick_unix_ms = unix_time_millis(); - state.local_webhook_endpoint = None; - persistence::record_event( - paths, - &mut state, - "rocmd", - "info", - "daemon_stop", - "rocmd automation supervisor stopped", - None, - )?; - state.write(paths)?; - if let Some(task) = local_webhook_task { - task.abort(); - } - Ok(()) -} - -fn print_status(paths: &AppPaths) -> Result<()> { - let config = RocmCliConfig::load(paths).unwrap_or_default(); - println!("rocmd status"); - println!(" config dir: {}", paths.config_dir.display()); - println!(" data dir: {}", paths.data_dir.display()); - println!(" policy: on-demand by default, persistent only with background features"); - println!( - " automations desired: {}", - if config.automation_daemon_enabled() { - "enabled" - } else { - "disabled" - } - ); - match AutomationRuntimeState::load(paths)? { - Some(state) => { - println!( - " automations runtime: {} pid={} last_tick_unix_ms={}", - if state.running { "running" } else { "stopped" }, - state.daemon_pid, - state.last_tick_unix_ms - ); - println!( - " local_webhook_endpoint: {}", - state - .local_webhook_endpoint - .as_deref() - .unwrap_or("disabled") - ); - for watcher in state - .active_watchers - .into_iter() - .filter(|watcher| watcher.enabled) - { - println!( - " watcher {} mode={} last_event={}", - watcher.id, - watcher.mode.as_str(), - watcher.last_event.as_deref().unwrap_or("") - ); - } - } - None => println!(" automations runtime: inactive"), - } - println!( - " automation events: {}", - paths.automation_events_path().display() - ); - println!(" audit events: {}", paths.audit_events_path().display()); - - let records = persistence::load_managed_services(paths)?; - if records.is_empty() { - println!(" services: none"); - return Ok(()); - } - - for record in records { - println!( - " service {} engine={} status={} endpoint={}", - record.service_id, record.engine, record.status, record.endpoint_url - ); - } - - Ok(()) -} - -fn parse_gpu_indices_arg(value: Option<&str>) -> Result> { - let Some(raw) = value else { - return Ok(Vec::new()); - }; - match rocm_engine_protocol::GpuSelection::parse_cli_value(raw).map_err(anyhow::Error::msg)? { - rocm_engine_protocol::GpuSelection::Auto => Ok(Vec::new()), - rocm_engine_protocol::GpuSelection::Index(index) => Ok(vec![index]), - } -} - -#[allow(clippy::too_many_arguments)] -fn supervise_service( - paths: &AppPaths, - service_id: String, - engine: String, - model_ref: String, - canonical_model_id: String, - runtime_id: Option, - env_id: Option, - host: String, - port: u16, - device_policy: String, - gpu: Option, - engine_recipe_json: Option, -) -> Result<()> { - paths.ensure()?; - fs::create_dir_all(paths.engine_logs_dir(&engine))?; - fs::create_dir_all(paths.engine_state_dir(&engine))?; - fs::create_dir_all(paths.services_dir())?; - - let gpu_indices = parse_gpu_indices_arg(gpu.as_deref())?; - let _ = daemon_binary_path(); - - let mut record = ManagedServiceRecord::new( - paths, - service_id, - engine.clone(), - model_ref, - canonical_model_id.clone(), - host, - port, - "managed", - std::process::id(), - runtime_id.clone(), - env_id.clone(), - Some(device_policy.clone()), - ); - record.gpu_indices = gpu_indices; - record.engine_recipe_json = engine_recipe_json.clone(); - // Carried over from whatever is on disk. `ManagedServiceRecord::new` starts - // this false, so rebuilding a record here without restoring it would not - // just skip the check now — it would write the weakened record back and - // disarm every later `rocm services restart` as well. - // - // Propagated, not defaulted. This read arms the guard below, so it is not - // best-effort the way an identical-looking call feeding a printed warning - // would be. `load_managed_services` already *skips* unparseable records, so - // an `Err` here is a real I/O failure — and a missing directory is `Ok` - // anyway. Swallowing it would say "no service ever required a key", the - // key-file fallback is false precisely when a service has been stopped, and - // the weakened record would then be written back at the bottom of this - // function. That is the outcome the comment above says must not happen. - let previously_required = persistence::load_managed_services(paths) - .context( - "could not read the service registry to check whether this service requires an \ - endpoint API key; refusing to recover it rather than assume it does not", - )? - .iter() - .any(|existing| existing.service_id == record.service_id && existing.requires_api_key); - // Only what the registry recorded. The `|| key-file-is-present` clause that - // used to be here re-derived the flag the same way `spawn_managed_engine_child` - // did, and was wrong for the same reason: a public bind always has a key file - // whether or not auth was ever demanded, so recovery re-armed this on services - // that never asked for it and refused them with the wrong remediation. - record.requires_api_key = previously_required; - // Refuse a keyless public respawn before the manifest write, so a refused - // attempt leaves the recorded restart_count and timestamps intact instead of - // clobbering them with a record no live process will ever back. The spawn - // site below re-checks against the key actually threaded onto the command. - common::ensure_public_service_has_endpoint_key( - &record.host, - rocm_engine_protocol::endpoint_key_file_if_present(paths, &record.service_id) - .and_then(|path| rocm_engine_protocol::endpoint_api_key_file_if_valid(&path)) - .is_some(), - record.requires_api_key, - )?; - record.write()?; - - let log_file = fs::File::create(&record.log_path) - .with_context(|| format!("failed to create {}", record.log_path.display()))?; - let log_file_err = log_file - .try_clone() - .context("failed to clone service log file handle")?; - - let rocm_binary = - std::env::current_exe().context("failed to resolve current rocm executable path")?; - let mut command = ProcessCommand::new(rocm_binary); - command - .args(engine_serve_http_args( - &engine, - &record.service_id, - &canonical_model_id, - &record.host, - record.port, - &device_policy, - &record.gpu_indices, - runtime_id.as_deref(), - env_id.as_deref(), - engine_recipe_json.as_deref(), - &record.engine_state_path, - )) - .stdin(Stdio::null()) - .stdout(Stdio::from(log_file)) - .stderr(Stdio::from(log_file_err)); - // Re-thread the endpoint key file (public bind only) onto the engine child, - // same as the initial `rocm serve` spawn. This path also runs on daemon - // recovery (`restart_managed_service` re-execs `rocmd supervise`), so - // without this a previously-authenticated public service would come back - // up anonymous after a crash/recover cycle. - // If the key is gone the child would listen on the recorded public host with - // no auth, so fail closed instead — an unreachable service is recoverable, - // an anonymous public one is not. - let endpoint_key_applied = - common::apply_endpoint_key_env(&mut command, paths, &record.service_id); - common::ensure_public_service_has_endpoint_key( - &record.host, - endpoint_key_applied, - record.requires_api_key, - )?; - let mut child = command - .spawn() - .with_context(|| format!("failed to spawn engine supervisor child for {engine}"))?; - - record.engine_pid = Some(child.id()); - record.status = "running".to_owned(); - record.write()?; - - // Clone the fields the poller reads so the `on_phase` closure can borrow - // `record` mutably to persist each startup-phase transition to disk. - let ready_engine = record.engine.clone(); - let ready_service_id = record.service_id.clone(); - let ready_log_path = record.log_path.clone(); - let became_ready = wait_for_service_ready( - paths, - &ready_engine, - &ready_service_id, - &ready_log_path, - Duration::from_mins(3), - |phase| { - record.startup_phase = Some(phase.to_owned()); - let _ = record.write(); - }, - ); - if became_ready { - record.status = "ready".to_owned(); - // The phase only describes the coming-up window; clear it once ready. - record.startup_phase = None; - record.write()?; - } - - let exit_status = child.wait().context("failed waiting for engine child")?; - record.status = if exit_status.success() { - "stopped".to_owned() - } else { - "failed".to_owned() - }; - record.write()?; - - if exit_status.success() { - Ok(()) - } else { - std::process::exit(exit_status.code().unwrap_or(1)); - } -} - -#[allow(clippy::too_many_arguments)] -fn engine_serve_http_args( - engine: &str, - service_id: &str, - canonical_model_id: &str, - host: &str, - port: u16, - device_policy: &str, - gpu_indices: &[u32], - runtime_id: Option<&str>, - env_id: Option<&str>, - engine_recipe_json: Option<&str>, - state_path: &Path, -) -> Vec { - let mut args = vec![ - "__engine-serve-http".to_owned(), - engine.to_owned(), - service_id.to_owned(), - canonical_model_id.to_owned(), - "--host".to_owned(), - host.to_owned(), - "--port".to_owned(), - port.to_string(), - "--device-policy".to_owned(), - device_policy.to_owned(), - ]; - if let Some(csv) = rocm_engine_protocol::gpu_indices_to_csv(gpu_indices) { - args.extend(["--gpu".to_owned(), csv]); - } - args.extend(common::optional_arg("--runtime-id", runtime_id)); - args.extend(common::optional_arg("--env-id", env_id)); - args.extend(common::optional_arg( - "--engine-recipe-json", - engine_recipe_json, - )); - args.extend(["--state-path".to_owned(), state_path.display().to_string()]); - args -} - -fn build_runtime_state( - config: &RocmCliConfig, - automations_enabled: bool, -) -> AutomationRuntimeState { - let now = unix_time_millis(); - let active_watchers = builtin_watchers() - .iter() - .map(|watcher| WatcherRuntimeSnapshot { - id: watcher.id.to_owned(), - enabled: config.watcher_enabled(watcher), - mode: config.effective_watcher_mode(watcher), - summary: watcher.summary.to_owned(), - last_event: None, - last_event_unix_ms: None, - }) - .collect(); - AutomationRuntimeState { - running: automations_enabled, - automations_enabled, - daemon_pid: std::process::id(), - started_at_unix_ms: now, - last_tick_unix_ms: now, - local_webhook_endpoint: None, - active_watchers, - } -} - fn reconcile_watcher_snapshots(config: &RocmCliConfig, state: &mut AutomationRuntimeState) { for watcher in builtin_watchers() { match state.watcher_mut(watcher.id) { @@ -2328,370 +1600,7 @@ fn detached_rocmd_command(rocmd_binary: &std::path::Path) -> ProcessCommand { #[cfg(test)] mod tests { use super::*; - use crate::test_support::{temp_app_paths, unique_test_root}; - - /// Drive `supervise_service` far enough to reach the key guard, and return - /// what it did. - /// - /// The guard sits before the manifest write and well before any spawn, so a - /// refusal returns without starting a process — which is what makes the real - /// call site testable at all. The arguments below are the shape a recovery - /// re-exec passes: a loopback bind, no GPU, no recipe. - /// - /// This exists because testing `ensure_public_service_has_endpoint_key` - /// directly with literal arguments cannot catch the defect that actually - /// happened twice in this crate's history — the guard being *wired up* with - /// the wrong value at its call site. - /// - /// Bounded, and the bound is the assertion. A guard that fails to refuse - /// does not return an error — it falls through to the engine spawn and - /// supervises a child that never exits, so an unbounded call would hang the - /// suite instead of failing it. Both callers below are regression tests for - /// a fail-*open*, which is exactly the shape that turns into a hang. - fn supervise_at_the_guard(paths: &AppPaths, service_id: &str) -> Result<()> { - let paths = paths.clone(); - let service_id = service_id.to_owned(); - let (sender, receiver) = std::sync::mpsc::channel(); - std::thread::spawn(move || { - let outcome = supervise_service( - &paths, - service_id, - "llamacpp".to_owned(), - "a-model".to_owned(), - "a-model".to_owned(), - None, - None, - "127.0.0.1".to_owned(), - 11434, - "gpu_required".to_owned(), - None, - None, - ); - let _ = sender.send(outcome.map_err(|error| format!("{error:#}"))); - }); - receiver - .recv_timeout(std::time::Duration::from_secs(30)) - .unwrap_or_else(|_| { - panic!( - "supervise_service did not return within 30s: the key guard let the call \ - through and it reached the engine spawn, which is the fail-open this test \ - exists to catch" - ) - }) - .map_err(anyhow::Error::msg) - } - - /// Write a service record into the registry the way a live service would - /// have left it behind. - fn seed_registry(paths: &AppPaths, service_id: &str, requires_api_key: bool) { - fs::create_dir_all(paths.services_dir()).unwrap(); - let mut record = ManagedServiceRecord::new( - paths, - service_id.to_owned(), - "llamacpp".to_owned(), - "a-model".to_owned(), - "a-model".to_owned(), - "127.0.0.1".to_owned(), - 11434, - "managed", - std::process::id(), - None, - None, - Some("gpu_required".to_owned()), - ); - record.requires_api_key = requires_api_key; - record.write().unwrap(); - } - - /// Drive `supervise_service` and report whether the key guard let it past. - /// - /// Decided on what the call *returns*, not on any file. The obvious - /// observable — the manifest appearing — is useless here, because - /// `seed_registry` has already written one, so polling for it passes - /// whatever the guard does. That mistake was made first and caught by - /// mutating the code the test claims to protect. - /// - /// A call the guard admits does not return: it carries on to the engine - /// spawn. So the guard's refusal is the only thing that comes back quickly, - /// and it is identified by its message rather than by the mere fact of an - /// error — a later, unrelated failure must not read as a refusal. - fn guard_admits(paths: &AppPaths, service_id: &str) -> bool { - let owned_paths = paths.clone(); - let owned_id = service_id.to_owned(); - let (sender, receiver) = std::sync::mpsc::channel(); - std::thread::spawn(move || { - let outcome = supervise_service( - &owned_paths, - owned_id, - "llamacpp".to_owned(), - "a-model".to_owned(), - "a-model".to_owned(), - None, - None, - "127.0.0.1".to_owned(), - 11434, - "gpu_required".to_owned(), - None, - None, - ); - let _ = sender.send(outcome.map_err(|error| format!("{error:#}"))); - }); - - match receiver.recv_timeout(std::time::Duration::from_secs(10)) { - // The key guard refused, by its own words. - Ok(Err(rendered)) if rendered.contains("without authentication") => false, - // Anything else means it got past the guard: it either finished, or - // failed later for a reason that is not this guard, or is still - // running because it reached the spawn. - _ => true, - } - } - - #[test] - fn a_service_that_never_required_a_key_is_not_refused_for_lacking_one() { - // The other direction of the guard, and the one no test covered. - // Hardcoding `record.requires_api_key = true` at the restore site passes - // every other test in this crate, because they all seed a service that - // *does* require a key. This is the case that catches it. - // - // Two records are seeded, not one: with a single record the - // `existing.service_id == record.service_id` half of the lookup does - // nothing, so dropping that comparison would go unnoticed and one - // service's requirement would leak onto another's. - let (root, paths) = temp_app_paths("supervise-no-key-needed"); - seed_registry(&paths, "svc-needs-key", true); - seed_registry(&paths, "svc-plain", false); - - assert!( - guard_admits(&paths, "svc-plain"), - "a loopback service that never asked for a key must not be refused for lacking one" - ); - - let _ = fs::remove_dir_all(&root); - } - - #[test] - fn supervising_a_service_that_required_a_key_refuses_when_the_key_is_gone() { - // The real call site, not the guard in isolation. `supervise_service` - // rebuilds the record with `ManagedServiceRecord::new`, which starts - // `requires_api_key` false, and restores it from the registry. Passing - // the wrong value here — a literal, or the freshly-built field before it - // is restored — is exactly the miswiring that shipped twice in this - // crate and that a literal-argument unit test cannot see. - let (root, paths) = temp_app_paths("supervise-requires-key"); - seed_registry(&paths, "svc-needs-key", true); - - let error = supervise_at_the_guard(&paths, "svc-needs-key") - .expect_err("a service that required a key must not be recovered without one"); - assert!( - format!("{error:#}").contains("without authentication"), - "{error:#}" - ); - - // And the refusal must not have weakened what is on disk. The guard runs - // before `record.write()` precisely so a refused attempt leaves the - // recorded requirement armed for the next attempt. - let stored = persistence::load_managed_services(&paths).unwrap(); - let stored = stored - .iter() - .find(|candidate| candidate.service_id == "svc-needs-key") - .expect("the seeded record must survive a refused recovery"); - assert!( - stored.requires_api_key, - "a refused recovery must not disarm the requirement" - ); - - let _ = fs::remove_dir_all(&root); - } - - #[test] - fn a_registry_that_cannot_be_read_refuses_recovery_rather_than_assuming_no_key() { - // `load_managed_services` already skips records it cannot parse, so an - // `Err` from it is a real I/O failure — and a missing directory is `Ok`. - // Defaulting it away therefore says "no service ever required a key", - // which is fail-open on an auth gate and, worse, gets written back. - // - // The failure is provoked portably: a directory named like a record - // makes the `fs::read` inside the loop fail rather than the read_dir. - let (root, paths) = temp_app_paths("supervise-unreadable-registry"); - fs::create_dir_all(paths.services_dir().join("not-a-record.json")).unwrap(); - - let error = supervise_at_the_guard(&paths, "svc-unknown") - .expect_err("an unreadable registry must refuse, not assume no key was required"); - let rendered = format!("{error:#}"); - assert!( - rendered.contains("could not read the service registry"), - "{rendered}" - ); - - // Nothing was written: a registry we could not read is not a registry we - // may add a weakened record to. - assert!( - !paths.service_manifest_path("svc-unknown").exists(), - "a refused recovery must not persist a record" - ); - - let _ = fs::remove_dir_all(&root); - } - - #[test] - fn last_cr_segment_keeps_final_progress_redraw() { - // A tqdm/HF-style in-place redraw collapses to its last segment. - assert_eq!( - last_cr_segment("Downloading: 10%\rDownloading: 55%\rDownloading: 100%"), - "Downloading: 100%" - ); - // A plain line is unchanged. - assert_eq!( - last_cr_segment("Loading model weights"), - "Loading model weights" - ); - } - - #[test] - fn classify_startup_phase_maps_engine_vocabulary() { - assert_eq!( - classify_startup_phase("Downloading shards: 100%"), - Some("downloading") - ); - assert_eq!( - classify_startup_phase("Fetching 12 files"), - Some("downloading") - ); - assert_eq!( - classify_startup_phase("INFO: Loading model weights took 4.2s"), - Some("loading") - ); - assert_eq!( - classify_startup_phase("llama_model_loader: loaded meta data"), - Some("loading") - ); - assert_eq!( - classify_startup_phase("Capturing CUDA graph shapes"), - Some("warmup") - ); - assert_eq!( - classify_startup_phase("Warming up the engine"), - Some("warmup") - ); - // Ordinary chatter carries no phase signal. - assert_eq!( - classify_startup_phase("Uvicorn running on http://..."), - None - ); - } - - #[test] - fn classify_startup_phase_emits_only_dashboard_known_tokens() { - // These tokens are the wire contract with the dashboard's - // `StartupPhase::from_token` (rocm-dash-core); emitting anything else - // would be silently dropped there. rocmd can't link that crate, so the - // contract is pinned here by literal. - for line in [ - "Downloading shards", - "Loading model weights", - "Capturing CUDA graph", - ] { - let token = classify_startup_phase(line).expect("line is a phase signal"); - assert!( - matches!(token, "downloading" | "loading" | "warmup"), - "token {token:?} must be one the dashboard understands" - ); - } - } - - #[test] - fn read_new_log_phase_advances_and_tracks_latest() { - use std::io::Write as _; - // Workspace-local test root (rooted at CARGO_MANIFEST_DIR, not the - // ambient temp dir) — same helper the other rocmd tests use. - let dir = unique_test_root(&format!("rocmd-phase-{}", std::process::id())); - let log = dir.join("svc.log"); - std::fs::write(&log, "boot\nDownloading shards: 100%\n").unwrap(); - - let mut pos = 0_u64; - assert_eq!(read_new_log_phase(&log, &mut pos), Some("downloading")); - // No new bytes → no phase, cursor unchanged. - let after_first = pos; - assert_eq!(read_new_log_phase(&log, &mut pos), None); - assert_eq!(pos, after_first); - - // Appending a later stage advances the phase. - let mut f = std::fs::OpenOptions::new().append(true).open(&log).unwrap(); - writeln!(f, "Loading model weights took 3s").unwrap(); - assert_eq!(read_new_log_phase(&log, &mut pos), Some("loading")); - - let _ = std::fs::remove_dir_all(&dir); - } - - #[test] - fn engine_serve_http_args_forward_engine_recipe_json() { - let engine_recipe_json = r#"{"contract_version":"0.1.0","engine":"vllm","required_flags":["--enable-auto-tool-choice"]}"#; - let args = engine_serve_http_args( - "vllm", - "svc-1", - "Qwen/Qwen3.5-4B", - "127.0.0.1", - 11435, - "gpu_required", - &[], - Some("therock-release:gfx120X-all"), - Some("env-1"), - Some(engine_recipe_json), - Path::new("state.json"), - ); - - assert!( - args.windows(2) - .any(|pair| pair[0] == "--engine-recipe-json" && pair[1] == engine_recipe_json) - ); - assert!( - args.windows(2).any(|pair| { - pair[0] == "--runtime-id" && pair[1] == "therock-release:gfx120X-all" - }) - ); - assert!( - args.windows(2) - .any(|pair| pair[0] == "--state-path" && pair[1] == "state.json") - ); - } - - #[test] - fn engine_serve_http_args_emit_gpu_indices_when_pinned() { - let args = engine_serve_http_args( - "vllm", - "svc-1", - "Qwen/Qwen3.5-4B", - "127.0.0.1", - 11435, - "gpu_required", - &[1], - None, - None, - None, - Path::new("state.json"), - ); - - assert!( - args.windows(2) - .any(|pair| pair[0] == "--gpu" && pair[1] == "1") - ); - - let auto = engine_serve_http_args( - "vllm", - "svc-1", - "Qwen/Qwen3.5-4B", - "127.0.0.1", - 11435, - "gpu_required", - &[], - None, - None, - None, - Path::new("state.json"), - ); - assert!(!auto.iter().any(|arg| arg == "--gpu")); - } + use crate::test_support::temp_app_paths; #[test] fn recovery_supervise_args_preserve_engine_recipe_json() { @@ -2745,14 +1654,6 @@ mod tests { ); } - #[tokio::test] - async fn local_webhook_requires_enabled_automation_loop() { - let (_root, paths) = temp_app_paths("local-webhook-requires-loop"); - let error = run_daemon(&paths, false, Some(0)).await.unwrap_err(); - - assert!(error.to_string().contains("requires --automations-enabled")); - } - #[test] fn event_collector_emits_schedule_tick_for_due_update() -> Result<()> { let (root, paths) = temp_app_paths("event-bus-schedule"); @@ -4249,120 +3150,6 @@ mod tests { Ok(()) } - #[test] - fn stop_managed_service_removes_endpoint_key_file() -> Result<()> { - let (root, paths) = temp_app_paths("stop-removes-endpoint-key"); - paths.ensure()?; - let current_pid = std::process::id(); - let service_id = "svc-endpoint-key-stop"; - let mut record = ManagedServiceRecord::new( - &paths, - service_id, - "vllm", - "qwen", - "Qwen/Qwen3.5", - "0.0.0.0", - 11435, - "managed", - current_pid, - None, - None, - None, - ); - record.engine_pid = Some(current_pid); - record.status = "ready".to_owned(); - record.write()?; - - let key_path = rocm_engine_protocol::endpoint_key_file_path(&paths, service_id); - fs::create_dir_all(paths.services_dir())?; - fs::write(&key_path, "secret-key")?; - assert!(key_path.exists()); - - let result = stop_managed_service(&paths, service_id); - // Observe the real filesystem state before the blanket temp-dir cleanup, - // otherwise remove_dir_all would delete the key file and mask a missing - // production cleanup (the regression this test guards). - let key_removed = !key_path.exists(); - fs::remove_dir_all(root).ok(); - - let value = result?; - assert_eq!( - value - .get("service") - .and_then(|service| service.get("status")) - .and_then(Value::as_str), - Some("stopped") - ); - assert!(key_removed, "endpoint key file must be removed after stop"); - Ok(()) - } - - #[test] - fn stop_managed_service_without_endpoint_key_file_succeeds() -> Result<()> { - let (root, paths) = temp_app_paths("stop-no-endpoint-key"); - paths.ensure()?; - let current_pid = std::process::id(); - let service_id = "svc-no-endpoint-key-stop"; - let mut record = ManagedServiceRecord::new( - &paths, - service_id, - "vllm", - "qwen", - "Qwen/Qwen3.5", - "127.0.0.1", - 11435, - "managed", - current_pid, - None, - None, - None, - ); - record.engine_pid = Some(current_pid); - record.status = "ready".to_owned(); - record.write()?; - - // Loopback service: no endpoint key file was ever written for it. - let key_path = rocm_engine_protocol::endpoint_key_file_path(&paths, service_id); - assert!(!key_path.exists()); - - let result = stop_managed_service(&paths, service_id); - let reloaded = load_service_record(&paths, service_id); - fs::remove_dir_all(root).ok(); - - let value = result?; - assert_eq!( - value - .get("service") - .and_then(|service| service.get("status")) - .and_then(Value::as_str), - Some("stopped") - ); - assert_eq!(reloaded?.status, "stopped"); - assert!(!key_path.exists()); - Ok(()) - } - - #[test] - fn stop_server_process_tree_discovers_descendants_before_parents() { - let output = "\ -10 1 -11 10 -12 11 -13 10 -20 1 -21 20 -"; - - assert_eq!( - descendant_pids_from_ps_output(output, &[10]), - vec![12, 11, 13] - ); - assert_eq!( - descendant_pids_from_ps_output(output, &[10, 20]), - vec![12, 11, 13, 21] - ); - } - fn test_watcher_snapshot( id: &str, mode: WatcherMode, @@ -4390,119 +3177,3 @@ mod tests { } } } - -#[cfg(unix)] -async fn shutdown_signal() { - let mut term = tokio::signal::unix::signal(tokio::signal::unix::SignalKind::terminate()) - .expect("failed to register SIGTERM handler"); - tokio::select! { - _ = tokio::signal::ctrl_c() => {} - _ = term.recv() => {} - } -} - -#[cfg(not(unix))] -async fn shutdown_signal() { - let _ = tokio::signal::ctrl_c().await; -} - -/// Keep only the final visible segment of a `\r`-redrawn progress line. -/// -/// Progress tools (pip, tqdm, Hugging Face) redraw a line in place with a bare -/// carriage return and no newline, so the segment after the last `\r` is its -/// final visible state. Lines without `\r` pass through unchanged. (Same -/// collapse rule the dashboard job console applies to streamed job output.) -fn last_cr_segment(line: &str) -> &str { - line.rsplit('\r').next().unwrap_or(line) -} - -/// Classify a single serve-log line into a coarse startup phase token -/// (`downloading`/`loading`/`warmup`), or `None` when the line carries no phase -/// signal. Case-insensitive substring match over the common vLLM / llama.cpp / -/// Hugging Face startup vocabulary. Checked warmup → loading → downloading so -/// the latest lifecycle stage a line mentions wins. -fn classify_startup_phase(line: &str) -> Option<&'static str> { - let lower = line.to_ascii_lowercase(); - if lower.contains("capturing cuda graph") - || lower.contains("capturing the model") - || lower.contains("warming up") - || lower.contains("warmup") - { - Some("warmup") - } else if lower.contains("loading weights") - || lower.contains("loading model") - || lower.contains("load_tensors") - || lower.contains("llama_model_loader") - || lower.contains("model loading took") - { - Some("loading") - } else if lower.contains("downloading") || lower.contains("fetching") { - Some("downloading") - } else { - None - } -} - -/// Read log bytes appended since `*pos`, advance `*pos`, and return the most -/// recent recognizable startup phase in that new output (later lines win, so a -/// download → load → warmup progression advances naturally). -/// -/// Best-effort: any I/O error (file not created yet, transient read) yields -/// `None`. A shrunk file (rotation/truncation) resets the cursor to the top. -fn read_new_log_phase(log_path: &Path, pos: &mut u64) -> Option<&'static str> { - let mut file = fs::File::open(log_path).ok()?; - let len = file.metadata().ok()?.len(); - if len < *pos { - *pos = 0; - } - if len == *pos { - return None; - } - file.seek(SeekFrom::Start(*pos)).ok()?; - let mut bytes = Vec::new(); - let read = file.read_to_end(&mut bytes).ok()?; - *pos += read as u64; - let text = String::from_utf8_lossy(&bytes); - let mut phase = None; - for line in text.lines() { - if let Some(found) = classify_startup_phase(last_cr_segment(line)) { - phase = Some(found); - } - } - phase -} - -fn engine_healthcheck_ready(paths: &AppPaths, engine: &str, service_id: &str) -> Result { - Ok(common::healthcheck_response_ready( - &common::engine_healthcheck_response(paths, engine, service_id)?, - )) -} - -/// Poll a freshly-spawned service until its healthcheck reports ready (or the -/// timeout elapses), tailing its log file meanwhile and reporting each coarse -/// startup phase transition via `on_phase`. -fn wait_for_service_ready( - paths: &AppPaths, - engine: &str, - service_id: &str, - log_path: &Path, - timeout: Duration, - mut on_phase: impl FnMut(&str), -) -> bool { - let start = std::time::Instant::now(); - let mut log_pos: u64 = 0; - let mut last_phase: Option<&'static str> = None; - while start.elapsed() < timeout { - if let Some(phase) = read_new_log_phase(log_path, &mut log_pos) - && last_phase != Some(phase) - { - last_phase = Some(phase); - on_phase(phase); - } - if engine_healthcheck_ready(paths, engine, service_id).unwrap_or(false) { - return true; - } - thread::sleep(Duration::from_millis(200)); - } - false -} diff --git a/apps/rocmd/src/mcp.rs b/apps/rocmd/src/mcp.rs index e6995aca0..aab2654db 100644 --- a/apps/rocmd/src/mcp.rs +++ b/apps/rocmd/src/mcp.rs @@ -741,7 +741,7 @@ pub(crate) fn handle_mcp_tool_call(paths: &AppPaths, params: &Value) -> Result Result { + let mut record = crate::persistence::load_managed_services(paths)? + .into_iter() + .find(|record| record.service_id == service_id) + .with_context(|| format!("managed service `{service_id}` not found"))?; + let mut signaled_pids = Vec::new(); + let mut skipped_pids = Vec::new(); + let mut root_pids = Vec::new(); + if let Some(engine_pid) = record.engine_pid + && engine_pid != 0 + { + if engine_pid == std::process::id() { + skipped_pids.push(engine_pid); + } else { + root_pids.push(engine_pid); + } + } + if record.supervisor_pid != 0 + && record.supervisor_pid != std::process::id() + && Some(record.supervisor_pid) != record.engine_pid + { + root_pids.push(record.supervisor_pid); + } + let mut pids_to_signal = descendant_pids_for_roots(&root_pids)?; + pids_to_signal.extend(root_pids); + let mut seen_pids = HashSet::new(); + for pid in pids_to_signal { + if !seen_pids.insert(pid) { + continue; + } + if terminate_process(pid)? { + signaled_pids.push(pid); + } else { + skipped_pids.push(pid); + } + } + let force_signaled_pids = force_terminate_remaining_processes(&signaled_pids)?; + record.status = "stopped".to_owned(); + record.write()?; + // Best-effort and idempotent: a missing key file is not an error, so this + // is safe to call unconditionally on every stop (including loopback + // services that never had a key, and repeated stops of an already-stopped + // service). Leaving the 0600 key file behind after stop would strand a + // plaintext secret on disk for a service that no longer exists. + let _ = std::fs::remove_file(rocm_engine_protocol::endpoint_key_file_path( + paths, service_id, + )); + Ok(json!({ + "service": record, + "signaled_pids": signaled_pids, + "force_signaled_pids": force_signaled_pids, + "skipped_pids": skipped_pids, + })) +} + +#[cfg(unix)] +fn descendant_pids_for_roots(root_pids: &[u32]) -> Result> { + if root_pids.is_empty() { + return Ok(Vec::new()); + } + let output = ProcessCommand::new("ps") + .args(["-eo", "pid=,ppid="]) + .stdin(Stdio::null()) + .stdout(Stdio::piped()) + .stderr(Stdio::piped()) + .output() + .context("failed to list process tree with ps")?; + if !output.status.success() { + let stderr = String::from_utf8_lossy(&output.stderr).trim().to_owned(); + let stdout = String::from_utf8_lossy(&output.stdout).trim().to_owned(); + bail!( + "failed to list process tree: {}", + if !stderr.is_empty() { + stderr + } else if !stdout.is_empty() { + stdout + } else { + format!("exit status {}", output.status) + } + ); + } + let output = String::from_utf8_lossy(&output.stdout); + Ok(descendant_pids_from_ps_output(&output, root_pids)) +} + +#[cfg(not(unix))] +fn descendant_pids_for_roots(_root_pids: &[u32]) -> Result> { + Ok(Vec::new()) +} + +#[cfg(any(unix, test))] +fn descendant_pids_from_ps_output(output: &str, root_pids: &[u32]) -> Vec { + fn append_descendants( + parent: u32, + processes: &[(u32, u32)], + seen: &mut HashSet, + output: &mut Vec, + ) { + for (pid, ppid) in processes { + if *ppid != parent || *pid == parent || !seen.insert(*pid) { + continue; + } + append_descendants(*pid, processes, seen, output); + output.push(*pid); + } + } + + let processes = output + .lines() + .filter_map(|line| { + let mut parts = line.split_whitespace(); + let pid = parts.next()?.parse::().ok()?; + let ppid = parts.next()?.parse::().ok()?; + Some((pid, ppid)) + }) + .collect::>(); + let mut seen = root_pids.iter().copied().collect::>(); + let mut descendants = Vec::new(); + for root in root_pids { + append_descendants(*root, &processes, &mut seen, &mut descendants); + } + descendants +} + +fn terminate_process(pid: u32) -> Result { + if pid == std::process::id() { + return Ok(false); + } + #[cfg(unix)] + { + let output = ProcessCommand::new("kill") + .arg("-TERM") + .arg(pid.to_string()) + .stdin(Stdio::null()) + .stdout(Stdio::piped()) + .stderr(Stdio::piped()) + .output() + .with_context(|| format!("failed to launch kill for pid {pid}"))?; + if output.status.success() { + Ok(true) + } else { + let stderr = String::from_utf8_lossy(&output.stderr).trim().to_owned(); + let stdout = String::from_utf8_lossy(&output.stdout).trim().to_owned(); + if stderr.contains("No such process") || stdout.contains("No such process") { + return Ok(false); + } + bail!( + "failed to signal pid {pid}: {}", + if !stderr.is_empty() { + stderr + } else if !stdout.is_empty() { + stdout + } else { + format!("exit status {}", output.status) + } + ) + } + } + #[cfg(windows)] + { + let output = ProcessCommand::new("taskkill") + .arg("/PID") + .arg(pid.to_string()) + .arg("/T") + .arg("/F") + .stdin(Stdio::null()) + .stdout(Stdio::piped()) + .stderr(Stdio::piped()) + .output() + .with_context(|| format!("failed to launch taskkill for pid {pid}"))?; + if output.status.success() { + Ok(true) + } else { + let stderr = String::from_utf8_lossy(&output.stderr).trim().to_owned(); + let stdout = String::from_utf8_lossy(&output.stdout).trim().to_owned(); + if stderr.contains("not found") || stdout.contains("not found") { + return Ok(false); + } + bail!( + "failed to stop pid {pid}: {}", + if !stderr.is_empty() { + stderr + } else if !stdout.is_empty() { + stdout + } else { + format!("exit status {}", output.status) + } + ) + } + } +} + +#[cfg(unix)] +fn force_terminate_remaining_processes(pids: &[u32]) -> Result> { + if pids.is_empty() { + return Ok(Vec::new()); + } + thread::sleep(Duration::from_millis(750)); + let mut force_signaled = Vec::new(); + for pid in pids { + if *pid == std::process::id() || !process_is_running(*pid)? { + continue; + } + let output = ProcessCommand::new("kill") + .arg("-KILL") + .arg(pid.to_string()) + .stdin(Stdio::null()) + .stdout(Stdio::piped()) + .stderr(Stdio::piped()) + .output() + .with_context(|| format!("failed to launch kill -KILL for pid {pid}"))?; + if output.status.success() { + force_signaled.push(*pid); + } else { + let stderr = String::from_utf8_lossy(&output.stderr).trim().to_owned(); + let stdout = String::from_utf8_lossy(&output.stdout).trim().to_owned(); + bail!( + "failed to force stop pid {pid}: {}", + if !stderr.is_empty() { + stderr + } else if !stdout.is_empty() { + stdout + } else { + format!("exit status {}", output.status) + } + ); + } + } + Ok(force_signaled) +} + +#[cfg(not(unix))] +fn force_terminate_remaining_processes(_pids: &[u32]) -> Result> { + Ok(Vec::new()) +} + +#[cfg(unix)] +fn process_is_running(pid: u32) -> Result { + let output = ProcessCommand::new("kill") + .arg("-0") + .arg(pid.to_string()) + .stdin(Stdio::null()) + .stdout(Stdio::piped()) + .stderr(Stdio::piped()) + .output() + .with_context(|| format!("failed to launch kill -0 for pid {pid}"))?; + Ok(output.status.success()) +} + +pub(crate) async fn run_daemon( + paths: &AppPaths, + automations_enabled: bool, + local_webhook_port: Option, +) -> Result<()> { + if local_webhook_port.is_some() && !automations_enabled { + bail!("local webhook source requires --automations-enabled"); + } + + let config = RocmCliConfig::load(paths)?; + let mut state = build_runtime_state(&config, automations_enabled); + let local_webhook = if let Some(port) = local_webhook_port { + Some(crate::webhook::start_local_webhook_source(port).await?) + } else { + None + }; + let (local_webhook_endpoint, mut local_webhook_receiver, local_webhook_task) = + match local_webhook { + Some(source) => ( + Some(source.endpoint), + Some(source.receiver), + Some(source.task), + ), + None => (None, None, None), + }; + state.local_webhook_endpoint = local_webhook_endpoint.clone(); + + println!("rocmd run"); + println!(" automations enabled: {automations_enabled}"); + println!( + " lifecycle: {}", + if automations_enabled { + "persistent" + } else { + "on-demand" + } + ); + println!(" config: {}", paths.config_path().display()); + println!(" state: {}", paths.automation_state_path().display()); + println!( + " local_webhook_endpoint: {}", + local_webhook_endpoint.as_deref().unwrap_or("disabled") + ); + let enabled_count = state + .active_watchers + .iter() + .filter(|watcher| watcher.enabled) + .count(); + println!(" enabled watchers: {enabled_count}"); + // This banner is the foreground-loop readiness contract used by callers and + // integration tests. Flush it before any persistent work so piped stdout on + // Windows cannot retain the line in a userspace buffer indefinitely. + io::stdout() + .flush() + .context("failed to flush rocmd run banner")?; + + if !automations_enabled { + println!( + " note: rerun with --automations-enabled to keep rocmd alive for watcher execution" + ); + return Ok(()); + } + + paths.ensure()?; + state.write(paths)?; + crate::persistence::record_event( + paths, + &mut state, + "rocmd", + "info", + "daemon_start", + "rocmd automation supervisor started", + None, + )?; + state.write(paths)?; + + crate::evaluate_watchers(paths, &config, &mut state)?; + state.last_tick_unix_ms = unix_time_millis(); + state.write(paths)?; + + let shutdown = shutdown_signal(); + tokio::pin!(shutdown); + + let mut ticker = time::interval(crate::WATCHER_TICK_INTERVAL); + ticker.set_missed_tick_behavior(MissedTickBehavior::Delay); + + loop { + tokio::select! { + _ = ticker.tick() => { + let config = RocmCliConfig::load(paths)?; + crate::reconcile_watcher_snapshots(&config, &mut state); + crate::evaluate_watchers(paths, &config, &mut state)?; + state.last_tick_unix_ms = unix_time_millis(); + state.write(paths)?; + } + event = crate::webhook::receive_local_webhook_event(&mut local_webhook_receiver) => { + if let Some(event) = event { + let config = RocmCliConfig::load(paths)?; + crate::reconcile_watcher_snapshots(&config, &mut state); + crate::persistence::record_event( + paths, + &mut state, + "rocmd", + "info", + "local_webhook_event", + &format!( + "received local webhook event kind={} watcher_hint={}; dispatching through existing watcher policy; webhook payload grants no new action", + event.kind, + event.watcher_hint.as_deref().unwrap_or("") + ), + event.service_id.clone(), + )?; + if let Err(error) = + crate::evaluate_watchers_for_events(paths, &config, &mut state, &[event]) + { + crate::persistence::record_event( + paths, + &mut state, + "rocmd", + "error", + "local_webhook_dispatch_failed", + &format!( + "local webhook event could not be dispatched through watcher policy: {error}" + ), + None, + )?; + } + state.last_tick_unix_ms = unix_time_millis(); + state.write(paths)?; + } else { + local_webhook_receiver = None; + state.local_webhook_endpoint = None; + crate::persistence::record_event( + paths, + &mut state, + "rocmd", + "warn", + "local_webhook_stopped", + "local webhook source stopped; automation daemon continues without webhook ingestion", + None, + )?; + state.write(paths)?; + } + } + () = &mut shutdown => { + break; + } + } + } + + state.running = false; + state.last_tick_unix_ms = unix_time_millis(); + state.local_webhook_endpoint = None; + crate::persistence::record_event( + paths, + &mut state, + "rocmd", + "info", + "daemon_stop", + "rocmd automation supervisor stopped", + None, + )?; + state.write(paths)?; + if let Some(task) = local_webhook_task { + task.abort(); + } + Ok(()) +} + +pub(crate) fn print_status(paths: &AppPaths) -> Result<()> { + let config = RocmCliConfig::load(paths).unwrap_or_default(); + println!("rocmd status"); + println!(" config dir: {}", paths.config_dir.display()); + println!(" data dir: {}", paths.data_dir.display()); + println!(" policy: on-demand by default, persistent only with background features"); + println!( + " automations desired: {}", + if config.automation_daemon_enabled() { + "enabled" + } else { + "disabled" + } + ); + match AutomationRuntimeState::load(paths)? { + Some(state) => { + println!( + " automations runtime: {} pid={} last_tick_unix_ms={}", + if state.running { "running" } else { "stopped" }, + state.daemon_pid, + state.last_tick_unix_ms + ); + println!( + " local_webhook_endpoint: {}", + state + .local_webhook_endpoint + .as_deref() + .unwrap_or("disabled") + ); + for watcher in state + .active_watchers + .into_iter() + .filter(|watcher| watcher.enabled) + { + println!( + " watcher {} mode={} last_event={}", + watcher.id, + watcher.mode.as_str(), + watcher.last_event.as_deref().unwrap_or("") + ); + } + } + None => println!(" automations runtime: inactive"), + } + println!( + " automation events: {}", + paths.automation_events_path().display() + ); + println!(" audit events: {}", paths.audit_events_path().display()); + + let records = crate::persistence::load_managed_services(paths)?; + if records.is_empty() { + println!(" services: none"); + return Ok(()); + } + + for record in records { + println!( + " service {} engine={} status={} endpoint={}", + record.service_id, record.engine, record.status, record.endpoint_url + ); + } + + Ok(()) +} + +fn parse_gpu_indices_arg(value: Option<&str>) -> Result> { + let Some(raw) = value else { + return Ok(Vec::new()); + }; + match rocm_engine_protocol::GpuSelection::parse_cli_value(raw).map_err(anyhow::Error::msg)? { + rocm_engine_protocol::GpuSelection::Auto => Ok(Vec::new()), + rocm_engine_protocol::GpuSelection::Index(index) => Ok(vec![index]), + } +} + +#[allow(clippy::too_many_arguments)] +pub(crate) fn supervise_service( + paths: &AppPaths, + service_id: String, + engine: String, + model_ref: String, + canonical_model_id: String, + runtime_id: Option, + env_id: Option, + host: String, + port: u16, + device_policy: String, + gpu: Option, + engine_recipe_json: Option, +) -> Result<()> { + paths.ensure()?; + fs::create_dir_all(paths.engine_logs_dir(&engine))?; + fs::create_dir_all(paths.engine_state_dir(&engine))?; + fs::create_dir_all(paths.services_dir())?; + + let gpu_indices = parse_gpu_indices_arg(gpu.as_deref())?; + let _ = daemon_binary_path(); + + let mut record = ManagedServiceRecord::new( + paths, + service_id, + engine.clone(), + model_ref, + canonical_model_id.clone(), + host, + port, + "managed", + std::process::id(), + runtime_id.clone(), + env_id.clone(), + Some(device_policy.clone()), + ); + record.gpu_indices = gpu_indices; + record.engine_recipe_json = engine_recipe_json.clone(); + // Carried over from whatever is on disk. `ManagedServiceRecord::new` starts + // this false, so rebuilding a record here without restoring it would not + // just skip the check now — it would write the weakened record back and + // disarm every later `rocm services restart` as well. + // + // Propagated, not defaulted. This read arms the guard below, so it is not + // best-effort the way an identical-looking call feeding a printed warning + // would be. `load_managed_services` already *skips* unparseable records, so + // an `Err` here is a real I/O failure — and a missing directory is `Ok` + // anyway. Swallowing it would say "no service ever required a key", the + // key-file fallback is false precisely when a service has been stopped, and + // the weakened record would then be written back at the bottom of this + // function. That is the outcome the comment above says must not happen. + let previously_required = crate::persistence::load_managed_services(paths) + .context( + "could not read the service registry to check whether this service requires an \ + endpoint API key; refusing to recover it rather than assume it does not", + )? + .iter() + .any(|existing| existing.service_id == record.service_id && existing.requires_api_key); + // Only what the registry recorded. The `|| key-file-is-present` clause that + // used to be here re-derived the flag the same way `spawn_managed_engine_child` + // did, and was wrong for the same reason: a public bind always has a key file + // whether or not auth was ever demanded, so recovery re-armed this on services + // that never asked for it and refused them with the wrong remediation. + record.requires_api_key = previously_required; + // Refuse a keyless public respawn before the manifest write, so a refused + // attempt leaves the recorded restart_count and timestamps intact instead of + // clobbering them with a record no live process will ever back. The spawn + // site below re-checks against the key actually threaded onto the command. + crate::common::ensure_public_service_has_endpoint_key( + &record.host, + rocm_engine_protocol::endpoint_key_file_if_present(paths, &record.service_id) + .and_then(|path| rocm_engine_protocol::endpoint_api_key_file_if_valid(&path)) + .is_some(), + record.requires_api_key, + )?; + record.write()?; + + let log_file = fs::File::create(&record.log_path) + .with_context(|| format!("failed to create {}", record.log_path.display()))?; + let log_file_err = log_file + .try_clone() + .context("failed to clone service log file handle")?; + + let rocm_binary = + std::env::current_exe().context("failed to resolve current rocm executable path")?; + let mut command = ProcessCommand::new(rocm_binary); + command + .args(engine_serve_http_args( + &engine, + &record.service_id, + &canonical_model_id, + &record.host, + record.port, + &device_policy, + &record.gpu_indices, + runtime_id.as_deref(), + env_id.as_deref(), + engine_recipe_json.as_deref(), + &record.engine_state_path, + )) + .stdin(Stdio::null()) + .stdout(Stdio::from(log_file)) + .stderr(Stdio::from(log_file_err)); + // Re-thread the endpoint key file (public bind only) onto the engine child, + // same as the initial `rocm serve` spawn. This path also runs on daemon + // recovery (`restart_managed_service` re-execs `rocmd supervise`), so + // without this a previously-authenticated public service would come back + // up anonymous after a crash/recover cycle. + // If the key is gone the child would listen on the recorded public host with + // no auth, so fail closed instead — an unreachable service is recoverable, + // an anonymous public one is not. + let endpoint_key_applied = + crate::common::apply_endpoint_key_env(&mut command, paths, &record.service_id); + crate::common::ensure_public_service_has_endpoint_key( + &record.host, + endpoint_key_applied, + record.requires_api_key, + )?; + let mut child = command + .spawn() + .with_context(|| format!("failed to spawn engine supervisor child for {engine}"))?; + + record.engine_pid = Some(child.id()); + record.status = "running".to_owned(); + record.write()?; + + // Clone the fields the poller reads so the `on_phase` closure can borrow + // `record` mutably to persist each startup-phase transition to disk. + let ready_engine = record.engine.clone(); + let ready_service_id = record.service_id.clone(); + let ready_log_path = record.log_path.clone(); + let became_ready = wait_for_service_ready( + paths, + &ready_engine, + &ready_service_id, + &ready_log_path, + Duration::from_mins(3), + |phase| { + record.startup_phase = Some(phase.to_owned()); + let _ = record.write(); + }, + ); + if became_ready { + record.status = "ready".to_owned(); + // The phase only describes the coming-up window; clear it once ready. + record.startup_phase = None; + record.write()?; + } + + let exit_status = child.wait().context("failed waiting for engine child")?; + record.status = if exit_status.success() { + "stopped".to_owned() + } else { + "failed".to_owned() + }; + record.write()?; + + if exit_status.success() { + Ok(()) + } else { + std::process::exit(exit_status.code().unwrap_or(1)); + } +} + +#[allow(clippy::too_many_arguments)] +fn engine_serve_http_args( + engine: &str, + service_id: &str, + canonical_model_id: &str, + host: &str, + port: u16, + device_policy: &str, + gpu_indices: &[u32], + runtime_id: Option<&str>, + env_id: Option<&str>, + engine_recipe_json: Option<&str>, + state_path: &Path, +) -> Vec { + let mut args = vec![ + "__engine-serve-http".to_owned(), + engine.to_owned(), + service_id.to_owned(), + canonical_model_id.to_owned(), + "--host".to_owned(), + host.to_owned(), + "--port".to_owned(), + port.to_string(), + "--device-policy".to_owned(), + device_policy.to_owned(), + ]; + if let Some(csv) = rocm_engine_protocol::gpu_indices_to_csv(gpu_indices) { + args.extend(["--gpu".to_owned(), csv]); + } + args.extend(crate::common::optional_arg("--runtime-id", runtime_id)); + args.extend(crate::common::optional_arg("--env-id", env_id)); + args.extend(crate::common::optional_arg( + "--engine-recipe-json", + engine_recipe_json, + )); + args.extend(["--state-path".to_owned(), state_path.display().to_string()]); + args +} + +fn build_runtime_state( + config: &RocmCliConfig, + automations_enabled: bool, +) -> AutomationRuntimeState { + let now = unix_time_millis(); + let active_watchers = builtin_watchers() + .iter() + .map(|watcher| WatcherRuntimeSnapshot { + id: watcher.id.to_owned(), + enabled: config.watcher_enabled(watcher), + mode: config.effective_watcher_mode(watcher), + summary: watcher.summary.to_owned(), + last_event: None, + last_event_unix_ms: None, + }) + .collect(); + AutomationRuntimeState { + running: automations_enabled, + automations_enabled, + daemon_pid: std::process::id(), + started_at_unix_ms: now, + last_tick_unix_ms: now, + local_webhook_endpoint: None, + active_watchers, + } +} + +#[cfg(unix)] +async fn shutdown_signal() { + let mut term = tokio::signal::unix::signal(tokio::signal::unix::SignalKind::terminate()) + .expect("failed to register SIGTERM handler"); + tokio::select! { + _ = tokio::signal::ctrl_c() => {} + _ = term.recv() => {} + } +} + +#[cfg(not(unix))] +async fn shutdown_signal() { + let _ = tokio::signal::ctrl_c().await; +} + +/// Keep only the final visible segment of a `\r`-redrawn progress line. +/// +/// Progress tools (pip, tqdm, Hugging Face) redraw a line in place with a bare +/// carriage return and no newline, so the segment after the last `\r` is its +/// final visible state. Lines without `\r` pass through unchanged. (Same +/// collapse rule the dashboard job console applies to streamed job output.) +fn last_cr_segment(line: &str) -> &str { + line.rsplit('\r').next().unwrap_or(line) +} + +/// Classify a single serve-log line into a coarse startup phase token +/// (`downloading`/`loading`/`warmup`), or `None` when the line carries no phase +/// signal. Case-insensitive substring match over the common vLLM / llama.cpp / +/// Hugging Face startup vocabulary. Checked warmup → loading → downloading so +/// the latest lifecycle stage a line mentions wins. +fn classify_startup_phase(line: &str) -> Option<&'static str> { + let lower = line.to_ascii_lowercase(); + if lower.contains("capturing cuda graph") + || lower.contains("capturing the model") + || lower.contains("warming up") + || lower.contains("warmup") + { + Some("warmup") + } else if lower.contains("loading weights") + || lower.contains("loading model") + || lower.contains("load_tensors") + || lower.contains("llama_model_loader") + || lower.contains("model loading took") + { + Some("loading") + } else if lower.contains("downloading") || lower.contains("fetching") { + Some("downloading") + } else { + None + } +} + +/// Read log bytes appended since `*pos`, advance `*pos`, and return the most +/// recent recognizable startup phase in that new output (later lines win, so a +/// download → load → warmup progression advances naturally). +/// +/// Best-effort: any I/O error (file not created yet, transient read) yields +/// `None`. A shrunk file (rotation/truncation) resets the cursor to the top. +fn read_new_log_phase(log_path: &Path, pos: &mut u64) -> Option<&'static str> { + use std::io::{Read, Seek, SeekFrom}; + let mut file = fs::File::open(log_path).ok()?; + let len = file.metadata().ok()?.len(); + if len < *pos { + *pos = 0; + } + if len == *pos { + return None; + } + file.seek(SeekFrom::Start(*pos)).ok()?; + let mut bytes = Vec::new(); + let read = file.read_to_end(&mut bytes).ok()?; + *pos += read as u64; + let text = String::from_utf8_lossy(&bytes); + let mut phase = None; + for line in text.lines() { + if let Some(found) = classify_startup_phase(last_cr_segment(line)) { + phase = Some(found); + } + } + phase +} + +fn engine_healthcheck_ready(paths: &AppPaths, engine: &str, service_id: &str) -> Result { + Ok(crate::common::healthcheck_response_ready( + &crate::common::engine_healthcheck_response(paths, engine, service_id)?, + )) +} + +/// Poll a freshly-spawned service until its healthcheck reports ready (or the +/// timeout elapses), tailing its log file meanwhile and reporting each coarse +/// startup phase transition via `on_phase`. +fn wait_for_service_ready( + paths: &AppPaths, + engine: &str, + service_id: &str, + log_path: &Path, + timeout: Duration, + mut on_phase: impl FnMut(&str), +) -> bool { + let start = std::time::Instant::now(); + let mut log_pos: u64 = 0; + let mut last_phase: Option<&'static str> = None; + while start.elapsed() < timeout { + if let Some(phase) = read_new_log_phase(log_path, &mut log_pos) + && last_phase != Some(phase) + { + last_phase = Some(phase); + on_phase(phase); + } + if engine_healthcheck_ready(paths, engine, service_id).unwrap_or(false) { + return true; + } + thread::sleep(Duration::from_millis(200)); + } + false +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::test_support::{temp_app_paths, unique_test_root}; + + /// Drive `supervise_service` far enough to reach the key guard, and return + /// what it did. + /// + /// The guard sits before the manifest write and well before any spawn, so a + /// refusal returns without starting a process — which is what makes the real + /// call site testable at all. The arguments below are the shape a recovery + /// re-exec passes: a loopback bind, no GPU, no recipe. + /// + /// This exists because testing `ensure_public_service_has_endpoint_key` + /// directly with literal arguments cannot catch the defect that actually + /// happened twice in this crate's history — the guard being *wired up* with + /// the wrong value at its call site. + /// + /// Bounded, and the bound is the assertion. A guard that fails to refuse + /// does not return an error — it falls through to the engine spawn and + /// supervises a child that never exits, so an unbounded call would hang the + /// suite instead of failing it. Both callers below are regression tests for + /// a fail-*open*, which is exactly the shape that turns into a hang. + fn supervise_at_the_guard(paths: &AppPaths, service_id: &str) -> Result<()> { + let paths = paths.clone(); + let service_id = service_id.to_owned(); + let (sender, receiver) = std::sync::mpsc::channel(); + std::thread::spawn(move || { + let outcome = supervise_service( + &paths, + service_id, + "llamacpp".to_owned(), + "a-model".to_owned(), + "a-model".to_owned(), + None, + None, + "127.0.0.1".to_owned(), + 11434, + "gpu_required".to_owned(), + None, + None, + ); + let _ = sender.send(outcome.map_err(|error| format!("{error:#}"))); + }); + receiver + .recv_timeout(std::time::Duration::from_secs(30)) + .unwrap_or_else(|_| { + panic!( + "supervise_service did not return within 30s: the key guard let the call \ + through and it reached the engine spawn, which is the fail-open this test \ + exists to catch" + ) + }) + .map_err(anyhow::Error::msg) + } + + /// Write a service record into the registry the way a live service would + /// have left it behind. + fn seed_registry(paths: &AppPaths, service_id: &str, requires_api_key: bool) { + fs::create_dir_all(paths.services_dir()).unwrap(); + let mut record = ManagedServiceRecord::new( + paths, + service_id.to_owned(), + "llamacpp".to_owned(), + "a-model".to_owned(), + "a-model".to_owned(), + "127.0.0.1".to_owned(), + 11434, + "managed", + std::process::id(), + None, + None, + Some("gpu_required".to_owned()), + ); + record.requires_api_key = requires_api_key; + record.write().unwrap(); + } + + /// Drive `supervise_service` and report whether the key guard let it past. + /// + /// Decided on what the call *returns*, not on any file. The obvious + /// observable — the manifest appearing — is useless here, because + /// `seed_registry` has already written one, so polling for it passes + /// whatever the guard does. That mistake was made first and caught by + /// mutating the code the test claims to protect. + /// + /// A call the guard admits does not return: it carries on to the engine + /// spawn. So the guard's refusal is the only thing that comes back quickly, + /// and it is identified by its message rather than by the mere fact of an + /// error — a later, unrelated failure must not read as a refusal. + fn guard_admits(paths: &AppPaths, service_id: &str) -> bool { + let owned_paths = paths.clone(); + let owned_id = service_id.to_owned(); + let (sender, receiver) = std::sync::mpsc::channel(); + std::thread::spawn(move || { + let outcome = supervise_service( + &owned_paths, + owned_id, + "llamacpp".to_owned(), + "a-model".to_owned(), + "a-model".to_owned(), + None, + None, + "127.0.0.1".to_owned(), + 11434, + "gpu_required".to_owned(), + None, + None, + ); + let _ = sender.send(outcome.map_err(|error| format!("{error:#}"))); + }); + + match receiver.recv_timeout(std::time::Duration::from_secs(10)) { + // The key guard refused, by its own words. + Ok(Err(rendered)) if rendered.contains("without authentication") => false, + // Anything else means it got past the guard: it either finished, or + // failed later for a reason that is not this guard, or is still + // running because it reached the spawn. + _ => true, + } + } + + #[test] + fn a_service_that_never_required_a_key_is_not_refused_for_lacking_one() { + // The other direction of the guard, and the one no test covered. + // Hardcoding `record.requires_api_key = true` at the restore site passes + // every other test in this crate, because they all seed a service that + // *does* require a key. This is the case that catches it. + // + // Two records are seeded, not one: with a single record the + // `existing.service_id == record.service_id` half of the lookup does + // nothing, so dropping that comparison would go unnoticed and one + // service's requirement would leak onto another's. + let (root, paths) = temp_app_paths("supervise-no-key-needed"); + seed_registry(&paths, "svc-needs-key", true); + seed_registry(&paths, "svc-plain", false); + + assert!( + guard_admits(&paths, "svc-plain"), + "a loopback service that never asked for a key must not be refused for lacking one" + ); + + let _ = fs::remove_dir_all(&root); + } + + #[test] + fn supervising_a_service_that_required_a_key_refuses_when_the_key_is_gone() { + // The real call site, not the guard in isolation. `supervise_service` + // rebuilds the record with `ManagedServiceRecord::new`, which starts + // `requires_api_key` false, and restores it from the registry. Passing + // the wrong value here — a literal, or the freshly-built field before it + // is restored — is exactly the miswiring that shipped twice in this + // crate and that a literal-argument unit test cannot see. + let (root, paths) = temp_app_paths("supervise-requires-key"); + seed_registry(&paths, "svc-needs-key", true); + + let error = supervise_at_the_guard(&paths, "svc-needs-key") + .expect_err("a service that required a key must not be recovered without one"); + assert!( + format!("{error:#}").contains("without authentication"), + "{error:#}" + ); + + // And the refusal must not have weakened what is on disk. The guard runs + // before `record.write()` precisely so a refused attempt leaves the + // recorded requirement armed for the next attempt. + let stored = crate::persistence::load_managed_services(&paths).unwrap(); + let stored = stored + .iter() + .find(|candidate| candidate.service_id == "svc-needs-key") + .expect("the seeded record must survive a refused recovery"); + assert!( + stored.requires_api_key, + "a refused recovery must not disarm the requirement" + ); + + let _ = fs::remove_dir_all(&root); + } + + #[test] + fn a_registry_that_cannot_be_read_refuses_recovery_rather_than_assuming_no_key() { + // `load_managed_services` already skips records it cannot parse, so an + // `Err` from it is a real I/O failure — and a missing directory is `Ok`. + // Defaulting it away therefore says "no service ever required a key", + // which is fail-open on an auth gate and, worse, gets written back. + // + // The failure is provoked portably: a directory named like a record + // makes the `fs::read` inside the loop fail rather than the read_dir. + let (root, paths) = temp_app_paths("supervise-unreadable-registry"); + fs::create_dir_all(paths.services_dir().join("not-a-record.json")).unwrap(); + + let error = supervise_at_the_guard(&paths, "svc-unknown") + .expect_err("an unreadable registry must refuse, not assume no key was required"); + let rendered = format!("{error:#}"); + assert!( + rendered.contains("could not read the service registry"), + "{rendered}" + ); + + // Nothing was written: a registry we could not read is not a registry we + // may add a weakened record to. + assert!( + !paths.service_manifest_path("svc-unknown").exists(), + "a refused recovery must not persist a record" + ); + + let _ = fs::remove_dir_all(&root); + } + + #[test] + fn last_cr_segment_keeps_final_progress_redraw() { + // A tqdm/HF-style in-place redraw collapses to its last segment. + assert_eq!( + last_cr_segment("Downloading: 10%\rDownloading: 55%\rDownloading: 100%"), + "Downloading: 100%" + ); + // A plain line is unchanged. + assert_eq!( + last_cr_segment("Loading model weights"), + "Loading model weights" + ); + } + + #[test] + fn classify_startup_phase_maps_engine_vocabulary() { + assert_eq!( + classify_startup_phase("Downloading shards: 100%"), + Some("downloading") + ); + assert_eq!( + classify_startup_phase("Fetching 12 files"), + Some("downloading") + ); + assert_eq!( + classify_startup_phase("INFO: Loading model weights took 4.2s"), + Some("loading") + ); + assert_eq!( + classify_startup_phase("llama_model_loader: loaded meta data"), + Some("loading") + ); + assert_eq!( + classify_startup_phase("Capturing CUDA graph shapes"), + Some("warmup") + ); + assert_eq!( + classify_startup_phase("Warming up the engine"), + Some("warmup") + ); + // Ordinary chatter carries no phase signal. + assert_eq!( + classify_startup_phase("Uvicorn running on http://..."), + None + ); + } + + #[test] + fn classify_startup_phase_emits_only_dashboard_known_tokens() { + // These tokens are the wire contract with the dashboard's + // `StartupPhase::from_token` (rocm-dash-core); emitting anything else + // would be silently dropped there. rocmd can't link that crate, so the + // contract is pinned here by literal. + for line in [ + "Downloading shards", + "Loading model weights", + "Capturing CUDA graph", + ] { + let token = classify_startup_phase(line).expect("line is a phase signal"); + assert!( + matches!(token, "downloading" | "loading" | "warmup"), + "token {token:?} must be one the dashboard understands" + ); + } + } + + #[test] + fn read_new_log_phase_advances_and_tracks_latest() { + use std::io::Write as _; + // Workspace-local test root (rooted at CARGO_MANIFEST_DIR, not the + // ambient temp dir) — same helper the other rocmd tests use. + let dir = unique_test_root(&format!("rocmd-phase-{}", std::process::id())); + let log = dir.join("svc.log"); + std::fs::write(&log, "boot\nDownloading shards: 100%\n").unwrap(); + + let mut pos = 0_u64; + assert_eq!(read_new_log_phase(&log, &mut pos), Some("downloading")); + // No new bytes → no phase, cursor unchanged. + let after_first = pos; + assert_eq!(read_new_log_phase(&log, &mut pos), None); + assert_eq!(pos, after_first); + + // Appending a later stage advances the phase. + let mut f = std::fs::OpenOptions::new().append(true).open(&log).unwrap(); + writeln!(f, "Loading model weights took 3s").unwrap(); + assert_eq!(read_new_log_phase(&log, &mut pos), Some("loading")); + + let _ = std::fs::remove_dir_all(&dir); + } + + #[test] + fn engine_serve_http_args_forward_engine_recipe_json() { + let engine_recipe_json = r#"{"contract_version":"0.1.0","engine":"vllm","required_flags":["--enable-auto-tool-choice"]}"#; + let args = engine_serve_http_args( + "vllm", + "svc-1", + "Qwen/Qwen3.5-4B", + "127.0.0.1", + 11435, + "gpu_required", + &[], + Some("therock-release:gfx120X-all"), + Some("env-1"), + Some(engine_recipe_json), + Path::new("state.json"), + ); + + assert!( + args.windows(2) + .any(|pair| pair[0] == "--engine-recipe-json" && pair[1] == engine_recipe_json) + ); + assert!( + args.windows(2).any(|pair| { + pair[0] == "--runtime-id" && pair[1] == "therock-release:gfx120X-all" + }) + ); + assert!( + args.windows(2) + .any(|pair| pair[0] == "--state-path" && pair[1] == "state.json") + ); + } + + #[test] + fn engine_serve_http_args_emit_gpu_indices_when_pinned() { + let args = engine_serve_http_args( + "vllm", + "svc-1", + "Qwen/Qwen3.5-4B", + "127.0.0.1", + 11435, + "gpu_required", + &[1], + None, + None, + None, + Path::new("state.json"), + ); + + assert!( + args.windows(2) + .any(|pair| pair[0] == "--gpu" && pair[1] == "1") + ); + + let auto = engine_serve_http_args( + "vllm", + "svc-1", + "Qwen/Qwen3.5-4B", + "127.0.0.1", + 11435, + "gpu_required", + &[], + None, + None, + None, + Path::new("state.json"), + ); + assert!(!auto.iter().any(|arg| arg == "--gpu")); + } + + #[tokio::test] + async fn local_webhook_requires_enabled_automation_loop() { + let (_root, paths) = temp_app_paths("local-webhook-requires-loop"); + let error = run_daemon(&paths, false, Some(0)).await.unwrap_err(); + + assert!(error.to_string().contains("requires --automations-enabled")); + } + + #[test] + fn stop_managed_service_removes_endpoint_key_file() -> Result<()> { + let (root, paths) = temp_app_paths("stop-removes-endpoint-key"); + paths.ensure()?; + let current_pid = std::process::id(); + let service_id = "svc-endpoint-key-stop"; + let mut record = ManagedServiceRecord::new( + &paths, + service_id, + "vllm", + "qwen", + "Qwen/Qwen3.5", + "0.0.0.0", + 11435, + "managed", + current_pid, + None, + None, + None, + ); + record.engine_pid = Some(current_pid); + record.status = "ready".to_owned(); + record.write()?; + + let key_path = rocm_engine_protocol::endpoint_key_file_path(&paths, service_id); + fs::create_dir_all(paths.services_dir())?; + fs::write(&key_path, "secret-key")?; + assert!(key_path.exists()); + + let result = stop_managed_service(&paths, service_id); + // Observe the real filesystem state before the blanket temp-dir cleanup, + // otherwise remove_dir_all would delete the key file and mask a missing + // production cleanup (the regression this test guards). + let key_removed = !key_path.exists(); + fs::remove_dir_all(root).ok(); + + let value = result?; + assert_eq!( + value + .get("service") + .and_then(|service| service.get("status")) + .and_then(Value::as_str), + Some("stopped") + ); + assert!(key_removed, "endpoint key file must be removed after stop"); + Ok(()) + } + + #[test] + fn stop_managed_service_without_endpoint_key_file_succeeds() -> Result<()> { + let (root, paths) = temp_app_paths("stop-no-endpoint-key"); + paths.ensure()?; + let current_pid = std::process::id(); + let service_id = "svc-no-endpoint-key-stop"; + let mut record = ManagedServiceRecord::new( + &paths, + service_id, + "vllm", + "qwen", + "Qwen/Qwen3.5", + "127.0.0.1", + 11435, + "managed", + current_pid, + None, + None, + None, + ); + record.engine_pid = Some(current_pid); + record.status = "ready".to_owned(); + record.write()?; + + // Loopback service: no endpoint key file was ever written for it. + let key_path = rocm_engine_protocol::endpoint_key_file_path(&paths, service_id); + assert!(!key_path.exists()); + + let result = stop_managed_service(&paths, service_id); + let reloaded = crate::load_service_record(&paths, service_id); + fs::remove_dir_all(root).ok(); + + let value = result?; + assert_eq!( + value + .get("service") + .and_then(|service| service.get("status")) + .and_then(Value::as_str), + Some("stopped") + ); + assert_eq!(reloaded?.status, "stopped"); + assert!(!key_path.exists()); + Ok(()) + } + + #[test] + fn stop_server_process_tree_discovers_descendants_before_parents() { + let output = "\ +10 1 +11 10 +12 11 +13 10 +20 1 +21 20 +"; + + assert_eq!( + descendant_pids_from_ps_output(output, &[10]), + vec![12, 11, 13] + ); + assert_eq!( + descendant_pids_from_ps_output(output, &[10, 20]), + vec![12, 11, 13, 21] + ); + } +} diff --git a/docs/architecture.md b/docs/architecture.md index 7117b9085..59852f120 100644 --- a/docs/architecture.md +++ b/docs/architecture.md @@ -31,7 +31,7 @@ Subsystem modules already following full domain extraction (each owns its own ty ### `apps/rocmd` — background daemon -`lib.rs` modularization is in progress (ROCMAI-83, Phase 5 of EAI-7768's sequencing). Extracted so far: `persistence.rs` (`record_event`/`load_managed_services`, the automation-event/audit-log and managed-service-registry I/O shared across the daemon's sandbox, MCP, service-lifecycle, and watcher code), `common.rs` (helpers shared across ≥2 of those remaining clusters: GPU/amd-smi snapshotting, the bridge-snapshot diagnostic, `CommandCapture`/command-timeout plumbing including the shared `rocm`-subprocess capture helpers the sandbox and MCP clusters both call, and small arg/healthcheck/endpoint-key utilities), `webhook.rs` (the local webhook source: its axum routes, request validation, and watcher-kind allow-list), `cli.rs` (the `Cli`/`Command` clap definitions, `SandboxToolArg`/`SandboxToolPolicy`, and the top-level dispatch in `run_cli`/`run_bin_cli`/`run_from_args` — the crate's only two externally-consumed entry points are re-exported from here via `lib.rs`'s `pub use`), `sandbox.rs` (bubblewrap/native sandbox execution, atomic-write helpers, artifact prefetch policy gating, and the sandbox-tool result shaping for `check_updates`/`driver_plan`), and `mcp.rs` (the MCP stdio server, tool schema table, tool dispatch, and the `rocm`-subprocess capture/argv-building helpers behind the MCP tools). A helper earns a place in `common.rs` only once a second still-inline cluster calls it directly; a helper with exactly one caller stays in `lib.rs` next to that caller until its own cluster's extraction PR, even if it is conceptually similar to something that did move. Still pending: the service-lifecycle and watcher clusters themselves — each landing as its own PR. +`lib.rs` modularization is in progress (ROCMAI-83, Phase 5 of EAI-7768's sequencing). Extracted so far: `persistence.rs` (`record_event`/`load_managed_services`, the automation-event/audit-log and managed-service-registry I/O shared across the daemon's sandbox, MCP, service-lifecycle, and watcher code), `common.rs` (helpers shared across ≥2 of those remaining clusters: GPU/amd-smi snapshotting, the bridge-snapshot diagnostic, `CommandCapture`/command-timeout plumbing including the shared `rocm`-subprocess capture helpers the sandbox and MCP clusters both call, and small arg/healthcheck/endpoint-key utilities), `webhook.rs` (the local webhook source: its axum routes, request validation, and watcher-kind allow-list), `cli.rs` (the `Cli`/`Command` clap definitions, `SandboxToolArg`/`SandboxToolPolicy`, and the top-level dispatch in `run_cli`/`run_bin_cli`/`run_from_args` — the crate's only two externally-consumed entry points are re-exported from here via `lib.rs`'s `pub use`), `sandbox.rs` (bubblewrap/native sandbox execution, atomic-write helpers, artifact prefetch policy gating, and the sandbox-tool result shaping for `check_updates`/`driver_plan`), `mcp.rs` (the MCP stdio server, tool schema table, tool dispatch, and the `rocm`-subprocess capture/argv-building helpers behind the MCP tools), and `service.rs` (managed-service PID lifecycle/stop, `run_daemon`'s foreground loop, `supervise_service`'s spawn-and-recover path, and serve-log startup-phase polling). A helper earns a place in `common.rs` only once a second still-inline cluster calls it directly; a helper with exactly one caller stays in `lib.rs` next to that caller until its own cluster's extraction PR, even if it is conceptually similar to something that did move. Still pending: the watcher cluster itself — landing as its own PR. ### `crates/rocm-core` — core library From 6d292ddb54b49c59c42ce1e5f2e258bead417099 Mon Sep 17 00:00:00 2001 From: Jussi Elo Date: Thu, 1 Oct 2026 13:45:41 +0000 Subject: [PATCH 3/3] ROCMAI-83: extract watchers.rs from apps/rocmd/src/lib.rs MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Eighth and final PR of Phase 5 (rocmd modularization, ROCMAI-27): pull event collection/dispatch for all built-in watchers (TheRock update, GPU metrics/thermal-pressure, cache-warm, driver-upgrade, server-recover), managed-service recovery classification, and automation-proposal queuing into their own module. Non-contiguous in the source: most of the cluster is one block, but queue_proposal/queue_proposal_with_arguments/proposal_tool_for_action/ proposal_arguments_for_action sit sandwiched between record_event and load_managed_services (persistence.rs territory, stays in lib.rs on this branch), and detached_rocmd_command sits at the tail of the file right before the test module. All three sub-blocks moved in this PR. Watcher-only consts (SERVER_RECOVER_BACKOFF_MS, SERVER_TRANSIENT_STALE_MS, ENDPOINT_HEALTH_TIMEOUT, THEROCK_UPDATE_INTERVAL_MS, GPU_METRICS_INTERVAL_MS, and the GPU thermal/VRAM pressure thresholds) moved with it, since nothing else in lib.rs used them. No behavior change. Reaches back into items owned by sibling modules (record_event/load_managed_services via crate::persistence, run_sandbox_tool/SandboxToolArg/SandboxToolPolicy via crate::sandbox/crate::cli, update_check_message/ record_notification_audit via crate::sandbox, and the engine-healthcheck/endpoint-key/wait_for_port/optional_arg/ gather_gpu_snapshot_for_config family via crate::common) and one webhook-domain type (LocalWebhookEventRequest/ local_webhook_event_from_request via crate::webhook). 34 tests that exercise this module's own logic (event-collector/gpu/ cache-warm/driver-upgrade/server-recover dispatch, therock-update handling, watcher-policy mode mapping, recovery-reason classification) moved into watchers.rs's own #[cfg(test)] mod in this same PR. Tests that are fundamentally about Cli parsing, run_daemon, sandbox-tool dispatch shape, or the persistence-layer record_event itself stayed in lib.rs for their own PRs. This completes Phase 5 of the rocm-cli modularization effort (ROCMAI-27): lib.rs is now top-level glue -- module declarations and the two externally-consumed entry points re-exported from cli.rs. Stacked on rocmai-83-service (ROCMAI-83 Phase 5 batch, AGENTS.md §11): rebased watchers.rs's crate:: reaches onto all seven sibling modules now landed by the predecessor PRs (persistence.rs, common.rs, webhook.rs, cli.rs, sandbox.rs, mcp.rs, service.rs) -- this is the last PR in the stack, so no reach-back is left unresolved. Repointed service.rs's stale crate::-root calls (evaluate_watchers, evaluate_watchers_for_events, reconcile_watcher_snapshots, load_service_record) and webhook.rs's stale crate::payload_string call to watchers::, now that this module owns them; fixed sandbox.rs's `use crate::{ARTIFACT_PREFETCH_TIMEOUT, restart_managed_service}` to import restart_managed_service from crate::watchers instead. Deduped the test-only temp_app_paths/unique_test_root/workspace_test_artifact_dir trio in watchers.rs's test module in favor of the shared crate::test_support module. Updated docs/architecture.md: all eight `apps/rocmd` modules are now listed as extracted, with no "still pending" clause remaining. Signed-off-by: Jussi Elo --- apps/rocmd/src/lib.rs | 3123 +---------------------------------- apps/rocmd/src/sandbox.rs | 3 +- apps/rocmd/src/service.rs | 12 +- apps/rocmd/src/watchers.rs | 3144 ++++++++++++++++++++++++++++++++++++ apps/rocmd/src/webhook.rs | 6 +- docs/architecture.md | 2 +- 6 files changed, 3170 insertions(+), 3120 deletions(-) create mode 100644 apps/rocmd/src/watchers.rs diff --git a/apps/rocmd/src/lib.rs b/apps/rocmd/src/lib.rs index add9b110d..a7a147a29 100644 --- a/apps/rocmd/src/lib.rs +++ b/apps/rocmd/src/lib.rs @@ -12,3094 +12,24 @@ mod sandbox; mod service; #[cfg(test)] mod test_support; +mod watchers; mod webhook; pub use cli::{run_bin_cli, run_from_args}; -use anyhow::{Context, Result, bail}; -#[cfg(test)] -use rocm_core::AuditEventRecord; -#[cfg(test)] -use rocm_core::AutomationEventRecord; -use rocm_core::{ - AppPaths, AutomationProposalRecord, AutomationRuntimeState, AutomationTriggerEvent, - CodexBridgeGpuSnapshot, ManagedServiceRecord, RocmCliConfig, WatcherMode, - WatcherRuntimeSnapshot, append_automation_proposal, builtin_watchers, - resolve_model_recipe_artifact, unix_time_millis, -}; -use serde_json::Value; -use serde_json::json; -use std::fs; -use std::process::{Command as ProcessCommand, Stdio}; -use std::thread; -use std::time::Duration; - -const WATCHER_TICK_INTERVAL: Duration = Duration::from_secs(5); -const SERVER_RECOVER_BACKOFF_MS: u128 = 30_000; -const SERVER_TRANSIENT_STALE_MS: u128 = 5 * 60 * 1_000; -const ENDPOINT_HEALTH_TIMEOUT: Duration = Duration::from_millis(250); -const THEROCK_UPDATE_INTERVAL_MS: u128 = 6 * 60 * 60 * 1000; -const GPU_METRICS_INTERVAL_MS: u128 = 60 * 1000; -const GPU_THERMAL_HOTSPOT_PRESSURE_C: f64 = 95.0; -const GPU_THERMAL_MEMORY_PRESSURE_C: f64 = 95.0; -const GPU_MEMORY_VRAM_PRESSURE_PERCENT: f64 = 95.0; -const ARTIFACT_PREFETCH_TIMEOUT: Duration = Duration::from_mins(10); - -fn reconcile_watcher_snapshots(config: &RocmCliConfig, state: &mut AutomationRuntimeState) { - for watcher in builtin_watchers() { - match state.watcher_mut(watcher.id) { - Some(snapshot) => { - snapshot.enabled = config.watcher_enabled(watcher); - snapshot.mode = config.effective_watcher_mode(watcher); - snapshot.summary = watcher.summary.to_owned(); - } - None => state.active_watchers.push(WatcherRuntimeSnapshot { - id: watcher.id.to_owned(), - enabled: config.watcher_enabled(watcher), - mode: config.effective_watcher_mode(watcher), - summary: watcher.summary.to_owned(), - last_event: None, - last_event_unix_ms: None, - }), - } - } -} - -fn evaluate_watchers( - paths: &AppPaths, - config: &RocmCliConfig, - state: &mut AutomationRuntimeState, -) -> Result<()> { - let events = collect_automation_events(paths, config, state)?; - evaluate_watchers_for_events(paths, config, state, &events) -} - -fn collect_automation_events( - paths: &AppPaths, - config: &RocmCliConfig, - state: &AutomationRuntimeState, -) -> Result> { - collect_automation_events_with_gpu_snapshot(paths, state, || { - common::gather_gpu_snapshot_for_config(config) - }) -} - -fn collect_automation_events_with_gpu_snapshot( - paths: &AppPaths, - state: &AutomationRuntimeState, - gpu_snapshot: F, -) -> Result> -where - F: FnMut() -> CodexBridgeGpuSnapshot, -{ - let now = unix_time_millis(); - let mut events = Vec::new(); - - if therock_update_due(state, now) { - events.push(AutomationTriggerEvent { - at_unix_ms: now, - kind: "schedule.tick".to_owned(), - source: "scheduler".to_owned(), - watcher_hint: Some("therock-update".to_owned()), - service_id: None, - reason: Some("therock_update_interval_due".to_owned()), - payload: json!({ - "interval_ms": THEROCK_UPDATE_INTERVAL_MS, - }), - }); - } - - if server_recover_due(state, now) - && let Some((record, recovery_reason)) = find_recoverable_service(paths)? - { - let kind = service_recovery_event_kind(&recovery_reason); - events.push(AutomationTriggerEvent { - at_unix_ms: now, - kind: kind.to_owned(), - source: "managed_service".to_owned(), - watcher_hint: Some("server-recover".to_owned()), - service_id: Some(record.service_id.clone()), - reason: Some(recovery_reason.clone()), - payload: json!({ - "engine": record.engine, - "status": record.status, - "endpoint": record.endpoint_url, - "recovery_reason": recovery_reason, - }), - }); - } - - let gpu_metrics_due_now = gpu_metrics_due(state, now); - let gpu_thermal_protect_due_now = gpu_thermal_protect_due(state, now); - let snapshot = (gpu_metrics_due_now || gpu_thermal_protect_due_now).then(gpu_snapshot); - - if gpu_metrics_due_now { - let snapshot = snapshot - .as_ref() - .expect("GPU snapshot should be collected for due metrics"); - let available = snapshot.amd_smi_available && snapshot.monitor_snapshot.is_some(); - events.push(AutomationTriggerEvent { - at_unix_ms: now, - kind: if available { - "gpu.metrics".to_owned() - } else { - "gpu.metrics_unavailable".to_owned() - }, - source: "gpu_telemetry".to_owned(), - watcher_hint: Some("gpu-metrics".to_owned()), - service_id: None, - reason: if available { - Some("amd_smi_snapshot_available".to_owned()) - } else { - snapshot - .note - .clone() - .or_else(|| Some("amd_smi_snapshot_unavailable".to_owned())) - }, - payload: json!({ - "amd_smi_available": snapshot.amd_smi_available, - "static_available": snapshot.static_snapshot.is_some(), - "monitor_available": snapshot.monitor_snapshot.is_some(), - "summary": gpu_snapshot_summary(snapshot), - "interval_ms": GPU_METRICS_INTERVAL_MS, - }), - }); - } - - if gpu_thermal_protect_due_now && let Some(snapshot) = snapshot.as_ref() { - events.extend(gpu_pressure_events(now, snapshot)); - } - - Ok(events) -} - -fn evaluate_watchers_for_events( - paths: &AppPaths, - config: &RocmCliConfig, - state: &mut AutomationRuntimeState, - events: &[AutomationTriggerEvent], -) -> Result<()> { - for watcher in builtin_watchers() { - if !config.watcher_enabled(watcher) { - continue; - } - let mode = config.effective_watcher_mode(watcher); - match watcher.id { - "therock-update" => { - for event in events_for_watcher(events, watcher.id, "schedule.tick") { - handle_therock_update_event(paths, mode, state, event)?; - } - } - "server-recover" => { - for event in events_for_watcher(events, watcher.id, "service.") { - handle_server_recover_event(paths, mode, state, event)?; - } - } - "gpu-metrics" => { - for event in events_for_watcher(events, watcher.id, "gpu.") { - handle_gpu_metrics_event(paths, mode, state, event)?; - } - } - "gpu-thermal-protect" => { - for event in - events_for_watcher_exact(events, watcher.id, "gpu.thermal_pressure").chain( - events_for_watcher_exact(events, watcher.id, "gpu.memory_pressure"), - ) - { - handle_gpu_thermal_protect_event(paths, mode, state, event)?; - } - } - "cache-warm" => { - for event in events_for_watcher_exact(events, watcher.id, "cache.warm") { - handle_cache_warm_event(paths, mode, state, event)?; - } - } - "driver-upgrade" => { - for event in events_for_watcher_exact(events, watcher.id, "update.available") { - handle_driver_upgrade_event(paths, mode, state, event)?; - } - } - _ => {} - } - } - Ok(()) -} - -fn events_for_watcher<'a>( - events: &'a [AutomationTriggerEvent], - watcher_id: &str, - kind_prefix: &str, -) -> impl Iterator { - events.iter().filter(move |event| { - event.watcher_hint.as_deref() == Some(watcher_id) && event.kind.starts_with(kind_prefix) - }) -} - -fn events_for_watcher_exact<'a>( - events: &'a [AutomationTriggerEvent], - watcher_id: &str, - kind: &'static str, -) -> impl Iterator { - events.iter().filter(move |event| { - event.watcher_hint.as_deref() == Some(watcher_id) && event.kind == kind - }) -} - -fn handle_therock_update_event( - paths: &AppPaths, - mode: WatcherMode, - state: &mut AutomationRuntimeState, - event: &AutomationTriggerEvent, -) -> Result<()> { - handle_therock_update_event_with_runner(paths, mode, state, event, |paths| { - sandbox::run_sandbox_tool( - paths, - cli::SandboxToolArg::CheckUpdates, - None, - None, - None, - cli::SandboxToolPolicy::default(), - ) - }) -} - -fn handle_therock_update_event_with_runner( - paths: &AppPaths, - mode: WatcherMode, - state: &mut AutomationRuntimeState, - _event: &AutomationTriggerEvent, - update_runner: F, -) -> Result<()> -where - F: FnOnce(&AppPaths) -> Result, -{ - let policy = watcher_policy_action("therock-update", mode); - let action = match policy { - WatcherPolicyAction::Observe => "observe_schedule", - WatcherPolicyAction::QueueProposal => "queue_update_proposal", - WatcherPolicyAction::RunContained => "run_update_check", - }; - let message = match policy { - WatcherPolicyAction::Observe => { - "scheduled TheRock update check reminder emitted; run `rocm update` to inspect the selected channel" - } - WatcherPolicyAction::QueueProposal => { - "scheduled TheRock update check reminder emitted; queueing read-only update-check proposal for review" - } - WatcherPolicyAction::RunContained => { - "scheduled TheRock update check is approved for contained read-only execution" - } - }; - match policy { - WatcherPolicyAction::Observe | WatcherPolicyAction::QueueProposal => { - persistence::record_event( - paths, - state, - "therock-update", - "info", - action, - message, - None, - )?; - if policy == WatcherPolicyAction::QueueProposal { - queue_proposal( - paths, - "therock-update", - action, - "Check TheRock updates", - "Run `rocm update` to inspect available CLI, runtime, engine, and recipe updates before applying changes.", - None, - )?; - } - } - WatcherPolicyAction::RunContained => match update_runner(paths) { - Ok(output) => match restricted_check_updates_result(&output) { - Ok(result) => { - persistence::record_event( - paths, - state, - "therock-update", - if result.exit_status == 0 { - "info" - } else { - "error" - }, - action, - &format!( - "{message}; restricted check_updates status={}; {}", - result.status, - common::update_check_message(result.status) - ), - None, - )?; - if result.update_available { - record_update_available_notification(paths, state, result.status)?; - } - } - Err(error) => { - persistence::record_event( - paths, - state, - "therock-update", - "error", - "update_check_failed", - &format!( - "scheduled TheRock update check failed during contained restricted execution: {error}; no updates were applied" - ), - None, - )?; - } - }, - Err(error) => { - persistence::record_event( - paths, - state, - "therock-update", - "error", - "update_check_failed", - &format!( - "scheduled TheRock update check failed during contained read-only execution: {error}; no updates were applied" - ), - None, - )?; - } - }, - } - Ok(()) -} - -struct RestrictedCheckUpdatesResult<'a> { - status: &'a str, - update_available: bool, - exit_status: i64, -} - -fn restricted_check_updates_result(value: &Value) -> Result> { - let tool = value - .get("tool") - .and_then(Value::as_str) - .context("restricted update check did not report a tool name")?; - if tool != cli::SandboxToolArg::CheckUpdates.as_cli_value() { - bail!("restricted update check returned `{tool}`, expected `check_updates`"); - } - let status = value - .get("status") - .and_then(Value::as_str) - .unwrap_or("checked"); - let update_available = value - .get("update_available") - .and_then(Value::as_bool) - .unwrap_or(matches!(status, "update_available" | "repair_available")); - let exit_status = value - .get("exit_status") - .and_then(Value::as_i64) - .unwrap_or_else(|| i64::from(status == "error")); - Ok(RestrictedCheckUpdatesResult { - status, - update_available, - exit_status, - }) -} - -fn record_update_available_notification( - paths: &AppPaths, - state: &mut AutomationRuntimeState, - status: &str, -) -> Result<()> { - let message = if status == "repair_available" { - "A ROCm runtime repair is available because its package composition changed. Preview it before applying. No updates were applied." - } else { - "A ROCm runtime update is available. Preview it before applying. No updates were applied." - }; - persistence::record_event( - paths, - state, - "therock-update", - "info", - "notify_if_newer", - message, - None, - )?; - sandbox::record_notification_audit( - paths, - "watcher:therock-update", - "notify_if_newer", - Some("therock-update"), - message, - ) -} - -fn therock_update_due(state: &AutomationRuntimeState, now: u128) -> bool { - let Some(snapshot) = state - .active_watchers - .iter() - .find(|watcher| watcher.id == "therock-update" && watcher.enabled) - else { - return false; - }; - snapshot - .last_event_unix_ms - .is_none_or(|last_event| now.saturating_sub(last_event) >= THEROCK_UPDATE_INTERVAL_MS) -} - -fn gpu_metrics_due(state: &AutomationRuntimeState, now: u128) -> bool { - let Some(snapshot) = state - .active_watchers - .iter() - .find(|watcher| watcher.id == "gpu-metrics" && watcher.enabled) - else { - return false; - }; - snapshot - .last_event_unix_ms - .is_none_or(|last_event| now.saturating_sub(last_event) >= GPU_METRICS_INTERVAL_MS) -} - -fn gpu_thermal_protect_due(state: &AutomationRuntimeState, now: u128) -> bool { - let Some(snapshot) = state - .active_watchers - .iter() - .find(|watcher| watcher.id == "gpu-thermal-protect" && watcher.enabled) - else { - return false; - }; - snapshot - .last_event_unix_ms - .is_none_or(|last_event| now.saturating_sub(last_event) >= GPU_METRICS_INTERVAL_MS) -} - -fn handle_gpu_metrics_event( - paths: &AppPaths, - mode: WatcherMode, - state: &mut AutomationRuntimeState, - event: &AutomationTriggerEvent, -) -> Result<()> { - let summary = event - .payload - .get("summary") - .and_then(Value::as_str) - .unwrap_or("summary unavailable"); - let level = if event.kind == "gpu.metrics" { - "info" - } else { - "warn" - }; - let mode_note = match mode { - WatcherMode::Observe => "observe mode records telemetry only", - WatcherMode::Propose => { - "propose mode has no GPU mutation policy yet, so telemetry is recorded only" - } - WatcherMode::Contained => { - "contained mode has no GPU mutation policy yet, so telemetry is recorded only" - } - }; - let reason = event - .reason - .as_deref() - .filter(|value| !value.trim().is_empty()) - .unwrap_or("no detail"); - let source = match event.source.as_str() { - "gpu_telemetry" => "local amd-smi telemetry", - "local_webhook" => "local webhook", - other => other, - }; - - persistence::record_event( - paths, - state, - "gpu-metrics", - level, - "record_gpu_metrics", - &format!( - "GPU metrics event from {source}: {summary}; reason={reason}; {mode_note}; no mutating action was taken" - ), - None, - ) -} - -fn gpu_snapshot_summary(snapshot: &CodexBridgeGpuSnapshot) -> String { - let mut parts = Vec::new(); - parts.push(format!("amd_smi_available={}", snapshot.amd_smi_available)); - parts.push(format!( - "static_snapshot={}", - if snapshot.static_snapshot.is_some() { - "available" - } else { - "missing" - } - )); - parts.push(format!( - "monitor_snapshot={}", - if snapshot.monitor_snapshot.is_some() { - "available" - } else { - "missing" - } - )); - if let Some(count) = snapshot.static_snapshot.as_ref().and_then(gpu_data_count) { - parts.push(format!("gpu_count={count}")); - } - if let Some(note) = snapshot.note.as_deref() - && !note.trim().is_empty() - { - parts.push(format!("note={note}")); - } - parts.join(" ") -} - -fn gpu_data_count(value: &Value) -> Option { - value - .get("gpu_data") - .and_then(Value::as_array) - .map(Vec::len) -} - -#[derive(Debug, Clone, Copy)] -struct GpuPressureReading { - gpu_index: Option, - hotspot_temperature_c: Option, - memory_temperature_c: Option, - vram_percent: Option, -} - -fn gpu_pressure_events( - now: u128, - snapshot: &CodexBridgeGpuSnapshot, -) -> Vec { - let Some(monitor_snapshot) = snapshot.monitor_snapshot.as_ref() else { - return Vec::new(); - }; - monitor_entries(monitor_snapshot) - .into_iter() - .filter_map(gpu_pressure_reading) - .filter_map(|reading| gpu_pressure_event(now, reading)) - .collect() -} - -fn gpu_pressure_event(now: u128, reading: GpuPressureReading) -> Option { - let (kind, reason, metric_label, value, threshold) = if let Some(value) = - reading.hotspot_temperature_c - && value >= GPU_THERMAL_HOTSPOT_PRESSURE_C - { - ( - "gpu.thermal_pressure", - "hotspot_temperature_threshold", - "hotspot temperature", - value, - GPU_THERMAL_HOTSPOT_PRESSURE_C, - ) - } else if let Some(value) = reading.memory_temperature_c - && value >= GPU_THERMAL_MEMORY_PRESSURE_C - { - ( - "gpu.thermal_pressure", - "memory_temperature_threshold", - "memory temperature", - value, - GPU_THERMAL_MEMORY_PRESSURE_C, - ) - } else if let Some(value) = reading.vram_percent - && value >= GPU_MEMORY_VRAM_PRESSURE_PERCENT - { - ( - "gpu.memory_pressure", - "vram_pressure_threshold", - "VRAM use", - value, - GPU_MEMORY_VRAM_PRESSURE_PERCENT, - ) - } else { - return None; - }; - let gpu_label = reading - .gpu_index - .map_or_else(|| "the GPU".to_owned(), |gpu| format!("GPU {gpu}")); - let unit = if metric_label == "VRAM use" { - "%" - } else { - " C" - }; - let summary = format!( - "{gpu_label} {metric_label} is {}{} (limit {}{})", - display_metric(value), - unit, - display_metric(threshold), - unit - ); - Some(AutomationTriggerEvent { - at_unix_ms: now, - kind: kind.to_owned(), - source: "gpu_telemetry".to_owned(), - watcher_hint: Some("gpu-thermal-protect".to_owned()), - service_id: None, - reason: Some(reason.to_owned()), - payload: json!({ - "gpu": reading.gpu_index, - "hotspot_temperature_c": reading.hotspot_temperature_c, - "memory_temperature_c": reading.memory_temperature_c, - "vram_percent": reading.vram_percent, - "hotspot_threshold_c": GPU_THERMAL_HOTSPOT_PRESSURE_C, - "memory_temperature_threshold_c": GPU_THERMAL_MEMORY_PRESSURE_C, - "vram_threshold_percent": GPU_MEMORY_VRAM_PRESSURE_PERCENT, - "recommended_action": "stop_serving_load", - "summary": summary, - }), - }) -} - -fn monitor_entries(value: &Value) -> Vec<&Value> { - if let Some(entries) = value.as_array() { - return entries.iter().collect(); - } - value - .get("gpu_data") - .and_then(Value::as_array) - .map(|entries| entries.iter().collect()) - .unwrap_or_default() -} - -fn gpu_pressure_reading(entry: &Value) -> Option { - let reading = GpuPressureReading { - gpu_index: metric_u64(entry, &["gpu", "gpu_id", "gpu_index"]), - hotspot_temperature_c: metric_f64( - entry, - &[ - "hotspot_temperature", - "hotspot_temperature_c", - "temperature_hotspot", - ], - ), - memory_temperature_c: metric_f64( - entry, - &[ - "memory_temperature", - "memory_temperature_c", - "temperature_memory", - ], - ), - vram_percent: metric_f64( - entry, - &["vram_percent", "vram_usage_percent", "vram_used_percent"], - ), - }; - (reading.hotspot_temperature_c.is_some() - || reading.memory_temperature_c.is_some() - || reading.vram_percent.is_some()) - .then_some(reading) -} - -fn metric_f64(entry: &Value, keys: &[&str]) -> Option { - keys.iter() - .find_map(|key| entry.get(*key).and_then(value_as_metric_f64)) -} - -fn metric_u64(entry: &Value, keys: &[&str]) -> Option { - keys.iter() - .find_map(|key| entry.get(*key).and_then(value_as_metric_u64)) -} - -fn value_as_metric_f64(value: &Value) -> Option { - match value { - Value::Number(number) => number.as_f64(), - Value::String(text) => text.trim().parse::().ok(), - Value::Object(map) => map - .get("value") - .or_else(|| map.get("val")) - .and_then(value_as_metric_f64), - _ => None, - } -} - -fn value_as_metric_u64(value: &Value) -> Option { - match value { - Value::Number(number) => number.as_u64(), - Value::String(text) => text.trim().parse::().ok(), - Value::Object(map) => map - .get("value") - .or_else(|| map.get("val")) - .and_then(value_as_metric_u64), - _ => None, - } -} - -fn display_metric(value: f64) -> String { - if value.fract().abs() < f64::EPSILON { - format!("{value:.0}") - } else { - format!("{value:.1}") - } -} - -fn handle_gpu_thermal_protect_event( - paths: &AppPaths, - mode: WatcherMode, - state: &mut AutomationRuntimeState, - event: &AutomationTriggerEvent, -) -> Result<()> { - let summary = payload_string(&event.payload, "summary") - .unwrap_or_else(|| "GPU pressure is high".to_owned()); - let reason = event.reason.as_deref().unwrap_or("gpu_pressure_threshold"); - - if matches!(mode, WatcherMode::Observe) { - return persistence::record_event( - paths, - state, - "gpu-thermal-protect", - "warn", - "observe_gpu_pressure", - &format!( - "{summary}; observe mode records this only and does not stop any model server" - ), - event.service_id.clone(), - ); - } - - let Some(record) = resolve_gpu_pressure_service_target(paths, event)? else { - return persistence::record_event( - paths, - state, - "gpu-thermal-protect", - "warn", - "gpu_pressure_no_clear_target", - &format!( - "{summary}; rocm-cli did not choose a model server to stop because there was no single clear running managed server" - ), - None, - ); - }; - - if pending_stop_proposal_exists(paths, &record.service_id)? { - return persistence::record_event( - paths, - state, - "gpu-thermal-protect", - "info", - "stop_proposal_already_pending", - &format!( - "{summary}; a reviewed stop request is already waiting for {}", - record.service_id - ), - Some(record.service_id), - ); - } - - let action = "queue_stop_server_proposal"; - let mode_note = if matches!(mode, WatcherMode::Contained) { - "contained mode still asks before stopping anything" - } else { - "asking before stopping anything" - }; - let message = format!( - "{summary}; {mode_note}; selected managed server {} ({})", - record.service_id, record.endpoint_url - ); - persistence::record_event( - paths, - state, - "gpu-thermal-protect", - "warn", - action, - &message, - Some(record.service_id.clone()), - )?; - queue_proposal_with_arguments( - paths, - "gpu-thermal-protect", - action, - "Review GPU pressure stop", - &message, - Some(record.service_id.clone()), - json!({ - "service_id": record.service_id, - "model_ref": record.model_ref, - "canonical_model_id": record.canonical_model_id, - "endpoint_url": record.endpoint_url, - "engine": record.engine, - "pressure_kind": event.kind, - "pressure_reason": reason, - "pressure_summary": summary, - "gpu": event.payload.get("gpu").cloned().unwrap_or(Value::Null), - "hotspot_temperature_c": event.payload.get("hotspot_temperature_c").cloned().unwrap_or(Value::Null), - "memory_temperature_c": event.payload.get("memory_temperature_c").cloned().unwrap_or(Value::Null), - "vram_percent": event.payload.get("vram_percent").cloned().unwrap_or(Value::Null), - "hotspot_threshold_c": GPU_THERMAL_HOTSPOT_PRESSURE_C, - "memory_temperature_threshold_c": GPU_THERMAL_MEMORY_PRESSURE_C, - "vram_threshold_percent": GPU_MEMORY_VRAM_PRESSURE_PERCENT, - }), - ) -} - -fn resolve_gpu_pressure_service_target( - paths: &AppPaths, - event: &AutomationTriggerEvent, -) -> Result> { - if let Some(service_id) = event - .service_id - .as_deref() - .map(str::trim) - .filter(|value| !value.is_empty()) - { - let record = load_service_record(paths, service_id)?; - return Ok(active_pressure_target(record)); - } - - let active = persistence::load_managed_services(paths)? - .into_iter() - .filter_map(active_pressure_target) - .collect::>(); - if active.len() == 1 { - Ok(active.into_iter().next()) - } else { - Ok(None) - } -} - -fn active_pressure_target(record: ManagedServiceRecord) -> Option { - (record.mode == "managed" && matches!(record.status.as_str(), "ready" | "running")) - .then_some(record) -} - -fn pending_stop_proposal_exists(paths: &AppPaths, service_id: &str) -> Result { - Ok(rocm_core::load_recent_automation_proposals(paths, 100)? - .into_iter() - .any(|proposal| { - proposal.status == "pending" - && proposal.watcher_id == "gpu-thermal-protect" - && proposal.service_id.as_deref() == Some(service_id) - && (proposal.action == "queue_stop_server_proposal" - || proposal.tool.as_deref() == Some("stop_server")) - })) -} - -fn handle_cache_warm_event( - paths: &AppPaths, - mode: WatcherMode, - state: &mut AutomationRuntimeState, - event: &AutomationTriggerEvent, -) -> Result<()> { - handle_cache_warm_event_with_resolver(paths, mode, state, event, |artifact_ref| { - resolve_model_recipe_artifact(artifact_ref).map(|resolved| resolved.is_some()) - }) -} - -fn handle_cache_warm_event_with_resolver( - paths: &AppPaths, - mode: WatcherMode, - state: &mut AutomationRuntimeState, - event: &AutomationTriggerEvent, - mut artifact_exists: F, -) -> Result<()> -where - F: FnMut(&str) -> Result, -{ - let Some(artifact_ref) = payload_string(&event.payload, "artifact_ref") else { - return persistence::record_event( - paths, - state, - "cache-warm", - "warn", - "cache_warm_missing_artifact", - "cache warm event did not include artifact_ref; no prefetch proposal was queued", - None, - ); - }; - match artifact_exists(&artifact_ref) { - Ok(true) => {} - Ok(false) => { - return persistence::record_event( - paths, - state, - "cache-warm", - "warn", - "cache_warm_unknown_artifact", - &format!( - "cache warm requested unknown registry artifact {artifact_ref}; no prefetch proposal was queued" - ), - None, - ); - } - Err(error) => { - return persistence::record_event( - paths, - state, - "cache-warm", - "error", - "cache_warm_registry_error", - &format!( - "cache warm could not verify registry artifact {artifact_ref}: {error}; no prefetch proposal was queued" - ), - None, - ); - } - } - match mode { - WatcherMode::Observe => persistence::record_event( - paths, - state, - "cache-warm", - "info", - "observe_cache_warm_request", - &format!( - "observed cache warm request for {artifact_ref}; observe mode does not queue or download artifacts" - ), - None, - ), - WatcherMode::Propose | WatcherMode::Contained => { - let action = "queue_prefetch_proposal"; - let message = if matches!(mode, WatcherMode::Contained) { - format!( - "cache warm requested for {artifact_ref}; contained mode still queues a review because artifact downloads require explicit source-policy approval" - ) - } else { - format!( - "cache warm requested for {artifact_ref}; queueing a reviewed prefetch proposal" - ) - }; - persistence::record_event(paths, state, "cache-warm", "info", action, &message, None)?; - queue_proposal_with_arguments( - paths, - "cache-warm", - action, - "Prefetch model artifact", - &message, - None, - json!({ - "artifact_ref": artifact_ref, - }), - ) - } - } -} - -fn handle_driver_upgrade_event( - paths: &AppPaths, - mode: WatcherMode, - state: &mut AutomationRuntimeState, - event: &AutomationTriggerEvent, -) -> Result<()> { - handle_driver_upgrade_event_with_runner(paths, mode, state, event, |paths| { - sandbox::run_sandbox_tool( - paths, - cli::SandboxToolArg::DriverPlan, - None, - None, - None, - cli::SandboxToolPolicy::default(), - ) - }) -} - -fn handle_driver_upgrade_event_with_runner( - paths: &AppPaths, - mode: WatcherMode, - state: &mut AutomationRuntimeState, - event: &AutomationTriggerEvent, - driver_plan_runner: F, -) -> Result<()> -where - F: FnOnce(&AppPaths) -> Result, -{ - if payload_string(&event.payload, "component").as_deref() != Some("driver") { - return persistence::record_event( - paths, - state, - "driver-upgrade", - "warn", - "driver_upgrade_ignored_component", - "driver-upgrade event did not include payload.component=driver; no driver plan proposal was queued", - None, - ); - } - - match mode { - WatcherMode::Observe => persistence::record_event( - paths, - state, - "driver-upgrade", - "info", - "observe_driver_update", - "observed local driver update signal; observe mode does not queue or run a driver plan", - None, - ), - WatcherMode::Propose => { - let action = "prepare_driver_plan"; - let message = - "local driver update signal received; queueing a reviewed read-only driver plan"; - persistence::record_event( - paths, - state, - "driver-upgrade", - "info", - action, - message, - None, - )?; - queue_proposal( - paths, - "driver-upgrade", - action, - "Review driver install plan", - message, - None, - ) - } - WatcherMode::Contained => match driver_plan_runner(paths) { - Ok(output) => match restricted_driver_plan_result(&output) { - Ok(result) => persistence::record_event( - paths, - state, - "driver-upgrade", - if result.exit_status == 0 { - "info" - } else { - "error" - }, - "run_driver_plan", - &format!( - "local driver update signal received; contained restricted driver_plan status={}; no driver commands were executed", - result.status - ), - None, - ), - Err(error) => persistence::record_event( - paths, - state, - "driver-upgrade", - "error", - "driver_plan_failed", - &format!( - "local driver update signal received, but contained restricted driver_plan failed: {error}; no driver commands were executed" - ), - None, - ), - }, - Err(error) => persistence::record_event( - paths, - state, - "driver-upgrade", - "error", - "driver_plan_failed", - &format!( - "local driver update signal received, but contained restricted driver_plan failed: {error}; no driver commands were executed" - ), - None, - ), - }, - } -} - -struct RestrictedDriverPlanResult<'a> { - status: &'a str, - exit_status: i64, -} - -fn restricted_driver_plan_result(value: &Value) -> Result> { - let tool = value - .get("tool") - .and_then(Value::as_str) - .context("restricted driver plan did not report a tool name")?; - if tool != cli::SandboxToolArg::DriverPlan.as_cli_value() { - bail!("restricted driver plan returned `{tool}`, expected `driver_plan`"); - } - let status = value - .get("status") - .and_then(Value::as_str) - .unwrap_or("planned"); - let exit_status = value - .get("exit_status") - .and_then(Value::as_i64) - .unwrap_or_else(|| i64::from(status == "error")); - Ok(RestrictedDriverPlanResult { - status, - exit_status, - }) -} - -pub(crate) fn payload_string(payload: &Value, key: &str) -> Option { - payload - .get(key) - .and_then(Value::as_str) - .map(str::trim) - .filter(|value| !value.is_empty()) - .map(str::to_owned) -} - -#[cfg(test)] -fn evaluate_server_recover( - paths: &AppPaths, - mode: WatcherMode, - state: &mut AutomationRuntimeState, -) -> Result<()> { - let now = unix_time_millis(); - if !server_recover_due(state, now) { - return Ok(()); - } - - let Some((mut record, recovery_reason)) = find_recoverable_service(paths)? else { - return Ok(()); - }; - let kind = service_recovery_event_kind(&recovery_reason); - let event = AutomationTriggerEvent { - at_unix_ms: now, - kind: kind.to_owned(), - source: "managed_service".to_owned(), - watcher_hint: Some("server-recover".to_owned()), - service_id: Some(record.service_id.clone()), - reason: Some(recovery_reason), - payload: json!({ - "engine": record.engine.clone(), - "status": record.status.clone(), - "endpoint": record.endpoint_url.clone(), - }), - }; - handle_server_recover_event_with_record(paths, mode, state, &event, &mut record) -} - -fn handle_server_recover_event( - paths: &AppPaths, - mode: WatcherMode, - state: &mut AutomationRuntimeState, - event: &AutomationTriggerEvent, -) -> Result<()> { - let service_id = event - .service_id - .as_deref() - .context("server-recover event is missing service_id")?; - let mut record = load_service_record(paths, service_id)?; - if !service_record_matches_recovery_event(paths, &record, event) { - persistence::record_event( - paths, - state, - "server-recover", - "info", - "ignore_nonrecoverable_service", - &format!( - "managed service {} does not currently need recovery; restart not attempted", - record.service_id - ), - Some(record.service_id.clone()), - )?; - return Ok(()); - } - handle_server_recover_event_with_record(paths, mode, state, event, &mut record) -} - -fn service_record_matches_recovery_event( - paths: &AppPaths, - record: &ManagedServiceRecord, - event: &AutomationTriggerEvent, -) -> bool { - match event.kind.as_str() { - "service.manifest_recoverable" => { - manifest_service_recovery_reason(record, unix_time_millis()).is_some() - } - "service.endpoint_recoverable" => endpoint_service_recovery_reason(record).is_some(), - "service.healthcheck_recoverable" => { - common::engine_healthcheck_response(paths, &record.engine, &record.service_id) - .is_ok_and(|healthcheck| common::healthcheck_response_recoverable(&healthcheck)) - } - _ => false, - } -} - -fn handle_server_recover_event_with_record( - paths: &AppPaths, - mode: WatcherMode, - state: &mut AutomationRuntimeState, - event: &AutomationTriggerEvent, - record: &mut ManagedServiceRecord, -) -> Result<()> { - let now = unix_time_millis(); - let recovery_reason = event.reason.as_deref().unwrap_or("recoverable_event"); - let recovery_reason_display = display_recovery_reason(recovery_reason); - - match watcher_policy_action("server-recover", mode) { - WatcherPolicyAction::Observe => persistence::record_event( - paths, - state, - "server-recover", - "warn", - "observe_failure", - &format!( - "observed managed service {} needing recovery ({recovery_reason_display}); restart not attempted in observe mode", - record.service_id, - ), - Some(record.service_id.clone()), - ), - WatcherPolicyAction::QueueProposal => { - let message = format!( - "managed service {} needs recovery ({recovery_reason_display}); queueing restart proposal", - record.service_id, - ); - persistence::record_event( - paths, - state, - "server-recover", - "warn", - "queue_restart_proposal", - &message, - Some(record.service_id.clone()), - )?; - queue_proposal( - paths, - "server-recover", - "queue_restart_proposal", - "Restart managed service", - &message, - Some(record.service_id.clone()), - ) - } - WatcherPolicyAction::RunContained => { - if let Some(last_restart) = record.last_restart_unix_ms - && now.saturating_sub(last_restart) < SERVER_RECOVER_BACKOFF_MS - { - return Ok(()); - } - // A public service whose endpoint key is gone can never be recovered: - // the respawn guard in `supervise_service` refuses it by design. - // Report it and stop, rather than letting a permanent failure - // propagate out of `evaluate_watchers` and take the whole daemon — - // and every other watcher — down on each 30s recovery tick. - if let Err(error) = common::ensure_public_service_has_endpoint_key( - &record.host, - rocm_engine_protocol::endpoint_key_file_if_present(paths, &record.service_id) - .and_then(|path| rocm_engine_protocol::endpoint_api_key_file_if_valid(&path)) - .is_some(), - record.requires_api_key, - ) { - return persistence::record_event( - paths, - state, - "server-recover", - "error", - "restart_managed_service_refused", - &format!( - "cannot recover managed service {} on {}:{} after \ - {recovery_reason_display}: {error}", - record.service_id, record.host, record.port - ), - Some(record.service_id.clone()), - ); - } - restart_managed_service(paths, &mut *record)?; - persistence::record_event( - paths, - state, - "server-recover", - "info", - "restart_managed_service", - &format!( - "restarted managed service {} on {}:{} after {recovery_reason_display}", - record.service_id, record.host, record.port - ), - Some(record.service_id.clone()), - ) - } - } -} - -fn display_recovery_reason(reason: &str) -> String { - match reason { - "manifest_status_failed" => "manifest reports failed".to_owned(), - "manifest_status_exited" => "manifest reports exited".to_owned(), - "manifest_status_unreachable" => "manifest reports unreachable".to_owned(), - "manifest_status_starting_stale" => "service has been starting for too long".to_owned(), - "manifest_status_recovering_stale" => "service has been recovering for too long".to_owned(), - "endpoint_status_unreachable" => "endpoint port is unreachable".to_owned(), - other if other.starts_with("healthcheck_status_") => format!( - "engine healthcheck reports {}", - other.trim_start_matches("healthcheck_status_") - ), - other => other.replace('_', " "), - } -} - -fn service_recovery_event_kind(recovery_reason: &str) -> &'static str { - if recovery_reason.starts_with("healthcheck_status_") { - "service.healthcheck_recoverable" - } else if recovery_reason.starts_with("endpoint_status_") { - "service.endpoint_recoverable" - } else { - "service.manifest_recoverable" - } -} - -fn server_recover_due(state: &AutomationRuntimeState, now: u128) -> bool { - let Some(snapshot) = state - .active_watchers - .iter() - .find(|watcher| watcher.id == "server-recover" && watcher.enabled) - else { - return false; - }; - snapshot - .last_event_unix_ms - .is_none_or(|last_event| now.saturating_sub(last_event) >= SERVER_RECOVER_BACKOFF_MS) -} - -#[derive(Debug, Clone, Copy, Eq, PartialEq)] -enum WatcherPolicyAction { - Observe, - QueueProposal, - RunContained, -} - -const fn watcher_policy_action(watcher_id: &str, mode: WatcherMode) -> WatcherPolicyAction { - match (watcher_id, mode) { - (_, WatcherMode::Observe) => WatcherPolicyAction::Observe, - (_, WatcherMode::Propose) => WatcherPolicyAction::QueueProposal, - (_, WatcherMode::Contained) => WatcherPolicyAction::RunContained, - } -} - -fn find_recoverable_service(paths: &AppPaths) -> Result> { - let now = unix_time_millis(); - for record in persistence::load_managed_services(paths)? { - if record.mode != "managed" { - continue; - } - if let Some(reason) = manifest_service_recovery_reason(&record, now) { - return Ok(Some((record, reason))); - } - if matches!(record.status.as_str(), "ready" | "running") { - let Ok(healthcheck) = - common::engine_healthcheck_response(paths, &record.engine, &record.service_id) - else { - if let Some(reason) = endpoint_service_recovery_reason(&record) { - return Ok(Some((record, reason))); - } - continue; - }; - if common::healthcheck_response_recoverable(&healthcheck) { - return Ok(Some(( - record, - format!("healthcheck_status_{}", healthcheck.status), - ))); - } - if let Some(reason) = endpoint_service_recovery_reason(&record) { - return Ok(Some((record, reason))); - } - } - } - Ok(None) -} - -fn endpoint_service_recovery_reason(record: &ManagedServiceRecord) -> Option { - (!common::wait_for_port(&record.host, record.port, ENDPOINT_HEALTH_TIMEOUT)) - .then(|| "endpoint_status_unreachable".to_owned()) -} - -fn load_service_record(paths: &AppPaths, service_id: &str) -> Result { - rocm_core::ServiceId::new(service_id) - .with_context(|| format!("invalid managed service id `{service_id}`"))?; - let manifest_path = paths.service_manifest_path(service_id); - let bytes = fs::read(&manifest_path).with_context(|| { - format!( - "managed service `{service_id}` not found at {}", - manifest_path.display() - ) - })?; - let record = serde_json::from_slice::(&bytes) - .with_context(|| format!("failed to parse {}", manifest_path.display()))?; - if record.service_id != service_id { - bail!( - "managed service manifest {} contains service_id `{}`, expected `{service_id}`", - manifest_path.display(), - record.service_id - ); - } - Ok(record) -} - -fn manifest_service_recovery_reason( - record: &ManagedServiceRecord, - now_unix_ms: u128, -) -> Option { - match record.status.as_str() { - "failed" | "exited" | "unreachable" => Some(format!("manifest_status_{}", record.status)), - "starting" | "recovering" => { - let started_at = record - .last_restart_unix_ms - .unwrap_or(record.created_at_unix_ms); - (now_unix_ms.saturating_sub(started_at) >= SERVER_TRANSIENT_STALE_MS) - .then(|| format!("manifest_status_{}_stale", record.status)) - } - _ => None, - } -} - -fn restart_managed_service(_paths: &AppPaths, record: &mut ManagedServiceRecord) -> Result<()> { - let rocmd_binary = - std::env::current_exe().context("failed to resolve current rocmd executable path")?; - let log_file = fs::OpenOptions::new() - .create(true) - .append(true) - .open(&record.log_path) - .with_context(|| format!("failed to open {}", record.log_path.display()))?; - let log_file_err = log_file - .try_clone() - .context("failed to clone service log file handle")?; - - record.status = "recovering".to_owned(); - // Counts the restart and drops the previous run's inference verification. - // The respawned child writes a fresh record of its own, and "recovering" is - // outside the statuses that probe, so a stale verdict would not currently be - // acted on — but this record is written again below, after the spawn, and - // that write can land after the child's. Clearing here keeps the invariant - // true at the one site that reuses a record across restarts. - record.reset_for_restart(); - record.supervisor_pid = std::process::id(); - record.write()?; - - let mut child = detached_rocmd_command(&rocmd_binary) - .args(recovery_supervise_args(record)) - .stdin(Stdio::null()) - .stdout(Stdio::from(log_file)) - .stderr(Stdio::from(log_file_err)) - .spawn() - .context("failed to spawn recovery supervisor")?; - - record.supervisor_pid = child.id(); - record.write()?; - - thread::sleep(Duration::from_millis(200)); - if let Some(status) = child - .try_wait() - .context("failed to check recovery supervisor startup state")? - { - record.status = "failed".to_owned(); - record.write()?; - anyhow::bail!( - "recovery supervisor exited immediately with status {status}; inspect {}", - record.log_path.display() - ); - } - - Ok(()) -} - -fn recovery_supervise_args(record: &ManagedServiceRecord) -> Vec { - let mut args = vec![ - "supervise".to_owned(), - record.service_id.clone(), - "--engine".to_owned(), - record.engine.clone(), - "--model-ref".to_owned(), - record.model_ref.clone(), - "--canonical-model-id".to_owned(), - record.canonical_model_id.clone(), - "--host".to_owned(), - record.host.clone(), - "--port".to_owned(), - record.port.to_string(), - "--device-policy".to_owned(), - record - .device_policy - .as_deref() - .unwrap_or("gpu_required") - .to_owned(), - ]; - args.extend(common::optional_arg( - "--runtime-id", - record.runtime_id.as_deref(), - )); - args.extend(common::optional_arg("--env-id", record.env_id.as_deref())); - if let Some(csv) = rocm_engine_protocol::gpu_indices_to_csv(&record.gpu_indices) { - args.extend(["--gpu".to_owned(), csv]); - } - args.extend(common::optional_arg( - "--engine-recipe-json", - record.engine_recipe_json.as_deref(), - )); - args -} - -fn queue_proposal( - paths: &AppPaths, - watcher_id: &str, - action: &str, - title: &str, - message: &str, - service_id: Option, -) -> Result<()> { - queue_proposal_with_arguments( - paths, - watcher_id, - action, - title, - message, - service_id.clone(), - proposal_arguments_for_action(action, service_id.as_deref()), - ) -} - -fn queue_proposal_with_arguments( - paths: &AppPaths, - watcher_id: &str, - action: &str, - title: &str, - message: &str, - service_id: Option, - arguments: Value, -) -> Result<()> { - append_automation_proposal( - paths, - &AutomationProposalRecord { - at_unix_ms: unix_time_millis(), - proposal_id: String::new(), - watcher_id: watcher_id.to_owned(), - action: action.to_owned(), - title: title.to_owned(), - message: message.to_owned(), - status: "pending".to_owned(), - service_id, - tool: proposal_tool_for_action(action).map(str::to_owned), - arguments, - reviewed_at_unix_ms: None, - }, - ) -} - -fn proposal_tool_for_action(action: &str) -> Option<&'static str> { - match action { - "queue_restart_proposal" => Some("restart_server"), - "queue_stop_server_proposal" => Some("stop_server"), - "queue_update_proposal" => Some("check_updates"), - "queue_prefetch_proposal" => Some("prefetch_artifact"), - "prepare_driver_plan" => Some("driver_plan"), - _ => None, - } -} - -fn proposal_arguments_for_action(action: &str, service_id: Option<&str>) -> Value { - match action { - "queue_restart_proposal" => json!({ - "service_id": service_id, - }), - "queue_stop_server_proposal" => json!({ - "service_id": service_id, - }), - "queue_update_proposal" => json!({}), - "prepare_driver_plan" => json!({}), - _ => Value::Null, - } -} - -#[cfg(unix)] -fn detached_rocmd_command(rocmd_binary: &std::path::Path) -> ProcessCommand { - let mut command = ProcessCommand::new("setsid"); - command.arg(rocmd_binary); - command -} - -#[cfg(not(unix))] -fn detached_rocmd_command(rocmd_binary: &std::path::Path) -> ProcessCommand { - ProcessCommand::new(rocmd_binary) -} - -#[cfg(test)] -mod tests { - use super::*; - use crate::test_support::temp_app_paths; - - #[test] - fn recovery_supervise_args_preserve_engine_recipe_json() { - let (_root, paths) = temp_app_paths("recovery-engine-recipe"); - let mut record = ManagedServiceRecord::new( - &paths, - "svc-1", - "vllm", - "qwen", - "Qwen/Qwen3.5-4B", - "127.0.0.1", - 11435, - "managed", - 123, - Some("therock-release:gfx120X-all".to_owned()), - Some("env-1".to_owned()), - Some("gpu_required".to_owned()), - ); - let engine_recipe_json = r#"{"contract_version":"0.1.0","engine":"vllm","required_flags":["--enable-auto-tool-choice"]}"#; - record.engine_recipe_json = Some(engine_recipe_json.to_owned()); - - let args = recovery_supervise_args(&record); - - assert!( - args.windows(2) - .any(|pair| pair[0] == "--engine-recipe-json" && pair[1] == engine_recipe_json) - ); - assert!( - args.windows(2) - .any(|pair| { pair[0] == "--canonical-model-id" && pair[1] == "Qwen/Qwen3.5-4B" }) - ); - } - - #[test] - fn watcher_policy_maps_modes_to_decisions() { - assert_eq!( - watcher_policy_action("server-recover", WatcherMode::Observe), - WatcherPolicyAction::Observe - ); - assert_eq!( - watcher_policy_action("server-recover", WatcherMode::Propose), - WatcherPolicyAction::QueueProposal - ); - assert_eq!( - watcher_policy_action("server-recover", WatcherMode::Contained), - WatcherPolicyAction::RunContained - ); - assert_eq!( - watcher_policy_action("therock-update", WatcherMode::Contained), - WatcherPolicyAction::RunContained - ); - } - - #[test] - fn event_collector_emits_schedule_tick_for_due_update() -> Result<()> { - let (root, paths) = temp_app_paths("event-bus-schedule"); - let state = test_runtime_state(vec![test_watcher_snapshot( - "therock-update", - WatcherMode::Observe, - None, - )]); - - let events = collect_automation_events(&paths, &RocmCliConfig::default(), &state)?; - fs::remove_dir_all(root).ok(); - - let event = events - .iter() - .find(|event| event.watcher_hint.as_deref() == Some("therock-update")) - .expect("schedule tick event should be emitted"); - assert_eq!(event.kind, "schedule.tick"); - assert_eq!(event.source, "scheduler"); - assert_eq!(event.reason.as_deref(), Some("therock_update_interval_due")); - Ok(()) - } - - #[test] - fn event_collector_emits_recoverable_service_event() -> Result<()> { - let (root, paths) = temp_app_paths("event-bus-service"); - paths.ensure()?; - let mut failed = ManagedServiceRecord::new( - &paths, - "svc-failed", - "vllm", - "qwen", - "Qwen/Qwen3.5", - "127.0.0.1", - 11435, - "managed", - 123, - None, - None, - None, - ); - failed.status = "failed".to_owned(); - failed.write()?; - let state = test_runtime_state(vec![test_watcher_snapshot( - "server-recover", - WatcherMode::Propose, - None, - )]); - - let events = collect_automation_events(&paths, &RocmCliConfig::default(), &state)?; - fs::remove_dir_all(root).ok(); - - let event = events - .iter() - .find(|event| event.watcher_hint.as_deref() == Some("server-recover")) - .expect("recoverable service event should be emitted"); - assert_eq!(event.kind, "service.manifest_recoverable"); - assert_eq!(event.service_id.as_deref(), Some("svc-failed")); - assert_eq!(event.reason.as_deref(), Some("manifest_status_failed")); - Ok(()) - } - - #[test] - fn event_collector_emits_endpoint_recoverable_service_event() -> Result<()> { - let (root, paths) = temp_app_paths("event-bus-endpoint"); - paths.ensure()?; - let mut service = ManagedServiceRecord::new( - &paths, - "svc-endpoint", - "missing-engine", - "qwen", - "Qwen/Qwen3.5", - "127.0.0.1", - 1, - "managed", - 123, - None, - None, - None, - ); - service.status = "ready".to_owned(); - service.write()?; - let state = test_runtime_state(vec![test_watcher_snapshot( - "server-recover", - WatcherMode::Propose, - None, - )]); - - let events = collect_automation_events(&paths, &RocmCliConfig::default(), &state)?; - fs::remove_dir_all(root).ok(); - - let event = events - .iter() - .find(|event| event.watcher_hint.as_deref() == Some("server-recover")) - .expect("endpoint recoverable service event should be emitted"); - assert_eq!(event.kind, "service.endpoint_recoverable"); - assert_eq!(event.service_id.as_deref(), Some("svc-endpoint")); - assert_eq!(event.reason.as_deref(), Some("endpoint_status_unreachable")); - Ok(()) - } - - #[test] - fn event_collector_emits_gpu_metrics_event_when_enabled() -> Result<()> { - let (root, paths) = temp_app_paths("event-bus-gpu-metrics"); - let state = test_runtime_state(vec![test_watcher_snapshot( - "gpu-metrics", - WatcherMode::Observe, - None, - )]); - - let events = collect_automation_events_with_gpu_snapshot(&paths, &state, || { - CodexBridgeGpuSnapshot { - amd_smi_available: true, - static_snapshot: Some(json!({ - "gpu_data": [ - { "gpu": 0, "asic": { "market_name": "AMD Radeon Test" } } - ] - })), - monitor_snapshot: Some(json!({ "gpu_data": [] })), - note: None, - } - })?; - fs::remove_dir_all(root).ok(); - - let event = events - .iter() - .find(|event| event.watcher_hint.as_deref() == Some("gpu-metrics")) - .expect("gpu metrics event should be emitted"); - assert_eq!(event.kind, "gpu.metrics"); - assert_eq!(event.source, "gpu_telemetry"); - assert_eq!(event.reason.as_deref(), Some("amd_smi_snapshot_available")); - assert_eq!( - event - .payload - .get("monitor_available") - .and_then(Value::as_bool), - Some(true) - ); - assert!( - event - .payload - .get("summary") - .and_then(Value::as_str) - .is_some_and(|summary| summary.contains("gpu_count=1")) - ); - Ok(()) - } - - #[test] - fn event_collector_emits_gpu_thermal_pressure_event_when_enabled() -> Result<()> { - let (root, paths) = temp_app_paths("event-bus-gpu-thermal-pressure"); - let state = test_runtime_state(vec![test_watcher_snapshot( - "gpu-thermal-protect", - WatcherMode::Propose, - None, - )]); - - let events = collect_automation_events_with_gpu_snapshot(&paths, &state, || { - CodexBridgeGpuSnapshot { - amd_smi_available: true, - static_snapshot: None, - monitor_snapshot: Some(json!({ - "gpu_data": [ - { - "gpu": 0, - "hotspot_temperature": { "value": 96.0 }, - "memory_temperature": { "value": 88.0 }, - "vram_percent": { "value": 72.0 } - } - ] - })), - note: None, - } - })?; - fs::remove_dir_all(root).ok(); - - let event = events - .iter() - .find(|event| event.watcher_hint.as_deref() == Some("gpu-thermal-protect")) - .expect("thermal pressure event should be emitted"); - assert_eq!(event.kind, "gpu.thermal_pressure"); - assert_eq!(event.source, "gpu_telemetry"); - assert_eq!( - event.reason.as_deref(), - Some("hotspot_temperature_threshold") - ); - assert!( - event - .payload - .get("summary") - .and_then(Value::as_str) - .is_some_and(|summary| summary.contains("GPU 0 hotspot temperature is 96 C")) - ); - Ok(()) - } - - #[test] - fn event_collector_skips_gpu_pressure_below_thresholds() -> Result<()> { - let (root, paths) = temp_app_paths("event-bus-gpu-pressure-cool"); - let state = test_runtime_state(vec![test_watcher_snapshot( - "gpu-thermal-protect", - WatcherMode::Propose, - None, - )]); - - let events = collect_automation_events_with_gpu_snapshot(&paths, &state, || { - CodexBridgeGpuSnapshot { - amd_smi_available: true, - static_snapshot: None, - monitor_snapshot: Some(json!([ - { - "gpu": 0, - "hotspot_temperature": 80.0, - "memory_temperature": 82.0, - "vram_percent": 50.0 - } - ])), - note: None, - } - })?; - fs::remove_dir_all(root).ok(); - - assert!( - !events - .iter() - .any(|event| event.watcher_hint.as_deref() == Some("gpu-thermal-protect")) - ); - Ok(()) - } - - #[test] - fn gpu_metrics_event_records_read_only_status_without_proposal() -> Result<()> { - let (root, paths) = temp_app_paths("gpu-metrics-record"); - paths.ensure()?; - let mut state = test_runtime_state(vec![test_watcher_snapshot( - "gpu-metrics", - WatcherMode::Contained, - None, - )]); - let mut config = RocmCliConfig::default(); - let watcher = config.watcher_config_mut("gpu-metrics"); - watcher.enabled = true; - watcher.mode = Some(WatcherMode::Contained); - let events = vec![AutomationTriggerEvent { - at_unix_ms: 1, - kind: "gpu.metrics_unavailable".to_owned(), - source: "gpu_telemetry".to_owned(), - watcher_hint: Some("gpu-metrics".to_owned()), - service_id: None, - reason: Some("amd-smi missing".to_owned()), - payload: json!({ - "summary": "amd_smi_available=false static_snapshot=missing monitor_snapshot=missing", - }), - }]; - - evaluate_watchers_for_events(&paths, &config, &mut state, &events)?; - let event_text = fs::read_to_string(paths.automation_events_path())?; - let event = serde_json::from_str::(event_text.trim())?; - let proposals = rocm_core::load_recent_automation_proposals(&paths, 1)?; - fs::remove_dir_all(root).ok(); - - assert_eq!(event.watcher_id, "gpu-metrics"); - assert_eq!(event.action, "record_gpu_metrics"); - assert!(event.message.contains("telemetry is recorded only")); - assert!(event.message.contains("amd-smi missing")); - assert!(event.message.contains("no mutating action was taken")); - assert!(proposals.is_empty()); - Ok(()) - } - - #[test] - fn gpu_thermal_protect_propose_queues_reviewed_stop_for_one_running_service() -> Result<()> { - let (root, paths) = temp_app_paths("gpu-thermal-protect-propose"); - paths.ensure()?; - let mut record = ManagedServiceRecord::new( - &paths, - "svc-hot", - "vllm", - "tiny", - "Tiny/Test", - "127.0.0.1", - 11435, - "managed", - 123, - None, - None, - None, - ); - record.status = "ready".to_owned(); - record.write()?; - let mut state = test_runtime_state(vec![test_watcher_snapshot( - "gpu-thermal-protect", - WatcherMode::Propose, - None, - )]); - let mut config = RocmCliConfig::default(); - let watcher = config.watcher_config_mut("gpu-thermal-protect"); - watcher.enabled = true; - watcher.mode = Some(WatcherMode::Propose); - let event = AutomationTriggerEvent { - at_unix_ms: 1, - kind: "gpu.thermal_pressure".to_owned(), - source: "gpu_telemetry".to_owned(), - watcher_hint: Some("gpu-thermal-protect".to_owned()), - service_id: None, - reason: Some("hotspot_temperature_threshold".to_owned()), - payload: json!({ - "gpu": 0, - "summary": "GPU 0 hotspot temperature is 96 C (limit 95 C)", - "hotspot_temperature_c": 96.0, - }), - }; - - evaluate_watchers_for_events(&paths, &config, &mut state, &[event])?; - let event_text = fs::read_to_string(paths.automation_events_path())?; - let event = serde_json::from_str::(event_text.trim())?; - let proposals = rocm_core::load_recent_automation_proposals(&paths, 1)?; - let saved = load_service_record(&paths, "svc-hot")?; - fs::remove_dir_all(root).ok(); - - assert_eq!(event.watcher_id, "gpu-thermal-protect"); - assert_eq!(event.action, "queue_stop_server_proposal"); - assert_eq!(event.service_id.as_deref(), Some("svc-hot")); - assert!(event.message.contains("asking before stopping anything")); - assert_eq!(saved.status, "ready"); - assert_eq!(proposals.len(), 1); - assert_eq!(proposals[0].tool.as_deref(), Some("stop_server")); - assert_eq!(proposals[0].service_id.as_deref(), Some("svc-hot")); - assert_eq!( - proposals[0] - .arguments - .get("pressure_reason") - .and_then(Value::as_str), - Some("hotspot_temperature_threshold") - ); - Ok(()) - } - - #[test] - fn gpu_thermal_protect_contained_still_queues_reviewed_stop() -> Result<()> { - let (root, paths) = temp_app_paths("gpu-thermal-protect-contained"); - paths.ensure()?; - let mut record = ManagedServiceRecord::new( - &paths, - "svc-hot", - "vllm", - "qwen", - "Qwen/Test", - "127.0.0.1", - 11436, - "managed", - 123, - None, - None, - None, - ); - record.status = "running".to_owned(); - record.write()?; - let mut state = test_runtime_state(vec![test_watcher_snapshot( - "gpu-thermal-protect", - WatcherMode::Contained, - None, - )]); - let event = AutomationTriggerEvent { - at_unix_ms: 1, - kind: "gpu.memory_pressure".to_owned(), - source: "gpu_telemetry".to_owned(), - watcher_hint: Some("gpu-thermal-protect".to_owned()), - service_id: Some("svc-hot".to_owned()), - reason: Some("vram_pressure_threshold".to_owned()), - payload: json!({ - "summary": "GPU 0 VRAM use is 96% (limit 95%)", - "vram_percent": 96.0, - }), - }; - - handle_gpu_thermal_protect_event(&paths, WatcherMode::Contained, &mut state, &event)?; - let event_text = fs::read_to_string(paths.automation_events_path())?; - let event = serde_json::from_str::(event_text.trim())?; - let proposals = rocm_core::load_recent_automation_proposals(&paths, 1)?; - let saved = load_service_record(&paths, "svc-hot")?; - fs::remove_dir_all(root).ok(); - - assert_eq!(event.action, "queue_stop_server_proposal"); - assert!(event.message.contains("contained mode still asks")); - assert_eq!(saved.status, "running"); - assert_eq!(proposals.len(), 1); - assert_eq!(proposals[0].tool.as_deref(), Some("stop_server")); - Ok(()) - } - - #[test] - fn gpu_thermal_protect_observe_records_without_proposal() -> Result<()> { - let (root, paths) = temp_app_paths("gpu-thermal-protect-observe"); - paths.ensure()?; - let mut state = test_runtime_state(vec![test_watcher_snapshot( - "gpu-thermal-protect", - WatcherMode::Observe, - None, - )]); - let event = AutomationTriggerEvent { - at_unix_ms: 1, - kind: "gpu.thermal_pressure".to_owned(), - source: "gpu_telemetry".to_owned(), - watcher_hint: Some("gpu-thermal-protect".to_owned()), - service_id: None, - reason: Some("memory_temperature_threshold".to_owned()), - payload: json!({ - "summary": "GPU memory temperature is high", - }), - }; - - handle_gpu_thermal_protect_event(&paths, WatcherMode::Observe, &mut state, &event)?; - let event_text = fs::read_to_string(paths.automation_events_path())?; - let event = serde_json::from_str::(event_text.trim())?; - let proposals = rocm_core::load_recent_automation_proposals(&paths, 1)?; - fs::remove_dir_all(root).ok(); - - assert_eq!(event.action, "observe_gpu_pressure"); - assert!(event.message.contains("does not stop any model server")); - assert!(proposals.is_empty()); - Ok(()) - } - - #[test] - fn gpu_thermal_protect_ambiguous_services_records_no_action() -> Result<()> { - let (root, paths) = temp_app_paths("gpu-thermal-protect-ambiguous"); - paths.ensure()?; - for service_id in ["svc-a", "svc-b"] { - let mut record = ManagedServiceRecord::new( - &paths, - service_id, - "vllm", - "qwen", - "Qwen/Test", - "127.0.0.1", - 11435, - "managed", - 123, - None, - None, - None, - ); - record.status = "ready".to_owned(); - record.write()?; - } - let mut state = test_runtime_state(vec![test_watcher_snapshot( - "gpu-thermal-protect", - WatcherMode::Propose, - None, - )]); - let event = AutomationTriggerEvent { - at_unix_ms: 1, - kind: "gpu.thermal_pressure".to_owned(), - source: "gpu_telemetry".to_owned(), - watcher_hint: Some("gpu-thermal-protect".to_owned()), - service_id: None, - reason: Some("hotspot_temperature_threshold".to_owned()), - payload: json!({ - "summary": "GPU 0 hotspot temperature is 96 C (limit 95 C)", - }), - }; - - handle_gpu_thermal_protect_event(&paths, WatcherMode::Propose, &mut state, &event)?; - let event_text = fs::read_to_string(paths.automation_events_path())?; - let event = serde_json::from_str::(event_text.trim())?; - let proposals = rocm_core::load_recent_automation_proposals(&paths, 1)?; - fs::remove_dir_all(root).ok(); - - assert_eq!(event.action, "gpu_pressure_no_clear_target"); - assert!(event.message.contains("did not choose a model server")); - assert!(proposals.is_empty()); - Ok(()) - } - - #[test] - fn gpu_thermal_protect_does_not_duplicate_pending_stop_proposals() -> Result<()> { - let (root, paths) = temp_app_paths("gpu-thermal-protect-dedupe"); - paths.ensure()?; - let mut record = ManagedServiceRecord::new( - &paths, - "svc-hot", - "vllm", - "tiny", - "Tiny/Test", - "127.0.0.1", - 11435, - "managed", - 123, - None, - None, - None, - ); - record.status = "ready".to_owned(); - record.write()?; - let mut state = test_runtime_state(vec![test_watcher_snapshot( - "gpu-thermal-protect", - WatcherMode::Propose, - None, - )]); - let event = AutomationTriggerEvent { - at_unix_ms: 1, - kind: "gpu.thermal_pressure".to_owned(), - source: "gpu_telemetry".to_owned(), - watcher_hint: Some("gpu-thermal-protect".to_owned()), - service_id: Some("svc-hot".to_owned()), - reason: Some("hotspot_temperature_threshold".to_owned()), - payload: json!({ - "summary": "GPU 0 hotspot temperature is 96 C (limit 95 C)", - }), - }; - - handle_gpu_thermal_protect_event(&paths, WatcherMode::Propose, &mut state, &event)?; - handle_gpu_thermal_protect_event(&paths, WatcherMode::Propose, &mut state, &event)?; - let events = rocm_core::load_recent_automation_events(&paths, 2)?; - let proposals = rocm_core::load_recent_automation_proposals(&paths, 10)?; - fs::remove_dir_all(root).ok(); - - assert_eq!(proposals.len(), 1); - assert_eq!(proposals[0].tool.as_deref(), Some("stop_server")); - assert!( - events - .iter() - .any(|event| event.action == "stop_proposal_already_pending") - ); - Ok(()) - } - - #[test] - fn local_webhook_gpu_metrics_event_uses_existing_read_only_policy() -> Result<()> { - let (root, paths) = temp_app_paths("local-webhook-gpu-metrics"); - paths.ensure()?; - let mut state = test_runtime_state(vec![test_watcher_snapshot( - "gpu-metrics", - WatcherMode::Contained, - None, - )]); - let mut config = RocmCliConfig::default(); - let watcher = config.watcher_config_mut("gpu-metrics"); - watcher.enabled = true; - watcher.mode = Some(WatcherMode::Contained); - let event = webhook::local_webhook_event_from_request(webhook::LocalWebhookEventRequest { - watcher_hint: "gpu-metrics".to_owned(), - kind: "gpu.metrics".to_owned(), - service_id: None, - reason: Some("manual smoke".to_owned()), - payload: json!({ - "summary": "manual webhook probe", - "action": "restart_server", - "mode": "contained", - }), - })?; - - evaluate_watchers_for_events(&paths, &config, &mut state, &[event])?; - let event_text = fs::read_to_string(paths.automation_events_path())?; - let event = serde_json::from_str::(event_text.trim())?; - let proposals = rocm_core::load_recent_automation_proposals(&paths, 1)?; - fs::remove_dir_all(root).ok(); - - assert_eq!(event.watcher_id, "gpu-metrics"); - assert_eq!(event.action, "record_gpu_metrics"); - assert!(event.message.contains("from local webhook")); - assert!(event.message.contains("manual smoke")); - assert!(event.message.contains("no mutating action was taken")); - assert!(proposals.is_empty()); - Ok(()) - } - - #[test] - fn cache_warm_propose_mode_queues_prefetch_proposal() -> Result<()> { - let (root, paths) = temp_app_paths("cache-warm-propose"); - paths.ensure()?; - let mut state = test_runtime_state(vec![test_watcher_snapshot( - "cache-warm", - WatcherMode::Propose, - None, - )]); - let event = webhook::local_webhook_event_from_request(webhook::LocalWebhookEventRequest { - watcher_hint: "cache-warm".to_owned(), - kind: "cache.warm".to_owned(), - service_id: None, - reason: Some("idle window".to_owned()), - payload: json!({ - "artifact_ref": "Qwen/Test-1B#hf-main", - "tool": "restart_server", - "allow_artifact_download": true, - "artifact_max_bytes": 1024, - }), - })?; - - handle_cache_warm_event_with_resolver( - &paths, - WatcherMode::Propose, - &mut state, - &event, - |_| Ok(true), - )?; - let event_text = fs::read_to_string(paths.automation_events_path())?; - let event = serde_json::from_str::(event_text.trim())?; - let proposals = rocm_core::load_recent_automation_proposals(&paths, 1)?; - fs::remove_dir_all(root).ok(); - - assert_eq!(event.watcher_id, "cache-warm"); - assert_eq!(event.action, "queue_prefetch_proposal"); - assert_eq!(proposals.len(), 1); - assert_eq!(proposals[0].watcher_id, "cache-warm"); - assert_eq!(proposals[0].tool.as_deref(), Some("prefetch_artifact")); - assert_eq!( - proposals[0] - .arguments - .get("artifact_ref") - .and_then(Value::as_str), - Some("Qwen/Test-1B#hf-main") - ); - assert!( - proposals[0].arguments.get("tool").is_none(), - "webhook payload must not grant arbitrary tool choice" - ); - assert!( - proposals[0] - .arguments - .get("allow_artifact_download") - .is_none(), - "webhook payload must not grant source-policy approval" - ); - assert!( - proposals[0].arguments.get("artifact_max_bytes").is_none(), - "webhook payload must not grant download byte-limit approval" - ); - Ok(()) - } - - #[test] - fn cache_warm_unknown_artifact_does_not_queue_proposal() -> Result<()> { - let (root, paths) = temp_app_paths("cache-warm-unknown"); - paths.ensure()?; - let mut state = test_runtime_state(vec![test_watcher_snapshot( - "cache-warm", - WatcherMode::Propose, - None, - )]); - let event = AutomationTriggerEvent { - at_unix_ms: 1, - kind: "cache.warm".to_owned(), - source: "test".to_owned(), - watcher_hint: Some("cache-warm".to_owned()), - service_id: None, - reason: Some("idle window".to_owned()), - payload: json!({ - "artifact_ref": "missing#artifact", - }), - }; - - handle_cache_warm_event_with_resolver( - &paths, - WatcherMode::Propose, - &mut state, - &event, - |_| Ok(false), - )?; - let event_text = fs::read_to_string(paths.automation_events_path())?; - let event = serde_json::from_str::(event_text.trim())?; - let proposals = rocm_core::load_recent_automation_proposals(&paths, 1)?; - fs::remove_dir_all(root).ok(); - - assert_eq!(event.action, "cache_warm_unknown_artifact"); - assert!(event.message.contains("unknown registry artifact")); - assert!(proposals.is_empty()); - Ok(()) - } - - #[test] - fn cache_warm_contained_mode_still_requires_reviewed_source_policy() -> Result<()> { - let (root, paths) = temp_app_paths("cache-warm-contained"); - paths.ensure()?; - let mut state = test_runtime_state(vec![test_watcher_snapshot( - "cache-warm", - WatcherMode::Contained, - None, - )]); - let event = AutomationTriggerEvent { - at_unix_ms: 1, - kind: "cache.warm".to_owned(), - source: "test".to_owned(), - watcher_hint: Some("cache-warm".to_owned()), - service_id: None, - reason: Some("idle window".to_owned()), - payload: json!({ - "artifact_ref": "Qwen/Test-1B#hf-main", - }), - }; - - handle_cache_warm_event_with_resolver( - &paths, - WatcherMode::Contained, - &mut state, - &event, - |_| Ok(true), - )?; - let event_text = fs::read_to_string(paths.automation_events_path())?; - let event = serde_json::from_str::(event_text.trim())?; - let proposals = rocm_core::load_recent_automation_proposals(&paths, 1)?; - fs::remove_dir_all(root).ok(); - - assert_eq!(event.action, "queue_prefetch_proposal"); - assert!(event.message.contains("explicit source-policy approval")); - assert_eq!(proposals.len(), 1); - assert_eq!(proposals[0].tool.as_deref(), Some("prefetch_artifact")); - Ok(()) - } - - #[test] - fn driver_upgrade_propose_mode_queues_driver_plan_proposal() -> Result<()> { - let (root, paths) = temp_app_paths("driver-upgrade-propose"); - paths.ensure()?; - let mut state = test_runtime_state(vec![test_watcher_snapshot( - "driver-upgrade", - WatcherMode::Propose, - None, - )]); - let event = webhook::local_webhook_event_from_request(webhook::LocalWebhookEventRequest { - watcher_hint: "driver-upgrade".to_owned(), - kind: "update.available".to_owned(), - service_id: None, - reason: Some("driver version is newer".to_owned()), - payload: json!({ - "component": "driver", - "tool": "restart_server", - }), - })?; - - handle_driver_upgrade_event(&paths, WatcherMode::Propose, &mut state, &event)?; - let event_text = fs::read_to_string(paths.automation_events_path())?; - let event = serde_json::from_str::(event_text.trim())?; - let proposals = rocm_core::load_recent_automation_proposals(&paths, 1)?; - fs::remove_dir_all(root).ok(); - - assert_eq!(event.watcher_id, "driver-upgrade"); - assert_eq!(event.action, "prepare_driver_plan"); - assert_eq!(proposals.len(), 1); - assert_eq!(proposals[0].watcher_id, "driver-upgrade"); - assert_eq!(proposals[0].tool.as_deref(), Some("driver_plan")); - assert!( - proposals[0].arguments.get("tool").is_none(), - "webhook payload must not grant arbitrary tool choice" - ); - Ok(()) - } - - #[test] - fn driver_upgrade_contained_mode_runs_restricted_driver_plan() -> Result<()> { - let (root, paths) = temp_app_paths("driver-upgrade-contained"); - paths.ensure()?; - let mut state = test_runtime_state(vec![test_watcher_snapshot( - "driver-upgrade", - WatcherMode::Contained, - None, - )]); - let event = AutomationTriggerEvent { - at_unix_ms: 1, - kind: "update.available".to_owned(), - source: "test".to_owned(), - watcher_hint: Some("driver-upgrade".to_owned()), - service_id: None, - reason: Some("driver version is newer".to_owned()), - payload: json!({ - "component": "driver", - }), - }; - - handle_driver_upgrade_event_with_runner( - &paths, - WatcherMode::Contained, - &mut state, - &event, - |_paths| { - Ok(sandbox::sandbox_driver_plan_value(common::CommandCapture { - argv: vec![ - "rocm".to_owned(), - "install".to_owned(), - "driver".to_owned(), - "--dkms".to_owned(), - "--dry-run".to_owned(), - ], - exit_status: 0, - stdout: "driver install plan\n supported: true\n".to_owned(), - stderr: String::new(), - })) - }, - )?; - let event_text = fs::read_to_string(paths.automation_events_path())?; - let event = serde_json::from_str::(event_text.trim())?; - let proposals = rocm_core::load_recent_automation_proposals(&paths, 1)?; - fs::remove_dir_all(root).ok(); - - assert_eq!(event.action, "run_driver_plan"); - assert!( - event - .message - .contains("contained restricted driver_plan status=planned") - ); - assert!(event.message.contains("no driver commands were executed")); - assert!(proposals.is_empty()); - Ok(()) - } - - #[test] - fn driver_upgrade_contained_mode_requires_restricted_driver_plan_tool() -> Result<()> { - let (root, paths) = temp_app_paths("driver-upgrade-contained-tool"); - paths.ensure()?; - let mut state = test_runtime_state(vec![test_watcher_snapshot( - "driver-upgrade", - WatcherMode::Contained, - None, - )]); - let event = AutomationTriggerEvent { - at_unix_ms: 1, - kind: "update.available".to_owned(), - source: "test".to_owned(), - watcher_hint: Some("driver-upgrade".to_owned()), - service_id: None, - reason: Some("driver version is newer".to_owned()), - payload: json!({ - "component": "driver", - }), - }; - - handle_driver_upgrade_event_with_runner( - &paths, - WatcherMode::Contained, - &mut state, - &event, - |_paths| { - Ok(json!({ - "tool": "check_updates", - "status": "checked", - "mutating": false, - })) - }, - )?; - let event_text = fs::read_to_string(paths.automation_events_path())?; - let event = serde_json::from_str::(event_text.trim())?; - let proposals = rocm_core::load_recent_automation_proposals(&paths, 1)?; - fs::remove_dir_all(root).ok(); - - assert_eq!(event.action, "driver_plan_failed"); - assert!(event.message.contains("expected `driver_plan`")); - assert!(event.message.contains("no driver commands were executed")); - assert!(proposals.is_empty()); - Ok(()) - } - - #[test] - fn driver_upgrade_observe_mode_records_without_proposal() -> Result<()> { - let (root, paths) = temp_app_paths("driver-upgrade-observe"); - paths.ensure()?; - let mut state = test_runtime_state(vec![test_watcher_snapshot( - "driver-upgrade", - WatcherMode::Observe, - None, - )]); - let event = AutomationTriggerEvent { - at_unix_ms: 1, - kind: "update.available".to_owned(), - source: "test".to_owned(), - watcher_hint: Some("driver-upgrade".to_owned()), - service_id: None, - reason: Some("driver version is newer".to_owned()), - payload: json!({ - "component": "driver", - }), - }; - - handle_driver_upgrade_event(&paths, WatcherMode::Observe, &mut state, &event)?; - let event_text = fs::read_to_string(paths.automation_events_path())?; - let event = serde_json::from_str::(event_text.trim())?; - let proposals = rocm_core::load_recent_automation_proposals(&paths, 1)?; - fs::remove_dir_all(root).ok(); - - assert_eq!(event.action, "observe_driver_update"); - assert!(event.message.contains("does not queue or run")); - assert!(proposals.is_empty()); - Ok(()) - } - - #[test] - fn driver_upgrade_ignores_non_driver_component_without_proposal() -> Result<()> { - let (root, paths) = temp_app_paths("driver-upgrade-wrong-component"); - paths.ensure()?; - let mut state = test_runtime_state(vec![test_watcher_snapshot( - "driver-upgrade", - WatcherMode::Propose, - None, - )]); - let event = AutomationTriggerEvent { - at_unix_ms: 1, - kind: "update.available".to_owned(), - source: "test".to_owned(), - watcher_hint: Some("driver-upgrade".to_owned()), - service_id: None, - reason: Some("runtime version is newer".to_owned()), - payload: json!({ - "component": "runtime", - }), - }; - - handle_driver_upgrade_event(&paths, WatcherMode::Propose, &mut state, &event)?; - let event_text = fs::read_to_string(paths.automation_events_path())?; - let event = serde_json::from_str::(event_text.trim())?; - let proposals = rocm_core::load_recent_automation_proposals(&paths, 1)?; - fs::remove_dir_all(root).ok(); - - assert_eq!(event.action, "driver_upgrade_ignored_component"); - assert!(event.message.contains("payload.component=driver")); - assert!(proposals.is_empty()); - Ok(()) - } - - #[test] - fn event_dispatcher_preserves_server_recover_proposal_behavior() -> Result<()> { - let (root, paths) = temp_app_paths("event-bus-dispatch"); - paths.ensure()?; - let mut failed = ManagedServiceRecord::new( - &paths, - "svc-failed", - "vllm", - "qwen", - "Qwen/Qwen3.5", - "127.0.0.1", - 11435, - "managed", - 123, - None, - None, - None, - ); - failed.status = "failed".to_owned(); - failed.write()?; - let mut state = test_runtime_state(vec![test_watcher_snapshot( - "server-recover", - WatcherMode::Propose, - None, - )]); - let mut config = RocmCliConfig::default(); - let watcher = config.watcher_config_mut("server-recover"); - watcher.enabled = true; - watcher.mode = Some(WatcherMode::Propose); - let events = vec![AutomationTriggerEvent { - at_unix_ms: 1, - kind: "service.manifest_recoverable".to_owned(), - source: "managed_service".to_owned(), - watcher_hint: Some("server-recover".to_owned()), - service_id: Some("svc-failed".to_owned()), - reason: Some("manifest_status_failed".to_owned()), - payload: json!({}), - }]; - - evaluate_watchers_for_events(&paths, &config, &mut state, &events)?; - let proposals = rocm_core::load_recent_automation_proposals(&paths, 1)?; - fs::remove_dir_all(root).ok(); - - assert_eq!(proposals.len(), 1); - assert_eq!(proposals[0].watcher_id, "server-recover"); - assert_eq!(proposals[0].service_id.as_deref(), Some("svc-failed")); - assert_eq!(proposals[0].tool.as_deref(), Some("restart_server")); - assert!(proposals[0].message.contains("manifest reports failed")); - assert!(!proposals[0].message.contains("manifest_status_failed")); - Ok(()) - } - - #[test] - fn server_recover_local_webhook_does_not_restart_healthy_service() -> Result<()> { - let (root, paths) = temp_app_paths("server-recover-healthy-webhook"); - paths.ensure()?; - let mut healthy = ManagedServiceRecord::new( - &paths, - "svc-healthy", - "vllm", - "qwen", - "Qwen/Qwen3.5", - "127.0.0.1", - 11435, - "managed", - 123, - None, - None, - None, - ); - healthy.status = "ready".to_owned(); - healthy.write()?; - let mut state = test_runtime_state(vec![test_watcher_snapshot( - "server-recover", - WatcherMode::Propose, - None, - )]); - let mut config = RocmCliConfig::default(); - let watcher = config.watcher_config_mut("server-recover"); - watcher.enabled = true; - watcher.mode = Some(WatcherMode::Propose); - let event = webhook::local_webhook_event_from_request(webhook::LocalWebhookEventRequest { - watcher_hint: "server-recover".to_owned(), - kind: "service.manifest_recoverable".to_owned(), - service_id: Some("svc-healthy".to_owned()), - reason: Some("manual recovery smoke".to_owned()), - payload: json!({}), - })?; - - evaluate_watchers_for_events(&paths, &config, &mut state, &[event])?; - let event_text = fs::read_to_string(paths.automation_events_path())?; - let event = serde_json::from_str::(event_text.trim())?; - let proposals = rocm_core::load_recent_automation_proposals(&paths, 1)?; - let reloaded = load_service_record(&paths, "svc-healthy")?; - fs::remove_dir_all(root).ok(); - - assert_eq!(event.action, "ignore_nonrecoverable_service"); - assert!(event.message.contains("does not currently need recovery")); - assert!(proposals.is_empty()); - assert_eq!(reloaded.status, "ready"); - Ok(()) - } - - #[test] - fn recovery_reason_display_avoids_raw_status_tokens() { - assert_eq!( - display_recovery_reason("manifest_status_starting_stale"), - "service has been starting for too long" - ); - assert_eq!( - display_recovery_reason("healthcheck_status_unreachable"), - "engine healthcheck reports unreachable" - ); - assert_eq!( - display_recovery_reason("endpoint_status_unreachable"), - "endpoint port is unreachable" - ); - } - - #[test] - fn manifest_recovery_policy_covers_terminal_and_stale_transient_states() { - let (root, paths) = temp_app_paths("manifest-recovery-policy"); - let mut record = ManagedServiceRecord::new( - &paths, - "svc-stale", - "vllm", - "qwen", - "Qwen/Qwen3.5", - "127.0.0.1", - 11435, - "managed", - 123, - None, - None, - None, - ); - record.created_at_unix_ms = 1_000; - - record.status = "exited".to_owned(); - assert_eq!( - manifest_service_recovery_reason(&record, 1_001).as_deref(), - Some("manifest_status_exited") - ); - - record.status = "unreachable".to_owned(); - assert_eq!( - manifest_service_recovery_reason(&record, 1_001).as_deref(), - Some("manifest_status_unreachable") - ); - - record.status = "starting".to_owned(); - assert_eq!(manifest_service_recovery_reason(&record, 2_000), None); - assert_eq!( - manifest_service_recovery_reason(&record, 1_000 + SERVER_TRANSIENT_STALE_MS).as_deref(), - Some("manifest_status_starting_stale") - ); - - record.status = "recovering".to_owned(); - record.last_restart_unix_ms = Some(5_000); - assert_eq!(manifest_service_recovery_reason(&record, 6_000), None); - assert_eq!( - manifest_service_recovery_reason(&record, 5_000 + SERVER_TRANSIENT_STALE_MS).as_deref(), - Some("manifest_status_recovering_stale") - ); - fs::remove_dir_all(root).ok(); - } - - #[test] - fn find_recoverable_service_prefers_failed_managed_manifest() -> Result<()> { - let (root, paths) = temp_app_paths("recoverable-service"); - paths.ensure()?; - let mut failed = ManagedServiceRecord::new( - &paths, - "svc-failed", - "vllm", - "qwen", - "Qwen/Qwen3.5", - "127.0.0.1", - 11435, - "managed", - 123, - None, - None, - None, - ); - failed.status = "failed".to_owned(); - failed.write()?; - - let found = find_recoverable_service(&paths)?.expect("failed service should be found"); - fs::remove_dir_all(root).ok(); - assert_eq!(found.0.service_id, "svc-failed"); - assert_eq!(found.1, "manifest_status_failed"); - Ok(()) - } - - #[test] - fn find_recoverable_service_detects_stale_starting_manifest() -> Result<()> { - let (root, paths) = temp_app_paths("recoverable-stale-starting"); - paths.ensure()?; - let mut stale = ManagedServiceRecord::new( - &paths, - "svc-starting", - "vllm", - "qwen", - "Qwen/Qwen3.5", - "127.0.0.1", - 11435, - "managed", - 123, - None, - None, - None, - ); - stale.status = "starting".to_owned(); - stale.created_at_unix_ms = 0; - stale.write()?; - - let found = - find_recoverable_service(&paths)?.expect("stale starting service should recover"); - fs::remove_dir_all(root).ok(); - assert_eq!(found.0.service_id, "svc-starting"); - assert_eq!(found.1, "manifest_status_starting_stale"); - Ok(()) - } - - #[test] - fn server_recover_propose_mode_queues_restart_proposal() -> Result<()> { - let (root, paths) = temp_app_paths("server-recover-proposal"); - paths.ensure()?; - let mut record = ManagedServiceRecord::new( - &paths, - "svc-1", - "vllm", - "qwen", - "Qwen/Qwen3.5", - "127.0.0.1", - 11435, - "managed", - 123, - None, - None, - None, - ); - record.status = "failed".to_owned(); - record.write()?; - let mut state = AutomationRuntimeState { - running: true, - automations_enabled: true, - daemon_pid: 1, - started_at_unix_ms: 1, - last_tick_unix_ms: 1, - local_webhook_endpoint: None, - active_watchers: vec![WatcherRuntimeSnapshot { - id: "server-recover".to_owned(), - enabled: true, - mode: WatcherMode::Propose, - summary: "recover".to_owned(), - last_event: None, - last_event_unix_ms: None, - }], - }; - - evaluate_server_recover(&paths, WatcherMode::Propose, &mut state)?; - let proposals = rocm_core::load_recent_automation_proposals(&paths, 1)?; - fs::remove_dir_all(root).ok(); - - assert_eq!(proposals.len(), 1); - assert_eq!(proposals[0].watcher_id, "server-recover"); - assert_eq!(proposals[0].action, "queue_restart_proposal"); - assert_eq!(proposals[0].service_id.as_deref(), Some("svc-1")); - assert_eq!(proposals[0].status, "pending"); - Ok(()) - } - - #[test] - fn therock_update_contained_mode_runs_read_only_check_without_queueing() -> Result<()> { - let (root, paths) = temp_app_paths("therock-update-contained"); - paths.ensure()?; - let mut state = AutomationRuntimeState { - running: true, - automations_enabled: true, - daemon_pid: 1, - started_at_unix_ms: 1, - last_tick_unix_ms: 1, - local_webhook_endpoint: None, - active_watchers: vec![WatcherRuntimeSnapshot { - id: "therock-update".to_owned(), - enabled: true, - mode: WatcherMode::Contained, - summary: "check updates".to_owned(), - last_event: None, - last_event_unix_ms: None, - }], - }; - let event = AutomationTriggerEvent { - at_unix_ms: 42, - kind: "schedule.tick".to_owned(), - source: "scheduler".to_owned(), - watcher_hint: Some("therock-update".to_owned()), - service_id: None, - reason: Some("therock_update_interval_due".to_owned()), - payload: json!({ "interval_ms": THEROCK_UPDATE_INTERVAL_MS }), - }; - - handle_therock_update_event_with_runner( - &paths, - WatcherMode::Contained, - &mut state, - &event, - |_paths| { - Ok(sandbox::sandbox_check_updates_value( - common::CommandCapture { - argv: vec!["rocm".to_owned(), "update".to_owned()], - exit_status: 0, - stdout: "update\n managed runtimes: none\n".to_owned(), - stderr: String::new(), - }, - )) - }, - )?; - - let event_text = fs::read_to_string(paths.automation_events_path())?; - let event = serde_json::from_str::(event_text.trim())?; - let proposals = rocm_core::load_recent_automation_proposals(&paths, 1)?; - fs::remove_dir_all(root).ok(); - - assert_eq!(event.watcher_id, "therock-update"); - assert_eq!(event.action, "run_update_check"); - assert!(event.message.contains("contained read-only execution")); - assert!( - event - .message - .contains("restricted check_updates status=checked") - ); - assert!(event.message.contains("no updates were applied")); - assert!(!event.message.contains("fallback")); - assert!(proposals.is_empty()); - Ok(()) - } - - #[test] - fn therock_update_contained_mode_records_update_available_without_applying() -> Result<()> { - let (root, paths) = temp_app_paths("therock-update-contained-available"); - paths.ensure()?; - let mut state = AutomationRuntimeState { - running: true, - automations_enabled: true, - daemon_pid: 1, - started_at_unix_ms: 1, - last_tick_unix_ms: 1, - local_webhook_endpoint: None, - active_watchers: vec![WatcherRuntimeSnapshot { - id: "therock-update".to_owned(), - enabled: true, - mode: WatcherMode::Contained, - summary: "check updates".to_owned(), - last_event: None, - last_event_unix_ms: None, - }], - }; - let event = AutomationTriggerEvent { - at_unix_ms: 42, - kind: "schedule.tick".to_owned(), - source: "scheduler".to_owned(), - watcher_hint: Some("therock-update".to_owned()), - service_id: None, - reason: Some("therock_update_interval_due".to_owned()), - payload: json!({ "interval_ms": THEROCK_UPDATE_INTERVAL_MS }), - }; - - handle_therock_update_event_with_runner( - &paths, - WatcherMode::Contained, - &mut state, - &event, - |_paths| { - Ok(sandbox::sandbox_check_updates_value(common::CommandCapture { - argv: vec!["rocm".to_owned(), "update".to_owned()], - exit_status: 0, - stdout: "update\n runtime release-pip-gfx120x-all status=update_available installed=7.13.0 latest=7.14.0\n".to_owned(), - stderr: String::new(), - })) - }, - )?; - - let event_text = fs::read_to_string(paths.automation_events_path())?; - let events = event_text - .lines() - .map(serde_json::from_str::) - .collect::, _>>()?; - let audit_text = fs::read_to_string(paths.audit_events_path())?; - let proposals = rocm_core::load_recent_automation_proposals(&paths, 1)?; - fs::remove_dir_all(root).ok(); - - let update_check = events - .iter() - .find(|event| event.action == "run_update_check") - .expect("update check event should be recorded"); - assert_eq!(update_check.watcher_id, "therock-update"); - assert!( - update_check - .message - .contains("restricted check_updates status=update_available") - ); - assert!( - update_check - .message - .contains("a ROCm runtime update is available") - ); - assert!(update_check.message.contains("no updates were applied")); - assert!(!update_check.message.contains("fallback")); - let notification = events - .iter() - .find(|event| event.action == "notify_if_newer") - .expect("notify-if-newer event should be recorded"); - assert_eq!(notification.watcher_id, "therock-update"); - assert!( - notification - .message - .contains("ROCm runtime update is available") - ); - assert!(notification.message.contains("No updates were applied")); - assert!(audit_text.contains("\"category\":\"notification\"")); - assert!(audit_text.contains("\"action\":\"notify_if_newer\"")); - assert!(audit_text.contains("ROCm runtime update is available")); - assert!(proposals.is_empty()); - Ok(()) - } - - #[test] - fn therock_update_contained_mode_uses_restricted_check_updates_tool() -> Result<()> { - let (root, paths) = temp_app_paths("therock-update-contained-tool"); - paths.ensure()?; - let mut state = AutomationRuntimeState { - running: true, - automations_enabled: true, - daemon_pid: 1, - started_at_unix_ms: 1, - last_tick_unix_ms: 1, - local_webhook_endpoint: None, - active_watchers: vec![WatcherRuntimeSnapshot { - id: "therock-update".to_owned(), - enabled: true, - mode: WatcherMode::Contained, - summary: "check updates".to_owned(), - last_event: None, - last_event_unix_ms: None, - }], - }; - let event = AutomationTriggerEvent { - at_unix_ms: 42, - kind: "schedule.tick".to_owned(), - source: "scheduler".to_owned(), - watcher_hint: Some("therock-update".to_owned()), - service_id: None, - reason: Some("therock_update_interval_due".to_owned()), - payload: json!({ "interval_ms": THEROCK_UPDATE_INTERVAL_MS }), - }; - - handle_therock_update_event_with_runner( - &paths, - WatcherMode::Contained, - &mut state, - &event, - |_paths| { - Ok(json!({ - "tool": "examine_snapshot", - "status": "captured", - "mutating": false, - })) - }, - )?; - - let event_text = fs::read_to_string(paths.automation_events_path())?; - let event = serde_json::from_str::(event_text.trim())?; - fs::remove_dir_all(root).ok(); - - assert_eq!(event.watcher_id, "therock-update"); - assert_eq!(event.action, "update_check_failed"); - assert!(event.message.contains("expected `check_updates`")); - assert!(event.message.contains("no updates were applied")); - Ok(()) - } - - #[test] - fn therock_update_notify_if_newer_uses_restricted_notification_contract() -> Result<()> { - let (root, paths) = temp_app_paths("therock-update-notify-contract"); - paths.ensure()?; - let mut state = AutomationRuntimeState { - running: true, - automations_enabled: true, - daemon_pid: 1, - started_at_unix_ms: 1, - last_tick_unix_ms: 1, - local_webhook_endpoint: None, - active_watchers: vec![WatcherRuntimeSnapshot { - id: "therock-update".to_owned(), - enabled: true, - mode: WatcherMode::Contained, - summary: "check updates".to_owned(), - last_event: None, - last_event_unix_ms: None, - }], - }; - - record_update_available_notification(&paths, &mut state, "update_available")?; +use std::time::Duration; - let audit_text = fs::read_to_string(paths.audit_events_path())?; - let audit = audit_text - .lines() - .map(serde_json::from_str::) - .collect::, _>>()?; - fs::remove_dir_all(root).ok(); +const WATCHER_TICK_INTERVAL: Duration = Duration::from_secs(5); +const ARTIFACT_PREFETCH_TIMEOUT: Duration = Duration::from_mins(10); - let notification = audit - .iter() - .find(|event| event.category == "notification" && event.action == "notify_if_newer") - .expect("notify_if_newer audit should be recorded"); - assert_eq!(notification.category, "notification"); - assert_eq!(notification.actor, "watcher:therock-update"); - assert_eq!(notification.watcher_id.as_deref(), Some("therock-update")); - assert!( - notification - .message - .contains("ROCm runtime update is available") - ); - assert!(notification.message.contains("No updates were applied")); - Ok(()) - } +#[cfg(test)] +mod tests { + use super::*; + use crate::test_support::temp_app_paths; + use anyhow::Result; + use rocm_core::ManagedServiceRecord; + use serde_json::Value; + use std::fs; #[test] fn sandbox_tool_stop_server_updates_manifest_and_skips_current_pid() -> Result<()> { @@ -3132,7 +62,7 @@ mod tests { None, cli::SandboxToolPolicy::default(), )?; - let reloaded = load_service_record(&paths, "svc-current")?; + let reloaded = watchers::load_service_record(&paths, "svc-current")?; fs::remove_dir_all(root).ok(); assert_eq!(value.get("status").and_then(Value::as_str), Some("stopped")); @@ -3149,31 +79,4 @@ mod tests { ); Ok(()) } - - fn test_watcher_snapshot( - id: &str, - mode: WatcherMode, - last_event_unix_ms: Option, - ) -> WatcherRuntimeSnapshot { - WatcherRuntimeSnapshot { - id: id.to_owned(), - enabled: true, - mode, - summary: "test watcher".to_owned(), - last_event: None, - last_event_unix_ms, - } - } - - fn test_runtime_state(active_watchers: Vec) -> AutomationRuntimeState { - AutomationRuntimeState { - running: true, - automations_enabled: true, - daemon_pid: 1, - started_at_unix_ms: 1, - last_tick_unix_ms: 1, - local_webhook_endpoint: None, - active_watchers, - } - } } diff --git a/apps/rocmd/src/sandbox.rs b/apps/rocmd/src/sandbox.rs index 6d67801e2..d11402aa2 100644 --- a/apps/rocmd/src/sandbox.rs +++ b/apps/rocmd/src/sandbox.rs @@ -2,11 +2,12 @@ // // SPDX-License-Identifier: MIT +use crate::ARTIFACT_PREFETCH_TIMEOUT; use crate::cli::{SandboxToolArg, SandboxToolPolicy}; use crate::common::{self, CommandCapture}; use crate::persistence::load_managed_services; use crate::service::stop_managed_service; -use crate::{ARTIFACT_PREFETCH_TIMEOUT, restart_managed_service}; +use crate::watchers::restart_managed_service; use anyhow::{Context, Result, bail}; use rocm_core::{ AppPaths, AuditEventRecord, ExamineSummary, ModelRecipeArtifactRecord, append_audit_event, diff --git a/apps/rocmd/src/service.rs b/apps/rocmd/src/service.rs index 3d1f4eeaa..eb5f2b2bf 100644 --- a/apps/rocmd/src/service.rs +++ b/apps/rocmd/src/service.rs @@ -343,7 +343,7 @@ pub(crate) async fn run_daemon( )?; state.write(paths)?; - crate::evaluate_watchers(paths, &config, &mut state)?; + crate::watchers::evaluate_watchers(paths, &config, &mut state)?; state.last_tick_unix_ms = unix_time_millis(); state.write(paths)?; @@ -357,15 +357,15 @@ pub(crate) async fn run_daemon( tokio::select! { _ = ticker.tick() => { let config = RocmCliConfig::load(paths)?; - crate::reconcile_watcher_snapshots(&config, &mut state); - crate::evaluate_watchers(paths, &config, &mut state)?; + crate::watchers::reconcile_watcher_snapshots(&config, &mut state); + crate::watchers::evaluate_watchers(paths, &config, &mut state)?; state.last_tick_unix_ms = unix_time_millis(); state.write(paths)?; } event = crate::webhook::receive_local_webhook_event(&mut local_webhook_receiver) => { if let Some(event) = event { let config = RocmCliConfig::load(paths)?; - crate::reconcile_watcher_snapshots(&config, &mut state); + crate::watchers::reconcile_watcher_snapshots(&config, &mut state); crate::persistence::record_event( paths, &mut state, @@ -380,7 +380,7 @@ pub(crate) async fn run_daemon( event.service_id.clone(), )?; if let Err(error) = - crate::evaluate_watchers_for_events(paths, &config, &mut state, &[event]) + crate::watchers::evaluate_watchers_for_events(paths, &config, &mut state, &[event]) { crate::persistence::record_event( paths, @@ -1313,7 +1313,7 @@ mod tests { assert!(!key_path.exists()); let result = stop_managed_service(&paths, service_id); - let reloaded = crate::load_service_record(&paths, service_id); + let reloaded = crate::watchers::load_service_record(&paths, service_id); fs::remove_dir_all(root).ok(); let value = result?; diff --git a/apps/rocmd/src/watchers.rs b/apps/rocmd/src/watchers.rs new file mode 100644 index 000000000..bd3916b77 --- /dev/null +++ b/apps/rocmd/src/watchers.rs @@ -0,0 +1,3144 @@ +// Copyright © Advanced Micro Devices, Inc., or its affiliates. +// +// SPDX-License-Identifier: MIT + +use anyhow::{Context, Result, bail}; +use rocm_core::{ + AppPaths, AutomationProposalRecord, AutomationRuntimeState, AutomationTriggerEvent, + CodexBridgeGpuSnapshot, ManagedServiceRecord, RocmCliConfig, WatcherMode, + WatcherRuntimeSnapshot, append_automation_proposal, builtin_watchers, + resolve_model_recipe_artifact, unix_time_millis, +}; +use serde_json::Value; +use serde_json::json; +use std::fs; +use std::process::{Command as ProcessCommand, Stdio}; +use std::thread; +use std::time::Duration; + +const SERVER_RECOVER_BACKOFF_MS: u128 = 30_000; +const SERVER_TRANSIENT_STALE_MS: u128 = 5 * 60 * 1_000; +const ENDPOINT_HEALTH_TIMEOUT: Duration = Duration::from_millis(250); +const THEROCK_UPDATE_INTERVAL_MS: u128 = 6 * 60 * 60 * 1000; +const GPU_METRICS_INTERVAL_MS: u128 = 60 * 1000; +const GPU_THERMAL_HOTSPOT_PRESSURE_C: f64 = 95.0; +const GPU_THERMAL_MEMORY_PRESSURE_C: f64 = 95.0; +const GPU_MEMORY_VRAM_PRESSURE_PERCENT: f64 = 95.0; + +pub(crate) fn reconcile_watcher_snapshots( + config: &RocmCliConfig, + state: &mut AutomationRuntimeState, +) { + for watcher in builtin_watchers() { + match state.watcher_mut(watcher.id) { + Some(snapshot) => { + snapshot.enabled = config.watcher_enabled(watcher); + snapshot.mode = config.effective_watcher_mode(watcher); + snapshot.summary = watcher.summary.to_owned(); + } + None => state.active_watchers.push(WatcherRuntimeSnapshot { + id: watcher.id.to_owned(), + enabled: config.watcher_enabled(watcher), + mode: config.effective_watcher_mode(watcher), + summary: watcher.summary.to_owned(), + last_event: None, + last_event_unix_ms: None, + }), + } + } +} + +pub(crate) fn evaluate_watchers( + paths: &AppPaths, + config: &RocmCliConfig, + state: &mut AutomationRuntimeState, +) -> Result<()> { + let events = collect_automation_events(paths, config, state)?; + evaluate_watchers_for_events(paths, config, state, &events) +} + +fn collect_automation_events( + paths: &AppPaths, + config: &RocmCliConfig, + state: &AutomationRuntimeState, +) -> Result> { + collect_automation_events_with_gpu_snapshot(paths, state, || { + crate::common::gather_gpu_snapshot_for_config(config) + }) +} + +fn collect_automation_events_with_gpu_snapshot( + paths: &AppPaths, + state: &AutomationRuntimeState, + gpu_snapshot: F, +) -> Result> +where + F: FnMut() -> CodexBridgeGpuSnapshot, +{ + let now = unix_time_millis(); + let mut events = Vec::new(); + + if therock_update_due(state, now) { + events.push(AutomationTriggerEvent { + at_unix_ms: now, + kind: "schedule.tick".to_owned(), + source: "scheduler".to_owned(), + watcher_hint: Some("therock-update".to_owned()), + service_id: None, + reason: Some("therock_update_interval_due".to_owned()), + payload: json!({ + "interval_ms": THEROCK_UPDATE_INTERVAL_MS, + }), + }); + } + + if server_recover_due(state, now) + && let Some((record, recovery_reason)) = find_recoverable_service(paths)? + { + let kind = service_recovery_event_kind(&recovery_reason); + events.push(AutomationTriggerEvent { + at_unix_ms: now, + kind: kind.to_owned(), + source: "managed_service".to_owned(), + watcher_hint: Some("server-recover".to_owned()), + service_id: Some(record.service_id.clone()), + reason: Some(recovery_reason.clone()), + payload: json!({ + "engine": record.engine, + "status": record.status, + "endpoint": record.endpoint_url, + "recovery_reason": recovery_reason, + }), + }); + } + + let gpu_metrics_due_now = gpu_metrics_due(state, now); + let gpu_thermal_protect_due_now = gpu_thermal_protect_due(state, now); + let snapshot = (gpu_metrics_due_now || gpu_thermal_protect_due_now).then(gpu_snapshot); + + if gpu_metrics_due_now { + let snapshot = snapshot + .as_ref() + .expect("GPU snapshot should be collected for due metrics"); + let available = snapshot.amd_smi_available && snapshot.monitor_snapshot.is_some(); + events.push(AutomationTriggerEvent { + at_unix_ms: now, + kind: if available { + "gpu.metrics".to_owned() + } else { + "gpu.metrics_unavailable".to_owned() + }, + source: "gpu_telemetry".to_owned(), + watcher_hint: Some("gpu-metrics".to_owned()), + service_id: None, + reason: if available { + Some("amd_smi_snapshot_available".to_owned()) + } else { + snapshot + .note + .clone() + .or_else(|| Some("amd_smi_snapshot_unavailable".to_owned())) + }, + payload: json!({ + "amd_smi_available": snapshot.amd_smi_available, + "static_available": snapshot.static_snapshot.is_some(), + "monitor_available": snapshot.monitor_snapshot.is_some(), + "summary": gpu_snapshot_summary(snapshot), + "interval_ms": GPU_METRICS_INTERVAL_MS, + }), + }); + } + + if gpu_thermal_protect_due_now && let Some(snapshot) = snapshot.as_ref() { + events.extend(gpu_pressure_events(now, snapshot)); + } + + Ok(events) +} + +pub(crate) fn evaluate_watchers_for_events( + paths: &AppPaths, + config: &RocmCliConfig, + state: &mut AutomationRuntimeState, + events: &[AutomationTriggerEvent], +) -> Result<()> { + for watcher in builtin_watchers() { + if !config.watcher_enabled(watcher) { + continue; + } + let mode = config.effective_watcher_mode(watcher); + match watcher.id { + "therock-update" => { + for event in events_for_watcher(events, watcher.id, "schedule.tick") { + handle_therock_update_event(paths, mode, state, event)?; + } + } + "server-recover" => { + for event in events_for_watcher(events, watcher.id, "service.") { + handle_server_recover_event(paths, mode, state, event)?; + } + } + "gpu-metrics" => { + for event in events_for_watcher(events, watcher.id, "gpu.") { + handle_gpu_metrics_event(paths, mode, state, event)?; + } + } + "gpu-thermal-protect" => { + for event in + events_for_watcher_exact(events, watcher.id, "gpu.thermal_pressure").chain( + events_for_watcher_exact(events, watcher.id, "gpu.memory_pressure"), + ) + { + handle_gpu_thermal_protect_event(paths, mode, state, event)?; + } + } + "cache-warm" => { + for event in events_for_watcher_exact(events, watcher.id, "cache.warm") { + handle_cache_warm_event(paths, mode, state, event)?; + } + } + "driver-upgrade" => { + for event in events_for_watcher_exact(events, watcher.id, "update.available") { + handle_driver_upgrade_event(paths, mode, state, event)?; + } + } + _ => {} + } + } + Ok(()) +} + +fn events_for_watcher<'a>( + events: &'a [AutomationTriggerEvent], + watcher_id: &str, + kind_prefix: &str, +) -> impl Iterator { + events.iter().filter(move |event| { + event.watcher_hint.as_deref() == Some(watcher_id) && event.kind.starts_with(kind_prefix) + }) +} + +fn events_for_watcher_exact<'a>( + events: &'a [AutomationTriggerEvent], + watcher_id: &str, + kind: &'static str, +) -> impl Iterator { + events.iter().filter(move |event| { + event.watcher_hint.as_deref() == Some(watcher_id) && event.kind == kind + }) +} + +fn handle_therock_update_event( + paths: &AppPaths, + mode: WatcherMode, + state: &mut AutomationRuntimeState, + event: &AutomationTriggerEvent, +) -> Result<()> { + handle_therock_update_event_with_runner(paths, mode, state, event, |paths| { + crate::sandbox::run_sandbox_tool( + paths, + crate::cli::SandboxToolArg::CheckUpdates, + None, + None, + None, + crate::cli::SandboxToolPolicy::default(), + ) + }) +} + +fn handle_therock_update_event_with_runner( + paths: &AppPaths, + mode: WatcherMode, + state: &mut AutomationRuntimeState, + _event: &AutomationTriggerEvent, + update_runner: F, +) -> Result<()> +where + F: FnOnce(&AppPaths) -> Result, +{ + let policy = watcher_policy_action("therock-update", mode); + let action = match policy { + WatcherPolicyAction::Observe => "observe_schedule", + WatcherPolicyAction::QueueProposal => "queue_update_proposal", + WatcherPolicyAction::RunContained => "run_update_check", + }; + let message = match policy { + WatcherPolicyAction::Observe => { + "scheduled TheRock update check reminder emitted; run `rocm update` to inspect the selected channel" + } + WatcherPolicyAction::QueueProposal => { + "scheduled TheRock update check reminder emitted; queueing read-only update-check proposal for review" + } + WatcherPolicyAction::RunContained => { + "scheduled TheRock update check is approved for contained read-only execution" + } + }; + match policy { + WatcherPolicyAction::Observe | WatcherPolicyAction::QueueProposal => { + crate::persistence::record_event( + paths, + state, + "therock-update", + "info", + action, + message, + None, + )?; + if policy == WatcherPolicyAction::QueueProposal { + queue_proposal( + paths, + "therock-update", + action, + "Check TheRock updates", + "Run `rocm update` to inspect available CLI, runtime, engine, and recipe updates before applying changes.", + None, + )?; + } + } + WatcherPolicyAction::RunContained => match update_runner(paths) { + Ok(output) => match restricted_check_updates_result(&output) { + Ok(result) => { + crate::persistence::record_event( + paths, + state, + "therock-update", + if result.exit_status == 0 { + "info" + } else { + "error" + }, + action, + &format!( + "{message}; restricted check_updates status={}; {}", + result.status, + crate::common::update_check_message(result.status) + ), + None, + )?; + if result.update_available { + record_update_available_notification(paths, state, result.status)?; + } + } + Err(error) => { + crate::persistence::record_event( + paths, + state, + "therock-update", + "error", + "update_check_failed", + &format!( + "scheduled TheRock update check failed during contained restricted execution: {error}; no updates were applied" + ), + None, + )?; + } + }, + Err(error) => { + crate::persistence::record_event( + paths, + state, + "therock-update", + "error", + "update_check_failed", + &format!( + "scheduled TheRock update check failed during contained read-only execution: {error}; no updates were applied" + ), + None, + )?; + } + }, + } + Ok(()) +} + +struct RestrictedCheckUpdatesResult<'a> { + status: &'a str, + update_available: bool, + exit_status: i64, +} + +fn restricted_check_updates_result(value: &Value) -> Result> { + let tool = value + .get("tool") + .and_then(Value::as_str) + .context("restricted update check did not report a tool name")?; + if tool != crate::cli::SandboxToolArg::CheckUpdates.as_cli_value() { + bail!("restricted update check returned `{tool}`, expected `check_updates`"); + } + let status = value + .get("status") + .and_then(Value::as_str) + .unwrap_or("checked"); + let update_available = value + .get("update_available") + .and_then(Value::as_bool) + .unwrap_or(matches!(status, "update_available" | "repair_available")); + let exit_status = value + .get("exit_status") + .and_then(Value::as_i64) + .unwrap_or_else(|| i64::from(status == "error")); + Ok(RestrictedCheckUpdatesResult { + status, + update_available, + exit_status, + }) +} + +fn record_update_available_notification( + paths: &AppPaths, + state: &mut AutomationRuntimeState, + status: &str, +) -> Result<()> { + let message = if status == "repair_available" { + "A ROCm runtime repair is available because its package composition changed. Preview it before applying. No updates were applied." + } else { + "A ROCm runtime update is available. Preview it before applying. No updates were applied." + }; + crate::persistence::record_event( + paths, + state, + "therock-update", + "info", + "notify_if_newer", + message, + None, + )?; + crate::sandbox::record_notification_audit( + paths, + "watcher:therock-update", + "notify_if_newer", + Some("therock-update"), + message, + ) +} + +fn therock_update_due(state: &AutomationRuntimeState, now: u128) -> bool { + let Some(snapshot) = state + .active_watchers + .iter() + .find(|watcher| watcher.id == "therock-update" && watcher.enabled) + else { + return false; + }; + snapshot + .last_event_unix_ms + .is_none_or(|last_event| now.saturating_sub(last_event) >= THEROCK_UPDATE_INTERVAL_MS) +} + +fn gpu_metrics_due(state: &AutomationRuntimeState, now: u128) -> bool { + let Some(snapshot) = state + .active_watchers + .iter() + .find(|watcher| watcher.id == "gpu-metrics" && watcher.enabled) + else { + return false; + }; + snapshot + .last_event_unix_ms + .is_none_or(|last_event| now.saturating_sub(last_event) >= GPU_METRICS_INTERVAL_MS) +} + +fn gpu_thermal_protect_due(state: &AutomationRuntimeState, now: u128) -> bool { + let Some(snapshot) = state + .active_watchers + .iter() + .find(|watcher| watcher.id == "gpu-thermal-protect" && watcher.enabled) + else { + return false; + }; + snapshot + .last_event_unix_ms + .is_none_or(|last_event| now.saturating_sub(last_event) >= GPU_METRICS_INTERVAL_MS) +} + +fn handle_gpu_metrics_event( + paths: &AppPaths, + mode: WatcherMode, + state: &mut AutomationRuntimeState, + event: &AutomationTriggerEvent, +) -> Result<()> { + let summary = event + .payload + .get("summary") + .and_then(Value::as_str) + .unwrap_or("summary unavailable"); + let level = if event.kind == "gpu.metrics" { + "info" + } else { + "warn" + }; + let mode_note = match mode { + WatcherMode::Observe => "observe mode records telemetry only", + WatcherMode::Propose => { + "propose mode has no GPU mutation policy yet, so telemetry is recorded only" + } + WatcherMode::Contained => { + "contained mode has no GPU mutation policy yet, so telemetry is recorded only" + } + }; + let reason = event + .reason + .as_deref() + .filter(|value| !value.trim().is_empty()) + .unwrap_or("no detail"); + let source = match event.source.as_str() { + "gpu_telemetry" => "local amd-smi telemetry", + "local_webhook" => "local webhook", + other => other, + }; + + crate::persistence::record_event( + paths, + state, + "gpu-metrics", + level, + "record_gpu_metrics", + &format!( + "GPU metrics event from {source}: {summary}; reason={reason}; {mode_note}; no mutating action was taken" + ), + None, + ) +} + +fn gpu_snapshot_summary(snapshot: &CodexBridgeGpuSnapshot) -> String { + let mut parts = Vec::new(); + parts.push(format!("amd_smi_available={}", snapshot.amd_smi_available)); + parts.push(format!( + "static_snapshot={}", + if snapshot.static_snapshot.is_some() { + "available" + } else { + "missing" + } + )); + parts.push(format!( + "monitor_snapshot={}", + if snapshot.monitor_snapshot.is_some() { + "available" + } else { + "missing" + } + )); + if let Some(count) = snapshot.static_snapshot.as_ref().and_then(gpu_data_count) { + parts.push(format!("gpu_count={count}")); + } + if let Some(note) = snapshot.note.as_deref() + && !note.trim().is_empty() + { + parts.push(format!("note={note}")); + } + parts.join(" ") +} + +fn gpu_data_count(value: &Value) -> Option { + value + .get("gpu_data") + .and_then(Value::as_array) + .map(Vec::len) +} + +#[derive(Debug, Clone, Copy)] +struct GpuPressureReading { + gpu_index: Option, + hotspot_temperature_c: Option, + memory_temperature_c: Option, + vram_percent: Option, +} + +fn gpu_pressure_events( + now: u128, + snapshot: &CodexBridgeGpuSnapshot, +) -> Vec { + let Some(monitor_snapshot) = snapshot.monitor_snapshot.as_ref() else { + return Vec::new(); + }; + monitor_entries(monitor_snapshot) + .into_iter() + .filter_map(gpu_pressure_reading) + .filter_map(|reading| gpu_pressure_event(now, reading)) + .collect() +} + +fn gpu_pressure_event(now: u128, reading: GpuPressureReading) -> Option { + let (kind, reason, metric_label, value, threshold) = if let Some(value) = + reading.hotspot_temperature_c + && value >= GPU_THERMAL_HOTSPOT_PRESSURE_C + { + ( + "gpu.thermal_pressure", + "hotspot_temperature_threshold", + "hotspot temperature", + value, + GPU_THERMAL_HOTSPOT_PRESSURE_C, + ) + } else if let Some(value) = reading.memory_temperature_c + && value >= GPU_THERMAL_MEMORY_PRESSURE_C + { + ( + "gpu.thermal_pressure", + "memory_temperature_threshold", + "memory temperature", + value, + GPU_THERMAL_MEMORY_PRESSURE_C, + ) + } else if let Some(value) = reading.vram_percent + && value >= GPU_MEMORY_VRAM_PRESSURE_PERCENT + { + ( + "gpu.memory_pressure", + "vram_pressure_threshold", + "VRAM use", + value, + GPU_MEMORY_VRAM_PRESSURE_PERCENT, + ) + } else { + return None; + }; + let gpu_label = reading + .gpu_index + .map_or_else(|| "the GPU".to_owned(), |gpu| format!("GPU {gpu}")); + let unit = if metric_label == "VRAM use" { + "%" + } else { + " C" + }; + let summary = format!( + "{gpu_label} {metric_label} is {}{} (limit {}{})", + display_metric(value), + unit, + display_metric(threshold), + unit + ); + Some(AutomationTriggerEvent { + at_unix_ms: now, + kind: kind.to_owned(), + source: "gpu_telemetry".to_owned(), + watcher_hint: Some("gpu-thermal-protect".to_owned()), + service_id: None, + reason: Some(reason.to_owned()), + payload: json!({ + "gpu": reading.gpu_index, + "hotspot_temperature_c": reading.hotspot_temperature_c, + "memory_temperature_c": reading.memory_temperature_c, + "vram_percent": reading.vram_percent, + "hotspot_threshold_c": GPU_THERMAL_HOTSPOT_PRESSURE_C, + "memory_temperature_threshold_c": GPU_THERMAL_MEMORY_PRESSURE_C, + "vram_threshold_percent": GPU_MEMORY_VRAM_PRESSURE_PERCENT, + "recommended_action": "stop_serving_load", + "summary": summary, + }), + }) +} + +fn monitor_entries(value: &Value) -> Vec<&Value> { + if let Some(entries) = value.as_array() { + return entries.iter().collect(); + } + value + .get("gpu_data") + .and_then(Value::as_array) + .map(|entries| entries.iter().collect()) + .unwrap_or_default() +} + +fn gpu_pressure_reading(entry: &Value) -> Option { + let reading = GpuPressureReading { + gpu_index: metric_u64(entry, &["gpu", "gpu_id", "gpu_index"]), + hotspot_temperature_c: metric_f64( + entry, + &[ + "hotspot_temperature", + "hotspot_temperature_c", + "temperature_hotspot", + ], + ), + memory_temperature_c: metric_f64( + entry, + &[ + "memory_temperature", + "memory_temperature_c", + "temperature_memory", + ], + ), + vram_percent: metric_f64( + entry, + &["vram_percent", "vram_usage_percent", "vram_used_percent"], + ), + }; + (reading.hotspot_temperature_c.is_some() + || reading.memory_temperature_c.is_some() + || reading.vram_percent.is_some()) + .then_some(reading) +} + +fn metric_f64(entry: &Value, keys: &[&str]) -> Option { + keys.iter() + .find_map(|key| entry.get(*key).and_then(value_as_metric_f64)) +} + +fn metric_u64(entry: &Value, keys: &[&str]) -> Option { + keys.iter() + .find_map(|key| entry.get(*key).and_then(value_as_metric_u64)) +} + +fn value_as_metric_f64(value: &Value) -> Option { + match value { + Value::Number(number) => number.as_f64(), + Value::String(text) => text.trim().parse::().ok(), + Value::Object(map) => map + .get("value") + .or_else(|| map.get("val")) + .and_then(value_as_metric_f64), + _ => None, + } +} + +fn value_as_metric_u64(value: &Value) -> Option { + match value { + Value::Number(number) => number.as_u64(), + Value::String(text) => text.trim().parse::().ok(), + Value::Object(map) => map + .get("value") + .or_else(|| map.get("val")) + .and_then(value_as_metric_u64), + _ => None, + } +} + +fn display_metric(value: f64) -> String { + if value.fract().abs() < f64::EPSILON { + format!("{value:.0}") + } else { + format!("{value:.1}") + } +} + +fn handle_gpu_thermal_protect_event( + paths: &AppPaths, + mode: WatcherMode, + state: &mut AutomationRuntimeState, + event: &AutomationTriggerEvent, +) -> Result<()> { + let summary = payload_string(&event.payload, "summary") + .unwrap_or_else(|| "GPU pressure is high".to_owned()); + let reason = event.reason.as_deref().unwrap_or("gpu_pressure_threshold"); + + if matches!(mode, WatcherMode::Observe) { + return crate::persistence::record_event( + paths, + state, + "gpu-thermal-protect", + "warn", + "observe_gpu_pressure", + &format!( + "{summary}; observe mode records this only and does not stop any model server" + ), + event.service_id.clone(), + ); + } + + let Some(record) = resolve_gpu_pressure_service_target(paths, event)? else { + return crate::persistence::record_event( + paths, + state, + "gpu-thermal-protect", + "warn", + "gpu_pressure_no_clear_target", + &format!( + "{summary}; rocm-cli did not choose a model server to stop because there was no single clear running managed server" + ), + None, + ); + }; + + if pending_stop_proposal_exists(paths, &record.service_id)? { + return crate::persistence::record_event( + paths, + state, + "gpu-thermal-protect", + "info", + "stop_proposal_already_pending", + &format!( + "{summary}; a reviewed stop request is already waiting for {}", + record.service_id + ), + Some(record.service_id), + ); + } + + let action = "queue_stop_server_proposal"; + let mode_note = if matches!(mode, WatcherMode::Contained) { + "contained mode still asks before stopping anything" + } else { + "asking before stopping anything" + }; + let message = format!( + "{summary}; {mode_note}; selected managed server {} ({})", + record.service_id, record.endpoint_url + ); + crate::persistence::record_event( + paths, + state, + "gpu-thermal-protect", + "warn", + action, + &message, + Some(record.service_id.clone()), + )?; + queue_proposal_with_arguments( + paths, + "gpu-thermal-protect", + action, + "Review GPU pressure stop", + &message, + Some(record.service_id.clone()), + json!({ + "service_id": record.service_id, + "model_ref": record.model_ref, + "canonical_model_id": record.canonical_model_id, + "endpoint_url": record.endpoint_url, + "engine": record.engine, + "pressure_kind": event.kind, + "pressure_reason": reason, + "pressure_summary": summary, + "gpu": event.payload.get("gpu").cloned().unwrap_or(Value::Null), + "hotspot_temperature_c": event.payload.get("hotspot_temperature_c").cloned().unwrap_or(Value::Null), + "memory_temperature_c": event.payload.get("memory_temperature_c").cloned().unwrap_or(Value::Null), + "vram_percent": event.payload.get("vram_percent").cloned().unwrap_or(Value::Null), + "hotspot_threshold_c": GPU_THERMAL_HOTSPOT_PRESSURE_C, + "memory_temperature_threshold_c": GPU_THERMAL_MEMORY_PRESSURE_C, + "vram_threshold_percent": GPU_MEMORY_VRAM_PRESSURE_PERCENT, + }), + ) +} + +fn resolve_gpu_pressure_service_target( + paths: &AppPaths, + event: &AutomationTriggerEvent, +) -> Result> { + if let Some(service_id) = event + .service_id + .as_deref() + .map(str::trim) + .filter(|value| !value.is_empty()) + { + let record = load_service_record(paths, service_id)?; + return Ok(active_pressure_target(record)); + } + + let active = crate::persistence::load_managed_services(paths)? + .into_iter() + .filter_map(active_pressure_target) + .collect::>(); + if active.len() == 1 { + Ok(active.into_iter().next()) + } else { + Ok(None) + } +} + +fn active_pressure_target(record: ManagedServiceRecord) -> Option { + (record.mode == "managed" && matches!(record.status.as_str(), "ready" | "running")) + .then_some(record) +} + +fn pending_stop_proposal_exists(paths: &AppPaths, service_id: &str) -> Result { + Ok(rocm_core::load_recent_automation_proposals(paths, 100)? + .into_iter() + .any(|proposal| { + proposal.status == "pending" + && proposal.watcher_id == "gpu-thermal-protect" + && proposal.service_id.as_deref() == Some(service_id) + && (proposal.action == "queue_stop_server_proposal" + || proposal.tool.as_deref() == Some("stop_server")) + })) +} + +fn handle_cache_warm_event( + paths: &AppPaths, + mode: WatcherMode, + state: &mut AutomationRuntimeState, + event: &AutomationTriggerEvent, +) -> Result<()> { + handle_cache_warm_event_with_resolver(paths, mode, state, event, |artifact_ref| { + resolve_model_recipe_artifact(artifact_ref).map(|resolved| resolved.is_some()) + }) +} + +fn handle_cache_warm_event_with_resolver( + paths: &AppPaths, + mode: WatcherMode, + state: &mut AutomationRuntimeState, + event: &AutomationTriggerEvent, + mut artifact_exists: F, +) -> Result<()> +where + F: FnMut(&str) -> Result, +{ + let Some(artifact_ref) = payload_string(&event.payload, "artifact_ref") else { + return crate::persistence::record_event( + paths, + state, + "cache-warm", + "warn", + "cache_warm_missing_artifact", + "cache warm event did not include artifact_ref; no prefetch proposal was queued", + None, + ); + }; + match artifact_exists(&artifact_ref) { + Ok(true) => {} + Ok(false) => { + return crate::persistence::record_event( + paths, + state, + "cache-warm", + "warn", + "cache_warm_unknown_artifact", + &format!( + "cache warm requested unknown registry artifact {artifact_ref}; no prefetch proposal was queued" + ), + None, + ); + } + Err(error) => { + return crate::persistence::record_event( + paths, + state, + "cache-warm", + "error", + "cache_warm_registry_error", + &format!( + "cache warm could not verify registry artifact {artifact_ref}: {error}; no prefetch proposal was queued" + ), + None, + ); + } + } + match mode { + WatcherMode::Observe => crate::persistence::record_event( + paths, + state, + "cache-warm", + "info", + "observe_cache_warm_request", + &format!( + "observed cache warm request for {artifact_ref}; observe mode does not queue or download artifacts" + ), + None, + ), + WatcherMode::Propose | WatcherMode::Contained => { + let action = "queue_prefetch_proposal"; + let message = if matches!(mode, WatcherMode::Contained) { + format!( + "cache warm requested for {artifact_ref}; contained mode still queues a review because artifact downloads require explicit source-policy approval" + ) + } else { + format!( + "cache warm requested for {artifact_ref}; queueing a reviewed prefetch proposal" + ) + }; + crate::persistence::record_event( + paths, + state, + "cache-warm", + "info", + action, + &message, + None, + )?; + queue_proposal_with_arguments( + paths, + "cache-warm", + action, + "Prefetch model artifact", + &message, + None, + json!({ + "artifact_ref": artifact_ref, + }), + ) + } + } +} + +fn handle_driver_upgrade_event( + paths: &AppPaths, + mode: WatcherMode, + state: &mut AutomationRuntimeState, + event: &AutomationTriggerEvent, +) -> Result<()> { + handle_driver_upgrade_event_with_runner(paths, mode, state, event, |paths| { + crate::sandbox::run_sandbox_tool( + paths, + crate::cli::SandboxToolArg::DriverPlan, + None, + None, + None, + crate::cli::SandboxToolPolicy::default(), + ) + }) +} + +fn handle_driver_upgrade_event_with_runner( + paths: &AppPaths, + mode: WatcherMode, + state: &mut AutomationRuntimeState, + event: &AutomationTriggerEvent, + driver_plan_runner: F, +) -> Result<()> +where + F: FnOnce(&AppPaths) -> Result, +{ + if payload_string(&event.payload, "component").as_deref() != Some("driver") { + return crate::persistence::record_event( + paths, + state, + "driver-upgrade", + "warn", + "driver_upgrade_ignored_component", + "driver-upgrade event did not include payload.component=driver; no driver plan proposal was queued", + None, + ); + } + + match mode { + WatcherMode::Observe => crate::persistence::record_event( + paths, + state, + "driver-upgrade", + "info", + "observe_driver_update", + "observed local driver update signal; observe mode does not queue or run a driver plan", + None, + ), + WatcherMode::Propose => { + let action = "prepare_driver_plan"; + let message = + "local driver update signal received; queueing a reviewed read-only driver plan"; + crate::persistence::record_event( + paths, + state, + "driver-upgrade", + "info", + action, + message, + None, + )?; + queue_proposal( + paths, + "driver-upgrade", + action, + "Review driver install plan", + message, + None, + ) + } + WatcherMode::Contained => match driver_plan_runner(paths) { + Ok(output) => match restricted_driver_plan_result(&output) { + Ok(result) => crate::persistence::record_event( + paths, + state, + "driver-upgrade", + if result.exit_status == 0 { + "info" + } else { + "error" + }, + "run_driver_plan", + &format!( + "local driver update signal received; contained restricted driver_plan status={}; no driver commands were executed", + result.status + ), + None, + ), + Err(error) => crate::persistence::record_event( + paths, + state, + "driver-upgrade", + "error", + "driver_plan_failed", + &format!( + "local driver update signal received, but contained restricted driver_plan failed: {error}; no driver commands were executed" + ), + None, + ), + }, + Err(error) => crate::persistence::record_event( + paths, + state, + "driver-upgrade", + "error", + "driver_plan_failed", + &format!( + "local driver update signal received, but contained restricted driver_plan failed: {error}; no driver commands were executed" + ), + None, + ), + }, + } +} + +struct RestrictedDriverPlanResult<'a> { + status: &'a str, + exit_status: i64, +} + +fn restricted_driver_plan_result(value: &Value) -> Result> { + let tool = value + .get("tool") + .and_then(Value::as_str) + .context("restricted driver plan did not report a tool name")?; + if tool != crate::cli::SandboxToolArg::DriverPlan.as_cli_value() { + bail!("restricted driver plan returned `{tool}`, expected `driver_plan`"); + } + let status = value + .get("status") + .and_then(Value::as_str) + .unwrap_or("planned"); + let exit_status = value + .get("exit_status") + .and_then(Value::as_i64) + .unwrap_or_else(|| i64::from(status == "error")); + Ok(RestrictedDriverPlanResult { + status, + exit_status, + }) +} + +pub(crate) fn payload_string(payload: &Value, key: &str) -> Option { + payload + .get(key) + .and_then(Value::as_str) + .map(str::trim) + .filter(|value| !value.is_empty()) + .map(str::to_owned) +} + +#[cfg(test)] +fn evaluate_server_recover( + paths: &AppPaths, + mode: WatcherMode, + state: &mut AutomationRuntimeState, +) -> Result<()> { + let now = unix_time_millis(); + if !server_recover_due(state, now) { + return Ok(()); + } + + let Some((mut record, recovery_reason)) = find_recoverable_service(paths)? else { + return Ok(()); + }; + let kind = service_recovery_event_kind(&recovery_reason); + let event = AutomationTriggerEvent { + at_unix_ms: now, + kind: kind.to_owned(), + source: "managed_service".to_owned(), + watcher_hint: Some("server-recover".to_owned()), + service_id: Some(record.service_id.clone()), + reason: Some(recovery_reason), + payload: json!({ + "engine": record.engine.clone(), + "status": record.status.clone(), + "endpoint": record.endpoint_url.clone(), + }), + }; + handle_server_recover_event_with_record(paths, mode, state, &event, &mut record) +} + +fn handle_server_recover_event( + paths: &AppPaths, + mode: WatcherMode, + state: &mut AutomationRuntimeState, + event: &AutomationTriggerEvent, +) -> Result<()> { + let service_id = event + .service_id + .as_deref() + .context("server-recover event is missing service_id")?; + let mut record = load_service_record(paths, service_id)?; + if !service_record_matches_recovery_event(paths, &record, event) { + crate::persistence::record_event( + paths, + state, + "server-recover", + "info", + "ignore_nonrecoverable_service", + &format!( + "managed service {} does not currently need recovery; restart not attempted", + record.service_id + ), + Some(record.service_id.clone()), + )?; + return Ok(()); + } + handle_server_recover_event_with_record(paths, mode, state, event, &mut record) +} + +fn service_record_matches_recovery_event( + paths: &AppPaths, + record: &ManagedServiceRecord, + event: &AutomationTriggerEvent, +) -> bool { + match event.kind.as_str() { + "service.manifest_recoverable" => { + manifest_service_recovery_reason(record, unix_time_millis()).is_some() + } + "service.endpoint_recoverable" => endpoint_service_recovery_reason(record).is_some(), + "service.healthcheck_recoverable" => { + crate::common::engine_healthcheck_response(paths, &record.engine, &record.service_id) + .is_ok_and(|healthcheck| { + crate::common::healthcheck_response_recoverable(&healthcheck) + }) + } + _ => false, + } +} + +fn handle_server_recover_event_with_record( + paths: &AppPaths, + mode: WatcherMode, + state: &mut AutomationRuntimeState, + event: &AutomationTriggerEvent, + record: &mut ManagedServiceRecord, +) -> Result<()> { + let now = unix_time_millis(); + let recovery_reason = event.reason.as_deref().unwrap_or("recoverable_event"); + let recovery_reason_display = display_recovery_reason(recovery_reason); + + match watcher_policy_action("server-recover", mode) { + WatcherPolicyAction::Observe => crate::persistence::record_event( + paths, + state, + "server-recover", + "warn", + "observe_failure", + &format!( + "observed managed service {} needing recovery ({recovery_reason_display}); restart not attempted in observe mode", + record.service_id, + ), + Some(record.service_id.clone()), + ), + WatcherPolicyAction::QueueProposal => { + let message = format!( + "managed service {} needs recovery ({recovery_reason_display}); queueing restart proposal", + record.service_id, + ); + crate::persistence::record_event( + paths, + state, + "server-recover", + "warn", + "queue_restart_proposal", + &message, + Some(record.service_id.clone()), + )?; + queue_proposal( + paths, + "server-recover", + "queue_restart_proposal", + "Restart managed service", + &message, + Some(record.service_id.clone()), + ) + } + WatcherPolicyAction::RunContained => { + if let Some(last_restart) = record.last_restart_unix_ms + && now.saturating_sub(last_restart) < SERVER_RECOVER_BACKOFF_MS + { + return Ok(()); + } + // A public service whose endpoint key is gone can never be recovered: + // the respawn guard in `supervise_service` refuses it by design. + // Report it and stop, rather than letting a permanent failure + // propagate out of `evaluate_watchers` and take the whole daemon — + // and every other watcher — down on each 30s recovery tick. + if let Err(error) = crate::common::ensure_public_service_has_endpoint_key( + &record.host, + rocm_engine_protocol::endpoint_key_file_if_present(paths, &record.service_id) + .and_then(|path| rocm_engine_protocol::endpoint_api_key_file_if_valid(&path)) + .is_some(), + record.requires_api_key, + ) { + return crate::persistence::record_event( + paths, + state, + "server-recover", + "error", + "restart_managed_service_refused", + &format!( + "cannot recover managed service {} on {}:{} after \ + {recovery_reason_display}: {error}", + record.service_id, record.host, record.port + ), + Some(record.service_id.clone()), + ); + } + restart_managed_service(paths, &mut *record)?; + crate::persistence::record_event( + paths, + state, + "server-recover", + "info", + "restart_managed_service", + &format!( + "restarted managed service {} on {}:{} after {recovery_reason_display}", + record.service_id, record.host, record.port + ), + Some(record.service_id.clone()), + ) + } + } +} + +fn display_recovery_reason(reason: &str) -> String { + match reason { + "manifest_status_failed" => "manifest reports failed".to_owned(), + "manifest_status_exited" => "manifest reports exited".to_owned(), + "manifest_status_unreachable" => "manifest reports unreachable".to_owned(), + "manifest_status_starting_stale" => "service has been starting for too long".to_owned(), + "manifest_status_recovering_stale" => "service has been recovering for too long".to_owned(), + "endpoint_status_unreachable" => "endpoint port is unreachable".to_owned(), + other if other.starts_with("healthcheck_status_") => format!( + "engine healthcheck reports {}", + other.trim_start_matches("healthcheck_status_") + ), + other => other.replace('_', " "), + } +} + +fn service_recovery_event_kind(recovery_reason: &str) -> &'static str { + if recovery_reason.starts_with("healthcheck_status_") { + "service.healthcheck_recoverable" + } else if recovery_reason.starts_with("endpoint_status_") { + "service.endpoint_recoverable" + } else { + "service.manifest_recoverable" + } +} + +fn server_recover_due(state: &AutomationRuntimeState, now: u128) -> bool { + let Some(snapshot) = state + .active_watchers + .iter() + .find(|watcher| watcher.id == "server-recover" && watcher.enabled) + else { + return false; + }; + snapshot + .last_event_unix_ms + .is_none_or(|last_event| now.saturating_sub(last_event) >= SERVER_RECOVER_BACKOFF_MS) +} + +#[derive(Debug, Clone, Copy, Eq, PartialEq)] +enum WatcherPolicyAction { + Observe, + QueueProposal, + RunContained, +} + +const fn watcher_policy_action(watcher_id: &str, mode: WatcherMode) -> WatcherPolicyAction { + match (watcher_id, mode) { + (_, WatcherMode::Observe) => WatcherPolicyAction::Observe, + (_, WatcherMode::Propose) => WatcherPolicyAction::QueueProposal, + (_, WatcherMode::Contained) => WatcherPolicyAction::RunContained, + } +} + +fn find_recoverable_service(paths: &AppPaths) -> Result> { + let now = unix_time_millis(); + for record in crate::persistence::load_managed_services(paths)? { + if record.mode != "managed" { + continue; + } + if let Some(reason) = manifest_service_recovery_reason(&record, now) { + return Ok(Some((record, reason))); + } + if matches!(record.status.as_str(), "ready" | "running") { + let Ok(healthcheck) = crate::common::engine_healthcheck_response( + paths, + &record.engine, + &record.service_id, + ) else { + if let Some(reason) = endpoint_service_recovery_reason(&record) { + return Ok(Some((record, reason))); + } + continue; + }; + if crate::common::healthcheck_response_recoverable(&healthcheck) { + return Ok(Some(( + record, + format!("healthcheck_status_{}", healthcheck.status), + ))); + } + if let Some(reason) = endpoint_service_recovery_reason(&record) { + return Ok(Some((record, reason))); + } + } + } + Ok(None) +} + +fn endpoint_service_recovery_reason(record: &ManagedServiceRecord) -> Option { + (!crate::common::wait_for_port(&record.host, record.port, ENDPOINT_HEALTH_TIMEOUT)) + .then(|| "endpoint_status_unreachable".to_owned()) +} + +pub(crate) fn load_service_record( + paths: &AppPaths, + service_id: &str, +) -> Result { + rocm_core::ServiceId::new(service_id) + .with_context(|| format!("invalid managed service id `{service_id}`"))?; + let manifest_path = paths.service_manifest_path(service_id); + let bytes = fs::read(&manifest_path).with_context(|| { + format!( + "managed service `{service_id}` not found at {}", + manifest_path.display() + ) + })?; + let record = serde_json::from_slice::(&bytes) + .with_context(|| format!("failed to parse {}", manifest_path.display()))?; + if record.service_id != service_id { + bail!( + "managed service manifest {} contains service_id `{}`, expected `{service_id}`", + manifest_path.display(), + record.service_id + ); + } + Ok(record) +} + +fn manifest_service_recovery_reason( + record: &ManagedServiceRecord, + now_unix_ms: u128, +) -> Option { + match record.status.as_str() { + "failed" | "exited" | "unreachable" => Some(format!("manifest_status_{}", record.status)), + "starting" | "recovering" => { + let started_at = record + .last_restart_unix_ms + .unwrap_or(record.created_at_unix_ms); + (now_unix_ms.saturating_sub(started_at) >= SERVER_TRANSIENT_STALE_MS) + .then(|| format!("manifest_status_{}_stale", record.status)) + } + _ => None, + } +} + +pub(crate) fn restart_managed_service( + _paths: &AppPaths, + record: &mut ManagedServiceRecord, +) -> Result<()> { + let rocmd_binary = + std::env::current_exe().context("failed to resolve current rocmd executable path")?; + let log_file = fs::OpenOptions::new() + .create(true) + .append(true) + .open(&record.log_path) + .with_context(|| format!("failed to open {}", record.log_path.display()))?; + let log_file_err = log_file + .try_clone() + .context("failed to clone service log file handle")?; + + record.status = "recovering".to_owned(); + // Counts the restart and drops the previous run's inference verification. + // The respawned child writes a fresh record of its own, and "recovering" is + // outside the statuses that probe, so a stale verdict would not currently be + // acted on — but this record is written again below, after the spawn, and + // that write can land after the child's. Clearing here keeps the invariant + // true at the one site that reuses a record across restarts. + record.reset_for_restart(); + record.supervisor_pid = std::process::id(); + record.write()?; + + let mut child = detached_rocmd_command(&rocmd_binary) + .args(recovery_supervise_args(record)) + .stdin(Stdio::null()) + .stdout(Stdio::from(log_file)) + .stderr(Stdio::from(log_file_err)) + .spawn() + .context("failed to spawn recovery supervisor")?; + + record.supervisor_pid = child.id(); + record.write()?; + + thread::sleep(Duration::from_millis(200)); + if let Some(status) = child + .try_wait() + .context("failed to check recovery supervisor startup state")? + { + record.status = "failed".to_owned(); + record.write()?; + anyhow::bail!( + "recovery supervisor exited immediately with status {status}; inspect {}", + record.log_path.display() + ); + } + + Ok(()) +} + +fn recovery_supervise_args(record: &ManagedServiceRecord) -> Vec { + let mut args = vec![ + "supervise".to_owned(), + record.service_id.clone(), + "--engine".to_owned(), + record.engine.clone(), + "--model-ref".to_owned(), + record.model_ref.clone(), + "--canonical-model-id".to_owned(), + record.canonical_model_id.clone(), + "--host".to_owned(), + record.host.clone(), + "--port".to_owned(), + record.port.to_string(), + "--device-policy".to_owned(), + record + .device_policy + .as_deref() + .unwrap_or("gpu_required") + .to_owned(), + ]; + args.extend(crate::common::optional_arg( + "--runtime-id", + record.runtime_id.as_deref(), + )); + args.extend(crate::common::optional_arg( + "--env-id", + record.env_id.as_deref(), + )); + if let Some(csv) = rocm_engine_protocol::gpu_indices_to_csv(&record.gpu_indices) { + args.extend(["--gpu".to_owned(), csv]); + } + args.extend(crate::common::optional_arg( + "--engine-recipe-json", + record.engine_recipe_json.as_deref(), + )); + args +} + +fn queue_proposal( + paths: &AppPaths, + watcher_id: &str, + action: &str, + title: &str, + message: &str, + service_id: Option, +) -> Result<()> { + queue_proposal_with_arguments( + paths, + watcher_id, + action, + title, + message, + service_id.clone(), + proposal_arguments_for_action(action, service_id.as_deref()), + ) +} + +fn queue_proposal_with_arguments( + paths: &AppPaths, + watcher_id: &str, + action: &str, + title: &str, + message: &str, + service_id: Option, + arguments: Value, +) -> Result<()> { + append_automation_proposal( + paths, + &AutomationProposalRecord { + at_unix_ms: unix_time_millis(), + proposal_id: String::new(), + watcher_id: watcher_id.to_owned(), + action: action.to_owned(), + title: title.to_owned(), + message: message.to_owned(), + status: "pending".to_owned(), + service_id, + tool: proposal_tool_for_action(action).map(str::to_owned), + arguments, + reviewed_at_unix_ms: None, + }, + ) +} + +fn proposal_tool_for_action(action: &str) -> Option<&'static str> { + match action { + "queue_restart_proposal" => Some("restart_server"), + "queue_stop_server_proposal" => Some("stop_server"), + "queue_update_proposal" => Some("check_updates"), + "queue_prefetch_proposal" => Some("prefetch_artifact"), + "prepare_driver_plan" => Some("driver_plan"), + _ => None, + } +} + +fn proposal_arguments_for_action(action: &str, service_id: Option<&str>) -> Value { + match action { + "queue_restart_proposal" => json!({ + "service_id": service_id, + }), + "queue_stop_server_proposal" => json!({ + "service_id": service_id, + }), + "queue_update_proposal" => json!({}), + "prepare_driver_plan" => json!({}), + _ => Value::Null, + } +} + +#[cfg(unix)] +fn detached_rocmd_command(rocmd_binary: &std::path::Path) -> ProcessCommand { + let mut command = ProcessCommand::new("setsid"); + command.arg(rocmd_binary); + command +} + +#[cfg(not(unix))] +fn detached_rocmd_command(rocmd_binary: &std::path::Path) -> ProcessCommand { + ProcessCommand::new(rocmd_binary) +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::test_support::temp_app_paths; + use rocm_core::{AuditEventRecord, AutomationEventRecord}; + + fn test_watcher_snapshot( + id: &str, + mode: WatcherMode, + last_event_unix_ms: Option, + ) -> WatcherRuntimeSnapshot { + WatcherRuntimeSnapshot { + id: id.to_owned(), + enabled: true, + mode, + summary: "test watcher".to_owned(), + last_event: None, + last_event_unix_ms, + } + } + + fn test_runtime_state(active_watchers: Vec) -> AutomationRuntimeState { + AutomationRuntimeState { + running: true, + automations_enabled: true, + daemon_pid: 1, + started_at_unix_ms: 1, + last_tick_unix_ms: 1, + local_webhook_endpoint: None, + active_watchers, + } + } + + #[test] + fn event_collector_emits_schedule_tick_for_due_update() -> Result<()> { + let (root, paths) = temp_app_paths("event-bus-schedule"); + let state = test_runtime_state(vec![test_watcher_snapshot( + "therock-update", + WatcherMode::Observe, + None, + )]); + + let events = collect_automation_events(&paths, &RocmCliConfig::default(), &state)?; + fs::remove_dir_all(root).ok(); + + let event = events + .iter() + .find(|event| event.watcher_hint.as_deref() == Some("therock-update")) + .expect("schedule tick event should be emitted"); + assert_eq!(event.kind, "schedule.tick"); + assert_eq!(event.source, "scheduler"); + assert_eq!(event.reason.as_deref(), Some("therock_update_interval_due")); + Ok(()) + } + + #[test] + fn event_collector_emits_recoverable_service_event() -> Result<()> { + let (root, paths) = temp_app_paths("event-bus-service"); + paths.ensure()?; + let mut failed = ManagedServiceRecord::new( + &paths, + "svc-failed", + "vllm", + "qwen", + "Qwen/Qwen3.5", + "127.0.0.1", + 11435, + "managed", + 123, + None, + None, + None, + ); + failed.status = "failed".to_owned(); + failed.write()?; + let state = test_runtime_state(vec![test_watcher_snapshot( + "server-recover", + WatcherMode::Propose, + None, + )]); + + let events = collect_automation_events(&paths, &RocmCliConfig::default(), &state)?; + fs::remove_dir_all(root).ok(); + + let event = events + .iter() + .find(|event| event.watcher_hint.as_deref() == Some("server-recover")) + .expect("recoverable service event should be emitted"); + assert_eq!(event.kind, "service.manifest_recoverable"); + assert_eq!(event.service_id.as_deref(), Some("svc-failed")); + assert_eq!(event.reason.as_deref(), Some("manifest_status_failed")); + Ok(()) + } + + #[test] + fn event_collector_emits_endpoint_recoverable_service_event() -> Result<()> { + let (root, paths) = temp_app_paths("event-bus-endpoint"); + paths.ensure()?; + let mut service = ManagedServiceRecord::new( + &paths, + "svc-endpoint", + "missing-engine", + "qwen", + "Qwen/Qwen3.5", + "127.0.0.1", + 1, + "managed", + 123, + None, + None, + None, + ); + service.status = "ready".to_owned(); + service.write()?; + let state = test_runtime_state(vec![test_watcher_snapshot( + "server-recover", + WatcherMode::Propose, + None, + )]); + + let events = collect_automation_events(&paths, &RocmCliConfig::default(), &state)?; + fs::remove_dir_all(root).ok(); + + let event = events + .iter() + .find(|event| event.watcher_hint.as_deref() == Some("server-recover")) + .expect("endpoint recoverable service event should be emitted"); + assert_eq!(event.kind, "service.endpoint_recoverable"); + assert_eq!(event.service_id.as_deref(), Some("svc-endpoint")); + assert_eq!(event.reason.as_deref(), Some("endpoint_status_unreachable")); + Ok(()) + } + + #[test] + fn event_collector_emits_gpu_metrics_event_when_enabled() -> Result<()> { + let (root, paths) = temp_app_paths("event-bus-gpu-metrics"); + let state = test_runtime_state(vec![test_watcher_snapshot( + "gpu-metrics", + WatcherMode::Observe, + None, + )]); + + let events = collect_automation_events_with_gpu_snapshot(&paths, &state, || { + CodexBridgeGpuSnapshot { + amd_smi_available: true, + static_snapshot: Some(json!({ + "gpu_data": [ + { "gpu": 0, "asic": { "market_name": "AMD Radeon Test" } } + ] + })), + monitor_snapshot: Some(json!({ "gpu_data": [] })), + note: None, + } + })?; + fs::remove_dir_all(root).ok(); + + let event = events + .iter() + .find(|event| event.watcher_hint.as_deref() == Some("gpu-metrics")) + .expect("gpu metrics event should be emitted"); + assert_eq!(event.kind, "gpu.metrics"); + assert_eq!(event.source, "gpu_telemetry"); + assert_eq!(event.reason.as_deref(), Some("amd_smi_snapshot_available")); + assert_eq!( + event + .payload + .get("monitor_available") + .and_then(Value::as_bool), + Some(true) + ); + assert!( + event + .payload + .get("summary") + .and_then(Value::as_str) + .is_some_and(|summary| summary.contains("gpu_count=1")) + ); + Ok(()) + } + + #[test] + fn event_collector_emits_gpu_thermal_pressure_event_when_enabled() -> Result<()> { + let (root, paths) = temp_app_paths("event-bus-gpu-thermal-pressure"); + let state = test_runtime_state(vec![test_watcher_snapshot( + "gpu-thermal-protect", + WatcherMode::Propose, + None, + )]); + + let events = collect_automation_events_with_gpu_snapshot(&paths, &state, || { + CodexBridgeGpuSnapshot { + amd_smi_available: true, + static_snapshot: None, + monitor_snapshot: Some(json!({ + "gpu_data": [ + { + "gpu": 0, + "hotspot_temperature": { "value": 96.0 }, + "memory_temperature": { "value": 88.0 }, + "vram_percent": { "value": 72.0 } + } + ] + })), + note: None, + } + })?; + fs::remove_dir_all(root).ok(); + + let event = events + .iter() + .find(|event| event.watcher_hint.as_deref() == Some("gpu-thermal-protect")) + .expect("thermal pressure event should be emitted"); + assert_eq!(event.kind, "gpu.thermal_pressure"); + assert_eq!(event.source, "gpu_telemetry"); + assert_eq!( + event.reason.as_deref(), + Some("hotspot_temperature_threshold") + ); + assert!( + event + .payload + .get("summary") + .and_then(Value::as_str) + .is_some_and(|summary| summary.contains("GPU 0 hotspot temperature is 96 C")) + ); + Ok(()) + } + + #[test] + fn event_collector_skips_gpu_pressure_below_thresholds() -> Result<()> { + let (root, paths) = temp_app_paths("event-bus-gpu-pressure-cool"); + let state = test_runtime_state(vec![test_watcher_snapshot( + "gpu-thermal-protect", + WatcherMode::Propose, + None, + )]); + + let events = collect_automation_events_with_gpu_snapshot(&paths, &state, || { + CodexBridgeGpuSnapshot { + amd_smi_available: true, + static_snapshot: None, + monitor_snapshot: Some(json!([ + { + "gpu": 0, + "hotspot_temperature": 80.0, + "memory_temperature": 82.0, + "vram_percent": 50.0 + } + ])), + note: None, + } + })?; + fs::remove_dir_all(root).ok(); + + assert!( + !events + .iter() + .any(|event| event.watcher_hint.as_deref() == Some("gpu-thermal-protect")) + ); + Ok(()) + } + + #[test] + fn gpu_metrics_event_records_read_only_status_without_proposal() -> Result<()> { + let (root, paths) = temp_app_paths("gpu-metrics-record"); + paths.ensure()?; + let mut state = test_runtime_state(vec![test_watcher_snapshot( + "gpu-metrics", + WatcherMode::Contained, + None, + )]); + let mut config = RocmCliConfig::default(); + let watcher = config.watcher_config_mut("gpu-metrics"); + watcher.enabled = true; + watcher.mode = Some(WatcherMode::Contained); + let events = vec![AutomationTriggerEvent { + at_unix_ms: 1, + kind: "gpu.metrics_unavailable".to_owned(), + source: "gpu_telemetry".to_owned(), + watcher_hint: Some("gpu-metrics".to_owned()), + service_id: None, + reason: Some("amd-smi missing".to_owned()), + payload: json!({ + "summary": "amd_smi_available=false static_snapshot=missing monitor_snapshot=missing", + }), + }]; + + evaluate_watchers_for_events(&paths, &config, &mut state, &events)?; + let event_text = fs::read_to_string(paths.automation_events_path())?; + let event = serde_json::from_str::(event_text.trim())?; + let proposals = rocm_core::load_recent_automation_proposals(&paths, 1)?; + fs::remove_dir_all(root).ok(); + + assert_eq!(event.watcher_id, "gpu-metrics"); + assert_eq!(event.action, "record_gpu_metrics"); + assert!(event.message.contains("telemetry is recorded only")); + assert!(event.message.contains("amd-smi missing")); + assert!(event.message.contains("no mutating action was taken")); + assert!(proposals.is_empty()); + Ok(()) + } + + #[test] + fn gpu_thermal_protect_propose_queues_reviewed_stop_for_one_running_service() -> Result<()> { + let (root, paths) = temp_app_paths("gpu-thermal-protect-propose"); + paths.ensure()?; + let mut record = ManagedServiceRecord::new( + &paths, + "svc-hot", + "vllm", + "tiny", + "Tiny/Test", + "127.0.0.1", + 11435, + "managed", + 123, + None, + None, + None, + ); + record.status = "ready".to_owned(); + record.write()?; + let mut state = test_runtime_state(vec![test_watcher_snapshot( + "gpu-thermal-protect", + WatcherMode::Propose, + None, + )]); + let mut config = RocmCliConfig::default(); + let watcher = config.watcher_config_mut("gpu-thermal-protect"); + watcher.enabled = true; + watcher.mode = Some(WatcherMode::Propose); + let event = AutomationTriggerEvent { + at_unix_ms: 1, + kind: "gpu.thermal_pressure".to_owned(), + source: "gpu_telemetry".to_owned(), + watcher_hint: Some("gpu-thermal-protect".to_owned()), + service_id: None, + reason: Some("hotspot_temperature_threshold".to_owned()), + payload: json!({ + "gpu": 0, + "summary": "GPU 0 hotspot temperature is 96 C (limit 95 C)", + "hotspot_temperature_c": 96.0, + }), + }; + + evaluate_watchers_for_events(&paths, &config, &mut state, &[event])?; + let event_text = fs::read_to_string(paths.automation_events_path())?; + let event = serde_json::from_str::(event_text.trim())?; + let proposals = rocm_core::load_recent_automation_proposals(&paths, 1)?; + let saved = load_service_record(&paths, "svc-hot")?; + fs::remove_dir_all(root).ok(); + + assert_eq!(event.watcher_id, "gpu-thermal-protect"); + assert_eq!(event.action, "queue_stop_server_proposal"); + assert_eq!(event.service_id.as_deref(), Some("svc-hot")); + assert!(event.message.contains("asking before stopping anything")); + assert_eq!(saved.status, "ready"); + assert_eq!(proposals.len(), 1); + assert_eq!(proposals[0].tool.as_deref(), Some("stop_server")); + assert_eq!(proposals[0].service_id.as_deref(), Some("svc-hot")); + assert_eq!( + proposals[0] + .arguments + .get("pressure_reason") + .and_then(Value::as_str), + Some("hotspot_temperature_threshold") + ); + Ok(()) + } + + #[test] + fn gpu_thermal_protect_contained_still_queues_reviewed_stop() -> Result<()> { + let (root, paths) = temp_app_paths("gpu-thermal-protect-contained"); + paths.ensure()?; + let mut record = ManagedServiceRecord::new( + &paths, + "svc-hot", + "vllm", + "qwen", + "Qwen/Test", + "127.0.0.1", + 11436, + "managed", + 123, + None, + None, + None, + ); + record.status = "running".to_owned(); + record.write()?; + let mut state = test_runtime_state(vec![test_watcher_snapshot( + "gpu-thermal-protect", + WatcherMode::Contained, + None, + )]); + let event = AutomationTriggerEvent { + at_unix_ms: 1, + kind: "gpu.memory_pressure".to_owned(), + source: "gpu_telemetry".to_owned(), + watcher_hint: Some("gpu-thermal-protect".to_owned()), + service_id: Some("svc-hot".to_owned()), + reason: Some("vram_pressure_threshold".to_owned()), + payload: json!({ + "summary": "GPU 0 VRAM use is 96% (limit 95%)", + "vram_percent": 96.0, + }), + }; + + handle_gpu_thermal_protect_event(&paths, WatcherMode::Contained, &mut state, &event)?; + let event_text = fs::read_to_string(paths.automation_events_path())?; + let event = serde_json::from_str::(event_text.trim())?; + let proposals = rocm_core::load_recent_automation_proposals(&paths, 1)?; + let saved = load_service_record(&paths, "svc-hot")?; + fs::remove_dir_all(root).ok(); + + assert_eq!(event.action, "queue_stop_server_proposal"); + assert!(event.message.contains("contained mode still asks")); + assert_eq!(saved.status, "running"); + assert_eq!(proposals.len(), 1); + assert_eq!(proposals[0].tool.as_deref(), Some("stop_server")); + Ok(()) + } + + #[test] + fn gpu_thermal_protect_observe_records_without_proposal() -> Result<()> { + let (root, paths) = temp_app_paths("gpu-thermal-protect-observe"); + paths.ensure()?; + let mut state = test_runtime_state(vec![test_watcher_snapshot( + "gpu-thermal-protect", + WatcherMode::Observe, + None, + )]); + let event = AutomationTriggerEvent { + at_unix_ms: 1, + kind: "gpu.thermal_pressure".to_owned(), + source: "gpu_telemetry".to_owned(), + watcher_hint: Some("gpu-thermal-protect".to_owned()), + service_id: None, + reason: Some("memory_temperature_threshold".to_owned()), + payload: json!({ + "summary": "GPU memory temperature is high", + }), + }; + + handle_gpu_thermal_protect_event(&paths, WatcherMode::Observe, &mut state, &event)?; + let event_text = fs::read_to_string(paths.automation_events_path())?; + let event = serde_json::from_str::(event_text.trim())?; + let proposals = rocm_core::load_recent_automation_proposals(&paths, 1)?; + fs::remove_dir_all(root).ok(); + + assert_eq!(event.action, "observe_gpu_pressure"); + assert!(event.message.contains("does not stop any model server")); + assert!(proposals.is_empty()); + Ok(()) + } + + #[test] + fn gpu_thermal_protect_ambiguous_services_records_no_action() -> Result<()> { + let (root, paths) = temp_app_paths("gpu-thermal-protect-ambiguous"); + paths.ensure()?; + for service_id in ["svc-a", "svc-b"] { + let mut record = ManagedServiceRecord::new( + &paths, + service_id, + "vllm", + "qwen", + "Qwen/Test", + "127.0.0.1", + 11435, + "managed", + 123, + None, + None, + None, + ); + record.status = "ready".to_owned(); + record.write()?; + } + let mut state = test_runtime_state(vec![test_watcher_snapshot( + "gpu-thermal-protect", + WatcherMode::Propose, + None, + )]); + let event = AutomationTriggerEvent { + at_unix_ms: 1, + kind: "gpu.thermal_pressure".to_owned(), + source: "gpu_telemetry".to_owned(), + watcher_hint: Some("gpu-thermal-protect".to_owned()), + service_id: None, + reason: Some("hotspot_temperature_threshold".to_owned()), + payload: json!({ + "summary": "GPU 0 hotspot temperature is 96 C (limit 95 C)", + }), + }; + + handle_gpu_thermal_protect_event(&paths, WatcherMode::Propose, &mut state, &event)?; + let event_text = fs::read_to_string(paths.automation_events_path())?; + let event = serde_json::from_str::(event_text.trim())?; + let proposals = rocm_core::load_recent_automation_proposals(&paths, 1)?; + fs::remove_dir_all(root).ok(); + + assert_eq!(event.action, "gpu_pressure_no_clear_target"); + assert!(event.message.contains("did not choose a model server")); + assert!(proposals.is_empty()); + Ok(()) + } + + #[test] + fn gpu_thermal_protect_does_not_duplicate_pending_stop_proposals() -> Result<()> { + let (root, paths) = temp_app_paths("gpu-thermal-protect-dedupe"); + paths.ensure()?; + let mut record = ManagedServiceRecord::new( + &paths, + "svc-hot", + "vllm", + "tiny", + "Tiny/Test", + "127.0.0.1", + 11435, + "managed", + 123, + None, + None, + None, + ); + record.status = "ready".to_owned(); + record.write()?; + let mut state = test_runtime_state(vec![test_watcher_snapshot( + "gpu-thermal-protect", + WatcherMode::Propose, + None, + )]); + let event = AutomationTriggerEvent { + at_unix_ms: 1, + kind: "gpu.thermal_pressure".to_owned(), + source: "gpu_telemetry".to_owned(), + watcher_hint: Some("gpu-thermal-protect".to_owned()), + service_id: Some("svc-hot".to_owned()), + reason: Some("hotspot_temperature_threshold".to_owned()), + payload: json!({ + "summary": "GPU 0 hotspot temperature is 96 C (limit 95 C)", + }), + }; + + handle_gpu_thermal_protect_event(&paths, WatcherMode::Propose, &mut state, &event)?; + handle_gpu_thermal_protect_event(&paths, WatcherMode::Propose, &mut state, &event)?; + let events = rocm_core::load_recent_automation_events(&paths, 2)?; + let proposals = rocm_core::load_recent_automation_proposals(&paths, 10)?; + fs::remove_dir_all(root).ok(); + + assert_eq!(proposals.len(), 1); + assert_eq!(proposals[0].tool.as_deref(), Some("stop_server")); + assert!( + events + .iter() + .any(|event| event.action == "stop_proposal_already_pending") + ); + Ok(()) + } + + #[test] + fn local_webhook_gpu_metrics_event_uses_existing_read_only_policy() -> Result<()> { + let (root, paths) = temp_app_paths("local-webhook-gpu-metrics"); + paths.ensure()?; + let mut state = test_runtime_state(vec![test_watcher_snapshot( + "gpu-metrics", + WatcherMode::Contained, + None, + )]); + let mut config = RocmCliConfig::default(); + let watcher = config.watcher_config_mut("gpu-metrics"); + watcher.enabled = true; + watcher.mode = Some(WatcherMode::Contained); + let event = crate::webhook::local_webhook_event_from_request( + crate::webhook::LocalWebhookEventRequest { + watcher_hint: "gpu-metrics".to_owned(), + kind: "gpu.metrics".to_owned(), + service_id: None, + reason: Some("manual smoke".to_owned()), + payload: json!({ + "summary": "manual webhook probe", + "action": "restart_server", + "mode": "contained", + }), + }, + )?; + + evaluate_watchers_for_events(&paths, &config, &mut state, &[event])?; + let event_text = fs::read_to_string(paths.automation_events_path())?; + let event = serde_json::from_str::(event_text.trim())?; + let proposals = rocm_core::load_recent_automation_proposals(&paths, 1)?; + fs::remove_dir_all(root).ok(); + + assert_eq!(event.watcher_id, "gpu-metrics"); + assert_eq!(event.action, "record_gpu_metrics"); + assert!(event.message.contains("from local webhook")); + assert!(event.message.contains("manual smoke")); + assert!(event.message.contains("no mutating action was taken")); + assert!(proposals.is_empty()); + Ok(()) + } + + #[test] + fn cache_warm_propose_mode_queues_prefetch_proposal() -> Result<()> { + let (root, paths) = temp_app_paths("cache-warm-propose"); + paths.ensure()?; + let mut state = test_runtime_state(vec![test_watcher_snapshot( + "cache-warm", + WatcherMode::Propose, + None, + )]); + let event = crate::webhook::local_webhook_event_from_request( + crate::webhook::LocalWebhookEventRequest { + watcher_hint: "cache-warm".to_owned(), + kind: "cache.warm".to_owned(), + service_id: None, + reason: Some("idle window".to_owned()), + payload: json!({ + "artifact_ref": "Qwen/Test-1B#hf-main", + "tool": "restart_server", + "allow_artifact_download": true, + "artifact_max_bytes": 1024, + }), + }, + )?; + + handle_cache_warm_event_with_resolver( + &paths, + WatcherMode::Propose, + &mut state, + &event, + |_| Ok(true), + )?; + let event_text = fs::read_to_string(paths.automation_events_path())?; + let event = serde_json::from_str::(event_text.trim())?; + let proposals = rocm_core::load_recent_automation_proposals(&paths, 1)?; + fs::remove_dir_all(root).ok(); + + assert_eq!(event.watcher_id, "cache-warm"); + assert_eq!(event.action, "queue_prefetch_proposal"); + assert_eq!(proposals.len(), 1); + assert_eq!(proposals[0].watcher_id, "cache-warm"); + assert_eq!(proposals[0].tool.as_deref(), Some("prefetch_artifact")); + assert_eq!( + proposals[0] + .arguments + .get("artifact_ref") + .and_then(Value::as_str), + Some("Qwen/Test-1B#hf-main") + ); + assert!( + proposals[0].arguments.get("tool").is_none(), + "webhook payload must not grant arbitrary tool choice" + ); + assert!( + proposals[0] + .arguments + .get("allow_artifact_download") + .is_none(), + "webhook payload must not grant source-policy approval" + ); + assert!( + proposals[0].arguments.get("artifact_max_bytes").is_none(), + "webhook payload must not grant download byte-limit approval" + ); + Ok(()) + } + + #[test] + fn cache_warm_unknown_artifact_does_not_queue_proposal() -> Result<()> { + let (root, paths) = temp_app_paths("cache-warm-unknown"); + paths.ensure()?; + let mut state = test_runtime_state(vec![test_watcher_snapshot( + "cache-warm", + WatcherMode::Propose, + None, + )]); + let event = AutomationTriggerEvent { + at_unix_ms: 1, + kind: "cache.warm".to_owned(), + source: "test".to_owned(), + watcher_hint: Some("cache-warm".to_owned()), + service_id: None, + reason: Some("idle window".to_owned()), + payload: json!({ + "artifact_ref": "missing#artifact", + }), + }; + + handle_cache_warm_event_with_resolver( + &paths, + WatcherMode::Propose, + &mut state, + &event, + |_| Ok(false), + )?; + let event_text = fs::read_to_string(paths.automation_events_path())?; + let event = serde_json::from_str::(event_text.trim())?; + let proposals = rocm_core::load_recent_automation_proposals(&paths, 1)?; + fs::remove_dir_all(root).ok(); + + assert_eq!(event.action, "cache_warm_unknown_artifact"); + assert!(event.message.contains("unknown registry artifact")); + assert!(proposals.is_empty()); + Ok(()) + } + + #[test] + fn cache_warm_contained_mode_still_requires_reviewed_source_policy() -> Result<()> { + let (root, paths) = temp_app_paths("cache-warm-contained"); + paths.ensure()?; + let mut state = test_runtime_state(vec![test_watcher_snapshot( + "cache-warm", + WatcherMode::Contained, + None, + )]); + let event = AutomationTriggerEvent { + at_unix_ms: 1, + kind: "cache.warm".to_owned(), + source: "test".to_owned(), + watcher_hint: Some("cache-warm".to_owned()), + service_id: None, + reason: Some("idle window".to_owned()), + payload: json!({ + "artifact_ref": "Qwen/Test-1B#hf-main", + }), + }; + + handle_cache_warm_event_with_resolver( + &paths, + WatcherMode::Contained, + &mut state, + &event, + |_| Ok(true), + )?; + let event_text = fs::read_to_string(paths.automation_events_path())?; + let event = serde_json::from_str::(event_text.trim())?; + let proposals = rocm_core::load_recent_automation_proposals(&paths, 1)?; + fs::remove_dir_all(root).ok(); + + assert_eq!(event.action, "queue_prefetch_proposal"); + assert!(event.message.contains("explicit source-policy approval")); + assert_eq!(proposals.len(), 1); + assert_eq!(proposals[0].tool.as_deref(), Some("prefetch_artifact")); + Ok(()) + } + + #[test] + fn driver_upgrade_propose_mode_queues_driver_plan_proposal() -> Result<()> { + let (root, paths) = temp_app_paths("driver-upgrade-propose"); + paths.ensure()?; + let mut state = test_runtime_state(vec![test_watcher_snapshot( + "driver-upgrade", + WatcherMode::Propose, + None, + )]); + let event = crate::webhook::local_webhook_event_from_request( + crate::webhook::LocalWebhookEventRequest { + watcher_hint: "driver-upgrade".to_owned(), + kind: "update.available".to_owned(), + service_id: None, + reason: Some("driver version is newer".to_owned()), + payload: json!({ + "component": "driver", + "tool": "restart_server", + }), + }, + )?; + + handle_driver_upgrade_event(&paths, WatcherMode::Propose, &mut state, &event)?; + let event_text = fs::read_to_string(paths.automation_events_path())?; + let event = serde_json::from_str::(event_text.trim())?; + let proposals = rocm_core::load_recent_automation_proposals(&paths, 1)?; + fs::remove_dir_all(root).ok(); + + assert_eq!(event.watcher_id, "driver-upgrade"); + assert_eq!(event.action, "prepare_driver_plan"); + assert_eq!(proposals.len(), 1); + assert_eq!(proposals[0].watcher_id, "driver-upgrade"); + assert_eq!(proposals[0].tool.as_deref(), Some("driver_plan")); + assert!( + proposals[0].arguments.get("tool").is_none(), + "webhook payload must not grant arbitrary tool choice" + ); + Ok(()) + } + + #[test] + fn driver_upgrade_contained_mode_runs_restricted_driver_plan() -> Result<()> { + let (root, paths) = temp_app_paths("driver-upgrade-contained"); + paths.ensure()?; + let mut state = test_runtime_state(vec![test_watcher_snapshot( + "driver-upgrade", + WatcherMode::Contained, + None, + )]); + let event = AutomationTriggerEvent { + at_unix_ms: 1, + kind: "update.available".to_owned(), + source: "test".to_owned(), + watcher_hint: Some("driver-upgrade".to_owned()), + service_id: None, + reason: Some("driver version is newer".to_owned()), + payload: json!({ + "component": "driver", + }), + }; + + handle_driver_upgrade_event_with_runner( + &paths, + WatcherMode::Contained, + &mut state, + &event, + |_paths| { + Ok(crate::sandbox::sandbox_driver_plan_value( + crate::common::CommandCapture { + argv: vec![ + "rocm".to_owned(), + "install".to_owned(), + "driver".to_owned(), + "--dkms".to_owned(), + "--dry-run".to_owned(), + ], + exit_status: 0, + stdout: "driver install plan\n supported: true\n".to_owned(), + stderr: String::new(), + }, + )) + }, + )?; + let event_text = fs::read_to_string(paths.automation_events_path())?; + let event = serde_json::from_str::(event_text.trim())?; + let proposals = rocm_core::load_recent_automation_proposals(&paths, 1)?; + fs::remove_dir_all(root).ok(); + + assert_eq!(event.action, "run_driver_plan"); + assert!( + event + .message + .contains("contained restricted driver_plan status=planned") + ); + assert!(event.message.contains("no driver commands were executed")); + assert!(proposals.is_empty()); + Ok(()) + } + + #[test] + fn driver_upgrade_contained_mode_requires_restricted_driver_plan_tool() -> Result<()> { + let (root, paths) = temp_app_paths("driver-upgrade-contained-tool"); + paths.ensure()?; + let mut state = test_runtime_state(vec![test_watcher_snapshot( + "driver-upgrade", + WatcherMode::Contained, + None, + )]); + let event = AutomationTriggerEvent { + at_unix_ms: 1, + kind: "update.available".to_owned(), + source: "test".to_owned(), + watcher_hint: Some("driver-upgrade".to_owned()), + service_id: None, + reason: Some("driver version is newer".to_owned()), + payload: json!({ + "component": "driver", + }), + }; + + handle_driver_upgrade_event_with_runner( + &paths, + WatcherMode::Contained, + &mut state, + &event, + |_paths| { + Ok(json!({ + "tool": "check_updates", + "status": "checked", + "mutating": false, + })) + }, + )?; + let event_text = fs::read_to_string(paths.automation_events_path())?; + let event = serde_json::from_str::(event_text.trim())?; + let proposals = rocm_core::load_recent_automation_proposals(&paths, 1)?; + fs::remove_dir_all(root).ok(); + + assert_eq!(event.action, "driver_plan_failed"); + assert!(event.message.contains("expected `driver_plan`")); + assert!(event.message.contains("no driver commands were executed")); + assert!(proposals.is_empty()); + Ok(()) + } + + #[test] + fn driver_upgrade_observe_mode_records_without_proposal() -> Result<()> { + let (root, paths) = temp_app_paths("driver-upgrade-observe"); + paths.ensure()?; + let mut state = test_runtime_state(vec![test_watcher_snapshot( + "driver-upgrade", + WatcherMode::Observe, + None, + )]); + let event = AutomationTriggerEvent { + at_unix_ms: 1, + kind: "update.available".to_owned(), + source: "test".to_owned(), + watcher_hint: Some("driver-upgrade".to_owned()), + service_id: None, + reason: Some("driver version is newer".to_owned()), + payload: json!({ + "component": "driver", + }), + }; + + handle_driver_upgrade_event(&paths, WatcherMode::Observe, &mut state, &event)?; + let event_text = fs::read_to_string(paths.automation_events_path())?; + let event = serde_json::from_str::(event_text.trim())?; + let proposals = rocm_core::load_recent_automation_proposals(&paths, 1)?; + fs::remove_dir_all(root).ok(); + + assert_eq!(event.action, "observe_driver_update"); + assert!(event.message.contains("does not queue or run")); + assert!(proposals.is_empty()); + Ok(()) + } + + #[test] + fn driver_upgrade_ignores_non_driver_component_without_proposal() -> Result<()> { + let (root, paths) = temp_app_paths("driver-upgrade-wrong-component"); + paths.ensure()?; + let mut state = test_runtime_state(vec![test_watcher_snapshot( + "driver-upgrade", + WatcherMode::Propose, + None, + )]); + let event = AutomationTriggerEvent { + at_unix_ms: 1, + kind: "update.available".to_owned(), + source: "test".to_owned(), + watcher_hint: Some("driver-upgrade".to_owned()), + service_id: None, + reason: Some("runtime version is newer".to_owned()), + payload: json!({ + "component": "runtime", + }), + }; + + handle_driver_upgrade_event(&paths, WatcherMode::Propose, &mut state, &event)?; + let event_text = fs::read_to_string(paths.automation_events_path())?; + let event = serde_json::from_str::(event_text.trim())?; + let proposals = rocm_core::load_recent_automation_proposals(&paths, 1)?; + fs::remove_dir_all(root).ok(); + + assert_eq!(event.action, "driver_upgrade_ignored_component"); + assert!(event.message.contains("payload.component=driver")); + assert!(proposals.is_empty()); + Ok(()) + } + + #[test] + fn event_dispatcher_preserves_server_recover_proposal_behavior() -> Result<()> { + let (root, paths) = temp_app_paths("event-bus-dispatch"); + paths.ensure()?; + let mut failed = ManagedServiceRecord::new( + &paths, + "svc-failed", + "vllm", + "qwen", + "Qwen/Qwen3.5", + "127.0.0.1", + 11435, + "managed", + 123, + None, + None, + None, + ); + failed.status = "failed".to_owned(); + failed.write()?; + let mut state = test_runtime_state(vec![test_watcher_snapshot( + "server-recover", + WatcherMode::Propose, + None, + )]); + let mut config = RocmCliConfig::default(); + let watcher = config.watcher_config_mut("server-recover"); + watcher.enabled = true; + watcher.mode = Some(WatcherMode::Propose); + let events = vec![AutomationTriggerEvent { + at_unix_ms: 1, + kind: "service.manifest_recoverable".to_owned(), + source: "managed_service".to_owned(), + watcher_hint: Some("server-recover".to_owned()), + service_id: Some("svc-failed".to_owned()), + reason: Some("manifest_status_failed".to_owned()), + payload: json!({}), + }]; + + evaluate_watchers_for_events(&paths, &config, &mut state, &events)?; + let proposals = rocm_core::load_recent_automation_proposals(&paths, 1)?; + fs::remove_dir_all(root).ok(); + + assert_eq!(proposals.len(), 1); + assert_eq!(proposals[0].watcher_id, "server-recover"); + assert_eq!(proposals[0].service_id.as_deref(), Some("svc-failed")); + assert_eq!(proposals[0].tool.as_deref(), Some("restart_server")); + assert!(proposals[0].message.contains("manifest reports failed")); + assert!(!proposals[0].message.contains("manifest_status_failed")); + Ok(()) + } + + #[test] + fn server_recover_local_webhook_does_not_restart_healthy_service() -> Result<()> { + let (root, paths) = temp_app_paths("server-recover-healthy-webhook"); + paths.ensure()?; + let mut healthy = ManagedServiceRecord::new( + &paths, + "svc-healthy", + "vllm", + "qwen", + "Qwen/Qwen3.5", + "127.0.0.1", + 11435, + "managed", + 123, + None, + None, + None, + ); + healthy.status = "ready".to_owned(); + healthy.write()?; + let mut state = test_runtime_state(vec![test_watcher_snapshot( + "server-recover", + WatcherMode::Propose, + None, + )]); + let mut config = RocmCliConfig::default(); + let watcher = config.watcher_config_mut("server-recover"); + watcher.enabled = true; + watcher.mode = Some(WatcherMode::Propose); + let event = crate::webhook::local_webhook_event_from_request( + crate::webhook::LocalWebhookEventRequest { + watcher_hint: "server-recover".to_owned(), + kind: "service.manifest_recoverable".to_owned(), + service_id: Some("svc-healthy".to_owned()), + reason: Some("manual recovery smoke".to_owned()), + payload: json!({}), + }, + )?; + + evaluate_watchers_for_events(&paths, &config, &mut state, &[event])?; + let event_text = fs::read_to_string(paths.automation_events_path())?; + let event = serde_json::from_str::(event_text.trim())?; + let proposals = rocm_core::load_recent_automation_proposals(&paths, 1)?; + let reloaded = load_service_record(&paths, "svc-healthy")?; + fs::remove_dir_all(root).ok(); + + assert_eq!(event.action, "ignore_nonrecoverable_service"); + assert!(event.message.contains("does not currently need recovery")); + assert!(proposals.is_empty()); + assert_eq!(reloaded.status, "ready"); + Ok(()) + } + + #[test] + fn recovery_reason_display_avoids_raw_status_tokens() { + assert_eq!( + display_recovery_reason("manifest_status_starting_stale"), + "service has been starting for too long" + ); + assert_eq!( + display_recovery_reason("healthcheck_status_unreachable"), + "engine healthcheck reports unreachable" + ); + assert_eq!( + display_recovery_reason("endpoint_status_unreachable"), + "endpoint port is unreachable" + ); + } + #[test] + fn manifest_recovery_policy_covers_terminal_and_stale_transient_states() { + let (root, paths) = temp_app_paths("manifest-recovery-policy"); + let mut record = ManagedServiceRecord::new( + &paths, + "svc-stale", + "vllm", + "qwen", + "Qwen/Qwen3.5", + "127.0.0.1", + 11435, + "managed", + 123, + None, + None, + None, + ); + record.created_at_unix_ms = 1_000; + + record.status = "exited".to_owned(); + assert_eq!( + manifest_service_recovery_reason(&record, 1_001).as_deref(), + Some("manifest_status_exited") + ); + + record.status = "unreachable".to_owned(); + assert_eq!( + manifest_service_recovery_reason(&record, 1_001).as_deref(), + Some("manifest_status_unreachable") + ); + + record.status = "starting".to_owned(); + assert_eq!(manifest_service_recovery_reason(&record, 2_000), None); + assert_eq!( + manifest_service_recovery_reason(&record, 1_000 + SERVER_TRANSIENT_STALE_MS).as_deref(), + Some("manifest_status_starting_stale") + ); + + record.status = "recovering".to_owned(); + record.last_restart_unix_ms = Some(5_000); + assert_eq!(manifest_service_recovery_reason(&record, 6_000), None); + assert_eq!( + manifest_service_recovery_reason(&record, 5_000 + SERVER_TRANSIENT_STALE_MS).as_deref(), + Some("manifest_status_recovering_stale") + ); + fs::remove_dir_all(root).ok(); + } + + #[test] + fn find_recoverable_service_prefers_failed_managed_manifest() -> Result<()> { + let (root, paths) = temp_app_paths("recoverable-service"); + paths.ensure()?; + let mut failed = ManagedServiceRecord::new( + &paths, + "svc-failed", + "vllm", + "qwen", + "Qwen/Qwen3.5", + "127.0.0.1", + 11435, + "managed", + 123, + None, + None, + None, + ); + failed.status = "failed".to_owned(); + failed.write()?; + + let found = find_recoverable_service(&paths)?.expect("failed service should be found"); + fs::remove_dir_all(root).ok(); + assert_eq!(found.0.service_id, "svc-failed"); + assert_eq!(found.1, "manifest_status_failed"); + Ok(()) + } + + #[test] + fn find_recoverable_service_detects_stale_starting_manifest() -> Result<()> { + let (root, paths) = temp_app_paths("recoverable-stale-starting"); + paths.ensure()?; + let mut stale = ManagedServiceRecord::new( + &paths, + "svc-starting", + "vllm", + "qwen", + "Qwen/Qwen3.5", + "127.0.0.1", + 11435, + "managed", + 123, + None, + None, + None, + ); + stale.status = "starting".to_owned(); + stale.created_at_unix_ms = 0; + stale.write()?; + + let found = + find_recoverable_service(&paths)?.expect("stale starting service should recover"); + fs::remove_dir_all(root).ok(); + assert_eq!(found.0.service_id, "svc-starting"); + assert_eq!(found.1, "manifest_status_starting_stale"); + Ok(()) + } + + #[test] + fn server_recover_propose_mode_queues_restart_proposal() -> Result<()> { + let (root, paths) = temp_app_paths("server-recover-proposal"); + paths.ensure()?; + let mut record = ManagedServiceRecord::new( + &paths, + "svc-1", + "vllm", + "qwen", + "Qwen/Qwen3.5", + "127.0.0.1", + 11435, + "managed", + 123, + None, + None, + None, + ); + record.status = "failed".to_owned(); + record.write()?; + let mut state = AutomationRuntimeState { + running: true, + automations_enabled: true, + daemon_pid: 1, + started_at_unix_ms: 1, + last_tick_unix_ms: 1, + local_webhook_endpoint: None, + active_watchers: vec![WatcherRuntimeSnapshot { + id: "server-recover".to_owned(), + enabled: true, + mode: WatcherMode::Propose, + summary: "recover".to_owned(), + last_event: None, + last_event_unix_ms: None, + }], + }; + + evaluate_server_recover(&paths, WatcherMode::Propose, &mut state)?; + let proposals = rocm_core::load_recent_automation_proposals(&paths, 1)?; + fs::remove_dir_all(root).ok(); + + assert_eq!(proposals.len(), 1); + assert_eq!(proposals[0].watcher_id, "server-recover"); + assert_eq!(proposals[0].action, "queue_restart_proposal"); + assert_eq!(proposals[0].service_id.as_deref(), Some("svc-1")); + assert_eq!(proposals[0].status, "pending"); + Ok(()) + } + + #[test] + fn therock_update_contained_mode_runs_read_only_check_without_queueing() -> Result<()> { + let (root, paths) = temp_app_paths("therock-update-contained"); + paths.ensure()?; + let mut state = AutomationRuntimeState { + running: true, + automations_enabled: true, + daemon_pid: 1, + started_at_unix_ms: 1, + last_tick_unix_ms: 1, + local_webhook_endpoint: None, + active_watchers: vec![WatcherRuntimeSnapshot { + id: "therock-update".to_owned(), + enabled: true, + mode: WatcherMode::Contained, + summary: "check updates".to_owned(), + last_event: None, + last_event_unix_ms: None, + }], + }; + let event = AutomationTriggerEvent { + at_unix_ms: 42, + kind: "schedule.tick".to_owned(), + source: "scheduler".to_owned(), + watcher_hint: Some("therock-update".to_owned()), + service_id: None, + reason: Some("therock_update_interval_due".to_owned()), + payload: json!({ "interval_ms": THEROCK_UPDATE_INTERVAL_MS }), + }; + + handle_therock_update_event_with_runner( + &paths, + WatcherMode::Contained, + &mut state, + &event, + |_paths| { + Ok(crate::sandbox::sandbox_check_updates_value( + crate::common::CommandCapture { + argv: vec!["rocm".to_owned(), "update".to_owned()], + exit_status: 0, + stdout: "update\n managed runtimes: none\n".to_owned(), + stderr: String::new(), + }, + )) + }, + )?; + + let event_text = fs::read_to_string(paths.automation_events_path())?; + let event = serde_json::from_str::(event_text.trim())?; + let proposals = rocm_core::load_recent_automation_proposals(&paths, 1)?; + fs::remove_dir_all(root).ok(); + + assert_eq!(event.watcher_id, "therock-update"); + assert_eq!(event.action, "run_update_check"); + assert!(event.message.contains("contained read-only execution")); + assert!( + event + .message + .contains("restricted check_updates status=checked") + ); + assert!(event.message.contains("no updates were applied")); + assert!(!event.message.contains("fallback")); + assert!(proposals.is_empty()); + Ok(()) + } + + #[test] + fn therock_update_contained_mode_records_update_available_without_applying() -> Result<()> { + let (root, paths) = temp_app_paths("therock-update-contained-available"); + paths.ensure()?; + let mut state = AutomationRuntimeState { + running: true, + automations_enabled: true, + daemon_pid: 1, + started_at_unix_ms: 1, + last_tick_unix_ms: 1, + local_webhook_endpoint: None, + active_watchers: vec![WatcherRuntimeSnapshot { + id: "therock-update".to_owned(), + enabled: true, + mode: WatcherMode::Contained, + summary: "check updates".to_owned(), + last_event: None, + last_event_unix_ms: None, + }], + }; + let event = AutomationTriggerEvent { + at_unix_ms: 42, + kind: "schedule.tick".to_owned(), + source: "scheduler".to_owned(), + watcher_hint: Some("therock-update".to_owned()), + service_id: None, + reason: Some("therock_update_interval_due".to_owned()), + payload: json!({ "interval_ms": THEROCK_UPDATE_INTERVAL_MS }), + }; + + handle_therock_update_event_with_runner( + &paths, + WatcherMode::Contained, + &mut state, + &event, + |_paths| { + Ok(crate::sandbox::sandbox_check_updates_value(crate::common::CommandCapture { + argv: vec!["rocm".to_owned(), "update".to_owned()], + exit_status: 0, + stdout: "update\n runtime release-pip-gfx120x-all status=update_available installed=7.13.0 latest=7.14.0\n".to_owned(), + stderr: String::new(), + })) + }, + )?; + + let event_text = fs::read_to_string(paths.automation_events_path())?; + let events = event_text + .lines() + .map(serde_json::from_str::) + .collect::, _>>()?; + let audit_text = fs::read_to_string(paths.audit_events_path())?; + let proposals = rocm_core::load_recent_automation_proposals(&paths, 1)?; + fs::remove_dir_all(root).ok(); + + let update_check = events + .iter() + .find(|event| event.action == "run_update_check") + .expect("update check event should be recorded"); + assert_eq!(update_check.watcher_id, "therock-update"); + assert!( + update_check + .message + .contains("restricted check_updates status=update_available") + ); + assert!( + update_check + .message + .contains("a ROCm runtime update is available") + ); + assert!(update_check.message.contains("no updates were applied")); + assert!(!update_check.message.contains("fallback")); + let notification = events + .iter() + .find(|event| event.action == "notify_if_newer") + .expect("notify-if-newer event should be recorded"); + assert_eq!(notification.watcher_id, "therock-update"); + assert!( + notification + .message + .contains("ROCm runtime update is available") + ); + assert!(notification.message.contains("No updates were applied")); + assert!(audit_text.contains("\"category\":\"notification\"")); + assert!(audit_text.contains("\"action\":\"notify_if_newer\"")); + assert!(audit_text.contains("ROCm runtime update is available")); + assert!(proposals.is_empty()); + Ok(()) + } + + #[test] + fn therock_update_contained_mode_uses_restricted_check_updates_tool() -> Result<()> { + let (root, paths) = temp_app_paths("therock-update-contained-tool"); + paths.ensure()?; + let mut state = AutomationRuntimeState { + running: true, + automations_enabled: true, + daemon_pid: 1, + started_at_unix_ms: 1, + last_tick_unix_ms: 1, + local_webhook_endpoint: None, + active_watchers: vec![WatcherRuntimeSnapshot { + id: "therock-update".to_owned(), + enabled: true, + mode: WatcherMode::Contained, + summary: "check updates".to_owned(), + last_event: None, + last_event_unix_ms: None, + }], + }; + let event = AutomationTriggerEvent { + at_unix_ms: 42, + kind: "schedule.tick".to_owned(), + source: "scheduler".to_owned(), + watcher_hint: Some("therock-update".to_owned()), + service_id: None, + reason: Some("therock_update_interval_due".to_owned()), + payload: json!({ "interval_ms": THEROCK_UPDATE_INTERVAL_MS }), + }; + + handle_therock_update_event_with_runner( + &paths, + WatcherMode::Contained, + &mut state, + &event, + |_paths| { + Ok(json!({ + "tool": "examine_snapshot", + "status": "captured", + "mutating": false, + })) + }, + )?; + + let event_text = fs::read_to_string(paths.automation_events_path())?; + let event = serde_json::from_str::(event_text.trim())?; + fs::remove_dir_all(root).ok(); + + assert_eq!(event.watcher_id, "therock-update"); + assert_eq!(event.action, "update_check_failed"); + assert!(event.message.contains("expected `check_updates`")); + assert!(event.message.contains("no updates were applied")); + Ok(()) + } + + #[test] + fn therock_update_notify_if_newer_uses_restricted_notification_contract() -> Result<()> { + let (root, paths) = temp_app_paths("therock-update-notify-contract"); + paths.ensure()?; + let mut state = AutomationRuntimeState { + running: true, + automations_enabled: true, + daemon_pid: 1, + started_at_unix_ms: 1, + last_tick_unix_ms: 1, + local_webhook_endpoint: None, + active_watchers: vec![WatcherRuntimeSnapshot { + id: "therock-update".to_owned(), + enabled: true, + mode: WatcherMode::Contained, + summary: "check updates".to_owned(), + last_event: None, + last_event_unix_ms: None, + }], + }; + + record_update_available_notification(&paths, &mut state, "update_available")?; + + let audit_text = fs::read_to_string(paths.audit_events_path())?; + let audit = audit_text + .lines() + .map(serde_json::from_str::) + .collect::, _>>()?; + fs::remove_dir_all(root).ok(); + + let notification = audit + .iter() + .find(|event| event.category == "notification" && event.action == "notify_if_newer") + .expect("notify_if_newer audit should be recorded"); + assert_eq!(notification.category, "notification"); + assert_eq!(notification.actor, "watcher:therock-update"); + assert_eq!(notification.watcher_id.as_deref(), Some("therock-update")); + assert!( + notification + .message + .contains("ROCm runtime update is available") + ); + assert!(notification.message.contains("No updates were applied")); + Ok(()) + } + + #[test] + fn recovery_supervise_args_preserve_engine_recipe_json() { + let (_root, paths) = temp_app_paths("recovery-engine-recipe"); + let mut record = ManagedServiceRecord::new( + &paths, + "svc-1", + "vllm", + "qwen", + "Qwen/Qwen3.5-4B", + "127.0.0.1", + 11435, + "managed", + 123, + Some("therock-release:gfx120X-all".to_owned()), + Some("env-1".to_owned()), + Some("gpu_required".to_owned()), + ); + let engine_recipe_json = r#"{"contract_version":"0.1.0","engine":"vllm","required_flags":["--enable-auto-tool-choice"]}"#; + record.engine_recipe_json = Some(engine_recipe_json.to_owned()); + + let args = recovery_supervise_args(&record); + + assert!( + args.windows(2) + .any(|pair| pair[0] == "--engine-recipe-json" && pair[1] == engine_recipe_json) + ); + assert!( + args.windows(2) + .any(|pair| { pair[0] == "--canonical-model-id" && pair[1] == "Qwen/Qwen3.5-4B" }) + ); + } + + #[test] + fn watcher_policy_maps_modes_to_decisions() { + assert_eq!( + watcher_policy_action("server-recover", WatcherMode::Observe), + WatcherPolicyAction::Observe + ); + assert_eq!( + watcher_policy_action("server-recover", WatcherMode::Propose), + WatcherPolicyAction::QueueProposal + ); + assert_eq!( + watcher_policy_action("server-recover", WatcherMode::Contained), + WatcherPolicyAction::RunContained + ); + assert_eq!( + watcher_policy_action("therock-update", WatcherMode::Contained), + WatcherPolicyAction::RunContained + ); + } +} diff --git a/apps/rocmd/src/webhook.rs b/apps/rocmd/src/webhook.rs index d5807a66e..3c455b043 100644 --- a/apps/rocmd/src/webhook.rs +++ b/apps/rocmd/src/webhook.rs @@ -161,11 +161,13 @@ pub(crate) fn local_webhook_event_from_request( "server-recover" if service_id.unwrap_or_default().is_empty() => { bail!("server-recover webhook events require service_id"); } - "cache-warm" if crate::payload_string(&request.payload, "artifact_ref").is_none() => { + "cache-warm" + if crate::watchers::payload_string(&request.payload, "artifact_ref").is_none() => + { bail!("cache-warm webhook events require payload.artifact_ref"); } "driver-upgrade" - if crate::payload_string(&request.payload, "component").as_deref() + if crate::watchers::payload_string(&request.payload, "component").as_deref() != Some("driver") => { bail!("driver-upgrade webhook events require payload.component=driver"); diff --git a/docs/architecture.md b/docs/architecture.md index 59852f120..fa1e7eda8 100644 --- a/docs/architecture.md +++ b/docs/architecture.md @@ -31,7 +31,7 @@ Subsystem modules already following full domain extraction (each owns its own ty ### `apps/rocmd` — background daemon -`lib.rs` modularization is in progress (ROCMAI-83, Phase 5 of EAI-7768's sequencing). Extracted so far: `persistence.rs` (`record_event`/`load_managed_services`, the automation-event/audit-log and managed-service-registry I/O shared across the daemon's sandbox, MCP, service-lifecycle, and watcher code), `common.rs` (helpers shared across ≥2 of those remaining clusters: GPU/amd-smi snapshotting, the bridge-snapshot diagnostic, `CommandCapture`/command-timeout plumbing including the shared `rocm`-subprocess capture helpers the sandbox and MCP clusters both call, and small arg/healthcheck/endpoint-key utilities), `webhook.rs` (the local webhook source: its axum routes, request validation, and watcher-kind allow-list), `cli.rs` (the `Cli`/`Command` clap definitions, `SandboxToolArg`/`SandboxToolPolicy`, and the top-level dispatch in `run_cli`/`run_bin_cli`/`run_from_args` — the crate's only two externally-consumed entry points are re-exported from here via `lib.rs`'s `pub use`), `sandbox.rs` (bubblewrap/native sandbox execution, atomic-write helpers, artifact prefetch policy gating, and the sandbox-tool result shaping for `check_updates`/`driver_plan`), `mcp.rs` (the MCP stdio server, tool schema table, tool dispatch, and the `rocm`-subprocess capture/argv-building helpers behind the MCP tools), and `service.rs` (managed-service PID lifecycle/stop, `run_daemon`'s foreground loop, `supervise_service`'s spawn-and-recover path, and serve-log startup-phase polling). A helper earns a place in `common.rs` only once a second still-inline cluster calls it directly; a helper with exactly one caller stays in `lib.rs` next to that caller until its own cluster's extraction PR, even if it is conceptually similar to something that did move. Still pending: the watcher cluster itself — landing as its own PR. +Modularized (ROCMAI-83, Phase 5 of EAI-7768's sequencing): `persistence.rs` (`record_event`/`load_managed_services`, the automation-event/audit-log and managed-service-registry I/O shared across the daemon's sandbox, MCP, service-lifecycle, and watcher code), `common.rs` (helpers shared across ≥2 of the other clusters: GPU/amd-smi snapshotting, the bridge-snapshot diagnostic, `CommandCapture`/command-timeout plumbing including the shared `rocm`-subprocess capture helpers the sandbox and MCP clusters both call, and small arg/healthcheck/endpoint-key utilities), `webhook.rs` (the local webhook source: its axum routes, request validation, and watcher-kind allow-list), `cli.rs` (the `Cli`/`Command` clap definitions, `SandboxToolArg`/`SandboxToolPolicy`, the bridge-snapshot diagnostic printer, and the top-level dispatch in `run_cli`/`run_bin_cli`/`run_from_args`), `sandbox.rs` (bubblewrap/native sandbox execution, atomic-write helpers, artifact prefetch policy gating, and the sandbox-tool result shaping for `check_updates`/`driver_plan`), `mcp.rs` (the MCP stdio server, tool schema table, tool dispatch, and the `rocm`-subprocess capture/argv-building helpers behind the MCP tools), `service.rs` (managed-service PID lifecycle/stop, `run_daemon`'s foreground loop, `supervise_service`'s spawn-and-recover path, and serve-log startup-phase polling), and `watchers.rs` (event collection/dispatch for all built-in watchers — TheRock update, GPU metrics/thermal-pressure, cache-warm, driver-upgrade, server-recover — plus managed-service recovery classification and automation-proposal queuing). A helper earns a place in `common.rs` only once a second still-inline cluster calls it directly; a helper with exactly one caller stays next to that caller until its own cluster's extraction PR, even if it is conceptually similar to something that did move. `lib.rs` itself is now top-level glue: module declarations and the two externally-consumed entry points (`run_bin_cli`/`run_from_args`, re-exported from `cli.rs`). ### `crates/rocm-core` — core library