Skip to main content

vyre_driver_wgpu/
backend_impl.rs

1//! vyre_driver::VyreBackend implementation and core WgpuBackend methods.
2
3use crate::staging_reserve::{reserve_backend_vec, reserve_smallvec, reserve_vec};
4use crate::{AdapterRecoveryTarget, DispatchArena, WgpuBackend};
5use std::hash::{BuildHasherDefault, Hasher};
6use std::sync::{
7    atomic::{AtomicBool, Ordering},
8    Arc,
9};
10use std::time::Instant;
11use vyre_driver::persistent::PersistentThreadMode;
12use vyre_driver::speculate::SpeculationMode;
13use vyre_foundation::ir::Program;
14
15fn empty_batch_result_slots<T>(
16    len: usize,
17) -> Result<Vec<Option<Result<T, vyre_driver::BackendError>>>, vyre_driver::BackendError> {
18    let mut slots = Vec::new();
19    reserve_vec(
20        &mut slots,
21        len,
22        "WGPU backend",
23        "batch result slot",
24        "split the batch before dispatch",
25    )?;
26    slots.resize_with(len, || None);
27    Ok(slots)
28}
29
30fn finalize_batch_results<T>(
31    slots: Vec<Option<Result<T, vyre_driver::BackendError>>>,
32    missing_slot_message: &'static str,
33) -> Result<Vec<Result<T, vyre_driver::BackendError>>, vyre_driver::BackendError> {
34    let mut results = Vec::new();
35    reserve_vec(
36        &mut results,
37        slots.len(),
38        "WGPU backend",
39        "final batch result",
40        "split the batch before dispatch",
41    )?;
42    for slot in slots {
43        results.push(
44            slot.unwrap_or_else(|| Err(vyre_driver::BackendError::new(missing_slot_message))),
45        );
46    }
47    Ok(results)
48}
49
50fn elapsed_micros_u64(start: Instant, label: &str) -> Result<u64, vyre_driver::BackendError> {
51    u64::try_from(start.elapsed().as_micros()).map_err(|source| {
52        vyre_driver::BackendError::new(format!(
53            "{label} elapsed time cannot fit u64 microseconds: {source}. Fix: split or timeout the dispatch before telemetry overflows."
54        ))
55    })
56}
57
58impl WgpuBackend {
59    /// Adapter information selected for this backend instance.
60    #[must_use]
61    pub fn adapter_info(&self) -> &wgpu::AdapterInfo {
62        &self.adapter_info
63    }
64
65    /// Device limits for this backend instance.
66    #[must_use]
67    pub fn device_limits(&self) -> &wgpu::Limits {
68        &self.device_limits
69    }
70
71    /// Acquire the backend, probing adapters and returning a structured error
72    /// when no compatible GPU is found.
73    pub fn acquire() -> Result<Self, vyre_driver::BackendError> {
74        let ((device, queue), adapter_info, enabled_features) = crate::runtime::init_device()
75            .map_err(|error| {
76                let report = crate::runtime::device::adapter_probe_report();
77                vyre_driver::BackendError::new(format!(
78                    "no compatible GPU adapter found. Probed adapters: [{}].                  Missing features / limits: [{}]. Underlying error: {error}.                  Fix: install a compatible GPU driver and ensure a wgpu-supported backend                  (Vulkan, Metal, DX12) is available.",
79                    report.probed.join(", "),
80                    if report.missing.is_empty() {
81                        "none".to_string()
82                    } else {
83                        report.missing.join(", ")
84                    }
85                ))
86            })?;
87        let recovery_target = AdapterRecoveryTarget::Identity(
88            crate::runtime::device::AdapterIdentity::from_info(&adapter_info),
89        );
90        Self::from_device_queue(
91            device,
92            queue,
93            adapter_info,
94            enabled_features,
95            recovery_target,
96        )
97    }
98
99    /// Acquire a backend bound to a specific enumerable adapter index.
100    pub fn acquire_adapter(index: usize) -> Result<Self, vyre_driver::BackendError> {
101        let ((device, queue), adapter_info, enabled_features) =
102            crate::runtime::device::init_device_for_adapter(index)
103                .map_err(|error| vyre_driver::BackendError::new(error.to_string()))?;
104        Self::from_device_queue(
105            device,
106            queue,
107            adapter_info,
108            enabled_features,
109            AdapterRecoveryTarget::Index(index),
110        )
111    }
112
113    fn from_device_queue(
114        device: wgpu::Device,
115        queue: wgpu::Queue,
116        adapter_info: wgpu::AdapterInfo,
117        enabled_features: crate::runtime::device::EnabledFeatures,
118        recovery_target: AdapterRecoveryTarget,
119    ) -> Result<Self, vyre_driver::BackendError> {
120        let device_limits = device.limits();
121        let adapter_name = Arc::<str>::from(adapter_info.name.as_str());
122        let cache_tiers = vec![
123            crate::runtime::cache::CacheTier::try_new("hot", 1 << 24)?,
124            crate::runtime::cache::CacheTier::try_new("cold", 1 << 30)?,
125        ];
126        let persistent_pool = crate::buffer::BufferPool::with_tiering(
127            device.clone(),
128            queue.clone(),
129            &vyre_driver::DispatchConfig::default(),
130            cache_tiers,
131        )?;
132        let (pipeline_cache_entries, pipeline_cache_bytes) =
133            vyre_driver::pipeline::pipeline_cache_limits_from_env();
134        Ok(Self {
135            adapter_name,
136            adapter_info,
137            device_limits,
138            device_queue: Arc::new(arc_swap::ArcSwap::new(Arc::new((
139                device.clone(),
140                queue.clone(),
141            )))),
142            dispatch_arena: Arc::new(arc_swap::ArcSwap::from_pointee(DispatchArena::new(
143                device.clone(),
144                queue.clone(),
145                &vyre_driver::DispatchConfig::default(),
146            ))),
147            persistent_pool: Arc::new(arc_swap::ArcSwap::new(Arc::new(persistent_pool))),
148            pipeline_cache: Arc::new(
149                crate::runtime::cache::pipeline::LruPipelineCache::with_limits(
150                    pipeline_cache_entries,
151                    pipeline_cache_bytes,
152                ),
153            ),
154            wgsl_dispatch_pipeline_cache: Arc::new(dashmap::DashMap::with_hasher(
155                BuildHasherDefault::<rustc_hash::FxHasher>::default(),
156            )),
157            resident_pipeline_cache: Arc::new(dashmap::DashMap::with_hasher(BuildHasherDefault::<
158                rustc_hash::FxHasher,
159            >::default(
160            ))),
161            validation_cache: Arc::new(vyre_driver::validation::ValidationCache::default()),
162            shape_history: Arc::new(std::sync::Mutex::new(
163                vyre_driver::shape_prediction::ShapeHistory::new(),
164            )),
165            predicted_programs: Arc::new(dashmap::DashMap::with_hasher(BuildHasherDefault::<
166                rustc_hash::FxHasher,
167            >::default())),
168            bind_group_layout_cache: Arc::new(dashmap::DashMap::with_hasher(BuildHasherDefault::<
169                rustc_hash::FxHasher,
170            >::default(
171            ))),
172            resident_handles: Arc::new(dashmap::DashMap::with_hasher(BuildHasherDefault::<
173                rustc_hash::FxHasher,
174            >::default())),
175            device_lost: Arc::new(AtomicBool::new(false)),
176            enabled_features,
177            recovery_target,
178        })
179    }
180
181    pub(crate) fn current_device_queue(&self) -> Arc<(wgpu::Device, wgpu::Queue)> {
182        self.device_queue.load_full()
183    }
184
185    /// Consumer-visible snapshot of the live wgpu device + queue.
186    #[must_use]
187    pub fn device_queue(&self) -> Arc<(wgpu::Device, wgpu::Queue)> {
188        self.current_device_queue()
189    }
190
191    pub(crate) fn current_persistent_pool(&self) -> crate::buffer::BufferPool {
192        self.persistent_pool.load_full().as_ref().clone()
193    }
194
195    fn resident_pipeline_cache_key(
196        &self,
197        program: &Program,
198        config: &vyre_driver::DispatchConfig,
199    ) -> Result<(u64, u64, usize), vyre_driver::BackendError> {
200        let wire = program.to_wire().map_err(|source| {
201            vyre_driver::BackendError::new(format!(
202                "WGPU resident pipeline cache could not encode Program: {source}. Fix: validate the Program before resident dispatch."
203            ))
204        })?;
205        let mut program_hasher = rustc_hash::FxHasher::default();
206        program_hasher.write(&wire);
207        let mut config_hasher = rustc_hash::FxHasher::default();
208        config_hasher.write(format!("{config:?}").as_bytes());
209        Ok((program_hasher.finish(), config_hasher.finish(), wire.len()))
210    }
211
212    pub(crate) fn compile_resident_pipeline_cached(
213        &self,
214        program: &Program,
215        config: &vyre_driver::DispatchConfig,
216    ) -> Result<Arc<crate::pipeline::WgpuPipeline>, vyre_driver::BackendError> {
217        let key = self.resident_pipeline_cache_key(program, config)?;
218        if let Some(hit) = self.resident_pipeline_cache.get(&key) {
219            return Ok(hit.clone());
220        }
221        self.enforce_config_caps(config)?;
222        self.validate_with_cache(program)?;
223        let compiled = crate::pipeline::WgpuPipeline::compile_with_device_queue(
224            program,
225            config,
226            self.adapter_info.clone(),
227            self.enabled_features,
228            self.current_device_queue(),
229            self.dispatch_arena_snapshot(),
230            self.current_persistent_pool(),
231            self.pipeline_cache.clone(),
232            self.bind_group_layout_cache.clone(),
233        )?;
234        match self.resident_pipeline_cache.entry(key) {
235            dashmap::mapref::entry::Entry::Occupied(entry) => Ok(entry.get().clone()),
236            dashmap::mapref::entry::Entry::Vacant(entry) => {
237                entry.insert(compiled.clone());
238                Ok(compiled)
239            }
240        }
241    }
242
243    pub(crate) fn dispatch_arena_snapshot(&self) -> Arc<DispatchArena> {
244        self.dispatch_arena.load_full()
245    }
246
247    pub(crate) fn validate_with_cache(
248        &self,
249        program: &Program,
250    ) -> Result<(), vyre_driver::BackendError> {
251        self.validation_cache.get_or_validate_backend(program, self)
252    }
253
254    /// Test-only hook that marks the backend device as lost and invalidates
255    /// caches tied to the current device generation.
256    pub fn force_device_lost(&self) -> Result<(), vyre_driver::BackendError> {
257        self.device_lost.store(true, Ordering::Release);
258        self.pipeline_cache.clear();
259        self.wgsl_dispatch_pipeline_cache.clear();
260        self.bind_group_layout_cache.clear();
261        self.validation_cache.clear()?;
262        let device_queue = self.device_queue.load_full();
263        self.dispatch_arena.store(Arc::new(DispatchArena::new(
264            device_queue.0.clone(),
265            device_queue.1.clone(),
266            &vyre_driver::DispatchConfig::default(),
267        )));
268        Ok(())
269    }
270
271    /// Invalidate compiled pipeline artifacts selected by a rule-impact mask.
272    pub fn invalidate_impacted_pipeline_cache(
273        &self,
274        intervention_mask: &[u32],
275        rule_adj: &[u32],
276        state: &[u32],
277        join_rules: &[u32],
278        n: u32,
279        max_iterations: u32,
280        pipeline_lineage_cell: &[u32],
281        pipeline_keys: &[[u8; 32]],
282    ) -> Result<(), vyre_driver::BackendError> {
283        let final_impact_mask = vyre_driver::cache_invalidation::impacted_entries(
284            self,
285            intervention_mask,
286            rule_adj,
287            state,
288            join_rules,
289            n,
290            max_iterations,
291            pipeline_lineage_cell,
292        )
293        .map_err(|error| vyre_driver::BackendError::new(error.to_string()))?;
294        self.pipeline_cache
295            .invalidate_impacted(&final_impact_mask, pipeline_keys);
296        Ok(())
297    }
298
299    /// Convenience wrapper around [`Self::invalidate_impacted_pipeline_cache`]
300    pub fn invalidate_pipeline_cache_for_changed_op(
301        &self,
302        changed_op_handle: u32,
303        pipeline_lineage_cell: &[u32],
304        pipeline_keys: &[[u8; 32]],
305    ) -> Result<(), vyre_driver::BackendError> {
306        let n = 1u32;
307        let rule_adj = vec![1u32];
308        let intervention_mask = vec![1u32];
309        let state = vec![1u32];
310        let join_rules = vec![1u32];
311        let max_iterations = 1u32;
312        let mut normalized_lineage_cell = Vec::with_capacity(pipeline_lineage_cell.len());
313        normalized_lineage_cell.extend(pipeline_lineage_cell.iter().map(|&op| {
314            if op == changed_op_handle {
315                0
316            } else {
317                u32::MAX
318            }
319        }));
320        self.invalidate_impacted_pipeline_cache(
321            &intervention_mask,
322            &rule_adj,
323            &state,
324            &join_rules,
325            n,
326            max_iterations,
327            &normalized_lineage_cell,
328            pipeline_keys,
329        )
330    }
331
332    /// Invalidate disk-cached pipeline artifacts selected by a rule-impact mask.
333    pub fn invalidate_impacted_disk_cache(
334        &self,
335        intervention_mask: &[u32],
336        rule_adj: &[u32],
337        state: &[u32],
338        join_rules: &[u32],
339        n: u32,
340        max_iterations: u32,
341        pipeline_lineage_cell: &[u32],
342        cache_keys: &[String],
343    ) -> Result<(), vyre_driver::BackendError> {
344        crate::pipeline::disk_cache::invalidate_impacted(
345            self,
346            intervention_mask,
347            rule_adj,
348            state,
349            join_rules,
350            n,
351            max_iterations,
352            pipeline_lineage_cell,
353            cache_keys,
354        )
355        .map_err(|e| vyre_driver::BackendError::new(e.to_string()))
356    }
357
358    /// Create the backend if a GPU adapter is available.
359    #[must_use]
360    #[inline]
361    pub fn new() -> Result<Self, vyre_driver::BackendError> {
362        Self::acquire().map_err(|e| vyre_driver::BackendError::new(e.to_string()))
363    }
364
365    /// Process-wide shared backend handle.
366    pub fn shared() -> Result<Arc<Self>, vyre_driver::BackendError> {
367        static SHARED: std::sync::OnceLock<Result<Arc<WgpuBackend>, String>> =
368            std::sync::OnceLock::new();
369        match SHARED.get_or_init(|| Self::new().map(Arc::new).map_err(|e| e.to_string())) {
370            Ok(arc) => Ok(arc.clone()),
371            Err(msg) => Err(vyre_driver::BackendError::new(msg.clone())),
372        }
373    }
374
375    /// Dispatch borrowed inputs and visit each mapped output byte slice.
376    pub fn dispatch_borrowed_for_each_mapped_output<F>(
377        &self,
378        program: &Program,
379        inputs: &[&[u8]],
380        config: &vyre_driver::DispatchConfig,
381        visitor: F,
382    ) -> Result<(), vyre_driver::BackendError>
383    where
384        F: FnMut(usize, &[u8]) -> Result<(), vyre_driver::BackendError>,
385    {
386        let _span = tracing::trace_span!(
387            "vyre.dispatch_mapped_outputs",
388            backend = "wgpu",
389            inputs = inputs.len(),
390            label = tracing::field::Empty,
391        );
392        let _enter = _span.enter();
393        if let Some(label) = config.label.as_deref() {
394            _span.record("label", label);
395        }
396        let start = Instant::now();
397        self.dispatch_borrowed_async(program, inputs, config)?
398            .await_mapped_outputs(visitor)?;
399        tracing::trace!(
400            target: "vyre.dispatch",
401            elapsed_us = elapsed_micros_u64(start, "mapped-output dispatch")?,
402            inputs = inputs.len(),
403            "mapped-output dispatch completed"
404        );
405        Ok(())
406    }
407
408    /// Dispatch borrowed inputs and visit each mapped output as a typed POD slice.
409    pub fn dispatch_borrowed_for_each_pod_output<T, F>(
410        &self,
411        program: &Program,
412        inputs: &[&[u8]],
413        config: &vyre_driver::DispatchConfig,
414        mut visitor: F,
415    ) -> Result<(), vyre_driver::BackendError>
416    where
417        T: bytemuck::Pod,
418        F: FnMut(usize, &[T]) -> Result<(), vyre_driver::BackendError>,
419    {
420        self.dispatch_borrowed_for_each_mapped_output(program, inputs, config, |index, bytes| {
421            let typed = bytemuck::try_cast_slice::<u8, T>(bytes).map_err(|error| {
422                vyre_driver::BackendError::new(format!(
423                    "mapped output #{index} cannot be viewed as {}: {error}. Fix: set output_byte_range to a length and offset aligned for the requested POD type.",
424                    std::any::type_name::<T>()
425                ))
426            })?;
427            visitor(index, typed)
428        })
429    }
430
431    /// Enforce capability requirements declared in `config`.
432    pub(crate) fn enforce_config_caps(
433        &self,
434        config: &vyre_driver::DispatchConfig,
435    ) -> Result<(), vyre_driver::BackendError> {
436        if matches!(config.speculation, Some(SpeculationMode::Force))
437            && !<Self as vyre_driver::VyreBackend>::supports_speculation(self)
438        {
439            return Err(vyre_driver::BackendError::UnsupportedFeature {
440                name: "speculative dispatch".to_string(),
441                backend: <Self as vyre_driver::VyreBackend>::id(self).to_string(),
442            });
443        }
444        if matches!(config.persistent_thread, Some(PersistentThreadMode::Force))
445            && !<Self as vyre_driver::VyreBackend>::supports_persistent_thread_dispatch(self)
446        {
447            return Err(vyre_driver::BackendError::UnsupportedFeature {
448                name: "persistent-thread dispatch".to_string(),
449                backend: <Self as vyre_driver::VyreBackend>::id(self).to_string(),
450            });
451        }
452        Ok(())
453    }
454
455    /// Dispatch a real prefilter/confirm scan through the adaptive speculative path.
456    pub fn dispatch_speculative_prefilter_confirm<F>(
457        &self,
458        speculator: &vyre_driver::speculate::AdaptiveSpeculator,
459        plan: vyre_driver::speculate::SpeculativeDispatchPlan<'_>,
460        inputs: &[&[u8]],
461        config: &vyre_driver::DispatchConfig,
462        confirm_serial: F,
463    ) -> Result<vyre_driver::speculate::SpeculativeDispatchOutcome, vyre_driver::BackendError>
464    where
465        F: FnMut(
466            vyre_driver::OutputBuffers,
467        ) -> Result<vyre_driver::OutputBuffers, vyre_driver::BackendError>,
468    {
469        vyre_driver::speculate::dispatch_prefilter_confirm(
470            self,
471            speculator,
472            plan,
473            inputs,
474            config,
475            confirm_serial,
476        )
477    }
478
479    fn record_borrowed_batch_job(
480        &self,
481        program: &Program,
482        inputs: &[&[u8]],
483        config: &vyre_driver::DispatchConfig,
484        started: Instant,
485    ) -> Result<crate::engine::record_and_readback::RecordedDispatch, vyre_driver::BackendError>
486    {
487        self.enforce_config_caps(config)?;
488        self.validate_with_cache(program)?;
489        let pipeline = crate::pipeline::WgpuPipeline::compile_with_device_queue(
490            program,
491            config,
492            self.adapter_info.clone(),
493            self.enabled_features,
494            self.current_device_queue(),
495            self.dispatch_arena_snapshot(),
496            self.current_persistent_pool(),
497            self.pipeline_cache.clone(),
498            self.bind_group_layout_cache.clone(),
499        )?;
500        if let Some(deadline) = config.timeout {
501            let elapsed = started.elapsed();
502            if elapsed > deadline {
503                return Err(vyre_driver::BackendError::new(format!(
504                    "batch dispatch cancelled before GPU submission: took {elapsed:?}, budget {deadline:?}.                      Fix: raise DispatchConfig.timeout or split the program into smaller chunks."
505                )));
506            }
507        }
508        let workgroup_count = pipeline.workgroups_for_dispatch(config)?;
509        let dispatch_arena = self.dispatch_arena_snapshot();
510        crate::engine::record_and_readback::record_dispatch_unsubmitted(
511            crate::engine::record_and_readback::RecordAndReadback::for_dispatch(
512                &pipeline,
513                &dispatch_arena,
514                inputs,
515                workgroup_count,
516                config,
517                crate::async_dispatch::timestamp_profile_requested(config),
518                crate::engine::record_and_readback::DispatchLabels {
519                    bind_group: "vyre batch dispatch bind group",
520                    encoder: "vyre batch dispatch",
521                    compute: "vyre batch dispatch compute",
522                },
523            ),
524        )
525    }
526
527    /// Dispatch a batch of borrowed `(Program, inputs, config)` triples.
528    pub fn dispatch_borrowed_batch(
529        &self,
530        jobs: &[(&Program, &[&[u8]], &vyre_driver::DispatchConfig)],
531    ) -> Result<
532        Vec<Result<vyre_driver::OutputBuffers, vyre_driver::BackendError>>,
533        vyre_driver::BackendError,
534    > {
535        let _span = tracing::trace_span!(
536            "vyre.dispatch_borrowed_batch",
537            backend = "wgpu",
538            jobs = jobs.len(),
539        );
540        let _enter = _span.enter();
541
542        let mut results = empty_batch_result_slots(jobs.len())?;
543        let mut recorded = Vec::new();
544        reserve_backend_vec(&mut recorded, jobs.len(), "recorded dispatch")?;
545        let mut meta = Vec::new();
546        reserve_backend_vec(&mut meta, jobs.len(), "batch dispatch metadata")?;
547        for (index, (program, inputs, config)) in jobs.iter().enumerate() {
548            let started = Instant::now();
549            if program.is_explicit_noop() {
550                results[index] = Some(Ok(Vec::new()));
551                continue;
552            }
553            let command = self.record_borrowed_batch_job(program, inputs, config, started)?;
554            recorded.push(command);
555            meta.push((index, started, config.timeout));
556        }
557
558        let pending = crate::engine::record_and_readback::submit_recorded_batch(recorded)?;
559        for ((index, started, timeout), result) in meta
560            .into_iter()
561            .zip(crate::engine::record_and_readback::WgpuPendingReadback::await_many_owned(pending))
562        {
563            results[index] = Some(result.and_then(|outputs| {
564                if let Some(deadline) = timeout {
565                    let elapsed = started.elapsed();
566                    if elapsed > deadline {
567                        return Err(vyre_driver::BackendError::new(format!(
568                            "batch dispatch exceeded configured timeout: took {elapsed:?}, budget {deadline:?}.                              Fix: raise DispatchConfig.timeout or split the program into smaller chunks."
569                        )));
570                    }
571                }
572                Ok(outputs)
573            }));
574        }
575        finalize_batch_results(
576            results,
577            "internal batch dispatch result slot was not filled. Fix: keep batch recording metadata synchronized.",
578        )
579    }
580
581    /// Dispatch a borrowed batch and write each job's outputs into caller-owned per-job output buffers.
582    pub fn dispatch_borrowed_batch_into(
583        &self,
584        jobs: &[(&Program, &[&[u8]], &vyre_driver::DispatchConfig)],
585        outputs: &mut [vyre_driver::OutputBuffers],
586    ) -> Result<Vec<Result<(), vyre_driver::BackendError>>, vyre_driver::BackendError> {
587        if outputs.len() != jobs.len() {
588            return Err(vyre_driver::BackendError::new(format!(
589                "dispatch_borrowed_batch_into received {} output slots for {} jobs. Fix: pass exactly one OutputBuffers slot per job.",
590                outputs.len(),
591                jobs.len()
592            )));
593        }
594
595        let _span = tracing::trace_span!(
596            "vyre.dispatch_borrowed_batch_into",
597            backend = "wgpu",
598            jobs = jobs.len(),
599        );
600        let _enter = _span.enter();
601
602        let mut results = empty_batch_result_slots(jobs.len())?;
603        let mut recorded = Vec::new();
604        reserve_backend_vec(&mut recorded, jobs.len(), "recorded dispatch")?;
605        let mut meta = Vec::new();
606        reserve_backend_vec(&mut meta, jobs.len(), "batch-into dispatch metadata")?;
607        for (index, (program, inputs, config)) in jobs.iter().enumerate() {
608            let started = Instant::now();
609            if program.is_explicit_noop() {
610                outputs[index].clear();
611                results[index] = Some(Ok(()));
612                continue;
613            }
614            let command = self.record_borrowed_batch_job(program, inputs, config, started)?;
615            recorded.push(command);
616            meta.push((index, started, config.timeout));
617        }
618
619        let pending = crate::engine::record_and_readback::submit_recorded_batch(recorded)?;
620        let deadline =
621            crate::engine::record_and_readback::WgpuPendingReadback::wait_for_many(&pending);
622        for ((index, started, timeout), readback) in meta.into_iter().zip(pending) {
623            results[index] = Some(
624                readback
625                    .collect_after_submission_wait(&mut outputs[index], deadline)
626                    .and_then(|()| {
627                        if let Some(deadline) = timeout {
628                            let elapsed = started.elapsed();
629                            if elapsed > deadline {
630                                return Err(vyre_driver::BackendError::new(format!(
631                                    "batch dispatch exceeded configured timeout: took {elapsed:?}, budget {deadline:?}.                                      Fix: raise DispatchConfig.timeout or split the program into smaller chunks."
632                                )));
633                            }
634                        }
635                        Ok(())
636                    }),
637            );
638        }
639        finalize_batch_results(
640            results,
641            "internal batch-into dispatch result slot was not filled. Fix: keep batch recording metadata synchronized.",
642        )
643    }
644
645    /// Dispatch an owned batch of `(Program, inputs, config)` triples.
646    pub fn dispatch_batch(
647        &self,
648        jobs: &[(
649            vyre_foundation::ir::Program,
650            Vec<Vec<u8>>,
651            vyre_driver::DispatchConfig,
652        )],
653    ) -> Result<
654        Vec<Result<vyre_driver::OutputBuffers, vyre_driver::BackendError>>,
655        vyre_driver::BackendError,
656    > {
657        let mut borrowed_inputs = Vec::new();
658        reserve_backend_vec(&mut borrowed_inputs, jobs.len(), "borrowed input batch")?;
659        for (_, inputs, _) in jobs {
660            let mut borrowed = smallvec::SmallVec::<[&[u8]; 8]>::new();
661            reserve_smallvec(
662                &mut borrowed,
663                inputs.len(),
664                "WGPU backend",
665                "borrowed input slice reference",
666                "split the batch job before dispatch",
667            )?;
668            borrowed.extend(inputs.iter().map(Vec::as_slice));
669            borrowed_inputs.push(borrowed);
670        }
671        let mut borrowed_jobs = Vec::new();
672        reserve_backend_vec(&mut borrowed_jobs, jobs.len(), "borrowed dispatch job")?;
673        for ((program, _, config), inputs) in jobs.iter().zip(borrowed_inputs.iter()) {
674            borrowed_jobs.push((program, inputs.as_slice(), config));
675        }
676        self.dispatch_borrowed_batch(&borrowed_jobs)
677    }
678
679    /// Compile a program into a host-ingress wgpu stream.
680    #[allow(deprecated)]
681    pub fn compile_streaming(
682        &self,
683        program: &vyre_foundation::ir::Program,
684        config: vyre_driver::DispatchConfig,
685    ) -> Result<crate::engine::streaming::HostIngressStream, vyre_driver::BackendError> {
686        self.enforce_config_caps(&config)?;
687        let pipeline = crate::pipeline::WgpuPipeline::compile_with_device_queue(
688            program,
689            &config,
690            self.adapter_info.clone(),
691            self.enabled_features,
692            self.current_device_queue(),
693            self.dispatch_arena_snapshot(),
694            self.current_persistent_pool(),
695            self.pipeline_cache.clone(),
696            self.bind_group_layout_cache.clone(),
697        )?;
698        Ok(crate::engine::streaming::HostIngressStream::new(
699            (*pipeline).clone(),
700            config,
701        ))
702    }
703
704    /// Compile a program into a persistent pipeline.
705    pub fn compile_persistent(
706        &self,
707        program: &vyre_foundation::ir::Program,
708        config: &vyre_driver::DispatchConfig,
709    ) -> Result<Arc<crate::pipeline::WgpuPipeline>, vyre_driver::BackendError> {
710        self.enforce_config_caps(config)?;
711        crate::pipeline::WgpuPipeline::compile_with_device_queue(
712            program,
713            config,
714            self.adapter_info.clone(),
715            self.enabled_features,
716            self.current_device_queue(),
717            self.dispatch_arena_snapshot(),
718            self.current_persistent_pool(),
719            self.pipeline_cache.clone(),
720            self.bind_group_layout_cache.clone(),
721        )
722    }
723}
724
725/// Converts caller-owned input buffers into a [`smallvec::SmallVec`] of borrowed slices.
726///
727/// The wgpu backend's `dispatch_async` routes through this helper so staging reads from the
728/// caller's [`Vec`] allocations without cloning payload bytes - only slice references are
729/// collected into the vector. With more than eight inputs the [`SmallVec`] spills to heap
730/// storage while elements still alias the original buffers.
731#[allow(clippy::needless_lifetimes)]
732pub(crate) fn borrowed_slices_from_owned_inputs<'a>(
733    inputs: &'a [Vec<u8>],
734) -> smallvec::SmallVec<[&'a [u8]; 8]> {
735    let mut borrowed = smallvec::SmallVec::<[&'a [u8]; 8]>::with_capacity(inputs.len());
736    borrowed.extend(inputs.iter().map(Vec::as_slice));
737    borrowed
738}
739
740impl vyre_driver::VyreBackend for WgpuBackend {
741    fn id(&self) -> &'static str {
742        "wgpu"
743    }
744
745    fn version(&self) -> &'static str {
746        env!("CARGO_PKG_VERSION")
747    }
748
749    fn supported_ops(&self) -> &std::collections::HashSet<vyre_foundation::ir::OpId> {
750        vyre_driver::backend::validation::default_supported_ops_with_trap()
751    }
752
753    fn dispatch(
754        &self,
755        program: &Program,
756        inputs: &[Vec<u8>],
757        config: &vyre_driver::DispatchConfig,
758    ) -> Result<Vec<Vec<u8>>, vyre_driver::BackendError> {
759        let _span = tracing::trace_span!(
760            "vyre.dispatch",
761            backend = "wgpu",
762            inputs = inputs.len(),
763            label = tracing::field::Empty,
764        );
765        let _enter = _span.enter();
766        if let Some(label) = config.label.as_deref() {
767            _span.record("label", label);
768        }
769        let borrowed = borrowed_slices_from_owned_inputs(inputs);
770        let start = Instant::now();
771        let result = self
772            .dispatch_borrowed_async(program, &borrowed, config)?
773            .await_owned();
774        tracing::trace!(
775            target: "vyre.dispatch",
776            elapsed_us = elapsed_micros_u64(start, "borrowed-path dispatch")?,
777            inputs = inputs.len(),
778            "dispatch completed (borrowed-path; clone-free)"
779        );
780        result
781    }
782
783    fn dispatch_borrowed(
784        &self,
785        program: &Program,
786        inputs: &[&[u8]],
787        config: &vyre_driver::DispatchConfig,
788    ) -> Result<Vec<Vec<u8>>, vyre_driver::BackendError> {
789        let _span = tracing::trace_span!(
790            "vyre.dispatch",
791            backend = "wgpu",
792            inputs = inputs.len(),
793            label = tracing::field::Empty,
794        );
795        let _enter = _span.enter();
796        if let Some(label) = config.label.as_deref() {
797            _span.record("label", label);
798        }
799        let start = Instant::now();
800        let result = self
801            .dispatch_borrowed_async(program, inputs, config)?
802            .await_owned();
803        tracing::trace!(
804            target: "vyre.dispatch",
805            elapsed_us = elapsed_micros_u64(start, "dispatch")?,
806            inputs = inputs.len(),
807            "dispatch completed"
808        );
809        result
810    }
811
812    fn dispatch_borrowed_into(
813        &self,
814        program: &Program,
815        inputs: &[&[u8]],
816        config: &vyre_driver::DispatchConfig,
817        outputs: &mut vyre_driver::OutputBuffers,
818    ) -> Result<(), vyre_driver::BackendError> {
819        let _span = tracing::trace_span!(
820            "vyre.dispatch_into",
821            backend = "wgpu",
822            inputs = inputs.len(),
823            label = tracing::field::Empty,
824        );
825        let _enter = _span.enter();
826        if let Some(label) = config.label.as_deref() {
827            _span.record("label", label);
828        }
829        if vyre_driver::grid_sync::contains_grid_sync(program)
830            && !<Self as vyre_driver::VyreBackend>::supports_grid_sync(self)
831        {
832            return vyre_driver::grid_sync::dispatch_with_grid_sync_split_into(
833                self, program, inputs, config, outputs,
834            );
835        }
836        let start = Instant::now();
837        self.dispatch_borrowed_async(program, inputs, config)?
838            .await_into(outputs)?;
839        tracing::trace!(
840            target: "vyre.dispatch",
841            elapsed_us = elapsed_micros_u64(start, "dispatch into caller-owned outputs")?,
842            inputs = inputs.len(),
843            "dispatch completed into caller-owned outputs"
844        );
845        Ok(())
846    }
847
848    fn dispatch_borrowed_timed(
849        &self,
850        program: &Program,
851        inputs: &[&[u8]],
852        config: &vyre_driver::DispatchConfig,
853    ) -> Result<vyre_driver::TimedDispatchResult, vyre_driver::BackendError> {
854        let _span = tracing::trace_span!(
855            "vyre.dispatch_timed",
856            backend = "wgpu",
857            inputs = inputs.len(),
858            label = tracing::field::Empty,
859        );
860        let _enter = _span.enter();
861        if let Some(label) = config.label.as_deref() {
862            _span.record("label", label);
863        }
864        WgpuBackend::dispatch_borrowed_async_timed(self, program, inputs, config)?
865            .await_timed_owned()
866    }
867
868    fn dispatch_async(
869        &self,
870        program: &Program,
871        inputs: &[Vec<u8>],
872        config: &vyre_driver::DispatchConfig,
873    ) -> Result<Box<dyn vyre_driver::backend::PendingDispatch>, vyre_driver::BackendError> {
874        let _span = tracing::trace_span!(
875            "vyre.dispatch_async",
876            backend = "wgpu",
877            inputs = inputs.len(),
878            label = tracing::field::Empty,
879        );
880        let _enter = _span.enter();
881        if let Some(label) = config.label.as_deref() {
882            _span.record("label", label);
883        }
884
885        let borrowed = borrowed_slices_from_owned_inputs(inputs);
886        Ok(Box::new(WgpuBackend::dispatch_borrowed_async(
887            self, program, &borrowed, config,
888        )?))
889    }
890
891    fn dispatch_borrowed_async(
892        &self,
893        program: &Program,
894        inputs: &[&[u8]],
895        config: &vyre_driver::DispatchConfig,
896    ) -> Result<Box<dyn vyre_driver::backend::PendingDispatch>, vyre_driver::BackendError> {
897        Ok(Box::new(WgpuBackend::dispatch_borrowed_async(
898            self, program, inputs, config,
899        )?))
900    }
901
902    fn compile_native(
903        &self,
904        program: &Program,
905        config: &vyre_driver::DispatchConfig,
906    ) -> Result<Option<std::sync::Arc<dyn vyre_driver::CompiledPipeline>>, vyre_driver::BackendError>
907    {
908        self.enforce_config_caps(config)?;
909        self.validate_with_cache(program)?;
910        let cached = crate::pipeline::WgpuPipeline::compile_with_device_queue(
911            program,
912            config,
913            self.adapter_info.clone(),
914            self.enabled_features,
915            self.current_device_queue(),
916            self.dispatch_arena_snapshot(),
917            self.current_persistent_pool(),
918            self.pipeline_cache.clone(),
919            self.bind_group_layout_cache.clone(),
920        )?;
921        Ok(Some(cached))
922    }
923
924    fn allocate_device_buffer(
925        &self,
926        byte_len: usize,
927    ) -> Result<Box<dyn vyre_driver::DeviceBuffer>, vyre_driver::BackendError> {
928        self.allocate_wgpu_device_buffer(byte_len)
929    }
930
931    fn upload_device_buffer(
932        &self,
933        buffer: &mut dyn vyre_driver::DeviceBuffer,
934        bytes: &[u8],
935    ) -> Result<(), vyre_driver::BackendError> {
936        self.upload_wgpu_device_buffer(buffer, bytes)
937    }
938
939    fn download_device_buffer(
940        &self,
941        buffer: &dyn vyre_driver::DeviceBuffer,
942    ) -> Result<Vec<u8>, vyre_driver::BackendError> {
943        self.download_wgpu_device_buffer(buffer)
944    }
945
946    fn free_device_buffer(
947        &self,
948        buffer: Box<dyn vyre_driver::DeviceBuffer>,
949    ) -> Result<(), vyre_driver::BackendError> {
950        self.free_wgpu_device_buffer(buffer)
951    }
952
953    fn allocate_resident(
954        &self,
955        byte_len: usize,
956    ) -> Result<vyre_driver::Resource, vyre_driver::BackendError> {
957        crate::resident_resource::allocate_resident(self, byte_len)
958    }
959
960    fn upload_resident(
961        &self,
962        resource: &vyre_driver::Resource,
963        bytes: &[u8],
964    ) -> Result<(), vyre_driver::BackendError> {
965        crate::resident_upload::upload_resident(self, resource, bytes)
966    }
967
968    fn upload_resident_many(
969        &self,
970        uploads: &[(&vyre_driver::Resource, &[u8])],
971    ) -> Result<(), vyre_driver::BackendError> {
972        crate::resident_upload::upload_resident_many(self, uploads)
973    }
974
975    fn upload_resident_at(
976        &self,
977        resource: &vyre_driver::Resource,
978        dst_offset_bytes: usize,
979        bytes: &[u8],
980    ) -> Result<(), vyre_driver::BackendError> {
981        crate::resident_upload::upload_resident_at(self, resource, dst_offset_bytes, bytes)
982    }
983
984    fn upload_resident_at_many(
985        &self,
986        uploads: &[(&vyre_driver::Resource, usize, &[u8])],
987    ) -> Result<(), vyre_driver::BackendError> {
988        crate::resident_upload::upload_resident_at_many(self, uploads)
989    }
990
991    fn download_resident(
992        &self,
993        resource: &vyre_driver::Resource,
994    ) -> Result<Vec<u8>, vyre_driver::BackendError> {
995        crate::resident_download::download_resident(self, resource)
996    }
997
998    fn download_resident_into(
999        &self,
1000        resource: &vyre_driver::Resource,
1001        out: &mut Vec<u8>,
1002    ) -> Result<(), vyre_driver::BackendError> {
1003        crate::resident_download::download_resident_into(self, resource, out)
1004    }
1005
1006    fn download_resident_range(
1007        &self,
1008        resource: &vyre_driver::Resource,
1009        byte_offset: usize,
1010        byte_len: usize,
1011    ) -> Result<Vec<u8>, vyre_driver::BackendError> {
1012        crate::resident_download::download_resident_range(self, resource, byte_offset, byte_len)
1013    }
1014
1015    fn download_resident_range_into(
1016        &self,
1017        resource: &vyre_driver::Resource,
1018        byte_offset: usize,
1019        byte_len: usize,
1020        out: &mut Vec<u8>,
1021    ) -> Result<(), vyre_driver::BackendError> {
1022        crate::resident_download::download_resident_range_into(
1023            self,
1024            resource,
1025            byte_offset,
1026            byte_len,
1027            out,
1028        )
1029    }
1030
1031    fn download_resident_ranges_into(
1032        &self,
1033        ranges: &[(&vyre_driver::Resource, usize, usize)],
1034        outputs: &mut [&mut Vec<u8>],
1035    ) -> Result<(), vyre_driver::BackendError> {
1036        crate::resident_download::download_resident_ranges_into(self, ranges, outputs)
1037    }
1038
1039    fn free_resident(
1040        &self,
1041        resource: vyre_driver::Resource,
1042    ) -> Result<(), vyre_driver::BackendError> {
1043        crate::resident_resource::free_resident(self, resource)
1044    }
1045
1046    fn dispatch_resident_timed(
1047        &self,
1048        program: &Program,
1049        resources: &[vyre_driver::Resource],
1050        config: &vyre_driver::DispatchConfig,
1051    ) -> Result<vyre_driver::TimedDispatchResult, vyre_driver::BackendError> {
1052        crate::resident_dispatch::dispatch_resident_timed(self, program, resources, config)
1053    }
1054
1055    fn dispatch_with_device_buffers(
1056        &self,
1057        program: &Program,
1058        inputs: &[&dyn vyre_driver::DeviceBuffer],
1059        outputs: &mut [&mut dyn vyre_driver::DeviceBuffer],
1060        config: &vyre_driver::DispatchConfig,
1061    ) -> Result<(), vyre_driver::BackendError> {
1062        // Validate all buffers were allocated by us so the downcast
1063        // below cannot fail mid-loop after partial side-effects.
1064        vyre_driver::validate_buffer_ownership(self.id(), inputs.iter().copied())?;
1065        vyre_driver::validate_buffer_ownership(
1066            self.id(),
1067            outputs
1068                .iter()
1069                .map(|b| &**b as &dyn vyre_driver::DeviceBuffer),
1070        )?;
1071
1072        let resource_count = inputs.len().checked_add(outputs.len()).ok_or_else(|| {
1073            vyre_driver::BackendError::new(
1074                "resident dispatch resource count overflowed usize. Fix: split input/output resources before dispatch.",
1075            )
1076        })?;
1077        let mut resources =
1078            smallvec::SmallVec::<[vyre_driver::Resource; 8]>::with_capacity(resource_count);
1079        for buffer in inputs {
1080            let wgpu_buf = buffer
1081                .as_any()
1082                .downcast_ref::<crate::WgpuDeviceBuffer>()
1083                .ok_or_else(|| {
1084                    vyre_driver::BackendError::new(format!(
1085                        "Fix: dispatch_with_device_buffers expected WgpuDeviceBuffer inputs but got buffer owned by `{}`.",
1086                        buffer.backend_id()
1087                    ))
1088                })?;
1089            resources.push(vyre_driver::Resource::Resident(wgpu_buf.handle().id()));
1090        }
1091        for buffer in outputs.iter() {
1092            let backend_id = buffer.backend_id().to_string();
1093            let wgpu_buf = buffer
1094                .as_any()
1095                .downcast_ref::<crate::WgpuDeviceBuffer>()
1096                .ok_or_else(|| {
1097                    vyre_driver::BackendError::new(format!(
1098                        "Fix: dispatch_with_device_buffers expected WgpuDeviceBuffer outputs but got buffer owned by `{backend_id}`."
1099                    ))
1100                })?;
1101            resources.push(vyre_driver::Resource::Resident(wgpu_buf.handle().id()));
1102        }
1103
1104        let pipeline = self
1105            .compile_native(program, config)?
1106            .ok_or_else(|| {
1107                vyre_driver::BackendError::new(
1108                    "Fix: WgpuBackend::compile_native unexpectedly returned None for dispatch_with_device_buffers.",
1109                )
1110            })?;
1111        let _outputs = pipeline.dispatch_persistent_handles(&resources, config)?;
1112        Ok(())
1113    }
1114
1115    fn pipeline_cache_snapshot(&self) -> Option<vyre_driver::pipeline::PipelineCacheSnapshot> {
1116        Some(vyre_driver::pipeline::PipelineCacheSnapshot {
1117            hits: self.pipeline_cache.hits(),
1118            misses: self.pipeline_cache.misses(),
1119        })
1120    }
1121
1122    fn supports_subgroup_ops(&self) -> bool {
1123        crate::capabilities::supports_subgroup_ops(&self.enabled_features)
1124    }
1125
1126    fn supports_f16(&self) -> bool {
1127        false
1128    }
1129
1130    fn supports_bf16(&self) -> bool {
1131        false
1132    }
1133
1134    fn supports_tensor_cores(&self) -> bool {
1135        false
1136    }
1137
1138    fn supports_async_compute(&self) -> bool {
1139        false
1140    }
1141
1142    fn supports_indirect_dispatch(&self) -> bool {
1143        crate::capabilities::supports_indirect_dispatch(&self.adapter_info, &self.enabled_features)
1144    }
1145
1146    fn supports_speculation(&self) -> bool {
1147        false
1148    }
1149
1150    fn supports_persistent_thread_dispatch(&self) -> bool {
1151        false
1152    }
1153
1154    fn is_distributed(&self) -> bool {
1155        false
1156    }
1157
1158    fn max_workgroup_size(&self) -> [u32; 3] {
1159        self.enabled_features.max_workgroup_size
1160    }
1161
1162    fn max_compute_workgroups_per_dimension(&self) -> u32 {
1163        self.device_limits.max_compute_workgroups_per_dimension
1164    }
1165
1166    fn max_compute_invocations_per_workgroup(&self) -> u32 {
1167        self.device_limits.max_compute_invocations_per_workgroup
1168    }
1169
1170    fn subgroup_size(&self) -> Option<u32> {
1171        crate::capabilities::supports_subgroup_ops(&self.enabled_features)
1172            .then_some(self.enabled_features.min_subgroup_size)
1173    }
1174
1175    fn max_storage_buffer_bytes(&self) -> u64 {
1176        self.enabled_features.max_storage_buffer_binding_size
1177    }
1178
1179    fn device_profile(&self) -> vyre_driver::DeviceProfile {
1180        WgpuBackend::device_profile(self)
1181    }
1182
1183    fn flush(&self) -> Result<(), vyre_driver::BackendError> {
1184        let device_queue = self.current_device_queue();
1185        let submission = device_queue.1.submit(std::iter::empty());
1186        crate::runtime::device::poll_device_wait_for(&device_queue.0, submission)?;
1187        crate::pipeline::disk_cache::flush_disk_pipeline_cache()
1188    }
1189
1190    fn device_lost(&self) -> bool {
1191        self.device_lost.load(Ordering::Acquire)
1192    }
1193
1194    fn try_recover(&self) -> Result<(), vyre_driver::BackendError> {
1195        let ((device, queue), adapter_info, enabled) = match &self.recovery_target {
1196            AdapterRecoveryTarget::Index(index) => {
1197                crate::runtime::device::init_device_for_adapter(*index)
1198            }
1199            AdapterRecoveryTarget::Identity(identity) => {
1200                crate::runtime::device::init_device_for_adapter_identity(identity)
1201            }
1202        }
1203        .map_err(|error| vyre_driver::BackendError::new(error.to_string()))?;
1204        let device_limits = device.limits();
1205        let recovered_identity = crate::runtime::device::AdapterIdentity::from_info(&adapter_info);
1206        let original_identity =
1207            crate::runtime::device::AdapterIdentity::from_info(&self.adapter_info);
1208        if recovered_identity != original_identity {
1209            return Err(vyre_driver::BackendError::new(format!(
1210                "wgpu recovery selected a different adapter than the backend was constructed with. Original: {:?}; recovered: {:?}. Fix: construct a new backend for the new adapter instead of reusing device-local caches across adapter identities.",
1211                self.adapter_info, adapter_info
1212            )));
1213        }
1214        if device_limits != self.device_limits || enabled != self.enabled_features {
1215            return Err(vyre_driver::BackendError::new(
1216                "wgpu recovery selected the original adapter but feature or limit negotiation changed. Fix: construct a new backend so dispatch planning and pipeline caches are rebuilt against the new device contract.",
1217            ));
1218        }
1219        let cache_tiers = vec![
1220            crate::runtime::cache::CacheTier::try_new("hot", 1 << 24)?,
1221            crate::runtime::cache::CacheTier::try_new("cold", 1 << 30)?,
1222        ];
1223        let persistent_pool = crate::buffer::BufferPool::with_tiering(
1224            device.clone(),
1225            queue.clone(),
1226            &vyre_driver::DispatchConfig::default(),
1227            cache_tiers,
1228        )?;
1229        self.device_queue
1230            .store(Arc::new((device.clone(), queue.clone())));
1231        self.persistent_pool.store(Arc::new(persistent_pool));
1232        self.pipeline_cache.clear();
1233        self.wgsl_dispatch_pipeline_cache.clear();
1234        self.bind_group_layout_cache.clear();
1235        self.validation_cache.clear()?;
1236        self.dispatch_arena.store(Arc::new(DispatchArena::new(
1237            device.clone(),
1238            queue.clone(),
1239            &vyre_driver::DispatchConfig::default(),
1240        )));
1241        self.device_lost.store(false, Ordering::Release);
1242
1243        Ok(())
1244    }
1245}
1246
1247impl vyre_self_substrate::optimizer::dispatcher::OptimizerDispatcher for WgpuBackend {
1248    fn dispatch(
1249        &self,
1250        program: &Program,
1251        inputs: &[Vec<u8>],
1252        grid_override: Option<[u32; 3]>,
1253    ) -> Result<Vec<Vec<u8>>, vyre_self_substrate::optimizer::dispatcher::DispatchError> {
1254        let mut config = vyre_driver::DispatchConfig::default();
1255        config.grid_override = grid_override;
1256        vyre_driver::VyreBackend::dispatch(self, program, inputs, &config).map_err(|error| {
1257            vyre_self_substrate::optimizer::dispatcher::DispatchError::BackendError(
1258                error.to_string(),
1259            )
1260        })
1261    }
1262}
1263
1264#[cfg(test)]
1265mod borrowed_slice_conversion_tests {
1266    use super::{
1267        borrowed_slices_from_owned_inputs, empty_batch_result_slots, finalize_batch_results,
1268    };
1269
1270    #[test]
1271    fn dispatch_async_input_conversion_is_zero_copy_slice_refs() {
1272        let inputs = vec![vec![1u8, 2, 3], vec![4u8, 5]];
1273        let borrowed = borrowed_slices_from_owned_inputs(&inputs);
1274        assert_eq!(borrowed.len(), 2);
1275        assert_eq!(borrowed[0].as_ptr(), inputs[0].as_ptr());
1276        assert_eq!(borrowed[1].as_ptr(), inputs[1].as_ptr());
1277    }
1278
1279    #[test]
1280    fn nine_inputs_spill_smallvec_but_slices_alias_vecs() {
1281        let inputs: Vec<Vec<u8>> = (0..9).map(|i| vec![i as u8]).collect();
1282        let borrowed = borrowed_slices_from_owned_inputs(&inputs);
1283        assert_eq!(borrowed.len(), 9);
1284        for i in 0..9 {
1285            assert_eq!(
1286                borrowed[i].as_ptr(),
1287                inputs[i].as_ptr(),
1288                "slice {i} must reference the corresponding Vec buffer"
1289            );
1290        }
1291    }
1292
1293    #[test]
1294    fn generated_batch_result_finalization_preserves_success_error_and_missing_slots() {
1295        for case in 0..4096usize {
1296            let len = (case % 19) + 1;
1297            let mut slots = empty_batch_result_slots::<usize>(len)
1298                .expect("Fix: generated WGPU batch result test must reserve slots");
1299            for slot in 0..len {
1300                match (slot + case) % 7 {
1301                    0 => {}
1302                    1 => {
1303                        slots[slot] = Some(Err(vyre_driver::BackendError::new(format!(
1304                            "generated-error-{case}-{slot}"
1305                        ))));
1306                    }
1307                    _ => {
1308                        slots[slot] = Some(Ok(case * 100 + slot));
1309                    }
1310                }
1311            }
1312
1313            let finalized =
1314                finalize_batch_results(slots, "generated missing WGPU batch result slot")
1315                    .expect("Fix: generated WGPU batch finalization must reserve output results");
1316            assert_eq!(
1317                finalized.len(),
1318                len,
1319                "generated WGPU batch case {case} must preserve slot count"
1320            );
1321            for (slot, result) in finalized.into_iter().enumerate() {
1322                match (slot + case) % 7 {
1323                    0 => {
1324                        let error =
1325                            result.expect_err("Fix: missing generated batch slot must error");
1326                        assert!(
1327                            error
1328                                .to_string()
1329                                .contains("generated missing WGPU batch result slot"),
1330                            "Fix: missing generated batch slot must report the supplied invariant, got {error}"
1331                        );
1332                    }
1333                    1 => {
1334                        let error = result
1335                            .expect_err("Fix: explicit generated batch error must stay error");
1336                        assert!(
1337                            error
1338                                .to_string()
1339                                .contains(&format!("generated-error-{case}-{slot}")),
1340                            "Fix: explicit generated batch error must be preserved, got {error}"
1341                        );
1342                    }
1343                    _ => {
1344                        assert_eq!(
1345                            result.expect("Fix: generated batch success must stay success"),
1346                            case * 100 + slot
1347                        );
1348                    }
1349                }
1350            }
1351        }
1352    }
1353}