diff --git a/apps/rocmd/src/common.rs b/apps/rocmd/src/common.rs new file mode 100644 index 000000000..9bb7545b3 --- /dev/null +++ b/apps/rocmd/src/common.rs @@ -0,0 +1,657 @@ +// Copyright © Advanced Micro Devices, Inc., or its affiliates. +// +// SPDX-License-Identifier: MIT + +use anyhow::{Context, Result, bail}; +use rocm_core::{ + AppPaths, AutomationRuntimeState, CodexBridgeEngine, CodexBridgeGpuSnapshot, + CodexBridgeSnapshot, ExamineSummary, RocmCliConfig, default_engine_for_platform, + format_host_port, load_recent_automation_events, resolve_amd_smi_binary, unix_time_millis, +}; +#[cfg(test)] +use rocm_engine_protocol::EnginePluginDescriptor; +use rocm_engine_protocol::{ + EngineMethod, EngineRequestEnvelope, EngineResponseEnvelope, HealthcheckRequest, + HealthcheckResponse, +}; +use serde::Serialize; +use serde::de::DeserializeOwned; +use serde_json::Value; +use std::io::Write; +use std::net::{SocketAddr, TcpStream}; +#[cfg(test)] +use std::path::PathBuf; +use std::process::{Command as ProcessCommand, Stdio}; +use std::thread; +use std::time::Duration; + +const AMD_SMI_PROBE_TIMEOUT: Duration = Duration::from_secs(2); + +pub(crate) fn print_bridge_snapshot(paths: &AppPaths, pretty: bool) -> Result<()> { + let snapshot = build_bridge_snapshot(paths)?; + + if pretty { + println!( + "{}", + serde_json::to_string_pretty(&snapshot) + .context("failed to serialize bridge snapshot")? + ); + } else { + println!( + "{}", + serde_json::to_string(&snapshot).context("failed to serialize bridge snapshot")? + ); + } + + Ok(()) +} + +pub(crate) fn build_bridge_snapshot(paths: &AppPaths) -> Result { + let config = RocmCliConfig::load(paths).unwrap_or_default(); + Ok(CodexBridgeSnapshot { + protocol: "rocmd-codex-bridge-v0".to_owned(), + generated_at_unix_ms: unix_time_millis(), + examine: ExamineSummary::gather()?, + gpu: gather_gpu_snapshot_for_config(&config), + config, + automation_runtime: AutomationRuntimeState::load(paths)?, + recent_automation_events: load_recent_automation_events(paths, 32)?, + engines: bridge_engine_inventory(), + services: crate::persistence::load_managed_services(paths)?, + }) +} + +pub(crate) fn update_check_message(status: &str) -> &'static str { + match status { + "repair_available" => { + "ran read-only `rocm update`; a ROCm runtime repair is available because its package composition changed; no updates were applied" + } + "update_available" => { + "ran read-only `rocm update`; a ROCm runtime update is available; no updates were applied" + } + "error" => "read-only `rocm update` failed; no updates were applied", + _ => "ran read-only `rocm update`; no updates were applied", + } +} + +pub(crate) fn gather_gpu_snapshot() -> CodexBridgeGpuSnapshot { + let static_snapshot = match capture_amd_smi_json(&["static", "-a", "-g", "all", "--json"]) { + Ok(value) => Some(value), + Err(error) => { + return CodexBridgeGpuSnapshot { + amd_smi_available: false, + static_snapshot: None, + monitor_snapshot: None, + note: Some(error.to_string()), + }; + } + }; + + let monitor_snapshot = match capture_amd_smi_json(&[ + "monitor", "-p", "-t", "-u", "-m", "-v", "-g", "all", "--json", + ]) { + Ok(value) => Some(value), + Err(error) => { + return CodexBridgeGpuSnapshot { + amd_smi_available: true, + static_snapshot, + monitor_snapshot: None, + note: Some(error.to_string()), + }; + } + }; + + CodexBridgeGpuSnapshot { + amd_smi_available: true, + static_snapshot, + monitor_snapshot, + note: None, + } +} + +pub(crate) fn gather_gpu_snapshot_for_config(config: &RocmCliConfig) -> CodexBridgeGpuSnapshot { + if config.telemetry.local_inspection_enabled() { + gather_gpu_snapshot() + } else { + CodexBridgeGpuSnapshot { + amd_smi_available: false, + static_snapshot: None, + monitor_snapshot: None, + note: Some( + "gpu telemetry is disabled by rocm-cli config; no external reporting is implemented" + .to_owned(), + ), + } + } +} + +fn capture_amd_smi_json(args: &[&str]) -> Result { + let amd_smi_binary = resolve_amd_smi_binary(); + let mut command = ProcessCommand::new(&amd_smi_binary); + command + .args(args) + .stdin(Stdio::null()) + .stdout(Stdio::piped()) + .stderr(Stdio::piped()); + let output = run_command_with_timeout(command, AMD_SMI_PROBE_TIMEOUT) + .with_context(|| format!("failed to launch amd-smi {}", args.join(" ")))?; + + 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(); + anyhow::bail!( + "amd-smi {} failed: {}", + args.join(" "), + if !stderr.is_empty() { + stderr + } else if !stdout.is_empty() { + stdout + } else { + format!("exit status {}", output.status) + } + ); + } + + serde_json::from_slice(&output.stdout) + .with_context(|| format!("failed to parse amd-smi {} json", args.join(" "))) +} + +pub(crate) fn bridge_engine_inventory() -> Vec { + let default_engine = default_engine_for_platform(); + let current_exe = std::env::current_exe().ok(); + rocmd_engine_inventory() + .iter() + .map(|(id, summary)| CodexBridgeEngine { + id: (*id).to_owned(), + summary: (*summary).to_owned(), + default_for_platform: *id == default_engine, + installed_binary: true, + binary_path: current_exe.as_ref().map(|path| path.display().to_string()), + }) + .collect() +} + +const fn rocmd_engine_inventory() -> &'static [(&'static str, &'static str)] { + &[ + ( + "lemonade", + "embedded Lemonade server with ROCm llama.cpp backend", + ), + ( + "vllm", + "Linux/WSL ROCm GPU serving engine through external vLLM", + ), + ] +} + +#[cfg(test)] +pub(crate) fn find_engine_plugin_binary( + engine: &str, + plugin_dirs: I, +) -> Result> +where + I: IntoIterator, + P: AsRef, +{ + Ok(rocm_engine_protocol::discover_engine_plugins(plugin_dirs) + .context("failed to discover engine plugin binaries")? + .into_iter() + .find(|plugin: &EnginePluginDescriptor| plugin.id == engine) + .map(|plugin| plugin.executable_path)) +} + +#[derive(Debug)] +pub(crate) struct CommandCapture { + pub(crate) argv: Vec, + pub(crate) exit_status: i32, + pub(crate) stdout: String, + pub(crate) stderr: String, +} + +pub(crate) fn run_command_with_timeout( + mut command: ProcessCommand, + timeout: Duration, +) -> Result { + let mut child = command.spawn().context("failed to spawn child process")?; + let started = std::time::Instant::now(); + loop { + if child + .try_wait() + .context("failed to poll child process")? + .is_some() + { + return child + .wait_with_output() + .context("failed to collect child process output"); + } + if started.elapsed() >= timeout { + let _ = child.kill(); + let output = child + .wait_with_output() + .context("failed to collect timed-out child process output")?; + let stderr = String::from_utf8_lossy(&output.stderr).trim().to_owned(); + let stdout = String::from_utf8_lossy(&output.stdout).trim().to_owned(); + bail!( + "process exceeded {}s timeout: {}", + timeout.as_secs(), + if !stderr.is_empty() { + stderr + } else if !stdout.is_empty() { + stdout + } else { + "no output".to_owned() + } + ); + } + thread::sleep(Duration::from_millis(50)); + } +} + +pub(crate) 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]), + } +} + +pub(crate) fn optional_arg(flag: &str, value: Option<&str>) -> Vec { + match value { + Some(value) => vec![flag.to_owned(), value.to_owned()], + None => Vec::new(), + } +} + +pub(crate) fn engine_healthcheck_ready( + paths: &AppPaths, + engine: &str, + service_id: &str, +) -> Result { + Ok(healthcheck_response_ready(&engine_healthcheck_response( + paths, engine, service_id, + )?)) +} + +pub(crate) fn engine_healthcheck_response( + paths: &AppPaths, + engine: &str, + service_id: &str, +) -> Result { + engine_request::<_, HealthcheckResponse>( + paths, + engine, + service_id, + EngineMethod::Healthcheck, + &HealthcheckRequest { + service_id: service_id.to_owned(), + }, + ) +} + +pub(crate) fn healthcheck_response_ready(response: &HealthcheckResponse) -> bool { + response.status == "ready" && response.model_loaded +} + +pub(crate) fn healthcheck_response_recoverable(response: &HealthcheckResponse) -> bool { + matches!( + response.status.as_str(), + "failed" | "unreachable" | "exited" + ) +} + +/// Re-thread the endpoint key file (public bind only) onto an engine child's +/// environment, mirroring the initial `rocm serve` spawn. A loopback service +/// with no stored key leaves the command's environment untouched, matching the +/// unauthenticated default. +/// +/// Returns whether a key was applied. Callers that spawn a *listener* must +/// check it against the service's host — a public bind with no key would come +/// up anonymous (see [`ensure_public_service_has_endpoint_key`]). Callers that +/// only make a stdio plugin call can ignore it. +/// +/// The file must hold a *usable* key, not merely exist: the engine adapters +/// resolve it with `endpoint_api_key_from_file` and enforce nothing when that +/// yields `None`, so treating an empty or malformed file as "protected" would +/// let the guard pass while the endpoint served anonymously. +#[must_use] +pub(crate) fn apply_endpoint_key_env( + command: &mut ProcessCommand, + paths: &AppPaths, + service_id: &str, +) -> bool { + if let Some(key_file) = rocm_engine_protocol::endpoint_key_file_if_present(paths, service_id) + .and_then(|path| rocm_engine_protocol::endpoint_api_key_file_if_valid(&path)) + { + command.env(rocm_engine_protocol::ENDPOINT_API_KEY_FILE_ENV, key_file); + return true; + } + false +} + +/// Refuse to respawn a recorded service on a public host without its endpoint +/// API key, so daemon recovery cannot reopen a protected endpoint anonymously. +/// +/// Mirrors the guard of the same name in `rocm`; the shared +/// [`rocm_engine_protocol::is_public_bind_host`] keeps the two classifications +/// identical for a given `ManagedServiceRecord::host`. +pub(crate) fn ensure_public_service_has_endpoint_key( + host: &str, + key_present: bool, + requires_api_key: bool, +) -> Result<()> { + // The bind address is not the whole story. A service bound to loopback is + // only private until something republishes the port, and a tailnet publish + // outlives both the process and this daemon. The requirement is recorded on + // the service precisely so recovery can honour it without re-deriving it + // from an address that no longer answers the question. + if requires_api_key && !key_present { + bail!( + "refusing to recover a service that was launched with an endpoint API key but no \ + longer has one: it would come back up without authentication, and something \ + outside this machine may still be publishing its port. Relaunch it with \ + `rocm serve --require-api-key` to issue a new key." + ); + } + if rocm_engine_protocol::is_public_bind_host(host) && !key_present { + bail!( + "refusing to respawn a service bound to the public host `{host}` without an endpoint \ + API key: it would come back up without authentication. Relaunch it with \ + `rocm serve --host {host} --allow-public-bind` to issue a new key." + ); + } + Ok(()) +} + +pub(crate) fn engine_request( + paths: &AppPaths, + engine: &str, + service_id: &str, + method: EngineMethod, + request: &T, +) -> Result +where + T: Serialize, + R: DeserializeOwned, +{ + let envelope = EngineRequestEnvelope { + method, + payload: serde_json::to_value(request).context("failed to serialize engine request")?, + }; + let engine_binary = + std::env::current_exe().context("failed to resolve current rocm executable path")?; + let mut command = ProcessCommand::new(&engine_binary); + command + .arg("__engine-stdio") + .arg(engine) + .stdin(Stdio::piped()) + .stdout(Stdio::piped()) + .stderr(Stdio::piped()); + // A protected public endpoint enforces its key on every request, including + // this stdio-transported healthcheck — without the carrier the engine + // adapter's HTTP probe is unauthenticated and the endpoint looks + // unreachable, which would misclassify a healthy protected service as + // recoverable. + // + // No host to check here: this is a stdio plugin call, not a listener, so a + // missing key only weakens this probe rather than opening a port. + let _ = apply_endpoint_key_env(&mut command, paths, service_id); + let mut child = command.spawn().with_context(|| { + format!( + "failed to spawn engine stdio process {}", + engine_binary.display() + ) + })?; + + { + let stdin = child + .stdin + .as_mut() + .context("engine stdio child did not expose stdin")?; + serde_json::to_writer(&mut *stdin, &envelope).context("failed to write engine request")?; + stdin.write_all(b"\n")?; + } + + let output = child + .wait_with_output() + .context("failed waiting for engine stdio response")?; + if !output.status.success() { + let stderr = String::from_utf8_lossy(&output.stderr).trim().to_owned(); + if stderr.is_empty() { + bail!("engine stdio process exited with status {}", output.status); + } + bail!( + "engine stdio process exited with status {}: {}", + output.status, + stderr + ); + } + let envelope: EngineResponseEnvelope = + serde_json::from_slice(&output.stdout).context("failed to parse engine response")?; + if !envelope.ok { + let detail = envelope.error.map_or_else( + || "unknown engine error".to_owned(), + |error| format!("{}: {}", error.code, error.message), + ); + bail!(detail); + } + let data = envelope + .data + .context("engine response did not include data")?; + serde_json::from_value(data).context("failed to deserialize engine response data") +} + +pub(crate) fn wait_for_port(host: &str, port: u16, timeout: Duration) -> bool { + let address: SocketAddr = match format_host_port(host, port).parse() { + Ok(value) => value, + Err(_) => return false, + }; + + let start = std::time::Instant::now(); + while start.elapsed() < timeout { + if TcpStream::connect_timeout(&address, Duration::from_millis(200)).is_ok() { + return true; + } + thread::sleep(Duration::from_millis(200)); + } + false +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::test_support::temp_app_paths; + use rocm_core::engine_plugin_dirs; + use std::fs; + + #[test] + fn recovery_refuses_a_service_that_lost_a_key_it_was_launched_with() { + // The daemon keeps its own copy of this guard, and it only knew about + // public binds. A loopback service that something republishes — a + // tailnet publish outlives this daemon, let alone the process — would + // be recovered without authentication, and the rebuilt record would + // then disarm `rocm services restart` too. + for host in ["127.0.0.1", "localhost", "::1"] { + let error = ensure_public_service_has_endpoint_key(host, false, true) + .expect_err("a service launched with a key must not be recovered without one"); + assert!( + format!("{error:#}").contains("without authentication"), + "{error:#}" + ); + } + // With the key still present, recovery proceeds. + ensure_public_service_has_endpoint_key("127.0.0.1", true, true).unwrap(); + // And a service that never had one is untouched. + ensure_public_service_has_endpoint_key("127.0.0.1", false, false).unwrap(); + } + + #[test] + fn recovery_respawn_fails_closed_for_a_public_service_without_a_key() { + // Daemon recovery re-execs `rocmd supervise` for the recorded host. With + // the key gone the child would listen on that public host anonymously, + // so the spawn must be refused instead. + let error = ensure_public_service_has_endpoint_key("0.0.0.0", false, false).unwrap_err(); + assert!( + error.to_string().contains("without authentication"), + "{error:#}" + ); + + ensure_public_service_has_endpoint_key("0.0.0.0", true, false).unwrap(); + for host in ["127.0.0.1", "localhost", "::1"] { + ensure_public_service_has_endpoint_key(host, false, false) + .unwrap_or_else(|error| panic!("{host} must not require a key: {error:#}")); + } + } + + #[test] + fn apply_endpoint_key_env_sets_var_only_when_key_file_present() { + let (_root, paths) = temp_app_paths("apply-endpoint-key-env"); + let service_id = "svc-endpoint-key-env"; + + // Loopback service: no key file has ever been written, so the child's + // environment must be left untouched. + let mut command = ProcessCommand::new("true"); + assert!(!apply_endpoint_key_env(&mut command, &paths, service_id)); + assert!( + command + .get_envs() + .all(|(key, _)| key != rocm_engine_protocol::ENDPOINT_API_KEY_FILE_ENV), + "loopback service must not receive the endpoint key env var" + ); + + // Public service: once a key file exists at the deterministic path, + // it must be threaded onto the child so the engine's HTTP probe + // authenticates. + let key_path = rocm_engine_protocol::endpoint_key_file_path(&paths, service_id); + fs::create_dir_all(paths.services_dir()).unwrap(); + fs::write(&key_path, "secret-key").unwrap(); + let mut command = ProcessCommand::new("true"); + assert!(apply_endpoint_key_env(&mut command, &paths, service_id)); + let env_value = command + .get_envs() + .find_map(|(key, value)| { + (key == rocm_engine_protocol::ENDPOINT_API_KEY_FILE_ENV) + .then_some(value) + .flatten() + }) + .expect("endpoint key env var must be set once the key file exists"); + assert_eq!(env_value, key_path.as_os_str()); + + // An empty (or otherwise unusable) key file is not a key: the engine + // adapters would enforce nothing, so reporting `true` here would let the + // respawn guard pass and still open an anonymous public listener. + fs::write(&key_path, " \n").unwrap(); + let mut command = ProcessCommand::new("true"); + assert!(!apply_endpoint_key_env(&mut command, &paths, service_id)); + assert!( + command + .get_envs() + .all(|(key, _)| key != rocm_engine_protocol::ENDPOINT_API_KEY_FILE_ENV), + "an unusable key file must not be threaded onto the child" + ); + } + + #[test] + fn healthcheck_readiness_requires_ready_loaded_model() { + let ready = HealthcheckResponse { + status: "ready".to_owned(), + model_loaded: true, + device: "cuda".to_owned(), + uptime_sec: 1, + queue_depth: 0, + last_error: None, + tokens_per_sec: None, + }; + assert!(healthcheck_response_ready(&ready)); + + let mut loading = ready.clone(); + loading.status = "loading_model".to_owned(); + assert!(!healthcheck_response_ready(&loading)); + + let mut unloaded = ready; + unloaded.model_loaded = false; + assert!(!healthcheck_response_ready(&unloaded)); + } + + #[test] + fn healthcheck_recoverability_tracks_failed_endpoint_state() { + let mut response = HealthcheckResponse { + status: "ready".to_owned(), + model_loaded: true, + device: "cuda".to_owned(), + uptime_sec: 1, + queue_depth: 0, + last_error: None, + tokens_per_sec: None, + }; + assert!(!healthcheck_response_recoverable(&response)); + + response.status = "unreachable".to_owned(); + assert!(healthcheck_response_recoverable(&response)); + + response.status = "failed".to_owned(); + assert!(healthcheck_response_recoverable(&response)); + + response.status = "loading_model".to_owned(); + assert!(!healthcheck_response_recoverable(&response)); + + // The status the engines report for a model that is listed but has not + // yet served an inference request. Restarting it would kill a model + // mid-load and start the wait over. + response.status = "loading".to_owned(); + assert!(!healthcheck_response_recoverable(&response)); + } + + #[test] + fn healthcheck_readiness_withheld_while_the_model_only_lists() { + // What an engine reports once `/v1/models` answers but inference has not: + // not ready, so `rocm serve` keeps waiting instead of handing the caller + // an endpoint that will hang on its first request. + let listing_only = HealthcheckResponse { + status: "loading".to_owned(), + model_loaded: false, + device: "unknown".to_owned(), + uptime_sec: 1, + queue_depth: 0, + last_error: None, + tokens_per_sec: None, + }; + assert!(!healthcheck_response_ready(&listing_only)); + } + + #[test] + fn gpu_snapshot_respects_disabled_telemetry_policy() { + let mut config = RocmCliConfig::default(); + config.telemetry.mode = rocm_core::TELEMETRY_MODE_OFF.to_owned(); + + let snapshot = gather_gpu_snapshot_for_config(&config); + + assert!(!snapshot.amd_smi_available); + assert!(snapshot.static_snapshot.is_none()); + assert!(snapshot.monitor_snapshot.is_none()); + assert!( + snapshot + .note + .as_deref() + .is_some_and(|note| note.contains("disabled by rocm-cli config")) + ); + } + + #[test] + fn engine_plugin_discovery_finds_runtime_binary() -> Result<()> { + let (root, paths) = temp_app_paths("engine-plugin"); + let plugin_dir = paths.data_dir.join("engines").join("plugins"); + fs::create_dir_all(&plugin_dir)?; + let plugin_path = plugin_dir.join( + rocm_engine_protocol::platform_engine_plugin_binary_name("vllm"), + ); + fs::write(&plugin_path, "plugin")?; + + let discovered = find_engine_plugin_binary("vllm", engine_plugin_dirs(&paths))?; + fs::remove_dir_all(root).ok(); + + assert_eq!(discovered, Some(plugin_path)); + Ok(()) + } +} diff --git a/apps/rocmd/src/lib.rs b/apps/rocmd/src/lib.rs index ee107a0c4..d390df03f 100644 --- a/apps/rocmd/src/lib.rs +++ b/apps/rocmd/src/lib.rs @@ -4,6 +4,7 @@ #![allow(clippy::items_after_test_module)] +mod common; mod persistence; #[cfg(test)] mod test_support; @@ -17,25 +18,14 @@ use axum::{Json, Router}; use clap::{Parser, Subcommand, ValueEnum}; #[cfg(test)] use rocm_core::AutomationEventRecord; -#[cfg(test)] -use rocm_core::engine_plugin_dirs; use rocm_core::{ AppPaths, AuditEventRecord, AutomationProposalRecord, AutomationRuntimeState, - AutomationTriggerEvent, CodexBridgeEngine, CodexBridgeGpuSnapshot, CodexBridgeSnapshot, - DEFAULT_LOCAL_HOST, ExamineSummary, ManagedServiceRecord, ModelRecipeArtifactRecord, - RocmCliConfig, WatcherMode, WatcherRuntimeSnapshot, append_audit_event, - append_automation_proposal, builtin_watcher, builtin_watchers, daemon_binary_path, - default_engine_for_platform, format_host_port, load_recent_automation_events, - model_artifact_cache_status, resolve_amd_smi_binary, resolve_model_recipe_artifact, - unix_time_millis, -}; -#[cfg(test)] -use rocm_engine_protocol::EnginePluginDescriptor; -use rocm_engine_protocol::{ - EngineMethod, EngineRequestEnvelope, EngineResponseEnvelope, HealthcheckRequest, - HealthcheckResponse, + AutomationTriggerEvent, CodexBridgeGpuSnapshot, DEFAULT_LOCAL_HOST, ExamineSummary, + ManagedServiceRecord, ModelRecipeArtifactRecord, RocmCliConfig, WatcherMode, + WatcherRuntimeSnapshot, append_audit_event, append_automation_proposal, builtin_watcher, + builtin_watchers, daemon_binary_path, load_recent_automation_events, + model_artifact_cache_status, resolve_model_recipe_artifact, unix_time_millis, }; -use serde::de::DeserializeOwned; use serde::{Deserialize, Serialize}; use serde_json::Value; use serde_json::json; @@ -45,7 +35,6 @@ use std::collections::{HashSet, VecDeque}; use std::ffi::{OsStr, OsString}; use std::fs; use std::io::{self, BufRead, Read, Seek, SeekFrom, Write}; -use std::net::{SocketAddr, TcpStream}; use std::path::{Path, PathBuf}; use std::process::{Command as ProcessCommand, Stdio}; use std::thread; @@ -64,7 +53,6 @@ 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 AMD_SMI_PROBE_TIMEOUT: Duration = Duration::from_secs(2); const ARTIFACT_PREFETCH_TIMEOUT: Duration = Duration::from_mins(10); #[derive(Parser, Debug)] @@ -310,7 +298,7 @@ async fn run_cli(cli: Cli) -> Result<()> { print_status(&paths)?; } Command::BridgeSnapshot { pretty } => { - print_bridge_snapshot(&paths, pretty)?; + common::print_bridge_snapshot(&paths, pretty)?; } Command::SandboxRun { tool, @@ -387,40 +375,6 @@ async fn run_cli(cli: Cli) -> Result<()> { Ok(()) } -fn print_bridge_snapshot(paths: &AppPaths, pretty: bool) -> Result<()> { - let snapshot = build_bridge_snapshot(paths)?; - - if pretty { - println!( - "{}", - serde_json::to_string_pretty(&snapshot) - .context("failed to serialize bridge snapshot")? - ); - } else { - println!( - "{}", - serde_json::to_string(&snapshot).context("failed to serialize bridge snapshot")? - ); - } - - Ok(()) -} - -fn build_bridge_snapshot(paths: &AppPaths) -> Result { - let config = RocmCliConfig::load(paths).unwrap_or_default(); - Ok(CodexBridgeSnapshot { - protocol: "rocmd-codex-bridge-v0".to_owned(), - generated_at_unix_ms: unix_time_millis(), - examine: ExamineSummary::gather()?, - gpu: gather_gpu_snapshot_for_config(&config), - config, - automation_runtime: AutomationRuntimeState::load(paths)?, - recent_automation_events: load_recent_automation_events(paths, 32)?, - engines: bridge_engine_inventory(), - services: persistence::load_managed_services(paths)?, - }) -} - fn run_sandbox_runner( paths: &AppPaths, tool: SandboxToolArg, @@ -1197,11 +1151,11 @@ fn replace_file_windows(path: &Path, replacement: &Path) -> io::Result<()> { } } -fn sandbox_check_updates_value(output: CommandCapture) -> Value { +fn sandbox_check_updates_value(output: common::CommandCapture) -> Value { let status = update_check_status(&output); let update_available = output.exit_status == 0 && update_output_reports_update_available(&output.stdout); - let message = update_check_message(status); + let message = common::update_check_message(status); json!({ "tool": SandboxToolArg::CheckUpdates.as_cli_value(), "status": status, @@ -1215,7 +1169,7 @@ fn sandbox_check_updates_value(output: CommandCapture) -> Value { }) } -fn update_check_status(output: &CommandCapture) -> &'static str { +fn update_check_status(output: &common::CommandCapture) -> &'static str { if output.exit_status != 0 { "error" // A newer version outranks a composition repair: reporting the repair while @@ -1242,20 +1196,7 @@ fn update_output_reports_update_available(stdout: &str) -> bool { .any(|part| part == "update_available=true") } -fn update_check_message(status: &str) -> &'static str { - match status { - "repair_available" => { - "ran read-only `rocm update`; a ROCm runtime repair is available because its package composition changed; no updates were applied" - } - "update_available" => { - "ran read-only `rocm update`; a ROCm runtime update is available; no updates were applied" - } - "error" => "read-only `rocm update` failed; no updates were applied", - _ => "ran read-only `rocm update`; no updates were applied", - } -} - -fn sandbox_driver_plan_value(output: CommandCapture) -> Value { +fn sandbox_driver_plan_value(output: common::CommandCapture) -> Value { let status = if output.exit_status == 0 { "planned" } else { @@ -1422,129 +1363,6 @@ fn command_available(name: &str) -> bool { .is_ok() } -fn gather_gpu_snapshot() -> CodexBridgeGpuSnapshot { - let static_snapshot = match capture_amd_smi_json(&["static", "-a", "-g", "all", "--json"]) { - Ok(value) => Some(value), - Err(error) => { - return CodexBridgeGpuSnapshot { - amd_smi_available: false, - static_snapshot: None, - monitor_snapshot: None, - note: Some(error.to_string()), - }; - } - }; - - let monitor_snapshot = match capture_amd_smi_json(&[ - "monitor", "-p", "-t", "-u", "-m", "-v", "-g", "all", "--json", - ]) { - Ok(value) => Some(value), - Err(error) => { - return CodexBridgeGpuSnapshot { - amd_smi_available: true, - static_snapshot, - monitor_snapshot: None, - note: Some(error.to_string()), - }; - } - }; - - CodexBridgeGpuSnapshot { - amd_smi_available: true, - static_snapshot, - monitor_snapshot, - note: None, - } -} - -fn gather_gpu_snapshot_for_config(config: &RocmCliConfig) -> CodexBridgeGpuSnapshot { - if config.telemetry.local_inspection_enabled() { - gather_gpu_snapshot() - } else { - CodexBridgeGpuSnapshot { - amd_smi_available: false, - static_snapshot: None, - monitor_snapshot: None, - note: Some( - "gpu telemetry is disabled by rocm-cli config; no external reporting is implemented" - .to_owned(), - ), - } - } -} - -fn capture_amd_smi_json(args: &[&str]) -> Result { - let amd_smi_binary = resolve_amd_smi_binary(); - let mut command = ProcessCommand::new(&amd_smi_binary); - command - .args(args) - .stdin(Stdio::null()) - .stdout(Stdio::piped()) - .stderr(Stdio::piped()); - let output = run_command_with_timeout(command, AMD_SMI_PROBE_TIMEOUT) - .with_context(|| format!("failed to launch amd-smi {}", args.join(" ")))?; - - 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(); - anyhow::bail!( - "amd-smi {} failed: {}", - args.join(" "), - if !stderr.is_empty() { - stderr - } else if !stdout.is_empty() { - stdout - } else { - format!("exit status {}", output.status) - } - ); - } - - serde_json::from_slice(&output.stdout) - .with_context(|| format!("failed to parse amd-smi {} json", args.join(" "))) -} - -fn bridge_engine_inventory() -> Vec { - let default_engine = default_engine_for_platform(); - let current_exe = std::env::current_exe().ok(); - rocmd_engine_inventory() - .iter() - .map(|(id, summary)| CodexBridgeEngine { - id: (*id).to_owned(), - summary: (*summary).to_owned(), - default_for_platform: *id == default_engine, - installed_binary: true, - binary_path: current_exe.as_ref().map(|path| path.display().to_string()), - }) - .collect() -} - -const fn rocmd_engine_inventory() -> &'static [(&'static str, &'static str)] { - &[ - ( - "lemonade", - "embedded Lemonade server with ROCm llama.cpp backend", - ), - ( - "vllm", - "Linux/WSL ROCm GPU serving engine through external vLLM", - ), - ] -} - -#[cfg(test)] -fn find_engine_plugin_binary(engine: &str, plugin_dirs: I) -> Result> -where - I: IntoIterator, - P: AsRef, -{ - Ok(rocm_engine_protocol::discover_engine_plugins(plugin_dirs) - .context("failed to discover engine plugin binaries")? - .into_iter() - .find(|plugin: &EnginePluginDescriptor| plugin.id == engine) - .map(|plugin| plugin.executable_path)) -} - fn run_mcp_server(paths: &AppPaths) -> Result<()> { let stdin = io::stdin(); let stdout = io::stdout(); @@ -2111,7 +1929,7 @@ fn handle_mcp_tool_call(paths: &AppPaths, params: &Value) -> Result { } } "bridge_snapshot" => { - let snapshot = build_bridge_snapshot(paths)?; + let snapshot = common::build_bridge_snapshot(paths)?; Ok(tool_success( format!( "Captured bridge snapshot for {} / {} with default engine `{}`.", @@ -2122,7 +1940,7 @@ fn handle_mcp_tool_call(paths: &AppPaths, params: &Value) -> Result { } "gpu_snapshot" => { let config = RocmCliConfig::load(paths).unwrap_or_default(); - let gpu = gather_gpu_snapshot_for_config(&config); + 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 { @@ -2133,7 +1951,7 @@ fn handle_mcp_tool_call(paths: &AppPaths, params: &Value) -> Result { Ok(tool_success(status.to_owned(), json!(gpu))) } "engines" => { - let engines = bridge_engine_inventory(); + let engines = common::bridge_engine_inventory(); Ok(tool_success( format!("Found {} engine entries.", engines.len()), json!({ "engines": engines }), @@ -2335,7 +2153,7 @@ fn tool_error(text: String, structured: Value) -> Value { }) } -fn tool_result_from_command(prefix: &str, output: CommandCapture, is_error: bool) -> Value { +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": [ @@ -2354,7 +2172,7 @@ fn tool_result_from_command(prefix: &str, output: CommandCapture, is_error: bool }) } -fn command_capture_text(output: &CommandCapture) -> String { +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() { @@ -2368,15 +2186,7 @@ fn command_capture_text(output: &CommandCapture) -> String { } } -#[derive(Debug)] -struct CommandCapture { - argv: Vec, - exit_status: i32, - stdout: String, - stderr: String, -} - -fn run_rocm_capture(args: &[&str]) -> Result { +fn run_rocm_capture(args: &[&str]) -> Result { let paths = AppPaths::discover()?; run_rocm_capture_for_paths(&paths, args, Duration::from_mins(2)) } @@ -2385,7 +2195,7 @@ fn run_rocm_capture_for_paths( paths: &AppPaths, args: &[&str], timeout: Duration, -) -> Result { +) -> Result { let rocm_binary = rocm_core::daemon_binary_path()?; let mut command = ProcessCommand::new(&rocm_binary); command @@ -2396,9 +2206,9 @@ fn run_rocm_capture_for_paths( .stdin(Stdio::null()) .stdout(Stdio::piped()) .stderr(Stdio::piped()); - let output = run_command_with_timeout(command, timeout) + let output = common::run_command_with_timeout(command, timeout) .with_context(|| format!("failed to run {}", rocm_binary.display()))?; - Ok(CommandCapture { + Ok(common::CommandCapture { argv: std::iter::once(rocm_binary.display().to_string()) .chain(args.iter().map(|value| (*value).to_owned())) .collect(), @@ -2408,45 +2218,6 @@ fn run_rocm_capture_for_paths( }) } -fn run_command_with_timeout( - mut command: ProcessCommand, - timeout: Duration, -) -> Result { - let mut child = command.spawn().context("failed to spawn child process")?; - let started = std::time::Instant::now(); - loop { - if child - .try_wait() - .context("failed to poll child process")? - .is_some() - { - return child - .wait_with_output() - .context("failed to collect child process output"); - } - if started.elapsed() >= timeout { - let _ = child.kill(); - let output = child - .wait_with_output() - .context("failed to collect timed-out child process output")?; - let stderr = String::from_utf8_lossy(&output.stderr).trim().to_owned(); - let stdout = String::from_utf8_lossy(&output.stdout).trim().to_owned(); - bail!( - "process exceeded {}s timeout: {}", - timeout.as_secs(), - if !stderr.is_empty() { - stderr - } else if !stdout.is_empty() { - stdout - } else { - "no output".to_owned() - } - ); - } - thread::sleep(Duration::from_millis(50)); - } -} - 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()))?; @@ -3248,7 +3019,7 @@ fn supervise_service( 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 gpu_indices = common::parse_gpu_indices_arg(gpu.as_deref())?; let _ = daemon_binary_path(); let mut record = ManagedServiceRecord::new( @@ -3297,7 +3068,7 @@ fn supervise_service( // 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. - ensure_public_service_has_endpoint_key( + 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)) @@ -3340,8 +3111,9 @@ fn supervise_service( // 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 = apply_endpoint_key_env(&mut command, paths, &record.service_id); - ensure_public_service_has_endpoint_key( + 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, @@ -3421,9 +3193,12 @@ fn engine_serve_http_args( if let Some(csv) = rocm_engine_protocol::gpu_indices_to_csv(gpu_indices) { args.extend(["--gpu".to_owned(), csv]); } - args.extend(optional_arg("--runtime-id", runtime_id)); - args.extend(optional_arg("--env-id", env_id)); - args.extend(optional_arg("--engine-recipe-json", engine_recipe_json)); + 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 } @@ -3682,7 +3457,7 @@ fn collect_automation_events( state: &AutomationRuntimeState, ) -> Result> { collect_automation_events_with_gpu_snapshot(paths, state, || { - gather_gpu_snapshot_for_config(config) + common::gather_gpu_snapshot_for_config(config) }) } @@ -3930,7 +3705,7 @@ where &format!( "{message}; restricted check_updates status={}; {}", result.status, - update_check_message(result.status) + common::update_check_message(result.status) ), None, )?; @@ -4797,8 +4572,8 @@ fn service_record_matches_recovery_event( } "service.endpoint_recoverable" => endpoint_service_recovery_reason(record).is_some(), "service.healthcheck_recoverable" => { - engine_healthcheck_response(paths, &record.engine, &record.service_id) - .is_ok_and(|healthcheck| healthcheck_response_recoverable(&healthcheck)) + common::engine_healthcheck_response(paths, &record.engine, &record.service_id) + .is_ok_and(|healthcheck| common::healthcheck_response_recoverable(&healthcheck)) } _ => false, } @@ -4862,7 +4637,7 @@ fn handle_server_recover_event_with_record( // 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) = ensure_public_service_has_endpoint_key( + 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)) @@ -4965,14 +4740,14 @@ fn find_recoverable_service(paths: &AppPaths) -> Result Result Option { - (!wait_for_port(&record.host, record.port, ENDPOINT_HEALTH_TIMEOUT)) + (!common::wait_for_port(&record.host, record.port, ENDPOINT_HEALTH_TIMEOUT)) .then(|| "endpoint_status_unreachable".to_owned()) } @@ -5101,12 +4876,15 @@ fn recovery_supervise_args(record: &ManagedServiceRecord) -> Vec { .unwrap_or("gpu_required") .to_owned(), ]; - args.extend(optional_arg("--runtime-id", record.runtime_id.as_deref())); - args.extend(optional_arg("--env-id", record.env_id.as_deref())); + 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(optional_arg( + args.extend(common::optional_arg( "--engine-recipe-json", record.engine_recipe_json.as_deref(), )); @@ -5204,27 +4982,6 @@ mod tests { use rocm_core::ModelRecipeArtifactSourcePolicyRecord; use std::path::PathBuf; - #[test] - fn recovery_refuses_a_service_that_lost_a_key_it_was_launched_with() { - // The daemon keeps its own copy of this guard, and it only knew about - // public binds. A loopback service that something republishes — a - // tailnet publish outlives this daemon, let alone the process — would - // be recovered without authentication, and the rebuilt record would - // then disarm `rocm services restart` too. - for host in ["127.0.0.1", "localhost", "::1"] { - let error = super::ensure_public_service_has_endpoint_key(host, false, true) - .expect_err("a service launched with a key must not be recovered without one"); - assert!( - format!("{error:#}").contains("without authentication"), - "{error:#}" - ); - } - // With the key still present, recovery proceeds. - super::ensure_public_service_has_endpoint_key("127.0.0.1", true, true).unwrap(); - // And a service that never had one is untouched. - super::ensure_public_service_has_endpoint_key("127.0.0.1", false, false).unwrap(); - } - #[test] fn remote_read_only_verbs_are_allowed_and_mutating_ones_are_not() { let allow = |args: &[&str]| { @@ -6057,24 +5814,6 @@ mod tests { Ok(()) } - #[test] - fn gpu_snapshot_respects_disabled_telemetry_policy() { - let mut config = RocmCliConfig::default(); - config.telemetry.mode = rocm_core::TELEMETRY_MODE_OFF.to_owned(); - - let snapshot = gather_gpu_snapshot_for_config(&config); - - assert!(!snapshot.amd_smi_available); - assert!(snapshot.static_snapshot.is_none()); - assert!(snapshot.monitor_snapshot.is_none()); - assert!( - snapshot - .note - .as_deref() - .is_some_and(|note| note.contains("disabled by rocm-cli config")) - ); - } - #[test] fn read_tail_lines_returns_last_lines_only() -> Result<()> { let path = unique_test_path(&format!( @@ -6089,23 +5828,6 @@ mod tests { Ok(()) } - #[test] - fn engine_plugin_discovery_finds_runtime_binary() -> Result<()> { - let (root, paths) = temp_app_paths("engine-plugin"); - let plugin_dir = paths.data_dir.join("engines").join("plugins"); - fs::create_dir_all(&plugin_dir)?; - let plugin_path = plugin_dir.join( - rocm_engine_protocol::platform_engine_plugin_binary_name("vllm"), - ); - fs::write(&plugin_path, "plugin")?; - - let discovered = find_engine_plugin_binary("vllm", engine_plugin_dirs(&paths))?; - fs::remove_dir_all(root).ok(); - - assert_eq!(discovered, Some(plugin_path)); - Ok(()) - } - #[test] fn launch_server_rejects_public_bind_without_ack() { let arguments = serde_json::Map::from_iter([ @@ -7443,7 +7165,7 @@ mod tests { &mut state, &event, |_paths| { - Ok(sandbox_driver_plan_value(CommandCapture { + Ok(sandbox_driver_plan_value(common::CommandCapture { argv: vec![ "rocm".to_owned(), "install".to_owned(), @@ -7704,140 +7426,6 @@ mod tests { ); } - #[test] - fn healthcheck_readiness_requires_ready_loaded_model() { - let ready = HealthcheckResponse { - status: "ready".to_owned(), - model_loaded: true, - device: "cuda".to_owned(), - uptime_sec: 1, - queue_depth: 0, - last_error: None, - tokens_per_sec: None, - }; - assert!(healthcheck_response_ready(&ready)); - - let mut loading = ready.clone(); - loading.status = "loading_model".to_owned(); - assert!(!healthcheck_response_ready(&loading)); - - let mut unloaded = ready; - unloaded.model_loaded = false; - assert!(!healthcheck_response_ready(&unloaded)); - } - - #[test] - fn healthcheck_recoverability_tracks_failed_endpoint_state() { - let mut response = HealthcheckResponse { - status: "ready".to_owned(), - model_loaded: true, - device: "cuda".to_owned(), - uptime_sec: 1, - queue_depth: 0, - last_error: None, - tokens_per_sec: None, - }; - assert!(!healthcheck_response_recoverable(&response)); - - response.status = "unreachable".to_owned(); - assert!(healthcheck_response_recoverable(&response)); - - response.status = "failed".to_owned(); - assert!(healthcheck_response_recoverable(&response)); - - response.status = "loading_model".to_owned(); - assert!(!healthcheck_response_recoverable(&response)); - - // The status the engines report for a model that is listed but has not - // yet served an inference request. Restarting it would kill a model - // mid-load and start the wait over. - response.status = "loading".to_owned(); - assert!(!healthcheck_response_recoverable(&response)); - } - - #[test] - fn healthcheck_readiness_withheld_while_the_model_only_lists() { - // What an engine reports once `/v1/models` answers but inference has not: - // not ready, so `rocm serve` keeps waiting instead of handing the caller - // an endpoint that will hang on its first request. - let listing_only = HealthcheckResponse { - status: "loading".to_owned(), - model_loaded: false, - device: "unknown".to_owned(), - uptime_sec: 1, - queue_depth: 0, - last_error: None, - tokens_per_sec: None, - }; - assert!(!healthcheck_response_ready(&listing_only)); - } - - #[test] - fn apply_endpoint_key_env_sets_var_only_when_key_file_present() { - let (_root, paths) = temp_app_paths("apply-endpoint-key-env"); - let service_id = "svc-endpoint-key-env"; - - // Loopback service: no key file has ever been written, so the child's - // environment must be left untouched. - let mut command = ProcessCommand::new("true"); - assert!(!apply_endpoint_key_env(&mut command, &paths, service_id)); - assert!( - command - .get_envs() - .all(|(key, _)| key != rocm_engine_protocol::ENDPOINT_API_KEY_FILE_ENV), - "loopback service must not receive the endpoint key env var" - ); - - // Public service: once a key file exists at the deterministic path, - // it must be threaded onto the child so the engine's HTTP probe - // authenticates. - let key_path = rocm_engine_protocol::endpoint_key_file_path(&paths, service_id); - fs::create_dir_all(paths.services_dir()).unwrap(); - fs::write(&key_path, "secret-key").unwrap(); - let mut command = ProcessCommand::new("true"); - assert!(apply_endpoint_key_env(&mut command, &paths, service_id)); - let env_value = command - .get_envs() - .find_map(|(key, value)| { - (key == rocm_engine_protocol::ENDPOINT_API_KEY_FILE_ENV) - .then_some(value) - .flatten() - }) - .expect("endpoint key env var must be set once the key file exists"); - assert_eq!(env_value, key_path.as_os_str()); - - // An empty (or otherwise unusable) key file is not a key: the engine - // adapters would enforce nothing, so reporting `true` here would let the - // respawn guard pass and still open an anonymous public listener. - fs::write(&key_path, " \n").unwrap(); - let mut command = ProcessCommand::new("true"); - assert!(!apply_endpoint_key_env(&mut command, &paths, service_id)); - assert!( - command - .get_envs() - .all(|(key, _)| key != rocm_engine_protocol::ENDPOINT_API_KEY_FILE_ENV), - "an unusable key file must not be threaded onto the child" - ); - } - - #[test] - fn recovery_respawn_fails_closed_for_a_public_service_without_a_key() { - // Daemon recovery re-execs `rocmd supervise` for the recorded host. With - // the key gone the child would listen on that public host anonymously, - // so the spawn must be refused instead. - let error = ensure_public_service_has_endpoint_key("0.0.0.0", false, false).unwrap_err(); - assert!( - error.to_string().contains("without authentication"), - "{error:#}" - ); - - ensure_public_service_has_endpoint_key("0.0.0.0", true, false).unwrap(); - for host in ["127.0.0.1", "localhost", "::1"] { - ensure_public_service_has_endpoint_key(host, false, false) - .unwrap_or_else(|error| panic!("{host} must not require a key: {error:#}")); - } - } - #[test] fn manifest_recovery_policy_covers_terminal_and_stale_transient_states() { let (root, paths) = temp_app_paths("manifest-recovery-policy"); @@ -8029,7 +7617,7 @@ mod tests { &mut state, &event, |_paths| { - Ok(sandbox_check_updates_value(CommandCapture { + Ok(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(), @@ -8093,7 +7681,7 @@ mod tests { &mut state, &event, |_paths| { - Ok(sandbox_check_updates_value(CommandCapture { + Ok(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(), @@ -8249,7 +7837,7 @@ mod tests { #[test] fn sandbox_check_updates_value_is_read_only_and_preserves_output() { - let value = sandbox_check_updates_value(CommandCapture { + let value = sandbox_check_updates_value(common::CommandCapture { argv: vec!["rocm".to_owned(), "update".to_owned()], exit_status: 0, stdout: "update\n runtime release status=up_to_date\n".to_owned(), @@ -8278,7 +7866,7 @@ mod tests { #[test] fn sandbox_check_updates_value_marks_runtime_update_available() { - let value = sandbox_check_updates_value(CommandCapture { + let value = 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(), @@ -8317,7 +7905,7 @@ mod tests { #[test] fn sandbox_check_updates_value_marks_runtime_repair_available() { - let value = sandbox_check_updates_value(CommandCapture { + let value = sandbox_check_updates_value(common::CommandCapture { argv: vec!["rocm".to_owned(), "update".to_owned()], exit_status: 0, stdout: "update\n runtime release-wheel-multi-arch-7-14-0 status=repair_available installed=7.14.0 latest=7.14.0\n".to_owned(), @@ -8345,7 +7933,7 @@ mod tests { #[test] fn a_newer_version_outranks_a_composition_repair_in_the_watcher_report() { - let value = sandbox_check_updates_value(CommandCapture { + let value = sandbox_check_updates_value(common::CommandCapture { argv: vec!["rocm".to_owned(), "update".to_owned()], exit_status: 0, stdout: "update\n runtime old status=repair_available installed=7.14.0 latest=7.14.0\n runtime stale status=update_available installed=7.13.0 latest=7.14.0\n".to_owned(), @@ -8360,7 +7948,7 @@ mod tests { #[test] fn sandbox_driver_plan_value_is_read_only_and_preserves_output() { - let value = sandbox_driver_plan_value(CommandCapture { + let value = sandbox_driver_plan_value(common::CommandCapture { argv: vec![ "rocm".to_owned(), "install".to_owned(), @@ -9540,23 +9128,6 @@ async fn shutdown_signal() { let _ = tokio::signal::ctrl_c().await; } -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]), - } -} - -fn optional_arg(flag: &str, value: Option<&str>) -> Vec { - match value { - Some(value) => vec![flag.to_owned(), value.to_owned()], - None => Vec::new(), - } -} - /// 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 @@ -9644,197 +9215,7 @@ fn wait_for_service_ready( 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 -} - -fn engine_healthcheck_ready(paths: &AppPaths, engine: &str, service_id: &str) -> Result { - Ok(healthcheck_response_ready(&engine_healthcheck_response( - paths, engine, service_id, - )?)) -} - -fn engine_healthcheck_response( - paths: &AppPaths, - engine: &str, - service_id: &str, -) -> Result { - engine_request::<_, HealthcheckResponse>( - paths, - engine, - service_id, - EngineMethod::Healthcheck, - &HealthcheckRequest { - service_id: service_id.to_owned(), - }, - ) -} - -fn healthcheck_response_ready(response: &HealthcheckResponse) -> bool { - response.status == "ready" && response.model_loaded -} - -fn healthcheck_response_recoverable(response: &HealthcheckResponse) -> bool { - matches!( - response.status.as_str(), - "failed" | "unreachable" | "exited" - ) -} - -/// Re-thread the endpoint key file (public bind only) onto an engine child's -/// environment, mirroring the initial `rocm serve` spawn. A loopback service -/// with no stored key leaves the command's environment untouched, matching the -/// unauthenticated default. -/// -/// Returns whether a key was applied. Callers that spawn a *listener* must -/// check it against the service's host — a public bind with no key would come -/// up anonymous (see [`ensure_public_service_has_endpoint_key`]). Callers that -/// only make a stdio plugin call can ignore it. -/// -/// The file must hold a *usable* key, not merely exist: the engine adapters -/// resolve it with `endpoint_api_key_from_file` and enforce nothing when that -/// yields `None`, so treating an empty or malformed file as "protected" would -/// let the guard pass while the endpoint served anonymously. -#[must_use] -fn apply_endpoint_key_env( - command: &mut ProcessCommand, - paths: &AppPaths, - service_id: &str, -) -> bool { - if let Some(key_file) = rocm_engine_protocol::endpoint_key_file_if_present(paths, service_id) - .and_then(|path| rocm_engine_protocol::endpoint_api_key_file_if_valid(&path)) - { - command.env(rocm_engine_protocol::ENDPOINT_API_KEY_FILE_ENV, key_file); - return true; - } - false -} - -/// Refuse to respawn a recorded service on a public host without its endpoint -/// API key, so daemon recovery cannot reopen a protected endpoint anonymously. -/// -/// Mirrors the guard of the same name in `rocm`; the shared -/// [`rocm_engine_protocol::is_public_bind_host`] keeps the two classifications -/// identical for a given `ManagedServiceRecord::host`. -fn ensure_public_service_has_endpoint_key( - host: &str, - key_present: bool, - requires_api_key: bool, -) -> Result<()> { - // The bind address is not the whole story. A service bound to loopback is - // only private until something republishes the port, and a tailnet publish - // outlives both the process and this daemon. The requirement is recorded on - // the service precisely so recovery can honour it without re-deriving it - // from an address that no longer answers the question. - if requires_api_key && !key_present { - bail!( - "refusing to recover a service that was launched with an endpoint API key but no \ - longer has one: it would come back up without authentication, and something \ - outside this machine may still be publishing its port. Relaunch it with \ - `rocm serve --require-api-key` to issue a new key." - ); - } - if rocm_engine_protocol::is_public_bind_host(host) && !key_present { - bail!( - "refusing to respawn a service bound to the public host `{host}` without an endpoint \ - API key: it would come back up without authentication. Relaunch it with \ - `rocm serve --host {host} --allow-public-bind` to issue a new key." - ); - } - Ok(()) -} - -fn engine_request( - paths: &AppPaths, - engine: &str, - service_id: &str, - method: EngineMethod, - request: &T, -) -> Result -where - T: Serialize, - R: DeserializeOwned, -{ - let envelope = EngineRequestEnvelope { - method, - payload: serde_json::to_value(request).context("failed to serialize engine request")?, - }; - let engine_binary = - std::env::current_exe().context("failed to resolve current rocm executable path")?; - let mut command = ProcessCommand::new(&engine_binary); - command - .arg("__engine-stdio") - .arg(engine) - .stdin(Stdio::piped()) - .stdout(Stdio::piped()) - .stderr(Stdio::piped()); - // A protected public endpoint enforces its key on every request, including - // this stdio-transported healthcheck — without the carrier the engine - // adapter's HTTP probe is unauthenticated and the endpoint looks - // unreachable, which would misclassify a healthy protected service as - // recoverable. - // - // No host to check here: this is a stdio plugin call, not a listener, so a - // missing key only weakens this probe rather than opening a port. - let _ = apply_endpoint_key_env(&mut command, paths, service_id); - let mut child = command.spawn().with_context(|| { - format!( - "failed to spawn engine stdio process {}", - engine_binary.display() - ) - })?; - - { - let stdin = child - .stdin - .as_mut() - .context("engine stdio child did not expose stdin")?; - serde_json::to_writer(&mut *stdin, &envelope).context("failed to write engine request")?; - stdin.write_all(b"\n")?; - } - - let output = child - .wait_with_output() - .context("failed waiting for engine stdio response")?; - if !output.status.success() { - let stderr = String::from_utf8_lossy(&output.stderr).trim().to_owned(); - if stderr.is_empty() { - bail!("engine stdio process exited with status {}", output.status); - } - bail!( - "engine stdio process exited with status {}: {}", - output.status, - stderr - ); - } - let envelope: EngineResponseEnvelope = - serde_json::from_slice(&output.stdout).context("failed to parse engine response")?; - if !envelope.ok { - let detail = envelope.error.map_or_else( - || "unknown engine error".to_owned(), - |error| format!("{}: {}", error.code, error.message), - ); - bail!(detail); - } - let data = envelope - .data - .context("engine response did not include data")?; - serde_json::from_value(data).context("failed to deserialize engine response data") -} - -fn wait_for_port(host: &str, port: u16, timeout: Duration) -> bool { - let address: SocketAddr = match format_host_port(host, port).parse() { - Ok(value) => value, - Err(_) => return false, - }; - - let start = std::time::Instant::now(); - while start.elapsed() < timeout { - if TcpStream::connect_timeout(&address, Duration::from_millis(200)).is_ok() { + if common::engine_healthcheck_ready(paths, engine, service_id).unwrap_or(false) { return true; } thread::sleep(Duration::from_millis(200)); diff --git a/docs/architecture.md b/docs/architecture.md index b07b2f52c..b64e9d564 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). Still pending: a shared-helpers module for code used across ≥2 of those remaining clusters (GPU/amd-smi snapshotting, the bridge-snapshot diagnostic, small arg/healthcheck utilities), plus the CLI, sandbox, MCP, service-lifecycle, webhook, 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) and `common.rs` (helpers shared across ≥2 of those remaining clusters: GPU/amd-smi snapshotting, the bridge-snapshot diagnostic, `CommandCapture`/command-timeout plumbing, and small arg/healthcheck/endpoint-key utilities). Still pending: the CLI, sandbox, MCP, service-lifecycle, webhook, and watcher clusters themselves — each landing as its own PR. ### `crates/rocm-core` — core library