Skip to content
Open
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
83 changes: 82 additions & 1 deletion apps/rocm/src/main.rs
Original file line number Diff line number Diff line change
Expand Up @@ -3457,6 +3457,24 @@ fn driver_command(phase: DriverCommandPhase, command: &str) -> DriverPlanCommand
}
}

/// Resolve a `${VAR:-default}` shell parameter-expansion template to its
/// effective value for user-facing display: the value of `VAR` when it is set
/// and non-empty (matching the shell `:-` semantics), otherwise the default.
/// Anything that is not a recognized `${VAR:-default}` template is returned
/// unchanged.
fn resolve_repo_version(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();
};
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");
Expand All @@ -3476,7 +3494,11 @@ fn render_driver_install_plan(plan: &DriverInstallPlan, yes: bool, dry_run: bool
empty_as_unknown(&plan.version_id)
);
let _ = writeln!(output, " codename: {}", empty_as_unknown(&plan.codename));
let _ = writeln!(output, " repo_version: {}", plan.repo_version_expr);
let _ = writeln!(
output,
" repo_version: {}",
resolve_repo_version(&plan.repo_version_expr)
);
let _ = writeln!(output, " reason: {}", plan.reason);
if !plan.preflight_checks.is_empty() {
let _ = writeln!(output, " preflight_checks:");
Expand Down Expand Up @@ -24004,6 +24026,65 @@ VERSION_CODENAME=noble
assert!(rendered.contains("add --dkms"));
}

#[test]
fn resolve_repo_version_uses_default_when_env_unset() {
// A made-up variable name that nothing else sets keeps this race-free.
assert_eq!(
resolve_repo_version("${ROCM_CLI_TEST_UNSET_REPO_VERSION:-7.2.4}"),
"7.2.4"
);
}

#[test]
#[allow(unsafe_code)] // std::env::set_var/remove_var are unsafe in edition 2024
fn resolve_repo_version_prefers_env_value_when_set() {
// Uses a unique per-test variable name so it cannot race with other
// tests reading the real ROCM_CLI_AMDGPU_VERSION variable.
let var = "ROCM_CLI_TEST_REPO_VERSION_OVERRIDE";
unsafe {
std::env::set_var(var, "9.9.9");
}
let resolved = resolve_repo_version(&format!("${{{var}:-7.2.4}}"));
unsafe {
std::env::remove_var(var);
}
assert_eq!(resolved, "9.9.9");
}

#[test]
#[allow(unsafe_code)] // std::env::set_var/remove_var are unsafe in edition 2024
fn resolve_repo_version_treats_empty_env_as_unset() {
let var = "ROCM_CLI_TEST_REPO_VERSION_EMPTY";
unsafe {
std::env::set_var(var, "");
}
let resolved = resolve_repo_version(&format!("${{{var}:-7.2.4}}"));
unsafe {
std::env::remove_var(var);
}
assert_eq!(resolved, "7.2.4");
}

#[test]
fn resolve_repo_version_passes_through_non_template() {
assert_eq!(resolve_repo_version("7.2.4"), "7.2.4");
}

#[test]
fn driver_plan_dry_run_repo_version_line_is_resolved() {
// 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);
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 os_release = r#"
Expand Down
Loading