use crate::network::ssh_transport::{SshFallbackPolicy, SshTransport};
pub const MIN_NATIVE_VERSION: (u32, u32) = (0, 22);
#[derive(Debug, Clone, PartialEq, Eq, Default)]
pub struct ProbeOutcomes {
pub native: ProbeResult,
pub nvidia_smi: ProbeResult,
pub rocm_smi: ProbeResult,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub enum ProbeResult {
#[default]
NotAttempted,
NotAvailable,
Available,
}
pub fn select_transport(outcomes: &ProbeOutcomes, policy: &SshFallbackPolicy) -> SshTransport {
if outcomes.native == ProbeResult::Available {
return SshTransport::Native;
}
if policy.try_nvidia_smi && outcomes.nvidia_smi == ProbeResult::Available {
return SshTransport::NvidiaSmi;
}
if policy.try_rocm_smi && outcomes.rocm_smi == ProbeResult::Available {
return SshTransport::RocmSmi;
}
SshTransport::Unsupported
}
pub fn native_supported(version_line: &str) -> bool {
let Some(v) = extract_version(version_line) else {
return false;
};
v >= MIN_NATIVE_VERSION
}
fn extract_version(line: &str) -> Option<(u32, u32)> {
let token = line
.split_whitespace()
.last()
.unwrap_or("")
.trim_start_matches('v');
let mut parts = token.split('.');
let major = parts.next()?.parse::<u32>().ok()?;
let minor = parts
.next()
.map(|s| s.split(['-', '+']).next().unwrap_or(s))
.and_then(|s| s.parse::<u32>().ok())
.unwrap_or(0);
Some((major, minor))
}
#[cfg(test)]
mod tests {
use super::*;
fn policy_all() -> SshFallbackPolicy {
SshFallbackPolicy {
try_nvidia_smi: true,
try_rocm_smi: true,
}
}
fn policy_none() -> SshFallbackPolicy {
SshFallbackPolicy::default()
}
#[test]
fn selects_native_when_available() {
let outcomes = ProbeOutcomes {
native: ProbeResult::Available,
nvidia_smi: ProbeResult::Available, rocm_smi: ProbeResult::NotAttempted,
};
assert_eq!(
select_transport(&outcomes, &policy_all()),
SshTransport::Native
);
}
#[test]
fn falls_back_to_nvidia_when_native_absent() {
let outcomes = ProbeOutcomes {
native: ProbeResult::NotAvailable,
nvidia_smi: ProbeResult::Available,
rocm_smi: ProbeResult::NotAttempted,
};
assert_eq!(
select_transport(&outcomes, &policy_all()),
SshTransport::NvidiaSmi
);
}
#[test]
fn falls_back_to_rocm_when_nvidia_unavailable() {
let outcomes = ProbeOutcomes {
native: ProbeResult::NotAvailable,
nvidia_smi: ProbeResult::NotAvailable,
rocm_smi: ProbeResult::Available,
};
assert_eq!(
select_transport(&outcomes, &policy_all()),
SshTransport::RocmSmi
);
}
#[test]
fn unsupported_when_no_probes_work() {
let outcomes = ProbeOutcomes {
native: ProbeResult::NotAvailable,
nvidia_smi: ProbeResult::NotAvailable,
rocm_smi: ProbeResult::NotAvailable,
};
assert_eq!(
select_transport(&outcomes, &policy_all()),
SshTransport::Unsupported
);
}
#[test]
fn policy_none_skips_fallbacks_even_when_available() {
let outcomes = ProbeOutcomes {
native: ProbeResult::NotAvailable,
nvidia_smi: ProbeResult::Available,
rocm_smi: ProbeResult::Available,
};
assert_eq!(
select_transport(&outcomes, &policy_none()),
SshTransport::Unsupported
);
}
#[test]
fn policy_only_rocm_ignores_nvidia() {
let policy = SshFallbackPolicy {
try_nvidia_smi: false,
try_rocm_smi: true,
};
let outcomes = ProbeOutcomes {
native: ProbeResult::NotAvailable,
nvidia_smi: ProbeResult::Available,
rocm_smi: ProbeResult::Available,
};
assert_eq!(select_transport(&outcomes, &policy), SshTransport::RocmSmi);
}
#[test]
fn native_supported_accepts_exact_min_version() {
assert!(native_supported("all-smi 0.22.0"));
assert!(native_supported("all-smi 0.22.1"));
}
#[test]
fn native_supported_rejects_older_versions() {
assert!(!native_supported("all-smi 0.21.5"));
assert!(!native_supported("all-smi 0.20.1"));
}
#[test]
fn native_supported_accepts_newer_major() {
assert!(native_supported("all-smi 1.0.0"));
assert!(native_supported("all-smi 0.23.0"));
}
#[test]
fn native_supported_handles_v_prefix() {
assert!(native_supported("all-smi v0.22.0"));
}
#[test]
fn native_supported_rejects_garbage() {
assert!(!native_supported(""));
assert!(!native_supported("not a version"));
assert!(!native_supported("all-smi unknown"));
}
#[test]
fn native_supported_handles_prerelease_suffix() {
assert!(native_supported("all-smi 0.22.0-alpha.1"));
}
}