Skip to main content

vyre_driver_wgpu/runtime/device/
selector.rs

1//! Adapter selection + enumeration (C5 refactor).
2//!
3//! The legacy [`super::device::cached_device`] singleton picks the
4//! first adapter `wgpu::Instance::request_adapter` returns  -  fine for
5//! a single-GPU dev box, useless for multi-GPU servers that need to
6//! choose a specific device by vendor, index, or power preference.
7//!
8//! This module ships the explicit selection API:
9//!
10//! * [`enumerate_adapters`]  -  list every adapter wgpu reports.
11//! * [`AdapterCriteria`]  -  match by device type, vendor, name
12//!   substring, or power preference.
13//! * [`select_adapter`]  -  pick one matching the criteria (returns
14//!   the first match; callers wanting all matches iterate
15//!   [`enumerate_adapters`] themselves).
16//! * [`init_device_for_adapter`]  -  build a device+queue bound to the
17//!   chosen adapter.
18//! * `VYRE_ADAPTER_INDEX`  -  env override used by the backend
19//!   auto-picker to route programs to a specific device without
20//!   patching code.
21//!
22//! The legacy `cached_device()` still serves the default case: one
23//! singleton device, first compatible adapter. Callers that want
24//! multi-GPU now select an adapter by index before constructing a
25//! device/queue pair.
26
27use vyre_driver::BackendError;
28
29type Result<T, E = BackendError> = std::result::Result<T, E>;
30
31use super::reserve_probe_vec;
32
33/// Stable adapter identity used for deterministic recovery.
34#[derive(Clone, Debug, Eq, PartialEq)]
35pub(crate) struct AdapterIdentity {
36    name: String,
37    vendor: u32,
38    device: u32,
39    device_type: wgpu::DeviceType,
40    driver: String,
41    driver_info: String,
42    backend: wgpu::Backend,
43}
44
45impl AdapterIdentity {
46    pub(crate) fn from_info(info: &wgpu::AdapterInfo) -> Self {
47        Self {
48            name: info.name.clone(),
49            vendor: info.vendor,
50            device: info.device,
51            device_type: info.device_type,
52            driver: info.driver.clone(),
53            driver_info: info.driver_info.clone(),
54            backend: info.backend,
55        }
56    }
57
58    fn matches(&self, info: &wgpu::AdapterInfo) -> bool {
59        self.name == info.name
60            && self.vendor == info.vendor
61            && self.device == info.device
62            && self.device_type == info.device_type
63            && self.driver == info.driver
64            && self.driver_info == info.driver_info
65            && self.backend == info.backend
66    }
67}
68
69/// Criteria used by [`select_adapter`].
70#[derive(Debug, Default, Clone)]
71pub struct AdapterCriteria {
72    /// Prefer an adapter whose `device_type` matches.
73    pub device_type: Option<wgpu::DeviceType>,
74    /// Prefer an adapter whose vendor id matches.
75    pub vendor: Option<u32>,
76    /// Prefer an adapter whose name contains this substring
77    /// (case-insensitive).
78    pub name_contains: Option<String>,
79    /// Prefer an adapter with this power policy.
80    pub power: Option<wgpu::PowerPreference>,
81}
82
83/// Human-readable adapter probe details for GPU acquisition failures.
84#[derive(Clone, Debug, Default, Eq, PartialEq)]
85pub struct AdapterProbeReport {
86    /// Adapters visible to wgpu during the centralized probe.
87    pub probed: Vec<String>,
88    /// Feature, limit, or device-request reasons that prevented use.
89    pub missing: Vec<String>,
90}
91
92impl AdapterCriteria {
93    /// Build criteria for a high-performance discrete GPU.
94    #[must_use]
95    pub fn high_performance() -> Self {
96        Self {
97            device_type: Some(wgpu::DeviceType::DiscreteGpu),
98            power: Some(wgpu::PowerPreference::HighPerformance),
99            ..Self::default()
100        }
101    }
102
103    /// Build criteria for a low-power integrated GPU (laptop
104    /// battery savings).
105    #[must_use]
106    pub fn low_power() -> Self {
107        Self {
108            device_type: Some(wgpu::DeviceType::IntegratedGpu),
109            power: Some(wgpu::PowerPreference::LowPower),
110            ..Self::default()
111        }
112    }
113}
114
115/// List every adapter the wgpu instance reports.
116#[must_use]
117pub fn enumerate_adapters() -> Vec<wgpu::AdapterInfo> {
118    match try_enumerate_adapters() {
119        Ok(adapters) => adapters,
120        Err(error) => {
121            // Law 10: a probe failure is NOT "no GPUs present". Surface it
122            // loudly so callers do not read an empty vec as a device-free host.
123            tracing::error!(
124                %error,
125                "adapter enumeration probe failed; reporting zero adapters is a probe error, not an absence of GPUs"
126            );
127            Vec::new()
128        }
129    }
130}
131
132/// List every adapter the wgpu instance reports with fallible metadata staging.
133///
134/// # Errors
135///
136/// Returns `BackendError` when probe-result metadata cannot be reserved.
137pub(crate) fn try_enumerate_adapters() -> Result<Vec<wgpu::AdapterInfo>> {
138    let instance = wgpu::Instance::default();
139    let adapters = instance.enumerate_adapters(wgpu::Backends::all());
140    let mut infos = Vec::new();
141    reserve_probe_vec(&mut infos, adapters.len(), "adapter enumeration metadata")?;
142    infos.extend(adapters.iter().map(wgpu::Adapter::get_info));
143    Ok(infos)
144}
145
146/// Report whether the centralized adapter probe can see at least one real GPU.
147#[must_use]
148pub fn has_real_gpu_adapter() -> bool {
149    let instance = wgpu::Instance::default();
150    instance
151        .enumerate_adapters(wgpu::Backends::all())
152        .iter()
153        .any(|adapter| crate::capabilities::is_real_gpu(&adapter.get_info()))
154}
155
156/// Re-open the live wgpu adapter matching a previously selected adapter info.
157///
158/// Tests and capability probes use this instead of directly constructing their
159/// own `wgpu::Instance` so adapter identity and failure diagnostics stay in the
160/// runtime device contract.
161///
162/// # Errors
163///
164/// Returns `BackendError` when the adapter is no longer visible.
165pub fn adapter_for_info(expected: &wgpu::AdapterInfo) -> Result<wgpu::Adapter> {
166    let instance = wgpu::Instance::default();
167    let adapters = instance.enumerate_adapters(wgpu::Backends::all());
168    let mut probed = Vec::new();
169    reserve_probe_vec(&mut probed, adapters.len(), "adapter recovery probe report")?;
170    for adapter in adapters {
171        let candidate = adapter.get_info();
172        if adapter_info_matches(&candidate, expected) {
173            return Ok(adapter);
174        }
175        probed.push(format!(
176            "{} ({:?}, backend={:?}, vendor={:08x}, device={:08x})",
177            candidate.name,
178            candidate.device_type,
179            candidate.backend,
180            candidate.vendor,
181            candidate.device
182        ));
183    }
184
185    Err(BackendError::new(format!(
186        "selected adapter `{}` ({:?}, backend={:?}, vendor={:08x}, device={:08x}) is no longer enumerable. Probed adapters: [{}]. Fix: repair GPU visibility or reacquire the WGPU backend.",
187        expected.name,
188        expected.device_type,
189        expected.backend,
190        expected.vendor,
191        expected.device,
192        probed.join(", ")
193    )))
194}
195
196fn adapter_info_matches(candidate: &wgpu::AdapterInfo, expected: &wgpu::AdapterInfo) -> bool {
197    candidate.name == expected.name
198        && candidate.vendor == expected.vendor
199        && candidate.device == expected.device
200        && candidate.device_type == expected.device_type
201        && candidate.driver == expected.driver
202        && candidate.driver_info == expected.driver_info
203        && candidate.backend == expected.backend
204}
205
206/// Build the centralized adapter diagnostic report used by acquisition errors.
207#[must_use]
208pub fn adapter_probe_report() -> AdapterProbeReport {
209    let instance = wgpu::Instance::default();
210    let adapters = instance.enumerate_adapters(wgpu::Backends::all());
211    let mut report = AdapterProbeReport {
212        probed: Vec::new(),
213        missing: Vec::new(),
214    };
215
216    for adapter in adapters {
217        let info = adapter.get_info();
218        report.probed.push(format!(
219            "{} ({:?}, backend={:?})",
220            info.name, info.device_type, info.backend
221        ));
222        if matches!(
223            info.device_type,
224            wgpu::DeviceType::Cpu | wgpu::DeviceType::Other
225        ) {
226            continue;
227        }
228        if !adapter.features().contains(wgpu::Features::TIMESTAMP_QUERY) {
229            report.missing.push("TIMESTAMP_QUERY".to_string());
230        }
231        if !adapter
232            .features()
233            .contains(wgpu::Features::TIMESTAMP_QUERY_INSIDE_ENCODERS)
234        {
235            report
236                .missing
237                .push("TIMESTAMP_QUERY_INSIDE_ENCODERS".to_string());
238        }
239        let adapter_limits = adapter.limits();
240        if let Err(error) = pollster::block_on(adapter.request_device(&wgpu::DeviceDescriptor {
241            label: Some("vyre probe"),
242            required_features: wgpu::Features::empty(),
243            required_limits: wgpu::Limits {
244                max_storage_buffers_per_shader_stage:
245                    adapter_limits.max_storage_buffers_per_shader_stage,
246                ..wgpu::Limits::default()
247            },
248            memory_hints: wgpu::MemoryHints::default(),
249            trace: wgpu::Trace::Off,
250        })) {
251            report
252                .missing
253                .push(format!("device request failed on {}: {error}", info.name));
254        }
255    }
256
257    report
258}
259
260/// Select the first adapter matching `criteria`. Returns its index
261/// into [`enumerate_adapters`] plus its info.
262///
263/// # Errors
264///
265/// Returns `BackendError` when no adapter matches.
266pub fn select_adapter(criteria: &AdapterCriteria) -> Result<(usize, wgpu::AdapterInfo)> {
267    let instance = wgpu::Instance::default();
268    let adapters = instance.enumerate_adapters(wgpu::Backends::all());
269    for (idx, adapter) in adapters.iter().enumerate() {
270        let info = adapter.get_info();
271        if adapter_is_selectable(&info, criteria) {
272            return Ok((idx, info));
273        }
274    }
275    Err(BackendError::new(format!(
276        "no real GPU adapter matches criteria {criteria:?}. Fix: loosen the criteria or install drivers exposing the requested GPU class."
277    )))
278}
279
280fn adapter_is_selectable(info: &wgpu::AdapterInfo, criteria: &AdapterCriteria) -> bool {
281    crate::capabilities::is_real_gpu(info) && adapter_matches(info, criteria)
282}
283
284fn adapter_matches(info: &wgpu::AdapterInfo, criteria: &AdapterCriteria) -> bool {
285    if let Some(ty) = criteria.device_type {
286        if info.device_type != ty {
287            return false;
288        }
289    }
290    if let Some(vendor) = criteria.vendor {
291        if info.vendor != vendor {
292            return false;
293        }
294    }
295    if let Some(needle) = &criteria.name_contains {
296        if !adapter_name_contains(&info.name, needle) {
297            return false;
298        }
299    }
300    true
301}
302
303fn adapter_name_contains(name: &str, needle: &str) -> bool {
304    if needle.is_empty() {
305        return true;
306    }
307    if name.is_ascii() && needle.is_ascii() {
308        return name
309            .as_bytes()
310            .windows(needle.len())
311            .any(|window| window.eq_ignore_ascii_case(needle.as_bytes()));
312    }
313    name.to_lowercase().contains(&needle.to_lowercase())
314}
315
316/// Initialize a device + queue bound to the adapter at `index`.
317///
318/// Pairs with [`enumerate_adapters`] / [`select_adapter`] to give
319/// callers full control over which GPU the backend binds to.
320///
321/// # Errors
322///
323/// Returns `BackendError` when `index` is out of range or device
324/// creation fails.
325pub fn init_device_for_adapter(
326    index: usize,
327) -> Result<(
328    (wgpu::Device, wgpu::Queue),
329    wgpu::AdapterInfo,
330    crate::runtime::device::EnabledFeatures,
331)> {
332    super::device::wait_for_gpu(acquire_gpu_for_adapter(index))
333}
334
335/// Recreate a device on the same adapter identity used by an existing backend.
336///
337/// # Errors
338///
339/// Returns `BackendError` when the adapter disappeared, no longer reports as a
340/// real GPU, or rejects device creation.
341pub(crate) fn init_device_for_adapter_identity(
342    identity: &AdapterIdentity,
343) -> Result<(
344    (wgpu::Device, wgpu::Queue),
345    wgpu::AdapterInfo,
346    crate::runtime::device::EnabledFeatures,
347)> {
348    super::device::wait_for_gpu(acquire_gpu_for_adapter_identity(identity))
349}
350
351async fn acquire_gpu_for_adapter_identity(
352    identity: &AdapterIdentity,
353) -> Result<(
354    (wgpu::Device, wgpu::Queue),
355    wgpu::AdapterInfo,
356    crate::runtime::device::EnabledFeatures,
357)> {
358    let instance = wgpu::Instance::default();
359    let adapters = instance.enumerate_adapters(wgpu::Backends::all());
360    for adapter in &adapters {
361        let info = adapter.get_info();
362        if identity.matches(&info) {
363            if !crate::capabilities::is_real_gpu(&info) {
364                return Err(BackendError::new(format!(
365                    "recovery target `{}` now reports device type {:?}, which is not a real GPU execution target. Fix: restore the original GPU adapter or construct a new backend for the changed adapter.",
366                    info.name, info.device_type
367                )));
368            }
369            return super::device::request_device_for_adapter(adapter, "vyre device (recovered)")
370                .await;
371        }
372    }
373
374    let mut probed = Vec::new();
375    reserve_probe_vec(
376        &mut probed,
377        adapters.len(),
378        "adapter identity recovery probe report",
379    )?;
380    probed.extend(adapters.iter().map(|adapter| {
381        let info = adapter.get_info();
382        format!(
383            "{} ({:?}, backend={:?}, vendor={:08x}, device={:08x})",
384            info.name, info.device_type, info.backend, info.vendor, info.device
385        )
386    }));
387    Err(BackendError::new(format!(
388        "original recovery adapter was not found. Target: {:?}. Probed adapters: [{}]. Fix: restore the original GPU or create a new WgpuBackend for the available adapter.",
389        identity,
390        probed.join(", ")
391    )))
392}
393
394/// Async variant of [`init_device_for_adapter`].
395///
396/// # Errors
397///
398/// Returns `BackendError` when `index` is out of range or device
399/// creation fails.
400pub async fn acquire_gpu_for_adapter(
401    index: usize,
402) -> Result<(
403    (wgpu::Device, wgpu::Queue),
404    wgpu::AdapterInfo,
405    crate::runtime::device::EnabledFeatures,
406)> {
407    let instance = wgpu::Instance::default();
408    let adapters = instance.enumerate_adapters(wgpu::Backends::all());
409    let adapter = adapters.get(index).ok_or_else(|| BackendError::new(format!(
410        "adapter index {index} out of range (saw {} adapters). Fix: call enumerate_adapters() first to see valid indices.",
411        adapters.len()
412    )))?;
413    let info = adapter.get_info();
414    if !crate::capabilities::is_real_gpu(&info) {
415        return Err(BackendError::new(format!(
416            "adapter index {index} resolved to `{}` with device type {:?}, which is not a real GPU execution target. Fix: choose a discrete, integrated, or virtual GPU adapter.",
417            info.name, info.device_type
418        )));
419    }
420    super::device::request_device_for_adapter(adapter, "vyre device (selected)").await
421}
422
423/// Read the `VYRE_ADAPTER_INDEX` env override. `None` when unset.
424///
425/// # Errors
426///
427/// Returns an actionable GPU configuration error when the env var is
428/// set but cannot be parsed. A typoed adapter override must not
429/// silently fall back to automatic GPU selection.
430#[must_use]
431pub fn adapter_index_from_env() -> Result<Option<usize>> {
432    adapter_index_from_raw(std::env::var("VYRE_ADAPTER_INDEX").ok().as_deref())
433}
434
435fn adapter_index_from_raw(raw: Option<&str>) -> Result<Option<usize>> {
436    let Some(raw) = raw else {
437        return Ok(None);
438    };
439    raw.parse::<usize>().map(Some).map_err(|error| BackendError::new(format!(
440        "VYRE_ADAPTER_INDEX={raw:?} is not a valid adapter index: {error}. Fix: set VYRE_ADAPTER_INDEX to a non-negative integer from enumerate_adapters(), or unset it for automatic GPU selection."
441    )))
442}
443
444#[cfg(test)]
445mod tests {
446    use super::*;
447
448    #[test]
449    fn enumerate_adapters_finds_required_gpu() {
450        let adapters = enumerate_adapters();
451        assert_ne!(adapters.len(), 0,
452            "Fix: WGPU adapter enumeration returned no adapters on a GPU-required release host; repair driver/runtime configuration instead of accepting a CPU-only environment."
453        );
454    }
455
456    #[test]
457    fn criteria_high_perf_has_discrete_preset() {
458        let c = AdapterCriteria::high_performance();
459        assert_eq!(c.device_type, Some(wgpu::DeviceType::DiscreteGpu));
460        assert_eq!(c.power, Some(wgpu::PowerPreference::HighPerformance));
461    }
462
463    #[test]
464    fn criteria_low_power_has_integrated_preset() {
465        let c = AdapterCriteria::low_power();
466        assert_eq!(c.device_type, Some(wgpu::DeviceType::IntegratedGpu));
467    }
468
469    #[test]
470    fn env_override_parses_valid_index() {
471        assert_eq!(adapter_index_from_raw(Some("3")).unwrap(), Some(3));
472    }
473
474    #[test]
475    fn env_override_rejects_garbage() {
476        let error = adapter_index_from_raw(Some("not-a-number"))
477            .expect_err("invalid VYRE_ADAPTER_INDEX must error");
478        assert!(
479            error.to_string().contains("VYRE_ADAPTER_INDEX"),
480            "Fix: invalid adapter-index errors must name the misconfigured env var"
481        );
482    }
483
484    #[test]
485    fn selection_rejects_cpu_adapters_before_device_acquisition() {
486        let cpu_info = wgpu::AdapterInfo {
487            name: "llvmpipe".to_string(),
488            vendor: 0,
489            device: 0,
490            device_type: wgpu::DeviceType::Cpu,
491            driver: "software".to_string(),
492            driver_info: "cpu".to_string(),
493            backend: wgpu::Backend::Vulkan,
494        };
495        let gpu_info = wgpu::AdapterInfo {
496            name: "RTX 5090".to_string(),
497            vendor: 0x10de,
498            device: 0x2c02,
499            device_type: wgpu::DeviceType::DiscreteGpu,
500            driver: "nvidia".to_string(),
501            driver_info: "gpu".to_string(),
502            backend: wgpu::Backend::Vulkan,
503        };
504        let criteria = AdapterCriteria::default();
505
506        assert!(
507            !adapter_is_selectable(&cpu_info, &criteria),
508            "Fix: adapter selection must never return CPU/Other devices for later fallback handling."
509        );
510        assert!(adapter_is_selectable(&gpu_info, &criteria));
511    }
512
513    #[test]
514    fn adapter_identity_matches_every_recovery_field() {
515        let info = wgpu::AdapterInfo {
516            name: "gpu-a".to_string(),
517            vendor: 0x10de,
518            device: 0x2684,
519            device_type: wgpu::DeviceType::DiscreteGpu,
520            driver: "nvidia".to_string(),
521            driver_info: "driver-a".to_string(),
522            backend: wgpu::Backend::Vulkan,
523        };
524        let identity = AdapterIdentity::from_info(&info);
525        assert!(identity.matches(&info));
526
527        let mut changed = info.clone();
528        changed.device = 0x2685;
529        assert!(
530            !identity.matches(&changed),
531            "Fix: recovery must not silently bind to a different physical adapter."
532        );
533    }
534
535    #[test]
536    fn adapter_name_contains_matches_ascii_without_lowercase_in_hot_path() {
537        assert!(adapter_name_contains("NVIDIA GeForce RTX 5090", "rtx"));
538        assert!(adapter_name_contains("NVIDIA GeForce RTX 5090", "RTX"));
539        assert!(!adapter_name_contains("NVIDIA GeForce RTX 5090", "radeon"));
540        assert!(adapter_name_contains("Mötley GPU", "mötley"));
541    }
542}