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() => {
234            String::from_utf8_lossy(&out.stdout).to_string()
235        }
236        _ => return Vec::new(),
237    };
238
239    for line in names.lines() {
240        let line = line.trim();
241        if line.is_empty() || line.eq_ignore_ascii_case("Name") {
242            continue;
243        }
244
245        let lower = line.to_lowercase();
246        if lower.contains("nvidia") {
247            if !vendors.contains(&GpuVendor::Nvidia) {
248                vendors.push(GpuVendor::Nvidia);
249            }
250        } else if lower.contains("amd") || lower.contains("radeon") || lower.contains("rx ") {
251            if !vendors.contains(&GpuVendor::Amd) {
252                vendors.push(GpuVendor::Amd);
253            }
254        } else if lower.contains("intel") {
255            if !vendors.contains(&GpuVendor::Intel) {
256                vendors.push(GpuVendor::Intel);
257            }
258        }
259    }
260
261    if vendors.is_empty() {
262        vendors.push(GpuVendor::Unknown);
263    }
264
265    vendors
266}
267
268/// Detect GPU models on Windows using wmic.
269#[cfg(target_os = "windows")]
270pub fn detect_gpu_models_windows() -> Vec<Option<String>> {
271    let output = std::process::Command::new("wmic")
272        .args(["path", "win32_VideoController", "get", "Name"])
273        .output();
274
275    let names = match output {
276        Ok(out) if out.status.success() => {
277            String::from_utf8_lossy(&out.stdout).to_string()
278        }
279        _ => return Vec::new(),
280    };
281
282    let mut models = Vec::new();
283    for line in names.lines() {
284        let line = line.trim();
285        if line.is_empty() || line.eq_ignore_ascii_case("Name") {
286            continue;
287        }
288        models.push(Some(line.to_string()));
289    }
290
291    models
292}
293
294/// Detect GPU vendors on macOS using system_profiler.
295#[cfg(target_os = "macos")]
296pub fn detect_gpu_vendors_macos() -> Vec<GpuVendor> {
297    let mut vendors = Vec::new();
298    let output = std::process::Command::new("system_profiler")
299        .args(["SPDisplaysDataType"])
300        .output();
301
302    let data = match output {
303        Ok(out) if out.status.success() => {
304            String::from_utf8_lossy(&out.stdout).to_string()
305        }
306        _ => return Vec::new(),
307    };
308
309    for line in data.lines() {
310        let trimmed = line.trim();
311        if !trimmed.contains(":") {
312            continue;
313        }
314
315        let gpu_name = trimmed.split(':').nth(1).unwrap_or("").trim();
316        let lower = gpu_name.to_lowercase();
317
318        if lower.contains("apple") && (lower.contains("m1") || lower.contains("m2") || lower.contains("m3") || lower.contains("m4") || lower.contains("apple gpu") || lower.contains("apple silicon")) {
319            if !vendors.contains(&GpuVendor::Apple) {
320                vendors.push(GpuVendor::Apple);
321            }
322        } else if lower.contains("nvidia") {
323            if !vendors.contains(&GpuVendor::Nvidia) {
324                vendors.push(GpuVendor::Nvidia);
325            }
326        } else if lower.contains("amd") || lower.contains("radeon") || lower.contains("firepro") {
327            if !vendors.contains(&GpuVendor::Amd) {
328                vendors.push(GpuVendor::Amd);
329            }
330        } else if lower.contains("intel") {
331            if !vendors.contains(&GpuVendor::Intel) {
332                vendors.push(GpuVendor::Intel);
333            }
334        }
335    }
336
337    if vendors.is_empty() {
338        vendors.push(GpuVendor::Unknown);
339    }
340
341    vendors
342}
343
344/// Detect GPU models on macOS using system_profiler.
345#[cfg(target_os = "macos")]
346pub fn detect_gpu_models_macos() -> Vec<Option<String>> {
347    let output = std::process::Command::new("system_profiler")
348        .args(["SPDisplaysDataType"])
349        .output();
350
351    let data = match output {
352        Ok(out) if out.status.success() => {
353            String::from_utf8_lossy(&out.stdout).to_string()
354        }
355        _ => return Vec::new(),
356    };
357
358    let mut models = Vec::new();
359    let mut in_gpu_section = false;
360
361    for line in data.lines() {
362        let trimmed = line.trim();
363
364        if trimmed.contains("Chipset Model") || trimmed.contains("GPU Name") {
365            in_gpu_section = true;
366            if let Some(name) = trimmed.split(':').nth(1) {
367                let name = name.trim();
368                if !name.is_empty() {
369                    models.push(Some(name.to_string()));
370                }
371            }
372        } else if in_gpu_section && trimmed.contains("Vendor") {
373            in_gpu_section = false;
374        } else if in_gpu_section && trimmed.is_empty() {
375            in_gpu_section = false;
376        }
377    }
378
379    models
380}
381
382/// Detect GPU vendors on Linux (wrapper for cfg visibility).
383#[cfg(target_os = "linux")]
384#[allow(dead_code)]
385pub fn detect_gpu_vendors_linux() -> Vec<GpuVendor> {
386    detect_gpu_vendors_linux_impl()
387}
388
389/// Detect GPU models on Linux (wrapper for cfg visibility).
390#[cfg(target_os = "linux")]
391#[allow(dead_code)]
392pub fn detect_gpu_models_linux() -> Vec<Option<String>> {
393    detect_gpu_models_linux_impl()
394}
395
396/// Detect GPU vendors using platform-specific methods.
397#[cfg(target_os = "linux")]
398pub fn detect_gpu_vendors() -> Vec<GpuVendor> {
399    detect_gpu_vendors_linux_impl()
400}
401
402/// Detect GPU models using platform-specific methods.
403#[cfg(target_os = "linux")]
404pub fn detect_gpu_models() -> Vec<Option<String>> {
405    detect_gpu_models_linux_impl()
406}
407
408/// Detect GPU vendors using platform-specific methods.
409#[cfg(target_os = "windows")]
410pub fn detect_gpu_vendors() -> Vec<GpuVendor> {
411    detect_gpu_vendors_windows()
412}
413
414/// Detect GPU models using platform-specific methods.
415#[cfg(target_os = "windows")]
416pub fn detect_gpu_models() -> Vec<Option<String>> {
417    detect_gpu_models_windows()
418}
419
420/// Detect GPU vendors using platform-specific methods.
421#[cfg(target_os = "macos")]
422pub fn detect_gpu_vendors() -> Vec<GpuVendor> {
423    detect_gpu_vendors_macos()
424}
425
426/// Detect GPU models using platform-specific methods.
427#[cfg(target_os = "macos")]
428pub fn detect_gpu_models() -> Vec<Option<String>> {
429    detect_gpu_models_macos()
430}
431
432#[cfg(test)]
433mod tests {
434    use super::*;
435
436    #[test]
437    fn test_parse_windows_nvidia() {
438        let input = "Name\nNVIDIA GeForce RTX 4090\n";
439        let vendors = parse_gpu_name_for_vendor(input);
440        assert!(vendors.contains(&GpuVendor::Nvidia));
441    }
442
443    #[test]
444    fn test_parse_windows_amd() {
445        let input = "Name\nAMD Radeon RX 7900 XTX\n";
446        let vendors = parse_gpu_name_for_vendor(input);
447        assert!(vendors.contains(&GpuVendor::Amd));
448    }
449
450    #[test]
451    fn test_parse_windows_intel() {
452        let input = "Name\nIntel(R) UHD Graphics 770\n";
453        let vendors = parse_gpu_name_for_vendor(input);
454        assert!(vendors.contains(&GpuVendor::Intel));
455    }
456
457    #[test]
458    fn test_parse_windows_radeon() {
459        let input = "Name\nAMD Radeon RX 6600\nName\nRadeon RX 580\n";
460        let vendors = parse_gpu_name_for_vendor(input);
461        assert!(vendors.contains(&GpuVendor::Amd));
462        assert_eq!(vendors.len(), 1);
463    }
464
465    #[test]
466    fn test_parse_windows_multiple_gpus() {
467        let input = "Name\nNVIDIA GeForce RTX 3080\nName\nIntel(R) UHD Graphics 750\n";
468        let vendors = parse_gpu_name_for_vendor(input);
469        assert!(vendors.contains(&GpuVendor::Nvidia));
470        assert!(vendors.contains(&GpuVendor::Intel));
471        assert_eq!(vendors.len(), 2);
472    }
473
474    #[test]
475    fn test_parse_windows_empty() {
476        let input = "Name\n\n";
477        let vendors = parse_gpu_name_for_vendor(input);
478        assert!(vendors.is_empty());
479    }
480
481    #[test]
482    fn test_parse_macos_apple_silicon() {
483        let input = "Chipset Model: Apple M2\nType: GPU\nBus: Built-In\n";
484        let vendors = parse_macos_gpu_output(input);
485        assert!(vendors.contains(&GpuVendor::Apple));
486    }
487
488    #[test]
489    fn test_parse_macos_amd() {
490        let input = "Chipset Model: AMD Radeon Pro 5500M\nType: GPU\nBus: PCIe\nVendor: AMD\n";
491        let vendors = parse_macos_gpu_output(input);
492        assert!(vendors.contains(&GpuVendor::Amd));
493    }
494
495    #[test]
496    fn test_parse_macos_nvidia() {
497        let input = "Chipset Model: NVIDIA GeForce GTX 775M\nType: GPU\nBus: PCIe\n";
498        let vendors = parse_macos_gpu_output(input);
499        assert!(vendors.contains(&GpuVendor::Nvidia));
500    }
501
502    #[test]
503    fn test_parse_macos_intel() {
504        let input = "Chipset Model: Intel Iris Pro\nType: GPU\nBus: Built-In\n";
505        let vendors = parse_macos_gpu_output(input);
506        assert!(vendors.contains(&GpuVendor::Intel));
507    }
508
509    #[test]
510    fn test_parse_macos_m3() {
511        let input = "Chipset Model: Apple M3 Max\nType: GPU\n";
512        let vendors = parse_macos_gpu_output(input);
513        assert!(vendors.contains(&GpuVendor::Apple));
514    }
515
516    #[test]
517    fn test_parse_macos_m4() {
518        let input = "Chipset Model: Apple M4 Pro\nType: GPU\n";
519        let vendors = parse_macos_gpu_output(input);
520        assert!(vendors.contains(&GpuVendor::Apple));
521    }
522
523    // Helper function to parse GPU names from wmic-like output
524    fn parse_gpu_name_for_vendor(input: &str) -> Vec<GpuVendor> {
525        let mut vendors = Vec::new();
526        for line in input.lines() {
527            let line = line.trim();
528            if line.is_empty() || line.eq_ignore_ascii_case("Name") {
529                continue;
530            }
531            let lower = line.to_lowercase();
532            if lower.contains("nvidia") {
533                if !vendors.contains(&GpuVendor::Nvidia) {
534                    vendors.push(GpuVendor::Nvidia);
535                }
536            } else if lower.contains("amd") || lower.contains("radeon") || lower.contains("rx ") {
537                if !vendors.contains(&GpuVendor::Amd) {
538                    vendors.push(GpuVendor::Amd);
539                }
540            } else if lower.contains("intel") {
541                if !vendors.contains(&GpuVendor::Intel) {
542                    vendors.push(GpuVendor::Intel);
543                }
544            }
545        }
546        vendors
547    }
548
549    // Helper function to parse GPU names from system_profiler output
550    fn parse_macos_gpu_output(input: &str) -> Vec<GpuVendor> {
551        let mut vendors = Vec::new();
552        for line in input.lines() {
553            let trimmed = line.trim();
554            if !trimmed.contains(":") {
555                continue;
556            }
557            let gpu_name = trimmed.split(':').nth(1).unwrap_or("").trim();
558            let lower = gpu_name.to_lowercase();
559            if lower.contains("apple") && (lower.contains("m1") || lower.contains("m2") || lower.contains("m3") || lower.contains("m4") || lower.contains("apple gpu") || lower.contains("apple silicon")) {
560                if !vendors.contains(&GpuVendor::Apple) {
561                    vendors.push(GpuVendor::Apple);
562                }
563            } else if lower.contains("nvidia") {
564                if !vendors.contains(&GpuVendor::Nvidia) {
565                    vendors.push(GpuVendor::Nvidia);
566                }
567            } else if lower.contains("amd") || lower.contains("radeon") || lower.contains("firepro") {
568                if !vendors.contains(&GpuVendor::Amd) {
569                    vendors.push(GpuVendor::Amd);
570                }
571            } else if lower.contains("intel") {
572                if !vendors.contains(&GpuVendor::Intel) {
573                    vendors.push(GpuVendor::Intel);
574                }
575            }
576        }
577        vendors
578    }
579}