#![allow(unused)]
use crate::config::model::DeviceType;
use log::{debug, info, warn};
use std::sync::Arc;
use std::sync::atomic::{AtomicBool, AtomicU64, Ordering};
use tokio::sync::RwLock;
const DEFAULT_OPENCL_VRAM_BYTES: u64 = 8 * 1024 * 1024 * 1024; const DEFAULT_OPENCL_COMPUTE_CAP: (u32, u32) = (5, 0);
const DEFAULT_ROCM_VRAM_BYTES: u64 = 16 * 1024 * 1024 * 1024; const DEFAULT_ROCM_COMPUTE_CAP: (u32, u32) = (9, 0);
const DEFAULT_DRIVER_VERSION: &str = "unknown";
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct AmdGpuInfo {
pub name: String,
pub device_id: u32,
pub vram_bytes: u64,
pub compute_capability: (u32, u32),
pub opencl_version: String,
pub roc_version: Option<String>,
pub driver_version: String,
pub is_available: bool,
}
impl Default for AmdGpuInfo {
fn default() -> Self {
Self {
name: "Unknown AMD GPU".to_string(),
device_id: 0,
vram_bytes: 0,
compute_capability: (0, 0),
opencl_version: "3.0".to_string(),
roc_version: None,
driver_version: "unknown".to_string(),
is_available: false,
}
}
}
#[derive(Debug)]
pub struct AmdDevice {
info: AmdGpuInfo,
device_type: DeviceType,
memory_used: AtomicU64,
memory_allocated: AtomicU64,
compute_units: u32,
max_work_group_size: usize,
max_work_item_dimensions: u32,
is_busy: AtomicBool,
}
impl Clone for AmdDevice {
fn clone(&self) -> Self {
Self {
info: self.info.clone(),
device_type: self.device_type.clone(),
memory_used: AtomicU64::new(self.memory_used.load(Ordering::Relaxed)),
memory_allocated: AtomicU64::new(self.memory_allocated.load(Ordering::Relaxed)),
compute_units: self.compute_units,
max_work_group_size: self.max_work_group_size,
max_work_item_dimensions: self.max_work_item_dimensions,
is_busy: AtomicBool::new(self.is_busy.load(Ordering::Relaxed)),
}
}
}
impl AmdDevice {
pub fn new(info: AmdGpuInfo) -> Self {
let compute_units = (info.vram_bytes / (1024 * 1024 * 1024) * 64).min(80) as u32;
Self {
info,
device_type: DeviceType::Amd,
memory_used: AtomicU64::new(0),
memory_allocated: AtomicU64::new(0),
compute_units,
max_work_group_size: 256,
max_work_item_dimensions: 3,
is_busy: AtomicBool::new(false),
}
}
pub fn from_opencl(index: usize) -> Option<Self> {
debug!("Attempting to detect AMD GPU via OpenCL at index {}", index);
warn!(
"AMD GPU (OpenCL) 使用 fallback 默认参数 ({}GB, compute {}.{}). \
OpenCL 无法查询真实 VRAM,建议通过 ROCm 路径获取准确信息",
DEFAULT_OPENCL_VRAM_BYTES / (1024 * 1024 * 1024),
DEFAULT_OPENCL_COMPUTE_CAP.0,
DEFAULT_OPENCL_COMPUTE_CAP.1,
);
let info = AmdGpuInfo {
name: format!("AMD GPU (OpenCL) - Device {}", index),
device_id: index as u32,
vram_bytes: DEFAULT_OPENCL_VRAM_BYTES,
compute_capability: DEFAULT_OPENCL_COMPUTE_CAP,
opencl_version: "3.0".to_string(),
roc_version: None,
driver_version: query_amd_driver_version(),
is_available: true,
};
Some(Self::new(info))
}
pub fn from_rocm(index: usize) -> Option<Self> {
debug!("Attempting to detect AMD GPU via ROCm at index {}", index);
if std::process::Command::new("rocm-smi")
.arg("--version")
.output()
.is_err()
{
debug!(
"rocm-smi not available, skipping ROCm device at index {}",
index
);
return None;
}
let vram = query_rocm_vram(index);
let info = AmdGpuInfo {
name: format!("AMD GPU (ROCm) - Device {}", index),
device_id: index as u32,
vram_bytes: vram,
compute_capability: DEFAULT_ROCM_COMPUTE_CAP,
opencl_version: "3.0".to_string(),
roc_version: query_rocm_version(),
driver_version: query_amd_driver_version(),
is_available: true,
};
Some(Self::new(info))
}
pub fn info(&self) -> &AmdGpuInfo {
&self.info
}
pub fn device_type(&self) -> DeviceType {
self.device_type.clone()
}
pub fn name(&self) -> &str {
&self.info.name
}
pub fn vram_bytes(&self) -> u64 {
self.info.vram_bytes
}
pub fn available_memory(&self) -> u64 {
self.info.vram_bytes - self.memory_used.load(Ordering::SeqCst)
}
pub fn memory_usage_percent(&self) -> f64 {
let used = self.memory_used.load(Ordering::SeqCst);
if self.info.vram_bytes == 0 {
0.0
} else {
(used as f64 / self.info.vram_bytes as f64) * 100.0
}
}
pub fn compute_units(&self) -> u32 {
self.compute_units
}
pub fn max_work_group_size(&self) -> usize {
self.max_work_group_size
}
pub fn allocate(&self, bytes: u64) -> bool {
let result =
self.memory_allocated
.fetch_update(Ordering::SeqCst, Ordering::SeqCst, |current| {
let new_allocated = current + bytes;
if new_allocated > self.info.vram_bytes {
None } else {
Some(new_allocated)
}
});
match result {
Ok(_old) => {
self.memory_used.store(
self.memory_allocated.load(Ordering::SeqCst),
Ordering::SeqCst,
);
true
}
Err(_) => {
log::warn!(
"GPU memory allocation failed: requested {} bytes, available {} bytes",
bytes,
self.available_memory()
);
false
}
}
}
pub fn deallocate(&self, bytes: u64) {
let current = self.memory_allocated.load(Ordering::SeqCst);
self.memory_allocated
.store(current.saturating_sub(bytes), Ordering::SeqCst);
self.memory_used.store(
self.memory_allocated.load(Ordering::SeqCst),
Ordering::SeqCst,
);
}
pub fn is_busy(&self) -> bool {
self.is_busy.load(Ordering::SeqCst)
}
pub fn set_busy(&self, busy: bool) {
self.is_busy.store(busy, Ordering::SeqCst);
}
pub fn supports_precision(&self, precision: &str) -> bool {
matches!(precision, "fp32" | "fp16" | "bf16")
}
pub fn supports_operation(&self, operation: &str) -> bool {
matches!(
operation,
"matrix_multiply" | "convolution" | "activation" | "normalization" | "reduction"
)
}
}
fn query_amd_driver_version() -> String {
if let Ok(version) = std::fs::read_to_string("/sys/module/amdgpu/version") {
let trimmed = version.trim().to_string();
if !trimmed.is_empty() {
info!("AMD 驱动版本 (sysfs): {}", trimmed);
return trimmed;
}
}
if let Ok(output) = std::process::Command::new("rocm-smi")
.arg("--showdriverversion")
.output()
&& output.status.success()
{
let text = String::from_utf8_lossy(&output.stdout).to_string();
for line in text.lines() {
if let Some(ver) = line.strip_prefix("Driver version:") {
let trimmed = ver.trim().to_string();
if !trimmed.is_empty() {
info!("AMD 驱动版本 (rocm-smi): {}", trimmed);
return trimmed;
}
}
}
}
debug!(
"AMD 驱动版本查询失败,使用 fallback: {}",
DEFAULT_DRIVER_VERSION
);
DEFAULT_DRIVER_VERSION.to_string()
}
fn query_rocm_vram(index: usize) -> u64 {
if let Ok(output) = std::process::Command::new("rocm-smi")
.args(["--showmeminfo", "vram", "--json"])
.output()
&& output.status.success()
{
let text = String::from_utf8_lossy(&output.stdout);
for line in text.lines() {
if line.contains("VRAM Total") || line.contains("Total Memory") {
if let Some(mib) = extract_number_from_line(line) {
let bytes = mib * 1024 * 1024;
info!("ROCm VRAM (rocm-smi, device {}): {} MB", index, mib);
return bytes;
}
}
}
}
warn!(
"ROCm VRAM 查询失败,使用 fallback 默认值: {}GB",
DEFAULT_ROCM_VRAM_BYTES / (1024 * 1024 * 1024)
);
DEFAULT_ROCM_VRAM_BYTES
}
fn query_rocm_version() -> Option<String> {
if let Ok(output) = std::process::Command::new("rocm-smi")
.arg("--showdriverversion")
.output()
&& output.status.success()
{
let text = String::from_utf8_lossy(&output.stdout);
for line in text.lines() {
if let Some(ver) = line.strip_prefix("ROCm version:") {
let trimmed = ver.trim().to_string();
if !trimmed.is_empty() {
info!("ROCm 版本: {}", trimmed);
return Some(trimmed);
}
}
}
}
None
}
fn extract_number_from_line(line: &str) -> Option<u64> {
let mut num_str = String::new();
let mut found_digits = false;
for c in line.chars() {
if c.is_ascii_digit() {
num_str.push(c);
found_digits = true;
} else if found_digits {
break;
}
}
if found_digits {
num_str.parse().ok()
} else {
None
}
}
pub struct AmdDeviceManager {
devices: Arc<RwLock<Vec<Arc<AmdDevice>>>>,
primary_device: Arc<RwLock<Option<usize>>>,
opencl_available: AtomicBool,
rocm_available: AtomicBool,
initialized: Arc<AtomicBool>,
}
impl Default for AmdDeviceManager {
fn default() -> Self {
Self::new()
}
}
impl AmdDeviceManager {
pub fn new() -> Self {
Self {
devices: Arc::new(RwLock::new(Vec::new())),
primary_device: Arc::new(RwLock::new(None)),
opencl_available: AtomicBool::new(false),
rocm_available: AtomicBool::new(false),
initialized: Arc::new(AtomicBool::new(false)),
}
}
pub async fn initialize(&self) -> Result<(), crate::error::VecboostError> {
if self.initialized.load(Ordering::SeqCst) {
return Ok(());
}
log::info!("Initializing AMD GPU device manager...");
let mut devices = self.devices.write().await;
devices.clear();
let mut opencl_found = false;
let mut rocm_found = false;
for i in 0..4 {
if let Some(device) = AmdDevice::from_rocm(i) {
log::info!(
"Found ROCm-compatible AMD GPU: {} with {} bytes VRAM",
device.name(),
device.vram_bytes()
);
devices.push(Arc::new(device));
rocm_found = true;
}
}
if !rocm_found {
for i in 0..4 {
if let Some(device) = AmdDevice::from_opencl(i) {
log::info!(
"Found OpenCL-compatible AMD GPU: {} with {} bytes VRAM",
device.name(),
device.vram_bytes()
);
devices.push(Arc::new(device));
opencl_found = true;
}
}
}
self.opencl_available.store(opencl_found, Ordering::SeqCst);
self.rocm_available.store(rocm_found, Ordering::SeqCst);
if !devices.is_empty() {
let mut primary = self.primary_device.write().await;
*primary = Some(0);
}
self.initialized.store(true, Ordering::SeqCst);
log::info!(
"AMD GPU initialization complete. Found {} device(s) (ROCm: {}, OpenCL: {})",
devices.len(),
rocm_found,
opencl_found
);
Ok(())
}
pub fn is_initialized(&self) -> bool {
self.initialized.load(Ordering::SeqCst)
}
pub async fn devices(&self) -> Vec<Arc<AmdDevice>> {
self.devices.read().await.clone()
}
pub async fn primary_device(&self) -> Option<Arc<AmdDevice>> {
let primary = self.primary_device.read().await;
let devices = self.devices.read().await;
match *primary {
Some(idx) if idx < devices.len() => Some(devices[idx].clone()),
_ => devices.first().cloned(),
}
}
pub fn is_opencl_available(&self) -> bool {
self.opencl_available.load(Ordering::SeqCst)
}
pub fn is_rocm_available(&self) -> bool {
self.rocm_available.load(Ordering::SeqCst)
}
pub async fn get_device(&self, index: usize) -> Option<Arc<AmdDevice>> {
let devices = self.devices.read().await;
devices.get(index).cloned()
}
pub async fn total_vram(&self) -> u64 {
let devices = self.devices.read().await;
devices.iter().map(|d| d.vram_bytes()).sum()
}
pub async fn available_vram(&self) -> u64 {
let devices = self.devices.read().await;
devices.iter().map(|d| d.available_memory()).sum()
}
pub async fn device_count(&self) -> usize {
self.devices.read().await.len()
}
pub async fn memory_usage_summary(&self) -> String {
let devices = self.devices.read().await;
let total_used: u64 = devices
.iter()
.map(|d| d.memory_used.load(Ordering::SeqCst))
.sum();
let total_vram: u64 = devices.iter().map(|d| d.vram_bytes()).sum();
format!(
"AMD GPU Memory: {} bytes used / {} bytes total ({:.1}%)",
total_used,
total_vram,
if total_vram > 0 {
(total_used as f64 / total_vram as f64) * 100.0
} else {
0.0
}
)
}
pub async fn set_primary(&self, index: usize) -> bool {
let devices = self.devices.read().await;
if index < devices.len() {
let mut primary = self.primary_device.write().await;
*primary = Some(index);
true
} else {
false
}
}
pub async fn reset(&self) {
let devices = self.devices.write().await;
for device in devices.iter() {
device.memory_allocated.store(0, Ordering::SeqCst);
device.memory_used.store(0, Ordering::SeqCst);
device.is_busy.store(false, Ordering::SeqCst);
}
}
}
pub async fn create_amd_device_manager() -> Result<AmdDeviceManager, crate::error::VecboostError> {
let manager = AmdDeviceManager::new();
manager.initialize().await?;
Ok(manager)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_amd_gpu_info_default() {
let info = AmdGpuInfo::default();
assert_eq!(info.name, "Unknown AMD GPU");
assert_eq!(info.device_id, 0);
assert_eq!(info.vram_bytes, 0);
assert_eq!(info.compute_capability, (0, 0));
assert_eq!(info.opencl_version, "3.0");
assert!(info.roc_version.is_none());
assert_eq!(info.driver_version, "unknown");
assert!(!info.is_available);
}
#[test]
fn test_amd_gpu_info_default_eq() {
let a = AmdGpuInfo::default();
let b = AmdGpuInfo::default();
assert_eq!(a, b);
}
fn make_info(vram_bytes: u64) -> AmdGpuInfo {
AmdGpuInfo {
name: "AMD Radeon RX 7900 XTX".to_string(),
device_id: 0x73BF,
vram_bytes,
compute_capability: (9, 0),
opencl_version: "3.0".to_string(),
roc_version: Some("6.0.0".to_string()),
driver_version: "24.0.0".to_string(),
is_available: true,
}
}
#[test]
fn test_amd_device_new_compute_units_cap_at_80() {
let info = make_info(64 * 1024 * 1024 * 1024);
let device = AmdDevice::new(info);
assert_eq!(device.compute_units(), 80);
assert_eq!(device.max_work_group_size(), 256);
}
#[test]
fn test_amd_device_new_small_vram_compute_units() {
let info = make_info(1024 * 1024 * 1024);
let device = AmdDevice::new(info);
assert_eq!(device.compute_units(), 64);
}
#[test]
fn test_amd_device_new_medium_vram_compute_units_capped() {
let info = make_info(4 * 1024 * 1024 * 1024);
let device = AmdDevice::new(info);
assert_eq!(device.compute_units(), 80);
}
#[test]
fn test_amd_device_new_zero_vram_compute_units() {
let info = make_info(0);
let device = AmdDevice::new(info);
assert_eq!(device.compute_units(), 0);
}
#[test]
fn test_amd_device_from_opencl() {
let device = AmdDevice::from_opencl(2).expect("from_opencl should return Some");
assert!(device.info().is_available);
assert_eq!(device.device_type(), DeviceType::Amd);
assert!(device.name().contains("OpenCL"));
assert!(device.name().contains("Device 2"));
assert_eq!(device.vram_bytes(), 8 * 1024 * 1024 * 1024);
}
#[test]
fn test_amd_device_from_rocm() {
let result = AmdDevice::from_rocm(1);
match result {
None => {} Some(device) => {
assert!(device.info().is_available);
assert_eq!(device.device_type(), DeviceType::Amd);
assert!(device.name().contains("ROCm"));
assert!(device.name().contains("Device 1"));
assert!(device.vram_bytes() > 0);
}
}
}
#[test]
fn test_amd_device_info_accessors() {
let info = make_info(16 * 1024 * 1024 * 1024);
let device = AmdDevice::new(info.clone());
assert_eq!(device.info().name, info.name);
assert_eq!(device.info().device_id, info.device_id);
assert_eq!(device.info().vram_bytes, info.vram_bytes);
assert_eq!(device.name(), info.name);
assert_eq!(device.vram_bytes(), info.vram_bytes);
}
#[test]
fn test_amd_device_available_memory_initial() {
let device = AmdDevice::new(make_info(1024));
assert_eq!(device.available_memory(), 1024);
}
#[test]
fn test_amd_device_memory_usage_percent_zero_vram() {
let device = AmdDevice::new(make_info(0));
assert_eq!(device.memory_usage_percent(), 0.0);
}
#[test]
fn test_amd_device_memory_usage_percent_after_allocate() {
let device = AmdDevice::new(make_info(1024));
assert!(device.allocate(256));
assert!((device.memory_usage_percent() - 25.0).abs() < 0.001);
}
#[test]
fn test_amd_device_allocate_success() {
let device = AmdDevice::new(make_info(1024));
assert!(device.allocate(512));
assert_eq!(device.available_memory(), 512);
}
#[test]
fn test_amd_device_allocate_exceeds_vram() {
let device = AmdDevice::new(make_info(1024));
assert!(device.allocate(512));
assert!(!device.allocate(1024));
}
#[test]
fn test_amd_device_deallocate() {
let device = AmdDevice::new(make_info(1024));
assert!(device.allocate(512));
device.deallocate(256);
assert_eq!(device.available_memory(), 768);
}
#[test]
fn test_amd_device_deallocate_saturating() {
let device = AmdDevice::new(make_info(1024));
device.deallocate(2048);
assert_eq!(device.available_memory(), 1024);
}
#[test]
fn test_amd_device_is_busy_set_busy() {
let device = AmdDevice::new(make_info(1024));
assert!(!device.is_busy());
device.set_busy(true);
assert!(device.is_busy());
device.set_busy(false);
assert!(!device.is_busy());
}
#[test]
fn test_amd_device_supports_precision() {
let device = AmdDevice::new(make_info(1024));
assert!(device.supports_precision("fp32"));
assert!(device.supports_precision("fp16"));
assert!(device.supports_precision("bf16"));
assert!(!device.supports_precision("int8"));
assert!(!device.supports_precision("fp64"));
}
#[test]
fn test_amd_device_supports_operation() {
let device = AmdDevice::new(make_info(1024));
assert!(device.supports_operation("matrix_multiply"));
assert!(device.supports_operation("convolution"));
assert!(device.supports_operation("activation"));
assert!(device.supports_operation("normalization"));
assert!(device.supports_operation("reduction"));
assert!(!device.supports_operation("unknown"));
}
#[test]
fn test_amd_device_clone_preserves_state() {
let device = AmdDevice::new(make_info(1024));
assert!(device.allocate(128));
device.set_busy(true);
let cloned = device.clone();
assert_eq!(cloned.name(), device.name());
assert_eq!(cloned.vram_bytes(), device.vram_bytes());
assert_eq!(cloned.compute_units(), device.compute_units());
assert_eq!(cloned.available_memory(), device.available_memory());
assert!(cloned.is_busy());
}
#[tokio::test]
async fn test_amd_device_manager_new_default() {
let manager = AmdDeviceManager::new();
assert!(!manager.is_initialized());
assert!(!manager.is_opencl_available());
assert!(!manager.is_rocm_available());
assert_eq!(manager.device_count().await, 0);
let default_mgr = AmdDeviceManager::default();
assert!(!default_mgr.is_initialized());
}
#[tokio::test]
async fn test_amd_device_manager_initialize_finds_rocm_devices() {
let manager = AmdDeviceManager::new();
manager
.initialize()
.await
.expect("initialize should succeed");
assert!(manager.is_initialized());
assert_eq!(manager.device_count().await, 4);
let total = manager.total_vram().await;
assert!(total > 0);
}
#[tokio::test]
async fn test_amd_device_manager_initialize_idempotent() {
let manager = AmdDeviceManager::new();
manager.initialize().await.unwrap();
let count_after_first = manager.device_count().await;
manager.initialize().await.unwrap();
let count_after_second = manager.device_count().await;
assert_eq!(count_after_first, count_after_second);
}
#[tokio::test]
async fn test_amd_device_manager_primary_device() {
let manager = AmdDeviceManager::new();
manager.initialize().await.unwrap();
let primary = manager.primary_device().await;
assert!(primary.is_some());
let name = primary.as_ref().unwrap().name();
assert!(name.contains("ROCm") || name.contains("OpenCL"));
}
#[tokio::test]
async fn test_amd_device_manager_primary_device_none_when_empty() {
let manager = AmdDeviceManager::new();
let primary = manager.primary_device().await;
assert!(primary.is_none());
}
#[tokio::test]
async fn test_amd_device_manager_get_device_in_range() {
let manager = AmdDeviceManager::new();
manager.initialize().await.unwrap();
let device = manager.get_device(1).await;
assert!(device.is_some());
assert!(device.as_ref().unwrap().name().contains("Device 1"));
}
#[tokio::test]
async fn test_amd_device_manager_get_device_out_of_range() {
let manager = AmdDeviceManager::new();
manager.initialize().await.unwrap();
let device = manager.get_device(100).await;
assert!(device.is_none());
}
#[tokio::test]
async fn test_amd_device_manager_available_vram() {
let manager = AmdDeviceManager::new();
manager.initialize().await.unwrap();
let available = manager.available_vram().await;
assert_eq!(available, manager.total_vram().await);
}
#[tokio::test]
async fn test_amd_device_manager_memory_usage_summary_empty() {
let manager = AmdDeviceManager::new();
let summary = manager.memory_usage_summary().await;
assert!(summary.contains("0 bytes used"));
assert!(summary.contains("0.0%"));
}
#[tokio::test]
async fn test_amd_device_manager_memory_usage_summary_with_devices() {
let manager = AmdDeviceManager::new();
manager.initialize().await.unwrap();
let summary = manager.memory_usage_summary().await;
assert!(summary.contains("AMD GPU Memory"));
assert!(summary.contains("0.0%"));
}
#[tokio::test]
async fn test_amd_device_manager_set_primary_valid() {
let manager = AmdDeviceManager::new();
manager.initialize().await.unwrap();
let count = manager.device_count().await;
if count > 0 {
assert!(manager.set_primary(0).await);
let primary = manager.primary_device().await;
assert!(primary.is_some());
}
}
#[tokio::test]
async fn test_amd_device_manager_set_primary_out_of_range() {
let manager = AmdDeviceManager::new();
manager.initialize().await.unwrap();
assert!(!manager.set_primary(100).await);
}
#[tokio::test]
async fn test_amd_device_manager_reset() {
let manager = AmdDeviceManager::new();
manager.initialize().await.unwrap();
let primary = manager.primary_device().await.unwrap();
assert!(primary.allocate(1024));
assert!(!primary.is_busy());
primary.set_busy(true);
manager.reset().await;
let primary_after = manager.primary_device().await.unwrap();
assert_eq!(primary_after.available_memory(), primary_after.vram_bytes());
assert!(!primary_after.is_busy());
}
#[tokio::test]
async fn test_amd_device_manager_devices_returns_clone() {
let manager = AmdDeviceManager::new();
manager.initialize().await.unwrap();
let devices = manager.devices().await;
assert!(!devices.is_empty());
}
#[tokio::test]
async fn test_create_amd_device_manager() {
let _path_guard = PATH_MUTEX.lock().await;
let manager = create_amd_device_manager()
.await
.expect("create_amd_device_manager should succeed");
assert!(manager.is_initialized());
assert!(manager.device_count().await > 0);
}
#[tokio::test]
async fn test_amd_device_manager_primary_device_index_out_of_range_falls_back() {
let _path_guard = PATH_MUTEX.lock().await;
let manager = AmdDeviceManager::new();
manager.initialize().await.unwrap();
manager.set_primary(100).await;
let primary = manager.primary_device().await;
if manager.device_count().await > 0 {
assert!(primary.is_some());
assert!(primary.as_ref().unwrap().name().contains("Device 0"));
} else {
assert!(
primary.is_none(),
"no devices enumerated → primary must be None"
);
}
}
#[test]
fn test_extract_number_from_line_basic() {
assert_eq!(extract_number_from_line("16384"), Some(16384));
}
#[test]
fn test_extract_number_from_line_with_prefix() {
assert_eq!(
extract_number_from_line("VRAM Total: 16384 MiB"),
Some(16384)
);
}
#[test]
fn test_extract_number_from_line_with_trailing_text() {
assert_eq!(extract_number_from_line("16384 MiB"), Some(16384));
}
#[test]
fn test_extract_number_from_line_no_digits() {
assert_eq!(extract_number_from_line("no numbers here"), None);
}
#[test]
fn test_extract_number_from_line_empty() {
assert_eq!(extract_number_from_line(""), None);
}
#[test]
fn test_extract_number_from_line_multiple_numbers() {
assert_eq!(extract_number_from_line("card0: 16384 8192"), Some(0));
assert_eq!(extract_number_from_line("16384 8192"), Some(16384));
}
#[test]
fn test_extract_number_from_line_large_number() {
assert_eq!(
extract_number_from_line("VRAM Total Memory (MiB): 16384"),
Some(16384)
);
}
#[test]
fn test_query_rocm_vram_fallback_when_rocm_smi_unavailable() {
let vram = query_rocm_vram(0);
if std::process::Command::new("rocm-smi").output().is_err() {
assert_eq!(vram, DEFAULT_ROCM_VRAM_BYTES);
}
}
#[test]
fn test_query_rocm_version_returns_none_without_rocm_smi() {
if std::process::Command::new("rocm-smi").output().is_err() {
assert_eq!(query_rocm_version(), None);
}
}
#[test]
fn test_query_amd_driver_version_fallback() {
let version = query_amd_driver_version();
let has_sysfs = std::fs::read_to_string("/sys/module/amdgpu/version")
.map(|v| !v.trim().is_empty())
.unwrap_or(false);
let has_rocm_smi = std::process::Command::new("rocm-smi")
.arg("--showdriverversion")
.output()
.map(|o| o.status.success())
.unwrap_or(false);
if !has_sysfs && !has_rocm_smi {
assert_eq!(version, DEFAULT_DRIVER_VERSION);
}
}
#[tokio::test]
async fn test_memory_usage_summary_with_allocation() {
let manager = AmdDeviceManager::new();
manager.initialize().await.unwrap();
let device = manager.get_device(0).await.unwrap();
assert!(device.allocate(1024 * 1024));
let summary = manager.memory_usage_summary().await;
assert!(summary.contains("AMD GPU Memory"));
assert!(summary.contains("bytes used"));
}
#[tokio::test]
async fn test_available_vram_after_allocation() {
let manager = AmdDeviceManager::new();
manager.initialize().await.unwrap();
let total_before = manager.total_vram().await;
let available_before = manager.available_vram().await;
assert_eq!(total_before, available_before);
let device = manager.get_device(0).await.unwrap();
assert!(device.allocate(1024));
let available_after = manager.available_vram().await;
assert_eq!(available_after, available_before - 1024);
}
#[test]
fn test_from_opencl_different_indices() {
let d0 = AmdDevice::from_opencl(0).unwrap();
let d3 = AmdDevice::from_opencl(3).unwrap();
assert!(d0.name().contains("Device 0"));
assert!(d3.name().contains("Device 3"));
assert_eq!(d0.info().device_id, 0);
assert_eq!(d3.info().device_id, 3);
}
#[test]
fn test_amd_device_allocate_concurrent_safety() {
let device = AmdDevice::new(make_info(1024));
assert!(device.allocate(600));
assert!(!device.allocate(600));
assert!(device.allocate(424));
}
#[test]
fn test_amd_device_memory_usage_percent_with_various_vram() {
let device = AmdDevice::new(make_info(2048));
assert_eq!(device.memory_usage_percent(), 0.0);
assert!(device.allocate(1024));
assert!((device.memory_usage_percent() - 50.0).abs() < 0.001);
}
fn create_mock_rocm_smi() -> tempfile::TempDir {
let temp_dir = tempfile::tempdir().unwrap();
let script_path = temp_dir.path().join("rocm-smi");
std::fs::write(
&script_path,
r#"#!/bin/bash
if echo "$@" | grep -q "showmeminfo"; then
echo " VRAM Total Memory (MiB): 16384"
elif echo "$@" | grep -q "showdriverversion"; then
echo "ROCm version: 6.0.0"
echo "Driver version: 24.0.0"
elif echo "$@" | grep -q "version"; then
echo "rocm-smi 6.0.0"
else
echo "AMD GPU"
fi
"#,
)
.unwrap();
#[cfg(unix)]
{
use std::os::unix::fs::PermissionsExt;
std::fs::set_permissions(&script_path, std::fs::Permissions::from_mode(0o755)).unwrap();
}
temp_dir
}
fn prepend_path(dir: &std::path::Path) {
let current = std::env::var("PATH").unwrap_or_default();
unsafe { std::env::set_var("PATH", format!("{}:{}", dir.display(), current)) };
}
fn restore_path(original: &str) {
unsafe { std::env::set_var("PATH", original) };
}
static PATH_MUTEX: tokio::sync::Mutex<()> = tokio::sync::Mutex::const_new(());
#[tokio::test]
async fn test_query_rocm_vram_with_mock() {
let _guard = PATH_MUTEX.lock().await;
let original_path = std::env::var("PATH").unwrap_or_default();
let temp = create_mock_rocm_smi();
prepend_path(temp.path());
let vram = query_rocm_vram(0);
assert_eq!(vram, 16384 * 1024 * 1024);
restore_path(&original_path);
}
#[tokio::test]
async fn test_query_rocm_version_with_mock() {
let _guard = PATH_MUTEX.lock().await;
let original_path = std::env::var("PATH").unwrap_or_default();
let temp = create_mock_rocm_smi();
prepend_path(temp.path());
let version = query_rocm_version();
assert_eq!(version, Some("6.0.0".to_string()));
restore_path(&original_path);
}
#[tokio::test]
async fn test_query_amd_driver_version_with_mock() {
let _guard = PATH_MUTEX.lock().await;
let original_path = std::env::var("PATH").unwrap_or_default();
let temp = create_mock_rocm_smi();
prepend_path(temp.path());
let has_sysfs = std::fs::read_to_string("/sys/module/amdgpu/version")
.map(|v| !v.trim().is_empty())
.unwrap_or(false);
let version = query_amd_driver_version();
if !has_sysfs {
assert_eq!(version, "24.0.0");
}
restore_path(&original_path);
}
#[tokio::test]
async fn test_from_rocm_with_mock() {
let _guard = PATH_MUTEX.lock().await;
let original_path = std::env::var("PATH").unwrap_or_default();
let temp = create_mock_rocm_smi();
prepend_path(temp.path());
let device = AmdDevice::from_rocm(0).expect("from_rocm should return Some with mock");
assert!(device.info().is_available);
assert_eq!(device.device_type(), DeviceType::Amd);
assert!(device.name().contains("ROCm"));
assert!(device.name().contains("Device 0"));
assert_eq!(device.vram_bytes(), 16384 * 1024 * 1024);
assert!(device.info().roc_version.is_some());
restore_path(&original_path);
}
#[tokio::test]
async fn test_initialize_with_mock_rocm_smi() {
let _guard = PATH_MUTEX.lock().await;
let original_path = std::env::var("PATH").unwrap_or_default();
let temp = create_mock_rocm_smi();
prepend_path(temp.path());
let manager = AmdDeviceManager::new();
manager.initialize().await.unwrap();
assert!(manager.is_initialized());
assert!(manager.is_rocm_available());
assert_eq!(manager.device_count().await, 4);
let total = manager.total_vram().await;
assert_eq!(total, 4 * 16384 * 1024 * 1024u64);
restore_path(&original_path);
}
#[tokio::test]
async fn test_create_amd_device_manager_with_mock() {
let _guard = PATH_MUTEX.lock().await;
let original_path = std::env::var("PATH").unwrap_or_default();
let temp = create_mock_rocm_smi();
prepend_path(temp.path());
let manager = create_amd_device_manager().await.unwrap();
assert!(manager.is_initialized());
assert!(manager.is_rocm_available());
restore_path(&original_path);
}
}