Skip to main content

llm_manager/backend/
hardware.rs

1use std::fs;
2use std::path::Path;
3
4/// Detected operating system platform.
5#[derive(Debug, Clone, Copy, PartialEq, Eq)]
6pub enum Platform {
7    Linux,
8    Windows,
9    Macos,
10}
11
12/// GPU vendors
13#[derive(Debug, Clone, Copy, PartialEq, Eq)]
14#[allow(dead_code)]
15pub enum GpuVendor {
16    Amd,
17    Nvidia,
18    Intel,
19    Apple,
20    Unknown,
21}
22
23/// Detect the current operating system platform.
24pub fn detect_platform() -> Platform {
25    match std::env::consts::OS {
26        "windows" => Platform::Windows,
27        "macos" => Platform::Macos,
28        _ => Platform::Linux,
29    }
30}
31
32/// Check if the current architecture is ARM64.
33pub fn is_arm64() -> bool {
34    cfg!(target_arch = "aarch64")
35}
36
37/// Get the platform as a string slice.
38pub fn platform_name(platform: Platform) -> &'static str {
39    match platform {
40        Platform::Linux => "linux",
41        Platform::Windows => "windows",
42        Platform::Macos => "macos",
43    }
44}
45
46/// Check if a backend variant is available on the given platform.
47pub fn backend_supported(backend: crate::models::Backend, platform: Platform) -> bool {
48    match platform {
49        Platform::Linux => backend.is_linux(),
50        Platform::Windows => backend.is_windows(),
51        Platform::Macos => backend.is_macos(),
52    }
53}
54
55/// Returns paths to all primary DRM card directories (card0, card1, ...).
56fn drm_card_paths() -> Vec<std::path::PathBuf> {
57    let drm_path = Path::new("/sys/class/drm");
58    if !drm_path.exists() {
59        return Vec::new();
60    }
61    fs::read_dir(drm_path)
62        .map(|entries| {
63            entries
64                .flatten()
65                .filter(|e| {
66                    let n = e.file_name();
67                    let s = n.to_string_lossy();
68                    s.starts_with("card") && !s.contains('-')
69                })
70                .map(|e| e.path())
71                .collect()
72        })
73        .unwrap_or_default()
74}
75
76/// Detect all GPU vendors by scanning /sys/class/drm/card*/device/vendor (Linux).
77/// Returns a Vec of unique vendors (preserves detection order, deduplicates).
78fn detect_gpu_vendors_linux_impl() -> Vec<GpuVendor> {
79    let mut vendors = Vec::new();
80    for card_path in drm_card_paths() {
81        let vendor_path = card_path.join("device/vendor");
82        if let Ok(vendor_id) = fs::read_to_string(vendor_path) {
83            let vendor_id = vendor_id.trim();
84            let vendor = match vendor_id {
85                "0x1002" => GpuVendor::Amd,
86                "0x10de" => GpuVendor::Nvidia,
87                "0x8086" => GpuVendor::Intel,
88                _ => continue,
89            };
90            if !vendors.contains(&vendor) {
91                vendors.push(vendor);
92            }
93        }
94    }
95
96    if vendors.is_empty() {
97        vendors.push(GpuVendor::Unknown);
98    }
99
100    vendors
101}
102
103/// Detect all GPU model names (one per GPU, Linux).
104/// For AMD GPUs, includes the GFX target version.
105fn detect_gpu_models_linux_impl() -> Vec<Option<String>> {
106    let card_paths = drm_card_paths();
107    if card_paths.is_empty() {
108        return Vec::new();
109    }
110
111    let amd_gfx_targets = detect_amd_gfx_targets();
112    let mut amd_card_idx: usize = 0;
113    let mut models = Vec::new();
114    for card_path in &card_paths {
115        let vendor_path = card_path.join("device/vendor");
116        if let Ok(vendor_id) = fs::read_to_string(vendor_path) {
117            let vendor_id = vendor_id.trim();
118            let vendor = match vendor_id {
119                "0x1002" => GpuVendor::Amd,
120                "0x10de" => GpuVendor::Nvidia,
121                "0x8086" => GpuVendor::Intel,
122                _ => continue,
123            };
124
125            let vendor_name = match vendor {
126                GpuVendor::Amd => "AMD",
127                GpuVendor::Nvidia => "NVIDIA",
128                GpuVendor::Intel => "Intel",
129                GpuVendor::Apple => continue,
130                GpuVendor::Unknown => continue,
131            };
132
133            if vendor == GpuVendor::Amd {
134                if let Some(gfx) = amd_gfx_targets.get(amd_card_idx % amd_gfx_targets.len()) {
135                    models.push(Some(format!("{} ({})", vendor_name, gfx)));
136                } else {
137                    models.push(Some(vendor_name.to_string()));
138                }
139                amd_card_idx += 1;
140            } else {
141                models.push(Some(vendor_name.to_string()));
142            }
143        }
144    }
145
146    models
147}
148
149/// Format a raw GFX target version value to a string (e.g. 110003 -> "gfx1103").
150/// Returns None for value 0 (CPU node).
151fn gfx_target_to_string(val: u32) -> Option<String> {
152    if val == 0 {
153        return None;
154    }
155    let major = val / 10000;
156    let minor = (val % 10000) / 100;
157    let stepping = val % 100;
158
159    if stepping > 0 {
160        Some(format!("gfx{}{}{}", major, minor, stepping))
161    } else {
162        Some(format!("gfx{}{}", major, minor))
163    }
164}
165
166/// Collect all unique, non-zero AMD GFX target versions from KFD nodes.
167/// Skips CPU nodes (gfx_target_version == 0).
168/// Returns deduplicated targets in detection order.
169pub fn detect_amd_gfx_targets() -> Vec<String> {
170    let kfd_path = Path::new("/sys/class/kfd/kfd/topology/nodes");
171    if !kfd_path.exists() {
172        return Vec::new();
173    }
174
175    let mut targets = Vec::new();
176    if let Ok(entries) = fs::read_dir(kfd_path) {
177        for entry in entries.flatten() {
178            let props_path = entry.path().join("properties");
179            if let Ok(props) = fs::read_to_string(props_path) {
180                for line in props.lines() {
181                    if line.starts_with("gfx_target_version")
182                        && let Some(val_str) = line.split_whitespace().last()
183                        && let Ok(val) = val_str.parse::<u32>()
184                        && let Some(gfx) = gfx_target_to_string(val)
185                    {
186                        if !targets.contains(&gfx) {
187                            targets.push(gfx);
188                        }
189                        break;
190                    }
191                }
192            }
193        }
194    }
195    targets
196}
197
198/// Detect AMD GFX target version (e.g. "gfx1100").
199/// Returns the first non-zero GFX target found, or None.
200pub fn detect_amd_gfx_target() -> Option<String> {
201    detect_amd_gfx_targets().into_iter().next()
202}
203
204/// Get the best Lemonade asset suffix for the detected AMD architecture
205pub fn get_lemonade_gfx_suffix(gfx: &str) -> &'static str {
206    if gfx.starts_with("gfx103") {
207        "gfx103X"
208    } else if gfx.starts_with("gfx110") {
209        "gfx110X"
210    } else if gfx == "gfx1150" {
211        "gfx1150"
212    } else if gfx == "gfx1151" {
213        "gfx1151"
214    } else if gfx.starts_with("gfx120") {
215        "gfx120X"
216    } else {
217        // Fallback to most common recent if unknown
218        "gfx110X"
219    }
220}
221
222// ── Platform-specific GPU detection ──────────────────────────────────
223
224/// Detect GPU vendors on Windows using wmic.
225#[cfg(target_os = "windows")]
226pub fn detect_gpu_vendors_windows() -> Vec<GpuVendor> {
227    let mut vendors = Vec::new();
228    let output = std::process::Command::new("wmic")
229        .args(["path", "win32_VideoController", "get", "Name"])
230        .output();
231
232    let names = match output {
233        Ok(out) if out.status.success() => String::from_utf8_lossy(&out.stdout).to_string(),
234        _ => return Vec::new(),
235    };
236
237    for line in names.lines() {
238        let line = line.trim();
239        if line.is_empty() || line.eq_ignore_ascii_case("Name") {
240            continue;
241        }
242
243        let lower = line.to_lowercase();
244        if lower.contains("nvidia") {
245            if !vendors.contains(&GpuVendor::Nvidia) {
246                vendors.push(GpuVendor::Nvidia);
247            }
248        } else if lower.contains("amd") || lower.contains("radeon") || lower.contains("rx ") {
249            if !vendors.contains(&GpuVendor::Amd) {
250                vendors.push(GpuVendor::Amd);
251            }
252        } else if lower.contains("intel") {
253            if !vendors.contains(&GpuVendor::Intel) {
254                vendors.push(GpuVendor::Intel);
255            }
256        }
257    }
258
259    if vendors.is_empty() {
260        vendors.push(GpuVendor::Unknown);
261    }
262
263    vendors
264}
265
266/// Detect GPU models on Windows using wmic.
267#[cfg(target_os = "windows")]
268pub fn detect_gpu_models_windows() -> Vec<Option<String>> {
269    let output = std::process::Command::new("wmic")
270        .args(["path", "win32_VideoController", "get", "Name"])
271        .output();
272
273    let names = match output {
274        Ok(out) if out.status.success() => String::from_utf8_lossy(&out.stdout).to_string(),
275        _ => return Vec::new(),
276    };
277
278    let mut models = Vec::new();
279    for line in names.lines() {
280        let line = line.trim();
281        if line.is_empty() || line.eq_ignore_ascii_case("Name") {
282            continue;
283        }
284        models.push(Some(line.to_string()));
285    }
286
287    models
288}
289
290/// Detect GPU vendors on macOS using system_profiler.
291#[cfg(target_os = "macos")]
292pub fn detect_gpu_vendors_macos() -> Vec<GpuVendor> {
293    let mut vendors = Vec::new();
294    let output = std::process::Command::new("system_profiler")
295        .args(["SPDisplaysDataType"])
296        .output();
297
298    let data = match output {
299        Ok(out) if out.status.success() => String::from_utf8_lossy(&out.stdout).to_string(),
300        _ => return Vec::new(),
301    };
302
303    for line in data.lines() {
304        let trimmed = line.trim();
305        if !trimmed.contains(":") {
306            continue;
307        }
308
309        let gpu_name = trimmed.split(':').nth(1).unwrap_or("").trim();
310        let lower = gpu_name.to_lowercase();
311
312        if lower.contains("apple")
313            && (lower.contains("m1")
314                || lower.contains("m2")
315                || lower.contains("m3")
316                || lower.contains("m4")
317                || lower.contains("apple gpu")
318                || lower.contains("apple silicon"))
319        {
320            if !vendors.contains(&GpuVendor::Apple) {
321                vendors.push(GpuVendor::Apple);
322            }
323        } else if lower.contains("nvidia") {
324            if !vendors.contains(&GpuVendor::Nvidia) {
325                vendors.push(GpuVendor::Nvidia);
326            }
327        } else if lower.contains("amd") || lower.contains("radeon") || lower.contains("firepro") {
328            if !vendors.contains(&GpuVendor::Amd) {
329                vendors.push(GpuVendor::Amd);
330            }
331        } else if lower.contains("intel") {
332            if !vendors.contains(&GpuVendor::Intel) {
333                vendors.push(GpuVendor::Intel);
334            }
335        }
336    }
337
338    if vendors.is_empty() {
339        vendors.push(GpuVendor::Unknown);
340    }
341
342    vendors
343}
344
345/// Detect GPU models on macOS using system_profiler.
346#[cfg(target_os = "macos")]
347pub fn detect_gpu_models_macos() -> Vec<Option<String>> {
348    let output = std::process::Command::new("system_profiler")
349        .args(["SPDisplaysDataType"])
350        .output();
351
352    let data = match output {
353        Ok(out) if out.status.success() => String::from_utf8_lossy(&out.stdout).to_string(),
354        _ => return Vec::new(),
355    };
356
357    let mut models = Vec::new();
358    let mut in_gpu_section = false;
359
360    for line in data.lines() {
361        let trimmed = line.trim();
362
363        if trimmed.contains("Chipset Model") || trimmed.contains("GPU Name") {
364            in_gpu_section = true;
365            if let Some(name) = trimmed.split(':').nth(1) {
366                let name = name.trim();
367                if !name.is_empty() {
368                    models.push(Some(name.to_string()));
369                }
370            }
371        } else if in_gpu_section && trimmed.contains("Vendor") {
372            in_gpu_section = false;
373        } else if in_gpu_section && trimmed.is_empty() {
374            in_gpu_section = false;
375        }
376    }
377
378    models
379}
380
381/// Detect GPU vendors using platform-specific methods.
382#[cfg(target_os = "linux")]
383pub fn detect_gpu_vendors() -> Vec<GpuVendor> {
384    detect_gpu_vendors_linux_impl()
385}
386
387/// Detect GPU models using platform-specific methods.
388#[cfg(target_os = "linux")]
389pub fn detect_gpu_models() -> Vec<Option<String>> {
390    detect_gpu_models_linux_impl()
391}
392
393/// Detect GPU vendors using platform-specific methods.
394#[cfg(target_os = "windows")]
395pub fn detect_gpu_vendors() -> Vec<GpuVendor> {
396    detect_gpu_vendors_windows()
397}
398
399/// Detect GPU models using platform-specific methods.
400#[cfg(target_os = "windows")]
401pub fn detect_gpu_models() -> Vec<Option<String>> {
402    detect_gpu_models_windows()
403}
404
405/// Detect GPU vendors using platform-specific methods.
406#[cfg(target_os = "macos")]
407pub fn detect_gpu_vendors() -> Vec<GpuVendor> {
408    detect_gpu_vendors_macos()
409}
410
411/// Detect GPU models using platform-specific methods.
412#[cfg(target_os = "macos")]
413pub fn detect_gpu_models() -> Vec<Option<String>> {
414    detect_gpu_models_macos()
415}
416
417#[cfg(test)]
418mod tests {
419    use super::*;
420
421    #[test]
422    fn test_parse_windows_nvidia() {
423        let input = "Name\nNVIDIA GeForce RTX 4090\n";
424        let vendors = parse_gpu_name_for_vendor(input);
425        assert!(vendors.contains(&GpuVendor::Nvidia));
426    }
427
428    #[test]
429    fn test_parse_windows_amd() {
430        let input = "Name\nAMD Radeon RX 7900 XTX\n";
431        let vendors = parse_gpu_name_for_vendor(input);
432        assert!(vendors.contains(&GpuVendor::Amd));
433    }
434
435    #[test]
436    fn test_parse_windows_intel() {
437        let input = "Name\nIntel(R) UHD Graphics 770\n";
438        let vendors = parse_gpu_name_for_vendor(input);
439        assert!(vendors.contains(&GpuVendor::Intel));
440    }
441
442    #[test]
443    fn test_parse_windows_radeon() {
444        let input = "Name\nAMD Radeon RX 6600\nName\nRadeon RX 580\n";
445        let vendors = parse_gpu_name_for_vendor(input);
446        assert!(vendors.contains(&GpuVendor::Amd));
447        assert_eq!(vendors.len(), 1);
448    }
449
450    #[test]
451    fn test_parse_windows_multiple_gpus() {
452        let input = "Name\nNVIDIA GeForce RTX 3080\nName\nIntel(R) UHD Graphics 750\n";
453        let vendors = parse_gpu_name_for_vendor(input);
454        assert!(vendors.contains(&GpuVendor::Nvidia));
455        assert!(vendors.contains(&GpuVendor::Intel));
456        assert_eq!(vendors.len(), 2);
457    }
458
459    #[test]
460    fn test_parse_windows_empty() {
461        let input = "Name\n\n";
462        let vendors = parse_gpu_name_for_vendor(input);
463        assert!(vendors.is_empty());
464    }
465
466    #[test]
467    fn test_parse_macos_apple_silicon() {
468        let input = "Chipset Model: Apple M2\nType: GPU\nBus: Built-In\n";
469        let vendors = parse_macos_gpu_output(input);
470        assert!(vendors.contains(&GpuVendor::Apple));
471    }
472
473    #[test]
474    fn test_parse_macos_amd() {
475        let input = "Chipset Model: AMD Radeon Pro 5500M\nType: GPU\nBus: PCIe\nVendor: AMD\n";
476        let vendors = parse_macos_gpu_output(input);
477        assert!(vendors.contains(&GpuVendor::Amd));
478    }
479
480    #[test]
481    fn test_parse_macos_nvidia() {
482        let input = "Chipset Model: NVIDIA GeForce GTX 775M\nType: GPU\nBus: PCIe\n";
483        let vendors = parse_macos_gpu_output(input);
484        assert!(vendors.contains(&GpuVendor::Nvidia));
485    }
486
487    #[test]
488    fn test_parse_macos_intel() {
489        let input = "Chipset Model: Intel Iris Pro\nType: GPU\nBus: Built-In\n";
490        let vendors = parse_macos_gpu_output(input);
491        assert!(vendors.contains(&GpuVendor::Intel));
492    }
493
494    #[test]
495    fn test_parse_macos_m3() {
496        let input = "Chipset Model: Apple M3 Max\nType: GPU\n";
497        let vendors = parse_macos_gpu_output(input);
498        assert!(vendors.contains(&GpuVendor::Apple));
499    }
500
501    #[test]
502    fn test_parse_macos_m4() {
503        let input = "Chipset Model: Apple M4 Pro\nType: GPU\n";
504        let vendors = parse_macos_gpu_output(input);
505        assert!(vendors.contains(&GpuVendor::Apple));
506    }
507
508    // Helper function to parse GPU names from wmic-like output
509    fn parse_gpu_name_for_vendor(input: &str) -> Vec<GpuVendor> {
510        let mut vendors = Vec::new();
511        for line in input.lines() {
512            let line = line.trim();
513            if line.is_empty() || line.eq_ignore_ascii_case("Name") {
514                continue;
515            }
516            let lower = line.to_lowercase();
517            if lower.contains("nvidia") {
518                if !vendors.contains(&GpuVendor::Nvidia) {
519                    vendors.push(GpuVendor::Nvidia);
520                }
521            } else if lower.contains("amd") || lower.contains("radeon") || lower.contains("rx ") {
522                if !vendors.contains(&GpuVendor::Amd) {
523                    vendors.push(GpuVendor::Amd);
524                }
525            } else if lower.contains("intel") && !vendors.contains(&GpuVendor::Intel) {
526                vendors.push(GpuVendor::Intel);
527            }
528        }
529        vendors
530    }
531
532    // Helper function to parse GPU names from system_profiler output
533    fn parse_macos_gpu_output(input: &str) -> Vec<GpuVendor> {
534        let mut vendors = Vec::new();
535        for line in input.lines() {
536            let trimmed = line.trim();
537            if !trimmed.contains(":") {
538                continue;
539            }
540            let gpu_name = trimmed.split(':').nth(1).unwrap_or("").trim();
541            let lower = gpu_name.to_lowercase();
542            if lower.contains("apple")
543                && (lower.contains("m1")
544                    || lower.contains("m2")
545                    || lower.contains("m3")
546                    || lower.contains("m4")
547                    || lower.contains("apple gpu")
548                    || lower.contains("apple silicon"))
549            {
550                if !vendors.contains(&GpuVendor::Apple) {
551                    vendors.push(GpuVendor::Apple);
552                }
553            } else if lower.contains("nvidia") {
554                if !vendors.contains(&GpuVendor::Nvidia) {
555                    vendors.push(GpuVendor::Nvidia);
556                }
557            } else if lower.contains("amd") || lower.contains("radeon") || lower.contains("firepro")
558            {
559                if !vendors.contains(&GpuVendor::Amd) {
560                    vendors.push(GpuVendor::Amd);
561                }
562            } else if lower.contains("intel") && !vendors.contains(&GpuVendor::Intel) {
563                vendors.push(GpuVendor::Intel);
564            }
565        }
566        vendors
567    }
568}