Skip to main content

virtio_accel_vulkan/
native.rs

1//! Native Vulkan backend over `ash`: the audited `Accelerator` implementation.
2//!
3//! `SAFETY.md` is the audit of record; every `unsafe` block below carries a local `SAFETY:` note.
4//! Each Vulkan handle has exactly one Rust owner with a `Drop` implementation, every `VkResult`
5//! is checked before an out-value is trusted, and `VK_ERROR_DEVICE_LOST` poisons the whole backend
6//! instance (ADR 0006). Completion is a nonblocking `vkGetFenceStatus` read: no worker thread, no
7//! callback, no foreign code ever owns Rust memory.
8//!
9//! Handles are deliberately neither `Send` nor `Sync` (`Rc` inside): Vulkan queues and command
10//! pools are externally synchronized objects, and the contract permits thread-affine providers.
11
12use std::cell::{Cell, RefCell};
13use std::collections::HashMap;
14use std::ffi::CStr;
15use std::ptr::NonNull;
16use std::rc::Rc;
17use std::sync::{Arc, Mutex, OnceLock};
18
19use ash::vk;
20use virtio_accel_core::{
21    Accelerator, AcceleratorClass, AccessMode, AllocatedBuffer, ArtifactRef, BackendError,
22    BindingRef, BufferDesc, BufferInfo, BufferProperties, BufferUsage, ByteSink, ByteSource,
23    Capabilities, ContextDesc, DeviceIdentity, DeviceInfo, DeviceLimits, EventState, MemoryDomain,
24    QueueDesc, ReleaseFailure, SubmitFailure, Timeout,
25};
26use virtio_accel_tosa::{CapabilityDescriptor, TosaCapabilityProvider};
27
28use crate::lower::{KernelSpec, LoweringError, ProgramPlan, SlotRole, Work, lower_tosa};
29use crate::nvfp4::lower_nvfp4;
30use crate::shader::{self, KernelKey};
31use crate::{InitError, REQUIRED_RESIDENT_BYTES};
32
33/// Maximal TOSA artifact bytes admitted before parsing (mirrors the other TOSA backends).
34const MAX_TOSA_ARTIFACT_BYTES: u64 = 256 * 1024 * 1024;
35
36/// Stable provider-owned error namespace for unmapped `VkResult` codes (`"VULK"`).
37const VULKAN_EXTERNAL_DOMAIN: u32 = 0x5655_4c4b;
38
39/// Bounded staging allocation for explicit transfers into and out of `MemoryDomain::Device`.
40const STAGING_BYTES: u64 = 4 * 1024 * 1024;
41
42/// How long a synchronous explicit transfer may take before the device is treated as lost.
43const TRANSFER_TIMEOUT_NS: u64 = 30_000_000_000;
44
45/// `maxMemoryAllocationCount` is assumed at the spec minimum (ADR 0005): every buffer here is one
46/// dedicated `VkDeviceMemory`, so the advertised aggregate buffer count plus one transient staging
47/// allocation must stay inside it.
48const ASSUMED_MAX_MEMORY_ALLOCATIONS: u32 = 4096;
49const MAX_CONTEXTS: u32 = 16;
50const MAX_BUFFERS_PER_CONTEXT: u32 = 190;
51/// Every program may own one arena allocation for its constants and intermediates.
52const MAX_PROGRAMS_PER_CONTEXT: u32 = 64;
53const MAX_QUEUES_PER_CONTEXT: u32 = 16;
54/// Ring depth per context: one (command buffer, fence, descriptor set) triple per outstanding
55/// event (ADR 0006).
56const RING_DEPTH: u32 = 64;
57const MAX_BINDINGS_PER_SUBMISSION: u32 = 16;
58
59/// Explicit transfers and constant uploads hold at most one transient staging allocation at a
60/// time, on top of the guest buffers and program arenas.
61const TRANSIENT_STAGING_ALLOCATIONS: u32 = 1;
62
63const _: () = assert!(
64    MAX_CONTEXTS * (MAX_BUFFERS_PER_CONTEXT + MAX_PROGRAMS_PER_CONTEXT)
65        + TRANSIENT_STAGING_ALLOCATIONS
66        <= ASSUMED_MAX_MEMORY_ALLOCATIONS,
67    "advertised buffers, program arenas, and the staging allocation must fit the assumed allocation count"
68);
69
70/// Preferred 1-D workgroup size of the grid-stride kernels, and the size used on a device whose
71/// `maxComputeWorkGroupInvocations` is only the specification minimum (128).
72const PREFERRED_WORKGROUP: u32 = 256;
73const FALLBACK_WORKGROUP: u32 = 128;
74/// Preferred register-tiled MATMUL tile (256 invocations, a 64 × 64 block over 8 KiB of shared
75/// slabs) and the fallback tile (64 invocations, 32 × 32). The streaming kernel for eight rows or
76/// fewer is a fixed 64-invocation 1-D workgroup with an 8 KiB reduction buffer.
77const PREFERRED_MATMUL_TILE: u32 = 16;
78const FALLBACK_MATMUL_TILE: u32 = 8;
79/// Bytes every `VkBuffer` size is rounded up to so byte-storage tensors can be addressed by
80/// whole words at their tail; the logical buffer size the guest sees is unchanged.
81const WORD_BYTES: u64 = 4;
82
83/// Kernel parameters fixed per device from its limits (ADR 0007).
84#[derive(Clone, Copy, Debug, PartialEq, Eq)]
85struct Tuning {
86    /// 1-D workgroup size of the elementwise, reduction, pooling, and copy kernels.
87    workgroup: u32,
88    /// Side of the square MATMUL tile.
89    matmul_tile: u32,
90    /// Length of the storage-buffer descriptor array: bound slots plus the program arena.
91    buffers: u32,
92    /// The device advertises the exact subgroup cooperative-matrix shape used by the NVFP4
93    /// projection kernel: FP16 8x16x16 with FP32 accumulation.
94    cooperative_nvfp4: bool,
95    /// Fixed 32-lane subgroups with arithmetic reductions remove shared-memory synchronization
96    /// from the scalar NVFP4 decode kernel.
97    subgroup_nvfp4: bool,
98}
99
100impl Tuning {
101    /// Derive the tuning, or `None` when the device cannot host even the smallest kernels.
102    fn from_limits(limits: &vk::PhysicalDeviceLimits) -> Option<Self> {
103        let invocations = limits.max_compute_work_group_invocations;
104        let size = limits.max_compute_work_group_size;
105        let workgroup = if invocations >= PREFERRED_WORKGROUP && size[0] >= PREFERRED_WORKGROUP {
106            PREFERRED_WORKGROUP
107        } else if invocations >= FALLBACK_WORKGROUP && size[0] >= FALLBACK_WORKGROUP {
108            FALLBACK_WORKGROUP
109        } else {
110            return None;
111        };
112        // Both MATMUL kernels must fit: the square one is `tile × tile` invocations, the
113        // streaming one a 1-D `STREAM_WORKGROUP`, and each declares its shared slabs.
114        let tile_fits = |tile: u32| {
115            invocations >= tile * tile
116                && size[0] >= tile.max(shader::STREAM_WORKGROUP)
117                && size[1] >= tile
118                && limits.max_compute_shared_memory_size >= shader::matmul_shared_bytes(tile)
119        };
120        let matmul_tile = if tile_fits(PREFERRED_MATMUL_TILE) {
121            PREFERRED_MATMUL_TILE
122        } else if tile_fits(FALLBACK_MATMUL_TILE) {
123            FALLBACK_MATMUL_TILE
124        } else {
125            return None;
126        };
127        // Every element of the descriptor array counts against both per-stage and per-set
128        // storage-buffer limits; at least one input, one output, and the arena must fit.
129        let descriptors = limits
130            .max_per_stage_descriptor_storage_buffers
131            .min(limits.max_descriptor_set_storage_buffers)
132            .min(MAX_BINDINGS_PER_SUBMISSION + 1);
133        if descriptors < 3 {
134            return None;
135        }
136        Some(Self {
137            workgroup,
138            matmul_tile,
139            buffers: descriptors,
140            cooperative_nvfp4: false,
141            subgroup_nvfp4: false,
142        })
143    }
144
145    /// Bindings a submission may carry: every descriptor but the arena's.
146    const fn max_bindings(self) -> u32 {
147        self.buffers - 1
148    }
149
150    fn key(self, kernel: KernelSpec) -> KernelKey {
151        match kernel {
152            KernelSpec::Nvfp4Matmul { cooperative } => KernelKey::Nvfp4Matmul {
153                buffers: self.buffers,
154                cooperative: cooperative && self.cooperative_nvfp4,
155                subgroup: !(cooperative && self.cooperative_nvfp4) && self.subgroup_nvfp4,
156            },
157            KernelSpec::Elementwise {
158                op,
159                float,
160                broadcast,
161            } => KernelKey::Elementwise {
162                op,
163                float,
164                broadcast,
165                workgroup: self.workgroup,
166                buffers: self.buffers,
167            },
168            KernelSpec::Reduce { op, float } => KernelKey::Reduce {
169                op,
170                float,
171                workgroup: self.workgroup,
172                buffers: self.buffers,
173            },
174            KernelSpec::Matmul { input, output } => KernelKey::Matmul {
175                input,
176                output,
177                tile: self.matmul_tile,
178                buffers: self.buffers,
179            },
180            KernelSpec::MatmulStream { rhs, output } => KernelKey::MatmulStream {
181                rhs,
182                output,
183                buffers: self.buffers,
184            },
185            KernelSpec::MaxPool { nan_mode, float } => KernelKey::MaxPool {
186                nan_mode,
187                float,
188                workgroup: self.workgroup,
189                buffers: self.buffers,
190            },
191            KernelSpec::Cast { input, output } => KernelKey::Cast {
192                input,
193                output,
194                workgroup: self.workgroup,
195                buffers: self.buffers,
196            },
197            KernelSpec::Move {
198                storage,
199                contiguous,
200            } => KernelKey::Move {
201                storage,
202                contiguous,
203                workgroup: self.workgroup,
204                buffers: self.buffers,
205            },
206        }
207    }
208
209    /// Workgroup counts for `work`, or `None` when they exceed the device's dispatch limits.
210    fn workgroups(
211        self,
212        work: Work,
213        kernel: KernelSpec,
214        limits: &vk::PhysicalDeviceLimits,
215    ) -> Option<[u32; 3]> {
216        let max = limits.max_compute_work_group_count;
217        match work {
218            Work::Linear(count) => Some([
219                shader::linear_workgroups(count, self.workgroup, max[0].max(1)),
220                1,
221                1,
222            ]),
223            Work::Matmul { m, n, batch } => {
224                let groups = shader::matmul_workgroups(m, n, batch, self.matmul_tile);
225                (groups[0] <= max[0] && groups[1] <= max[1] && groups[2] <= max[2])
226                    .then_some(groups)
227            }
228            Work::MatmulStream { n, batch } => {
229                let groups = shader::stream_matmul_workgroups(n, batch);
230                (groups[0] <= max[0] && groups[2] <= max[2]).then_some(groups)
231            }
232            Work::Nvfp4Matmul { m, n } => {
233                let cooperative = matches!(kernel, KernelSpec::Nvfp4Matmul { cooperative: true })
234                    && self.cooperative_nvfp4;
235                let subgroup = !cooperative && self.subgroup_nvfp4;
236                let columns = if cooperative {
237                    n.div_ceil(16)
238                } else if subgroup {
239                    n.div_ceil(4)
240                } else {
241                    n
242                };
243                let rows = if cooperative { m.div_ceil(8) } else { m };
244                let groups = [columns, rows, 1];
245                (groups[0] <= max[0] && groups[1] <= max[1]).then_some(groups)
246            }
247        }
248    }
249}
250
251const EXCLUSIVE_ACCESS: u64 = 1 << 63;
252
253/// The process-wide loader handle: one `dlopen` of the platform Vulkan loader.
254fn entry() -> Result<ash::Entry, InitError> {
255    static ENTRY: OnceLock<Result<ash::Entry, InitError>> = OnceLock::new();
256    ENTRY
257        .get_or_init(|| {
258            // SAFETY: loading the platform Vulkan loader runs its initializers exactly once per
259            // process under this `OnceLock`; nothing else in this crate loads it.
260            unsafe { ash::Entry::load() }.map_err(|_| InitError::RuntimeUnavailable)
261        })
262        .clone()
263}
264
265fn backend_error(result: vk::Result) -> BackendError {
266    match result {
267        vk::Result::ERROR_OUT_OF_HOST_MEMORY | vk::Result::ERROR_OUT_OF_DEVICE_MEMORY => {
268            BackendError::OutOfMemory
269        }
270        vk::Result::ERROR_TOO_MANY_OBJECTS => BackendError::ResourceLimit,
271        vk::Result::ERROR_DEVICE_LOST => BackendError::DeviceLost,
272        vk::Result::ERROR_MEMORY_MAP_FAILED => BackendError::Incompatible,
273        other => BackendError::External {
274            domain: VULKAN_EXTERNAL_DOMAIN,
275            code: i64::from(other.as_raw()),
276        },
277    }
278}
279
280/// Owned `VkInstance`; destroyed after every device that was created from it.
281struct Instance {
282    /// Kept so the loaded library outlives the instance created from it.
283    _entry: ash::Entry,
284    instance: ash::Instance,
285}
286
287impl Instance {
288    fn create() -> Result<Self, InitError> {
289        let entry = entry()?;
290        // SAFETY: querying the loader's instance version has no preconditions.
291        let loader_version = unsafe { entry.try_enumerate_instance_version() }
292            .ok()
293            .flatten()
294            .unwrap_or(vk::API_VERSION_1_0);
295        if loader_version < vk::API_VERSION_1_3 {
296            return Err(InitError::RuntimeUnavailable);
297        }
298        let application = vk::ApplicationInfo::default()
299            .application_name(c"virtio-accel-vulkan")
300            .api_version(vk::API_VERSION_1_3);
301        // MoltenVK and other portability drivers refuse enumeration unless the instance opts
302        // into the portability extension; requesting it is a no-op on conformant native drivers.
303        let extension_names = [c"VK_KHR_portability_enumeration".as_ptr()];
304        let info = vk::InstanceCreateInfo::default()
305            .application_info(&application)
306            .flags(vk::InstanceCreateFlags::ENUMERATE_PORTABILITY_KHR)
307            .enabled_extension_names(&extension_names);
308        // SAFETY: `info` and the structures it points to outlive the call; no layers are
309        // requested, and the extension name is a static literal whose pointer outlives the call.
310        let instance = match unsafe { entry.create_instance(&info, None) } {
311            Ok(instance) => instance,
312            Err(vk::Result::ERROR_INCOMPATIBLE_DRIVER) => {
313                return Err(InitError::RuntimeUnavailable);
314            }
315            Err(_) => return Err(InitError::InstanceCreationFailed),
316        };
317        Ok(Self {
318            _entry: entry,
319            instance,
320        })
321    }
322}
323
324impl Drop for Instance {
325    fn drop(&mut self) {
326        // SAFETY: this owner is dropped exactly once, after every `Shared` (and thus every
327        // device) created from it: `Shared` holds the `Instance` and destroys its device first.
328        unsafe { self.instance.destroy_instance(None) };
329    }
330}
331
332/// Everything probed about one physical device before it is opened.
333#[derive(Clone)]
334struct PhysicalDeviceRecord {
335    handle: vk::PhysicalDevice,
336    name: String,
337    device_type: vk::PhysicalDeviceType,
338    vendor_id: u32,
339    device_id: u32,
340    uuid: [u8; 16],
341    queue_family: u32,
342    limits: vk::PhysicalDeviceLimits,
343    memory: vk::PhysicalDeviceMemoryProperties,
344    buffer_device_address: bool,
345    /// `minImportedHostPointerAlignment` when the device offers `VK_EXT_external_memory_host`,
346    /// the one extension this backend enables, and only for
347    /// [`VulkanAccelerator::import_host_buffer`] (ADR 0013).
348    host_import_alignment: Option<u64>,
349    /// Vulkan 1.2 `timelineSemaphore`, enabled when reported, for host gates (ADR 0013).
350    timeline_semaphore: bool,
351    shader_float16: bool,
352    vulkan_memory_model: bool,
353    cooperative_nvfp4: bool,
354    tuning: Tuning,
355}
356
357impl PhysicalDeviceRecord {
358    /// Probe one device; `None` when it cannot host this backend (API floor, compute queue,
359    /// mandatory `synchronization2`).
360    fn probe(owner: &Instance, handle: vk::PhysicalDevice) -> Option<Self> {
361        let instance = &owner.instance;
362        let mut vulkan11 = vk::PhysicalDeviceVulkan11Properties::default();
363        let mut subgroup = vk::PhysicalDeviceSubgroupProperties::default();
364        let mut properties = vk::PhysicalDeviceProperties2::default()
365            .push_next(&mut vulkan11)
366            .push_next(&mut subgroup);
367        // SAFETY: `handle` was enumerated from `instance`; the chained structures are live locals.
368        unsafe { instance.get_physical_device_properties2(handle, &mut properties) };
369        let properties = properties.properties;
370        if properties.api_version < vk::API_VERSION_1_3 {
371            return None;
372        }
373        let mut tuning = Tuning::from_limits(&properties.limits)?;
374
375        let mut vulkan12 = vk::PhysicalDeviceVulkan12Features::default();
376        let mut vulkan13 = vk::PhysicalDeviceVulkan13Features::default();
377        let mut cooperative = vk::PhysicalDeviceCooperativeMatrixFeaturesKHR::default();
378        let mut features = vk::PhysicalDeviceFeatures2::default()
379            .push_next(&mut vulkan12)
380            .push_next(&mut vulkan13)
381            .push_next(&mut cooperative);
382        // SAFETY: as above; the feature chain is fully initialized before the call.
383        unsafe { instance.get_physical_device_features2(handle, &mut features) };
384        if vulkan13.synchronization2 == vk::FALSE {
385            return None;
386        }
387
388        // SAFETY: `handle` is a live physical device of `instance`.
389        let families = unsafe { instance.get_physical_device_queue_family_properties(handle) };
390        // A compute-only family keeps this backend's work off the graphics queue when the device
391        // offers one; otherwise the first compute-capable family serves.
392        let compute_only = families.iter().position(|family| {
393            family.queue_flags.contains(vk::QueueFlags::COMPUTE)
394                && !family.queue_flags.contains(vk::QueueFlags::GRAPHICS)
395        });
396        let any_compute = families
397            .iter()
398            .position(|family| family.queue_flags.contains(vk::QueueFlags::COMPUTE));
399        let queue_family = u32::try_from(compute_only.or(any_compute)?).ok()?;
400
401        // SAFETY: `handle` is a live physical device of `instance`.
402        let memory = unsafe { instance.get_physical_device_memory_properties(handle) };
403        let host_import_alignment = host_import_alignment(instance, handle);
404        let cooperative_nvfp4 = vulkan12.shader_float16 == vk::TRUE
405            && vulkan12.vulkan_memory_model == vk::TRUE
406            && cooperative.cooperative_matrix == vk::TRUE
407            && subgroup.subgroup_size == 32
408            && subgroup
409                .supported_stages
410                .contains(vk::ShaderStageFlags::COMPUTE)
411            && supports_extension(instance, handle, ash::khr::cooperative_matrix::NAME)
412            && supports_nvfp4_cooperative_matrix(owner, handle);
413        tuning.cooperative_nvfp4 = cooperative_nvfp4;
414        tuning.subgroup_nvfp4 = subgroup.subgroup_size == 32
415            && subgroup
416                .supported_stages
417                .contains(vk::ShaderStageFlags::COMPUTE)
418            && subgroup
419                .supported_operations
420                .contains(vk::SubgroupFeatureFlags::ARITHMETIC);
421        let name = CStr::from_bytes_until_nul(bytemuck_i8_to_u8(&properties.device_name))
422            .map(|name| name.to_string_lossy().into_owned())
423            .unwrap_or_else(|_| String::from("vulkan-device"));
424        Some(Self {
425            handle,
426            name,
427            device_type: properties.device_type,
428            vendor_id: properties.vendor_id,
429            device_id: properties.device_id,
430            uuid: vulkan11.device_uuid,
431            queue_family,
432            limits: properties.limits,
433            memory,
434            buffer_device_address: vulkan12.buffer_device_address == vk::TRUE,
435            host_import_alignment,
436            timeline_semaphore: vulkan12.timeline_semaphore == vk::TRUE,
437            shader_float16: vulkan12.shader_float16 == vk::TRUE,
438            vulkan_memory_model: vulkan12.vulkan_memory_model == vk::TRUE,
439            cooperative_nvfp4,
440            tuning,
441        })
442    }
443
444    /// Preference order when no device was named: discrete, integrated, virtual, CPU, other.
445    fn rank(&self) -> u8 {
446        match self.device_type {
447            vk::PhysicalDeviceType::DISCRETE_GPU => 0,
448            vk::PhysicalDeviceType::INTEGRATED_GPU => 1,
449            vk::PhysicalDeviceType::VIRTUAL_GPU => 2,
450            vk::PhysicalDeviceType::CPU => 3,
451            _ => 4,
452        }
453    }
454
455    fn class(&self) -> AcceleratorClass {
456        if self.device_type == vk::PhysicalDeviceType::CPU {
457            AcceleratorClass::OTHER
458        } else {
459            AcceleratorClass::GPU
460        }
461    }
462}
463
464/// `minImportedHostPointerAlignment`, when `handle` offers `VK_EXT_external_memory_host`.
465fn host_import_alignment(instance: &ash::Instance, handle: vk::PhysicalDevice) -> Option<u64> {
466    // SAFETY: `handle` is a live physical device of `instance`; no layer is named.
467    let extensions = unsafe { instance.enumerate_device_extension_properties(handle) }.ok()?;
468    let name = ash::ext::external_memory_host::NAME;
469    extensions
470        .iter()
471        .any(|extension| extension.extension_name_as_c_str() == Ok(name))
472        .then_some(())?;
473    let mut host = vk::PhysicalDeviceExternalMemoryHostPropertiesEXT::default();
474    let mut properties = vk::PhysicalDeviceProperties2::default().push_next(&mut host);
475    // SAFETY: as above; the chained structure belongs to an extension the device reported.
476    unsafe { instance.get_physical_device_properties2(handle, &mut properties) };
477    let alignment = host.min_imported_host_pointer_alignment;
478    alignment.is_power_of_two().then_some(alignment)
479}
480
481fn supports_extension(instance: &ash::Instance, handle: vk::PhysicalDevice, name: &CStr) -> bool {
482    // SAFETY: `handle` is a live physical device of `instance`.
483    unsafe { instance.enumerate_device_extension_properties(handle) }.is_ok_and(|extensions| {
484        extensions
485            .iter()
486            .any(|extension| extension.extension_name_as_c_str() == Ok(name))
487    })
488}
489
490fn supports_nvfp4_cooperative_matrix(owner: &Instance, handle: vk::PhysicalDevice) -> bool {
491    let extension = ash::khr::cooperative_matrix::Instance::new(&owner._entry, &owner.instance);
492    // SAFETY: the caller checked that the physical device advertises the extension and `handle`
493    // belongs to this live instance.
494    let Ok(properties) =
495        (unsafe { extension.get_physical_device_cooperative_matrix_properties(handle) })
496    else {
497        return false;
498    };
499    properties.iter().any(|property| {
500        property.scope == vk::ScopeKHR::SUBGROUP
501            && property.m_size == 8
502            && property.n_size == 16
503            && property.k_size == 16
504            && property.a_type == vk::ComponentTypeKHR::FLOAT16
505            && property.b_type == vk::ComponentTypeKHR::FLOAT16
506            && property.c_type == vk::ComponentTypeKHR::FLOAT32
507            && property.result_type == vk::ComponentTypeKHR::FLOAT32
508            && property.saturating_accumulation == vk::FALSE
509    })
510}
511
512/// View a driver-filled `c_char` name array as bytes for `CStr` parsing.
513fn bytemuck_i8_to_u8(name: &[std::ffi::c_char; 256]) -> &[u8; 256] {
514    // SAFETY: `c_char` and `u8` have identical size and alignment; the array is plain data.
515    unsafe { &*(name as *const [std::ffi::c_char; 256]).cast::<[u8; 256]>() }
516}
517
518fn enumerate(instance: &Instance) -> Result<Vec<PhysicalDeviceRecord>, InitError> {
519    // SAFETY: enumeration on a live instance has no other preconditions.
520    let handles = unsafe { instance.instance.enumerate_physical_devices() }
521        .map_err(|_| InitError::DeviceEnumerationFailed)?;
522    Ok(handles
523        .into_iter()
524        .filter_map(|handle| PhysicalDeviceRecord::probe(instance, handle))
525        .collect())
526}
527
528/// The memory type chosen for each advertised domain (ADR 0005 memory-domain map).
529#[derive(Clone, Copy, Debug)]
530struct MemoryPlan {
531    /// `HOST_VISIBLE | HOST_COHERENT`, preferring system memory and a host-cached type.
532    host: u32,
533    /// `DEVICE_LOCAL`, preferring a type the host cannot see; absent on devices without one.
534    device: Option<u32>,
535    /// `DEVICE_LOCAL | HOST_VISIBLE | HOST_COHERENT`: ReBAR or UMA, never assumed.
536    shared: Option<u32>,
537}
538
539/// Backend options a host may set when opening a device. The defaults are what
540/// [`VulkanAccelerator::new`] and [`VulkanAccelerator::with_device`] use.
541#[derive(Clone, Copy, Debug, PartialEq, Eq)]
542pub struct VulkanOptions {
543    /// On a device with a single memory heap — an integrated GPU, Apple silicon, a software
544    /// ICD — there is no second memory for the `Device` domain to be local to, so its
545    /// allocations take a host-visible device-local type and `write_buffer`/`read_buffer` are a
546    /// mapped copy rather than a staged copy through a transient buffer and a GPU transfer
547    /// (ADR 0012). `false` keeps the discrete-GPU plan (a non-host-visible type, staged
548    /// transfers) on such devices too; the staging path's tests use it.
549    pub map_unified_device_memory: bool,
550}
551
552impl Default for VulkanOptions {
553    fn default() -> Self {
554        Self {
555            map_unified_device_memory: true,
556        }
557    }
558}
559
560impl MemoryPlan {
561    /// Choose one type per domain among those `buffer_type_mask` (the `memoryTypeBits` a
562    /// storage buffer of this backend reports) permits. With `unified`, the `Device` domain
563    /// prefers a host-visible type (see [`VulkanOptions::map_unified_device_memory`]).
564    fn select(
565        memory: &vk::PhysicalDeviceMemoryProperties,
566        buffer_type_mask: u32,
567        unified: bool,
568    ) -> Option<Self> {
569        let types = &memory.memory_types[..memory.memory_type_count as usize];
570        let usable = |index: usize, flags: vk::MemoryPropertyFlags| {
571            buffer_type_mask & (1 << index) != 0
572                && !flags.intersects(
573                    vk::MemoryPropertyFlags::PROTECTED
574                        | vk::MemoryPropertyFlags::LAZILY_ALLOCATED
575                        | vk::MemoryPropertyFlags::DEVICE_COHERENT_AMD
576                        | vk::MemoryPropertyFlags::RDMA_CAPABLE_NV,
577                )
578        };
579        let pick = |required: vk::MemoryPropertyFlags,
580                    score: fn(vk::MemoryPropertyFlags) -> u8|
581         -> Option<u32> {
582            types
583                .iter()
584                .enumerate()
585                .filter(|(index, ty)| {
586                    usable(*index, ty.property_flags) && ty.property_flags.contains(required)
587                })
588                .max_by_key(|(_, ty)| score(ty.property_flags))
589                .and_then(|(index, _)| u32::try_from(index).ok())
590        };
591        let host_coherent =
592            vk::MemoryPropertyFlags::HOST_VISIBLE | vk::MemoryPropertyFlags::HOST_COHERENT;
593        Some(Self {
594            host: pick(host_coherent, |flags| {
595                u8::from(!flags.contains(vk::MemoryPropertyFlags::DEVICE_LOCAL)) * 2
596                    + u8::from(flags.contains(vk::MemoryPropertyFlags::HOST_CACHED))
597            })?,
598            device: pick(
599                vk::MemoryPropertyFlags::DEVICE_LOCAL,
600                if unified {
601                    |flags| {
602                        u8::from(flags.contains(
603                            vk::MemoryPropertyFlags::HOST_VISIBLE
604                                | vk::MemoryPropertyFlags::HOST_COHERENT,
605                        )) * 2
606                            + u8::from(flags.contains(vk::MemoryPropertyFlags::HOST_CACHED))
607                    }
608                } else {
609                    |flags| u8::from(!flags.contains(vk::MemoryPropertyFlags::HOST_VISIBLE))
610                },
611            ),
612            shared: pick(
613                vk::MemoryPropertyFlags::DEVICE_LOCAL | host_coherent,
614                |flags| u8::from(flags.contains(vk::MemoryPropertyFlags::HOST_CACHED)),
615            ),
616        })
617    }
618
619    /// Whether the chosen type for `domain` is host-visible and coherent, i.e. whether its
620    /// allocations are persistently mapped and transfers are a plain copy.
621    fn is_mapped(self, memory: &vk::PhysicalDeviceMemoryProperties, domain: MemoryDomain) -> bool {
622        self.for_domain(domain).is_some_and(|index| {
623            memory.memory_types[index as usize].property_flags.contains(
624                vk::MemoryPropertyFlags::HOST_VISIBLE | vk::MemoryPropertyFlags::HOST_COHERENT,
625            )
626        })
627    }
628
629    fn for_domain(self, domain: MemoryDomain) -> Option<u32> {
630        match domain {
631            MemoryDomain::Host => Some(self.host),
632            MemoryDomain::Device => self.device,
633            MemoryDomain::Shared => self.shared,
634        }
635    }
636
637    fn capabilities(self) -> Capabilities {
638        let mut capabilities = Capabilities::HOST_VISIBLE_MEMORY;
639        if self.device.is_some() {
640            capabilities |= Capabilities::DEVICE_LOCAL_MEMORY;
641        }
642        if self.shared.is_some() {
643            capabilities |= Capabilities::SHARED_MEMORY;
644        }
645        capabilities
646    }
647}
648
649/// Live provider resource totals for accounting hooks.
650#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
651pub struct LiveResources {
652    pub contexts: u64,
653    pub buffers: u64,
654    pub programs: u64,
655    pub queues: u64,
656    pub events: u64,
657}
658
659#[derive(Default)]
660struct Counters {
661    direct_binding_admissions: Cell<u64>,
662    explicit_transfer_bytes: Cell<u64>,
663    contexts: Cell<u64>,
664    buffers: Cell<u64>,
665    programs: Cell<u64>,
666    queues: Cell<u64>,
667    events: Cell<u64>,
668}
669
670fn increment(cell: &Cell<u64>, by: u64) {
671    cell.set(cell.get().saturating_add(by));
672}
673
674fn decrement(cell: &Cell<u64>) {
675    cell.set(cell.get().saturating_sub(1));
676}
677
678/// The opened device and everything shared by all handles of one backend instance.
679///
680/// Field order matters for teardown: the explicit `Drop` destroys device-level objects and the
681/// device, then the `instance` field's own `Drop` destroys the instance.
682struct Shared {
683    device: ash::Device,
684    physical: PhysicalDeviceRecord,
685    queue: vk::Queue,
686    set_layout: vk::DescriptorSetLayout,
687    pipeline_layout: vk::PipelineLayout,
688    /// Driver-side cache shared by every pipeline of this instance: programs selecting the same
689    /// kernel with the same specialization are compiled once (ADR 0007).
690    pipeline_cache: vk::PipelineCache,
691    /// Assembled kernel modules by variant; assembled once per instance, on first use.
692    modules: RefCell<HashMap<KernelKey, Rc<[u32]>>>,
693    memory_plan: MemoryPlan,
694    /// The `VK_EXT_external_memory_host` entry points, when the extension was enabled.
695    host_memory: Option<ash::ext::external_memory_host::Device>,
696    info: DeviceInfo,
697    /// Sticky device-loss flag: after `VK_ERROR_DEVICE_LOST` no entry point is re-entered except
698    /// destruction (ADR 0006).
699    poisoned: Cell<bool>,
700    counters: Counters,
701    /// Dropped last (see the struct documentation): destroys the instance after the device.
702    _instance: Instance,
703}
704
705impl Shared {
706    fn open(
707        instance: Instance,
708        physical: PhysicalDeviceRecord,
709        options: VulkanOptions,
710    ) -> Result<Rc<Self>, InitError> {
711        let priorities = [1.0_f32];
712        let queue_info = vk::DeviceQueueCreateInfo::default()
713            .queue_family_index(physical.queue_family)
714            .queue_priorities(&priorities);
715        let queue_infos = [queue_info];
716        // `synchronization2` is core in 1.3 but still an opt-in feature (ADR 0005);
717        // `bufferDeviceAddress` is enabled only to measure allocation alignment honestly.
718        let mut vulkan12 = vk::PhysicalDeviceVulkan12Features::default()
719            .buffer_device_address(physical.buffer_device_address)
720            .timeline_semaphore(physical.timeline_semaphore)
721            .shader_float16(physical.shader_float16)
722            .vulkan_memory_model(physical.vulkan_memory_model);
723        let mut vulkan13 = vk::PhysicalDeviceVulkan13Features::default().synchronization2(true);
724        let mut cooperative = vk::PhysicalDeviceCooperativeMatrixFeaturesKHR::default()
725            .cooperative_matrix(physical.cooperative_nvfp4);
726        // `VK_EXT_external_memory_host`, when offered, only for importing caller memory
727        // (ADR 0013); nothing else depends on it.
728        let mut extensions = Vec::with_capacity(2);
729        if physical.host_import_alignment.is_some() {
730            extensions.push(ash::ext::external_memory_host::NAME.as_ptr());
731        }
732        if physical.cooperative_nvfp4 {
733            extensions.push(ash::khr::cooperative_matrix::NAME.as_ptr());
734        }
735        let info = vk::DeviceCreateInfo::default()
736            .queue_create_infos(&queue_infos)
737            .enabled_extension_names(&extensions)
738            .push_next(&mut vulkan12)
739            .push_next(&mut vulkan13)
740            .push_next(&mut cooperative);
741        // SAFETY: `physical.handle` belongs to `instance.instance`; every pointed-to structure
742        // outlives the call; the requested features and extension were reported supported by the
743        // probe.
744        let device = unsafe {
745            instance
746                .instance
747                .create_device(physical.handle, &info, None)
748        }
749        .map_err(|_| InitError::DeviceCreationFailed)?;
750        // SAFETY: the queue family and index 0 were requested at device creation.
751        let queue = unsafe { device.get_device_queue(physical.queue_family, 0) };
752
753        // Which memory types a storage buffer of this backend may live in is a property of the
754        // buffer usage, not of the heap list alone (ANV exposes types buffers cannot use), so the
755        // memory-domain map is chosen against a probe buffer's `memoryTypeBits`.
756        let buffer_type_mask = match probe_buffer_type_mask(&device, &physical) {
757            Ok(mask) => mask,
758            Err(_) => {
759                // SAFETY: the device was created above and has no other objects yet.
760                unsafe { device.destroy_device(None) };
761                return Err(InitError::DeviceCreationFailed);
762            }
763        };
764        let unified = options.map_unified_device_memory && physical.memory.memory_heap_count == 1;
765        let Some(memory_plan) = MemoryPlan::select(&physical.memory, buffer_type_mask, unified)
766        else {
767            // SAFETY: as above.
768            unsafe { device.destroy_device(None) };
769            return Err(InitError::DeviceUnavailable);
770        };
771
772        // One descriptor: set 0, binding 0, an array of storage buffers. Elements `0..bindings`
773        // are the submission's bound slots and the last element is the program arena; kernels
774        // select operands by specialization constant (ADR 0007).
775        let bindings = [vk::DescriptorSetLayoutBinding::default()
776            .binding(0)
777            .descriptor_type(vk::DescriptorType::STORAGE_BUFFER)
778            .descriptor_count(physical.tuning.buffers)
779            .stage_flags(vk::ShaderStageFlags::COMPUTE)];
780        let layout_info = vk::DescriptorSetLayoutCreateInfo::default().bindings(&bindings);
781        // SAFETY: the device is live and `layout_info` points at live locals.
782        let set_layout = match unsafe { device.create_descriptor_set_layout(&layout_info, None) } {
783            Ok(layout) => layout,
784            Err(_) => {
785                // SAFETY: the device was created above and has no other objects yet.
786                unsafe { device.destroy_device(None) };
787                return Err(InitError::DeviceCreationFailed);
788            }
789        };
790        let set_layouts = [set_layout];
791        let pipeline_layout_info =
792            vk::PipelineLayoutCreateInfo::default().set_layouts(&set_layouts);
793        // SAFETY: the device and set layout are live.
794        let pipeline_layout =
795            match unsafe { device.create_pipeline_layout(&pipeline_layout_info, None) } {
796                Ok(layout) => layout,
797                Err(_) => {
798                    // SAFETY: both objects were created above and nothing references them.
799                    unsafe {
800                        device.destroy_descriptor_set_layout(set_layout, None);
801                        device.destroy_device(None);
802                    }
803                    return Err(InitError::DeviceCreationFailed);
804                }
805            };
806
807        // SAFETY: the device is live; an empty create info is a valid empty cache.
808        let pipeline_cache = match unsafe {
809            device.create_pipeline_cache(&vk::PipelineCacheCreateInfo::default(), None)
810        } {
811            Ok(cache) => cache,
812            Err(_) => {
813                // SAFETY: the three objects were created above and nothing references them.
814                unsafe {
815                    device.destroy_pipeline_layout(pipeline_layout, None);
816                    device.destroy_descriptor_set_layout(set_layout, None);
817                    device.destroy_device(None);
818                }
819                return Err(InitError::DeviceCreationFailed);
820            }
821        };
822
823        let info = device_info(&physical, memory_plan);
824        let host_memory = physical
825            .host_import_alignment
826            .map(|_| ash::ext::external_memory_host::Device::new(&instance.instance, &device));
827        Ok(Rc::new(Self {
828            device,
829            physical,
830            queue,
831            set_layout,
832            pipeline_layout,
833            pipeline_cache,
834            modules: RefCell::new(HashMap::new()),
835            memory_plan,
836            host_memory,
837            info,
838            poisoned: Cell::new(false),
839            counters: Counters::default(),
840            _instance: instance,
841        }))
842    }
843
844    /// Map a failed `VkResult`, latching device loss.
845    fn fail(&self, result: vk::Result) -> BackendError {
846        if result == vk::Result::ERROR_DEVICE_LOST {
847            self.poisoned.set(true);
848        }
849        backend_error(result)
850    }
851
852    fn check_live(&self) -> Result<(), BackendError> {
853        if self.poisoned.get() {
854            Err(BackendError::DeviceLost)
855        } else {
856            Ok(())
857        }
858    }
859
860    /// The assembled module for `key`, built on first use and shared by every program after.
861    fn module(&self, key: KernelKey) -> Rc<[u32]> {
862        Rc::clone(
863            self.modules
864                .borrow_mut()
865                .entry(key)
866                .or_insert_with(|| Rc::from(key.assemble())),
867        )
868    }
869
870    /// Block until the device is idle; used only on teardown paths that must not free memory a
871    /// pending submission may still touch.
872    fn wait_idle(&self) {
873        // SAFETY: the device is live; waiting has no other preconditions. Errors are latched.
874        if let Err(result) = unsafe { self.device.device_wait_idle() } {
875            self.fail(result);
876        }
877    }
878}
879
880impl Drop for Shared {
881    fn drop(&mut self) {
882        // SAFETY: every child object holds an `Rc<Shared>`, so this runs only after all of them
883        // were destroyed; the layouts and device are destroyed exactly once, then the `_instance`
884        // field drops and destroys the instance.
885        unsafe {
886            let _ = self.device.device_wait_idle();
887            self.device
888                .destroy_pipeline_cache(self.pipeline_cache, None);
889            self.device
890                .destroy_pipeline_layout(self.pipeline_layout, None);
891            self.device
892                .destroy_descriptor_set_layout(self.set_layout, None);
893            self.device.destroy_device(None);
894        }
895    }
896}
897
898fn device_info(physical: &PhysicalDeviceRecord, memory_plan: MemoryPlan) -> DeviceInfo {
899    let largest_heap = physical.memory.memory_heaps[..physical.memory.memory_heap_count as usize]
900        .iter()
901        .map(|heap| heap.size)
902        .max()
903        .unwrap_or(0);
904    // A storage-buffer descriptor cannot exceed `maxStorageBufferRange` (128 MiB on lavapipe),
905    // so no buffer may either: a larger allocation could never be bound directly.
906    // Sizes are rounded up to whole words at allocation, so the advertised bound is rounded
907    // down to keep every rounded descriptor range inside `maxStorageBufferRange`.
908    let max_buffer_bytes = (u64::from(physical.limits.max_storage_buffer_range).min(largest_heap)
909        / WORD_BYTES
910        * WORD_BYTES)
911        .max(WORD_BYTES);
912    DeviceInfo {
913        identity: DeviceIdentity {
914            uuid: physical.uuid,
915            class: physical.class(),
916            vendor_id: physical.vendor_id,
917            device_id: physical.device_id,
918        },
919        capabilities: memory_plan.capabilities(),
920        limits: DeviceLimits {
921            max_contexts: MAX_CONTEXTS,
922            max_buffers_per_context: MAX_BUFFERS_PER_CONTEXT,
923            max_programs_per_context: MAX_PROGRAMS_PER_CONTEXT,
924            max_queues_per_context: MAX_QUEUES_PER_CONTEXT,
925            max_events_per_context: RING_DEPTH,
926            max_bindings_per_submission: physical.tuning.max_bindings(),
927            max_buffer_bytes,
928            max_artifact_bytes: MAX_TOSA_ARTIFACT_BYTES,
929        },
930    }
931}
932
933/// One preallocated ring slot: claimed by exactly one live event at a time (ADR 0006).
934struct Slot {
935    command_buffer: vk::CommandBuffer,
936    fence: vk::Fence,
937    descriptor_set: vk::DescriptorSet,
938}
939
940/// Provider state of one context: the pools, the ring, and the synchronous transfer kit.
941struct ContextInner {
942    shared: Rc<Shared>,
943    id: u64,
944    command_pool: vk::CommandPool,
945    descriptor_pool: vk::DescriptorPool,
946    slots: Vec<Slot>,
947    /// Indices into `slots` not owned by a live event.
948    free_slots: RefCell<Vec<u16>>,
949    /// Command buffer and fence for blocking `write_buffer`/`read_buffer` staging copies.
950    transfer_command_buffer: vk::CommandBuffer,
951    transfer_fence: vk::Fence,
952}
953
954impl ContextInner {
955    fn create(shared: &Rc<Shared>, id: u64) -> Result<Rc<Self>, BackendError> {
956        let device = &shared.device;
957        let pool_info = vk::CommandPoolCreateInfo::default()
958            .flags(vk::CommandPoolCreateFlags::RESET_COMMAND_BUFFER)
959            .queue_family_index(shared.physical.queue_family);
960        // SAFETY: the device is live; `pool_info` is a live local.
961        let command_pool = unsafe { device.create_command_pool(&pool_info, None) }
962            .map_err(|result| shared.fail(result))?;
963        let mut partial = PartialContext {
964            shared,
965            command_pool,
966            descriptor_pool: vk::DescriptorPool::null(),
967            fences: Vec::new(),
968        };
969
970        let allocate_info = vk::CommandBufferAllocateInfo::default()
971            .command_pool(command_pool)
972            .level(vk::CommandBufferLevel::PRIMARY)
973            .command_buffer_count(RING_DEPTH + 1);
974        // SAFETY: the pool is live; the buffers are freed with the pool.
975        let command_buffers = unsafe { device.allocate_command_buffers(&allocate_info) }
976            .map_err(|result| shared.fail(result))?;
977
978        // Each set holds the whole descriptor array, so every set charges `buffers` descriptors.
979        let pool_sizes = [vk::DescriptorPoolSize {
980            ty: vk::DescriptorType::STORAGE_BUFFER,
981            descriptor_count: RING_DEPTH * shared.physical.tuning.buffers,
982        }];
983        let descriptor_pool_info = vk::DescriptorPoolCreateInfo::default()
984            .max_sets(RING_DEPTH)
985            .pool_sizes(&pool_sizes);
986        // SAFETY: the device is live; `descriptor_pool_info` is a live local.
987        partial.descriptor_pool =
988            unsafe { device.create_descriptor_pool(&descriptor_pool_info, None) }
989                .map_err(|result| shared.fail(result))?;
990        let set_layouts = vec![shared.set_layout; RING_DEPTH as usize];
991        let set_info = vk::DescriptorSetAllocateInfo::default()
992            .descriptor_pool(partial.descriptor_pool)
993            .set_layouts(&set_layouts);
994        // SAFETY: the pool was sized for exactly these sets; they are freed with the pool.
995        let descriptor_sets = unsafe { device.allocate_descriptor_sets(&set_info) }
996            .map_err(|result| shared.fail(result))?;
997
998        for _ in 0..=RING_DEPTH {
999            // SAFETY: the device is live; each fence is owned by this context and destroyed once.
1000            let fence = unsafe { device.create_fence(&vk::FenceCreateInfo::default(), None) }
1001                .map_err(|result| shared.fail(result))?;
1002            partial.fences.push(fence);
1003        }
1004
1005        let (transfer_command_buffer, ring_command_buffers) = command_buffers
1006            .split_last()
1007            .expect("RING_DEPTH + 1 buffers");
1008        let transfer_fence = partial.fences.pop().expect("RING_DEPTH + 1 fences");
1009        let slots = ring_command_buffers
1010            .iter()
1011            .zip(&partial.fences)
1012            .zip(&descriptor_sets)
1013            .map(|((command_buffer, fence), descriptor_set)| Slot {
1014                command_buffer: *command_buffer,
1015                fence: *fence,
1016                descriptor_set: *descriptor_set,
1017            })
1018            .collect::<Vec<_>>();
1019        let free_slots = (0..RING_DEPTH as u16).rev().collect();
1020        let descriptor_pool = partial.descriptor_pool;
1021        // Ownership transfers to the context; the partial guard must not destroy anything now.
1022        partial.fences.clear();
1023        partial.descriptor_pool = vk::DescriptorPool::null();
1024        partial.command_pool = vk::CommandPool::null();
1025        increment(&shared.counters.contexts, 1);
1026        Ok(Rc::new(Self {
1027            shared: Rc::clone(shared),
1028            id,
1029            command_pool,
1030            descriptor_pool,
1031            slots,
1032            free_slots: RefCell::new(free_slots),
1033            transfer_command_buffer: *transfer_command_buffer,
1034            transfer_fence,
1035        }))
1036    }
1037
1038    fn claim_slot(&self) -> Option<u16> {
1039        self.free_slots.borrow_mut().pop()
1040    }
1041
1042    fn release_slot(&self, slot: u16) {
1043        self.free_slots.borrow_mut().push(slot);
1044    }
1045
1046    /// Record, submit, and wait for one buffer-to-buffer copy on the transfer kit.
1047    fn blocking_copy(
1048        &self,
1049        source: vk::Buffer,
1050        destination: vk::Buffer,
1051        region: vk::BufferCopy,
1052        visibility: CopyVisibility,
1053    ) -> Result<(), BackendError> {
1054        let shared = &self.shared;
1055        let device = &shared.device;
1056        let command_buffer = self.transfer_command_buffer;
1057        let fence = self.transfer_fence;
1058        let begin = vk::CommandBufferBeginInfo::default()
1059            .flags(vk::CommandBufferUsageFlags::ONE_TIME_SUBMIT);
1060        let barrier = vk::MemoryBarrier2::default()
1061            .src_stage_mask(vk::PipelineStageFlags2::COPY)
1062            .src_access_mask(vk::AccessFlags2::TRANSFER_WRITE);
1063        let barrier = match visibility {
1064            CopyVisibility::HostRead => barrier
1065                .dst_stage_mask(vk::PipelineStageFlags2::HOST)
1066                .dst_access_mask(vk::AccessFlags2::HOST_READ),
1067            CopyVisibility::Device => barrier
1068                .dst_stage_mask(
1069                    vk::PipelineStageFlags2::COMPUTE_SHADER | vk::PipelineStageFlags2::COPY,
1070                )
1071                .dst_access_mask(
1072                    vk::AccessFlags2::SHADER_STORAGE_READ
1073                        | vk::AccessFlags2::SHADER_STORAGE_WRITE
1074                        | vk::AccessFlags2::TRANSFER_READ
1075                        | vk::AccessFlags2::TRANSFER_WRITE,
1076                ),
1077        };
1078        let barriers = [barrier];
1079        let dependency = vk::DependencyInfo::default().memory_barriers(&barriers);
1080        let submit_buffers =
1081            [vk::CommandBufferSubmitInfo::default().command_buffer(command_buffer)];
1082        let submits = [vk::SubmitInfo2::default().command_buffer_infos(&submit_buffers)];
1083        // SAFETY: the transfer command buffer and fence are used only by this synchronous method,
1084        // which waits for the fence before returning, so no prior use is still pending; both
1085        // buffers are live allocations of this context and the region was bounds-checked by the
1086        // caller. `begin_command_buffer` implicitly resets the buffer (pool flag).
1087        unsafe {
1088            device
1089                .reset_fences(&[fence])
1090                .map_err(|result| shared.fail(result))?;
1091            device
1092                .begin_command_buffer(command_buffer, &begin)
1093                .map_err(|result| shared.fail(result))?;
1094            device.cmd_copy_buffer(command_buffer, source, destination, &[region]);
1095            device.cmd_pipeline_barrier2(command_buffer, &dependency);
1096            device
1097                .end_command_buffer(command_buffer)
1098                .map_err(|result| shared.fail(result))?;
1099            device
1100                .queue_submit2(shared.queue, &submits, fence)
1101                .map_err(|result| shared.fail(result))?;
1102            match device.wait_for_fences(&[fence], true, TRANSFER_TIMEOUT_NS) {
1103                Ok(()) => Ok(()),
1104                Err(vk::Result::TIMEOUT) => {
1105                    // A bounded copy that never completes leaves the staging allocation in
1106                    // unknown device use: treat the device as lost rather than free it.
1107                    shared.poisoned.set(true);
1108                    Err(BackendError::DeviceLost)
1109                }
1110                Err(result) => Err(shared.fail(result)),
1111            }
1112        }
1113    }
1114}
1115
1116/// Who consumes the destination of a blocking copy, and therefore which barrier follows it.
1117/// Submissions carry no implicit memory dependency between one another, so every consumer of a
1118/// copied range is named explicitly.
1119#[derive(Clone, Copy, Debug, PartialEq, Eq)]
1120enum CopyVisibility {
1121    /// The host reads the destination through a mapping after the fence signals.
1122    HostRead,
1123    /// Later submissions read or write the destination: compute dispatches over a bound buffer
1124    /// or the arena, and further staging copies out of or into it.
1125    Device,
1126}
1127
1128/// Copy arbitrary source bytes through a staging buffer. Vulkan copies operate on whole words, so
1129/// partial first and last words are read-modify-written to preserve neighbouring logical bytes.
1130fn write_through_staging<S: ByteSource + ?Sized>(
1131    context: &ContextInner,
1132    destination: vk::Buffer,
1133    start: u64,
1134    data: &S,
1135    len: u64,
1136    staging: &mut Staging<'_>,
1137) -> Result<(), BackendError> {
1138    let mut done = 0_u64;
1139    let misalignment = start % WORD_BYTES;
1140    if misalignment != 0 {
1141        let prefix = (WORD_BYTES - misalignment).min(len);
1142        context.blocking_copy(
1143            destination,
1144            staging.raw.buffer,
1145            vk::BufferCopy {
1146                src_offset: start - misalignment,
1147                dst_offset: 0,
1148                size: WORD_BYTES,
1149            },
1150            CopyVisibility::HostRead,
1151        )?;
1152        let slice_start = usize::try_from(misalignment).map_err(|_| BackendError::OutOfBounds)?;
1153        let end = slice_start
1154            .checked_add(usize::try_from(prefix).map_err(|_| BackendError::OutOfBounds)?)
1155            .ok_or(BackendError::OutOfBounds)?;
1156        data.read_at(done, &mut staging.as_mut_slice()[slice_start..end])?;
1157        context.blocking_copy(
1158            staging.raw.buffer,
1159            destination,
1160            vk::BufferCopy {
1161                src_offset: 0,
1162                dst_offset: start - misalignment,
1163                size: WORD_BYTES,
1164            },
1165            CopyVisibility::Device,
1166        )?;
1167        done += prefix;
1168    }
1169    while len - done >= WORD_BYTES {
1170        let chunk = ((len - done).min(staging.bytes) / WORD_BYTES) * WORD_BYTES;
1171        let chunk_len = usize::try_from(chunk).map_err(|_| BackendError::OutOfBounds)?;
1172        data.read_at(done, &mut staging.as_mut_slice()[..chunk_len])?;
1173        context.blocking_copy(
1174            staging.raw.buffer,
1175            destination,
1176            vk::BufferCopy {
1177                src_offset: 0,
1178                dst_offset: start + done,
1179                size: chunk,
1180            },
1181            CopyVisibility::Device,
1182        )?;
1183        done += chunk;
1184    }
1185    if done < len {
1186        let tail = len - done;
1187        context.blocking_copy(
1188            destination,
1189            staging.raw.buffer,
1190            vk::BufferCopy {
1191                src_offset: start + done,
1192                dst_offset: 0,
1193                size: WORD_BYTES,
1194            },
1195            CopyVisibility::HostRead,
1196        )?;
1197        let tail_len = usize::try_from(tail).map_err(|_| BackendError::OutOfBounds)?;
1198        data.read_at(done, &mut staging.as_mut_slice()[..tail_len])?;
1199        context.blocking_copy(
1200            staging.raw.buffer,
1201            destination,
1202            vk::BufferCopy {
1203                src_offset: 0,
1204                dst_offset: start + done,
1205                size: WORD_BYTES,
1206            },
1207            CopyVisibility::Device,
1208        )?;
1209    }
1210    Ok(())
1211}
1212
1213/// Copy arbitrary bytes from a buffer through staging. Partial first and last words are copied in
1214/// full, but only their requested logical bytes are sent to the caller.
1215fn read_through_staging(
1216    context: &ContextInner,
1217    source: vk::Buffer,
1218    start: u64,
1219    data: &mut dyn ByteSink,
1220    len: u64,
1221    staging: &mut Staging<'_>,
1222) -> Result<(), BackendError> {
1223    let mut done = 0_u64;
1224    let misalignment = start % WORD_BYTES;
1225    if misalignment != 0 {
1226        let prefix = (WORD_BYTES - misalignment).min(len);
1227        context.blocking_copy(
1228            source,
1229            staging.raw.buffer,
1230            vk::BufferCopy {
1231                src_offset: start - misalignment,
1232                dst_offset: 0,
1233                size: WORD_BYTES,
1234            },
1235            CopyVisibility::HostRead,
1236        )?;
1237        let slice_start = usize::try_from(misalignment).map_err(|_| BackendError::OutOfBounds)?;
1238        let end = slice_start
1239            .checked_add(usize::try_from(prefix).map_err(|_| BackendError::OutOfBounds)?)
1240            .ok_or(BackendError::OutOfBounds)?;
1241        data.write_at(done, &staging.as_mut_slice()[slice_start..end])?;
1242        done += prefix;
1243    }
1244    while len - done >= WORD_BYTES {
1245        let chunk = ((len - done).min(staging.bytes) / WORD_BYTES) * WORD_BYTES;
1246        let chunk_len = usize::try_from(chunk).map_err(|_| BackendError::OutOfBounds)?;
1247        context.blocking_copy(
1248            source,
1249            staging.raw.buffer,
1250            vk::BufferCopy {
1251                src_offset: start + done,
1252                dst_offset: 0,
1253                size: chunk,
1254            },
1255            CopyVisibility::HostRead,
1256        )?;
1257        data.write_at(done, &staging.as_mut_slice()[..chunk_len])?;
1258        done += chunk;
1259    }
1260    if done < len {
1261        let tail = usize::try_from(len - done).map_err(|_| BackendError::OutOfBounds)?;
1262        context.blocking_copy(
1263            source,
1264            staging.raw.buffer,
1265            vk::BufferCopy {
1266                src_offset: start + done,
1267                dst_offset: 0,
1268                size: WORD_BYTES,
1269            },
1270            CopyVisibility::HostRead,
1271        )?;
1272        data.write_at(done, &staging.as_mut_slice()[..tail])?;
1273    }
1274    Ok(())
1275}
1276
1277/// Destroys partially created context objects if creation fails midway.
1278struct PartialContext<'a> {
1279    shared: &'a Shared,
1280    command_pool: vk::CommandPool,
1281    descriptor_pool: vk::DescriptorPool,
1282    fences: Vec<vk::Fence>,
1283}
1284
1285impl Drop for PartialContext<'_> {
1286    fn drop(&mut self) {
1287        let device = &self.shared.device;
1288        // SAFETY: each handle here was created by `ContextInner::create` and not yet handed to a
1289        // context; null handles are skipped, and destroying a pool frees its allocations.
1290        unsafe {
1291            for fence in self.fences.drain(..) {
1292                device.destroy_fence(fence, None);
1293            }
1294            if self.descriptor_pool != vk::DescriptorPool::null() {
1295                device.destroy_descriptor_pool(self.descriptor_pool, None);
1296            }
1297            if self.command_pool != vk::CommandPool::null() {
1298                device.destroy_command_pool(self.command_pool, None);
1299            }
1300        }
1301    }
1302}
1303
1304impl Drop for ContextInner {
1305    fn drop(&mut self) {
1306        let shared = &self.shared;
1307        // Every child holds an `Rc<ContextInner>`, so no event can still be pending here; the
1308        // idle wait is defense in depth for a poisoned or misused instance.
1309        if self.free_slots.borrow().len() != self.slots.len() {
1310            shared.wait_idle();
1311        }
1312        // SAFETY: the pools own their command buffers and descriptor sets; the fences were created
1313        // by this context. Each is destroyed exactly once.
1314        unsafe {
1315            for slot in &self.slots {
1316                shared.device.destroy_fence(slot.fence, None);
1317            }
1318            shared.device.destroy_fence(self.transfer_fence, None);
1319            shared
1320                .device
1321                .destroy_descriptor_pool(self.descriptor_pool, None);
1322            shared.device.destroy_command_pool(self.command_pool, None);
1323        }
1324        decrement(&shared.counters.contexts);
1325    }
1326}
1327
1328/// Vulkan context handle: pools plus the bounded submission ring.
1329pub struct VulkanContext {
1330    inner: Rc<ContextInner>,
1331}
1332
1333impl std::fmt::Debug for VulkanContext {
1334    fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
1335        formatter
1336            .debug_struct("VulkanContext")
1337            .field("id", &self.inner.id)
1338            .finish_non_exhaustive()
1339    }
1340}
1341
1342/// In-flight gate shared between a buffer and the events bound to it.
1343///
1344/// Zero: idle. `1..EXCLUSIVE_ACCESS`: that many read-only bindings in flight. `EXCLUSIVE_ACCESS`:
1345/// one writing binding in flight. Explicit transfers require zero.
1346#[derive(Default)]
1347struct BufferState {
1348    in_flight: Cell<u64>,
1349}
1350
1351/// One dedicated `VkBuffer` + `VkDeviceMemory`, persistently mapped unless device-local.
1352pub struct VulkanBuffer {
1353    context: Rc<ContextInner>,
1354    desc: BufferDesc,
1355    buffer: vk::Buffer,
1356    memory: vk::DeviceMemory,
1357    mapped: Option<NonNull<u8>>,
1358    /// The memory is the caller's, imported (ADR 0013): `mapped` is the caller's pointer, not a
1359    /// `vkMapMemory` mapping, and the bytes outlive this handle.
1360    imported: bool,
1361    state: Rc<BufferState>,
1362}
1363
1364impl std::fmt::Debug for VulkanBuffer {
1365    fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
1366        formatter
1367            .debug_struct("VulkanBuffer")
1368            .field("context", &self.context.id)
1369            .field("desc", &self.desc)
1370            .field("mapped", &self.mapped.is_some())
1371            .field("imported", &self.imported)
1372            .finish_non_exhaustive()
1373    }
1374}
1375
1376impl VulkanBuffer {
1377    fn in_flight(&self) -> u64 {
1378        self.state.in_flight.get()
1379    }
1380
1381    /// Pointer to `offset` inside the persistent mapping, when the buffer is mapped at all.
1382    fn mapped_at(&self, offset: usize) -> Option<*mut u8> {
1383        // SAFETY: callers validated `offset` (plus their length) against `desc.bytes()`, and the
1384        // mapping covers the whole allocation.
1385        self.mapped
1386            .map(|pointer| unsafe { pointer.as_ptr().add(offset) })
1387    }
1388}
1389
1390impl Drop for VulkanBuffer {
1391    fn drop(&mut self) {
1392        let shared = &self.context.shared;
1393        // The contract forbids dropping an in-flight buffer; if it happens anyway, never free
1394        // memory a submission may still address.
1395        if self.in_flight() != 0 {
1396            shared.wait_idle();
1397        }
1398        // SAFETY: this handle owns the mapping, buffer, and memory, all created together in
1399        // `allocate` (or `import`, which maps nothing) and released exactly once here, in the
1400        // reverse order. Freeing imported memory releases the device's claim on the caller's
1401        // pages, never the pages.
1402        unsafe {
1403            if self.mapped.is_some() && !self.imported {
1404                shared.device.unmap_memory(self.memory);
1405            }
1406            shared.device.destroy_buffer(self.buffer, None);
1407            shared.device.free_memory(self.memory, None);
1408        }
1409        decrement(&shared.counters.buffers);
1410    }
1411}
1412
1413/// Usage flags of every buffer this backend creates: directly bindable as a storage buffer, a
1414/// transfer source and destination, and device-addressable when alignment can be measured.
1415fn buffer_usage(physical: &PhysicalDeviceRecord) -> vk::BufferUsageFlags {
1416    let mut usage = vk::BufferUsageFlags::STORAGE_BUFFER
1417        | vk::BufferUsageFlags::TRANSFER_SRC
1418        | vk::BufferUsageFlags::TRANSFER_DST;
1419    if physical.buffer_device_address {
1420        usage |= vk::BufferUsageFlags::SHADER_DEVICE_ADDRESS;
1421    }
1422    usage
1423}
1424
1425/// The `memoryTypeBits` a buffer with this backend's usage reports on `device`.
1426fn probe_buffer_type_mask(
1427    device: &ash::Device,
1428    physical: &PhysicalDeviceRecord,
1429) -> Result<u32, vk::Result> {
1430    let buffer_info = vk::BufferCreateInfo::default()
1431        .size(4)
1432        .usage(buffer_usage(physical))
1433        .sharing_mode(vk::SharingMode::EXCLUSIVE);
1434    // SAFETY: the device is live; the probe buffer is never bound and is destroyed here.
1435    unsafe {
1436        let buffer = device.create_buffer(&buffer_info, None)?;
1437        let requirements = device.get_buffer_memory_requirements(buffer);
1438        device.destroy_buffer(buffer, None);
1439        Ok(requirements.memory_type_bits)
1440    }
1441}
1442
1443/// Owned raw buffer + memory pair used while an allocation is being assembled or staged.
1444struct RawAllocation<'a> {
1445    shared: &'a Shared,
1446    buffer: vk::Buffer,
1447    memory: vk::DeviceMemory,
1448    mapped: Option<NonNull<u8>>,
1449    allocation_bytes: u64,
1450    /// The smallest power-of-two alignment every measured address satisfied.
1451    measured_alignment: u64,
1452    memory_flags: vk::MemoryPropertyFlags,
1453}
1454
1455impl<'a> RawAllocation<'a> {
1456    /// Create a buffer, allocate dedicated memory of `memory_type`, bind at offset 0, and map it
1457    /// when `map` is set. Alignment is measured, never assumed.
1458    fn create(
1459        shared: &'a Shared,
1460        bytes: u64,
1461        memory_type: u32,
1462        map: bool,
1463    ) -> Result<Self, BackendError> {
1464        let device = &shared.device;
1465        // Whole words: byte-storage tensors are read and atomically written by word, so the
1466        // buffer behind any binding must extend to the word containing its last byte.
1467        let rounded = bytes
1468            .checked_add(WORD_BYTES - 1)
1469            .ok_or(BackendError::ResourceLimit)?
1470            / WORD_BYTES
1471            * WORD_BYTES;
1472        let buffer_info = vk::BufferCreateInfo::default()
1473            .size(rounded)
1474            .usage(buffer_usage(&shared.physical))
1475            .sharing_mode(vk::SharingMode::EXCLUSIVE);
1476        // SAFETY: the device is live and `buffer_info` is a live local.
1477        let buffer = unsafe { device.create_buffer(&buffer_info, None) }
1478            .map_err(|result| shared.fail(result))?;
1479        let mut raw = Self {
1480            shared,
1481            buffer,
1482            memory: vk::DeviceMemory::null(),
1483            mapped: None,
1484            allocation_bytes: 0,
1485            measured_alignment: 0,
1486            memory_flags: shared.physical.memory.memory_types[memory_type as usize].property_flags,
1487        };
1488
1489        // SAFETY: `buffer` is live.
1490        let requirements = unsafe { device.get_buffer_memory_requirements(buffer) };
1491        if requirements.memory_type_bits & (1 << memory_type) == 0 {
1492            return Err(BackendError::Incompatible);
1493        }
1494        let mut flags =
1495            vk::MemoryAllocateFlagsInfo::default().flags(vk::MemoryAllocateFlags::DEVICE_ADDRESS);
1496        let mut allocate_info = vk::MemoryAllocateInfo::default()
1497            .allocation_size(requirements.size)
1498            .memory_type_index(memory_type);
1499        if shared.physical.buffer_device_address {
1500            allocate_info = allocate_info.push_next(&mut flags);
1501        }
1502        // SAFETY: the device is live; the chained structures outlive the call.
1503        raw.memory = unsafe { device.allocate_memory(&allocate_info, None) }
1504            .map_err(|result| shared.fail(result))?;
1505        raw.allocation_bytes = requirements.size;
1506        // SAFETY: fresh buffer and memory; offset 0 satisfies every alignment requirement.
1507        unsafe { device.bind_buffer_memory(buffer, raw.memory, 0) }
1508            .map_err(|result| shared.fail(result))?;
1509
1510        let mut alignment = u64::MAX;
1511        if map {
1512            // SAFETY: the memory is host-visible (chosen by the memory plan), unmapped, and
1513            // mapping the whole allocation is always in range.
1514            let pointer = unsafe {
1515                device.map_memory(raw.memory, 0, vk::WHOLE_SIZE, vk::MemoryMapFlags::empty())
1516            }
1517            .map_err(|result| shared.fail(result))?;
1518            let pointer = NonNull::new(pointer.cast::<u8>()).ok_or(BackendError::Incompatible)?;
1519            raw.mapped = Some(pointer);
1520            alignment = alignment.min(address_alignment(pointer.as_ptr() as u64));
1521        }
1522        if shared.physical.buffer_device_address {
1523            let address_info = vk::BufferDeviceAddressInfo::default().buffer(buffer);
1524            // SAFETY: the feature is enabled, the buffer carries the device-address usage, and
1525            // its memory was allocated with the device-address flag.
1526            let address = unsafe { device.get_buffer_device_address(&address_info) };
1527            alignment = alignment.min(address_alignment(address));
1528        } else if !map {
1529            // Nothing observable to measure: the binding requirement is the only guarantee.
1530            alignment = alignment.min(requirements.alignment.max(1));
1531        }
1532        raw.measured_alignment = alignment;
1533        Ok(raw)
1534    }
1535
1536    /// Create a buffer over `len` bytes of caller memory at `pointer`, imported as a host
1537    /// allocation (`VK_EXT_external_memory_host`) into a host-coherent memory type, and bind it
1538    /// at offset 0. Nothing is mapped: the caller's pointer is the host's view.
1539    ///
1540    /// # Safety
1541    ///
1542    /// `pointer..pointer + len` is live host memory that stays allocated and mapped until the
1543    /// memory object is freed; `pointer` and `len` are multiples of the device's
1544    /// `minImportedHostPointerAlignment`.
1545    unsafe fn import(
1546        shared: &'a Shared,
1547        pointer: NonNull<u8>,
1548        len: u64,
1549    ) -> Result<Self, BackendError> {
1550        let device = &shared.device;
1551        let host = shared
1552            .host_memory
1553            .as_ref()
1554            .ok_or(BackendError::Unsupported)?;
1555        let handle_type = vk::ExternalMemoryHandleTypeFlags::HOST_ALLOCATION_EXT;
1556        let mut host_properties = vk::MemoryHostPointerPropertiesEXT::default();
1557        // SAFETY: the extension is enabled on this device, so the entry point is loaded; the
1558        // caller vouches for the pointer; the out-structure is a live local. `ash` 0.38 has no
1559        // wrapper for this command, so the raw pointer is called and its result checked.
1560        unsafe {
1561            (host.fp().get_memory_host_pointer_properties_ext)(
1562                host.device(),
1563                handle_type,
1564                pointer.as_ptr().cast(),
1565                &mut host_properties,
1566            )
1567        }
1568        .result()
1569        .map_err(|result| shared.fail(result))?;
1570        let mut external = vk::ExternalMemoryBufferCreateInfo::default().handle_types(handle_type);
1571        let buffer_info = vk::BufferCreateInfo::default()
1572            .size(len)
1573            .usage(buffer_usage(&shared.physical))
1574            .sharing_mode(vk::SharingMode::EXCLUSIVE)
1575            .push_next(&mut external);
1576        // SAFETY: the device is live and `buffer_info` and its chain are live locals.
1577        let buffer = unsafe { device.create_buffer(&buffer_info, None) }
1578            .map_err(|result| shared.fail(result))?;
1579        // SAFETY: `buffer` is live.
1580        let requirements = unsafe { device.get_buffer_memory_requirements(buffer) };
1581        let coherent =
1582            vk::MemoryPropertyFlags::HOST_VISIBLE | vk::MemoryPropertyFlags::HOST_COHERENT;
1583        let memory = &shared.physical.memory;
1584        // A host-coherent type both the pointer and the buffer allow, device-local first.
1585        let memory_type = (0..memory.memory_type_count)
1586            .filter(|&index| {
1587                host_properties.memory_type_bits & requirements.memory_type_bits & (1 << index) != 0
1588                    && memory.memory_types[index as usize]
1589                        .property_flags
1590                        .contains(coherent)
1591            })
1592            .min_by_key(|&index| {
1593                !memory.memory_types[index as usize]
1594                    .property_flags
1595                    .contains(vk::MemoryPropertyFlags::DEVICE_LOCAL)
1596            });
1597        let mut raw = Self {
1598            shared,
1599            buffer,
1600            memory: vk::DeviceMemory::null(),
1601            mapped: None,
1602            allocation_bytes: 0,
1603            measured_alignment: 0,
1604            memory_flags: vk::MemoryPropertyFlags::empty(),
1605        };
1606        let memory_type = memory_type.ok_or(BackendError::Incompatible)?;
1607        raw.memory_flags = memory.memory_types[memory_type as usize].property_flags;
1608        if requirements.size > len {
1609            return Err(BackendError::ResourceLimit);
1610        }
1611        let mut import = vk::ImportMemoryHostPointerInfoEXT::default()
1612            .handle_type(handle_type)
1613            .host_pointer(pointer.as_ptr().cast());
1614        let mut flags =
1615            vk::MemoryAllocateFlagsInfo::default().flags(vk::MemoryAllocateFlags::DEVICE_ADDRESS);
1616        let mut allocate_info = vk::MemoryAllocateInfo::default()
1617            .allocation_size(len)
1618            .memory_type_index(memory_type)
1619            .push_next(&mut import);
1620        if shared.physical.buffer_device_address {
1621            allocate_info = allocate_info.push_next(&mut flags);
1622        }
1623        // SAFETY: the device is live; the chained structures outlive the call; the caller
1624        // vouches that the range is live, aligned host memory, and the type admits the pointer.
1625        raw.memory = unsafe { device.allocate_memory(&allocate_info, None) }
1626            .map_err(|result| shared.fail(result))?;
1627        raw.allocation_bytes = len;
1628        // SAFETY: fresh buffer and memory; offset 0 satisfies every alignment requirement.
1629        unsafe { device.bind_buffer_memory(buffer, raw.memory, 0) }
1630            .map_err(|result| shared.fail(result))?;
1631        let mut alignment = address_alignment(pointer.as_ptr() as u64);
1632        if shared.physical.buffer_device_address {
1633            let address_info = vk::BufferDeviceAddressInfo::default().buffer(buffer);
1634            // SAFETY: the feature is enabled, the buffer carries the device-address usage, and
1635            // its memory was allocated with the device-address flag.
1636            let address = unsafe { device.get_buffer_device_address(&address_info) };
1637            alignment = alignment.min(address_alignment(address));
1638        }
1639        raw.measured_alignment = alignment;
1640        Ok(raw)
1641    }
1642
1643    /// Transfer ownership of the handles to a `VulkanBuffer`.
1644    fn into_parts(self) -> (vk::Buffer, vk::DeviceMemory, Option<NonNull<u8>>) {
1645        let parts = (self.buffer, self.memory, self.mapped);
1646        std::mem::forget(self);
1647        parts
1648    }
1649}
1650
1651impl Drop for RawAllocation<'_> {
1652    fn drop(&mut self) {
1653        let device = &self.shared.device;
1654        // SAFETY: the handles were created by `create` and are released exactly once; null
1655        // memory (allocation failed) is skipped by the loader-defined null-handle rule.
1656        unsafe {
1657            if self.mapped.is_some() {
1658                device.unmap_memory(self.memory);
1659            }
1660            device.destroy_buffer(self.buffer, None);
1661            if self.memory != vk::DeviceMemory::null() {
1662                device.free_memory(self.memory, None);
1663            }
1664        }
1665    }
1666}
1667
1668/// The largest power of two dividing `address`, capped so a zero address does not overflow.
1669fn address_alignment(address: u64) -> u64 {
1670    1_u64 << address.trailing_zeros().min(40)
1671}
1672
1673/// Bounded host-visible staging buffer for device-local transfers; one per explicit transfer.
1674struct Staging<'a> {
1675    raw: RawAllocation<'a>,
1676    bytes: u64,
1677}
1678
1679impl<'a> Staging<'a> {
1680    fn new(shared: &'a Shared, bytes: u64) -> Result<Self, BackendError> {
1681        let bytes = bytes
1682            .max(WORD_BYTES)
1683            .checked_add(WORD_BYTES - 1)
1684            .ok_or(BackendError::ResourceLimit)?
1685            / WORD_BYTES
1686            * WORD_BYTES;
1687        let raw = RawAllocation::create(shared, bytes, shared.memory_plan.host, true)?;
1688        Ok(Self { raw, bytes })
1689    }
1690
1691    fn as_mut_slice(&mut self) -> &mut [u8] {
1692        let pointer = self.raw.mapped.expect("staging is mapped");
1693        // SAFETY: the mapping covers `bytes` bytes of host-coherent memory owned by `self`; the
1694        // GPU never accesses it while this borrow is live (every copy is waited for).
1695        unsafe { std::slice::from_raw_parts_mut(pointer.as_ptr(), self.bytes as usize) }
1696    }
1697}
1698
1699/// Program-side in-flight count: pipelines stay alive until every submission using them retired.
1700#[derive(Default)]
1701struct ProgramState {
1702    in_flight: Cell<u32>,
1703}
1704
1705/// One recorded dispatch of a resident program.
1706struct Dispatch {
1707    pipeline: vk::Pipeline,
1708    workgroups: [u32; 3],
1709    barrier_before: bool,
1710}
1711
1712/// The program-owned arena: constants and intermediates in one dedicated allocation, bound as
1713/// the last element of the descriptor array. Never mapped; constants arrive through staging.
1714struct Arena {
1715    buffer: vk::Buffer,
1716    memory: vk::DeviceMemory,
1717    bytes: u64,
1718    /// The persistent mapping when the arena's memory type is host-visible (a unified-memory
1719    /// device): constants are then written by a plain copy rather than through staging.
1720    mapped: Option<NonNull<u8>>,
1721}
1722
1723/// Resident compute pipelines specialized for one admitted TOSA graph.
1724pub struct VulkanProgram {
1725    context: Rc<ContextInner>,
1726    dispatches: Vec<Dispatch>,
1727    arena: Option<Arena>,
1728    plan: ProgramPlan,
1729    state: Rc<ProgramState>,
1730}
1731
1732impl std::fmt::Debug for VulkanProgram {
1733    fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
1734        formatter
1735            .debug_struct("VulkanProgram")
1736            .field("context", &self.context.id)
1737            .field("dispatches", &self.dispatches.len())
1738            .field("arena_bytes", &self.plan.arena_bytes)
1739            .field("slots", &self.plan.slots.len())
1740            .finish_non_exhaustive()
1741    }
1742}
1743
1744impl VulkanProgram {
1745    /// Number of `vkCmdDispatch` calls one submission of this program records.
1746    pub fn dispatch_count(&self) -> usize {
1747        self.dispatches.len()
1748    }
1749
1750    /// Bytes of program-owned arena storage (constants plus intermediates).
1751    pub fn arena_bytes(&self) -> u64 {
1752        self.plan.arena_bytes
1753    }
1754}
1755
1756impl Drop for VulkanProgram {
1757    fn drop(&mut self) {
1758        let shared = &self.context.shared;
1759        if self.state.in_flight.get() != 0 {
1760            shared.wait_idle();
1761        }
1762        // SAFETY: this handle owns every pipeline and the arena, created in `load_program` and
1763        // destroyed exactly once here; no submission references them (in-flight count is zero
1764        // or the device was idled above).
1765        unsafe {
1766            for dispatch in &self.dispatches {
1767                shared.device.destroy_pipeline(dispatch.pipeline, None);
1768            }
1769            if let Some(arena) = &self.arena {
1770                shared.device.destroy_buffer(arena.buffer, None);
1771                shared.device.free_memory(arena.memory, None);
1772            }
1773        }
1774        decrement(&shared.counters.programs);
1775    }
1776}
1777
1778/// Vulkan execution queue handle. Every queue of a context feeds the device's one compute queue.
1779pub struct VulkanQueue {
1780    context: Rc<ContextInner>,
1781}
1782
1783impl std::fmt::Debug for VulkanQueue {
1784    fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
1785        formatter
1786            .debug_struct("VulkanQueue")
1787            .field("context", &self.context.id)
1788            .finish_non_exhaustive()
1789    }
1790}
1791
1792impl Drop for VulkanQueue {
1793    fn drop(&mut self) {
1794        decrement(&self.context.shared.counters.queues);
1795    }
1796}
1797
1798/// In-flight guard for one buffer bound to one submission.
1799struct Guard {
1800    state: Rc<BufferState>,
1801    exclusive: bool,
1802}
1803
1804impl Guard {
1805    fn acquire(state: &Rc<BufferState>, exclusive: bool) -> Result<Self, BackendError> {
1806        let current = state.in_flight.get();
1807        let next = if exclusive {
1808            if current != 0 {
1809                return Err(BackendError::Busy);
1810            }
1811            EXCLUSIVE_ACCESS
1812        } else {
1813            if current >= EXCLUSIVE_ACCESS - 1 {
1814                return Err(BackendError::Busy);
1815            }
1816            current + 1
1817        };
1818        state.in_flight.set(next);
1819        Ok(Self {
1820            state: Rc::clone(state),
1821            exclusive,
1822        })
1823    }
1824}
1825
1826impl Drop for Guard {
1827    fn drop(&mut self) {
1828        let current = self.state.in_flight.get();
1829        self.state.in_flight.set(if self.exclusive {
1830            debug_assert_eq!(current, EXCLUSIVE_ACCESS);
1831            0
1832        } else {
1833            debug_assert!((1..EXCLUSIVE_ACCESS).contains(&current));
1834            current - 1
1835        });
1836    }
1837}
1838
1839/// One submission: a claimed ring slot, its fence, and the guards it holds until terminal.
1840pub struct VulkanEvent {
1841    context: Rc<ContextInner>,
1842    slot: u16,
1843    program: Rc<ProgramState>,
1844    guards: RefCell<[Option<Guard>; MAX_BINDINGS_PER_SUBMISSION as usize]>,
1845    latched: Cell<Option<EventState>>,
1846    /// Set once the slot was returned to the ring (by `destroy_event` or `Drop`).
1847    released: Cell<bool>,
1848}
1849
1850impl std::fmt::Debug for VulkanEvent {
1851    fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
1852        formatter
1853            .debug_struct("VulkanEvent")
1854            .field("context", &self.context.id)
1855            .field("slot", &self.slot)
1856            .field("latched", &self.latched.get())
1857            .finish_non_exhaustive()
1858    }
1859}
1860
1861impl VulkanEvent {
1862    /// Publish the first terminal state. Guards are released strictly before the latch becomes
1863    /// observable so a caller seeing a terminal state can transfer buffer bytes immediately.
1864    fn latch(&self, state: EventState) -> EventState {
1865        for guard in self.guards.borrow_mut().iter_mut() {
1866            *guard = None;
1867        }
1868        let in_flight = self.program.in_flight.get();
1869        self.program.in_flight.set(in_flight.saturating_sub(1));
1870        self.latched.set(Some(state));
1871        state
1872    }
1873
1874    /// Nonblocking status read of the slot's fence (ADR 0006).
1875    fn poll(&self) -> Result<EventState, BackendError> {
1876        if let Some(state) = self.latched.get() {
1877            return Ok(state);
1878        }
1879        let shared = &self.context.shared;
1880        let fence = self.context.slots[self.slot as usize].fence;
1881        // SAFETY: the fence belongs to this event's claimed slot and was submitted exactly once
1882        // since its last reset; `vkGetFenceStatus` is a read-only query.
1883        match unsafe { shared.device.get_fence_status(fence) } {
1884            Ok(true) => Ok(self.latch(EventState::Complete)),
1885            Ok(false) => Ok(EventState::Pending),
1886            Err(vk::Result::ERROR_DEVICE_LOST) => {
1887                shared.poisoned.set(true);
1888                Ok(self.latch(EventState::Failed(BackendError::DeviceLost)))
1889            }
1890            Err(result) => Err(shared.fail(result)),
1891        }
1892    }
1893
1894    fn release(&self) {
1895        if self.released.replace(true) {
1896            return;
1897        }
1898        self.context.release_slot(self.slot);
1899        decrement(&self.context.shared.counters.events);
1900    }
1901}
1902
1903impl Drop for VulkanEvent {
1904    fn drop(&mut self) {
1905        if self.released.get() {
1906            return;
1907        }
1908        // Dropping a pending event outside `destroy_event` is a contract violation; still, never
1909        // return a slot whose command buffer may be executing: wait for its fence first.
1910        if self.latched.get().is_none() {
1911            let shared = &self.context.shared;
1912            let fence = self.context.slots[self.slot as usize].fence;
1913            // SAFETY: the fence is this slot's, submitted once; waiting has no preconditions.
1914            match unsafe { shared.device.wait_for_fences(&[fence], true, u64::MAX) } {
1915                Ok(()) => {
1916                    self.latch(EventState::Complete);
1917                }
1918                Err(result) => {
1919                    shared.fail(result);
1920                    self.latch(EventState::Failed(BackendError::DeviceLost));
1921                }
1922            }
1923        }
1924        self.release();
1925    }
1926}
1927
1928/// Vulkan backend instance bound to one physical device.
1929/// A timeline semaphore the host raises, which [`VulkanAccelerator::submit_after`] submissions
1930/// wait on (ADR 0013): work is queued before its inputs exist, and whoever produces them (a
1931/// thread finishing a storage read, say) releases it with a [`VulkanGateSignal`], without a
1932/// round trip through the thread that owns the backend.
1933///
1934/// Dropping the gate releases every submission still waiting on it, then waits for the device
1935/// to idle before destroying the semaphore, so no submission can wait forever.
1936pub struct VulkanHostGate {
1937    shared: Rc<Shared>,
1938    core: Arc<GateCore>,
1939    /// The highest value any submission was queued to wait for.
1940    awaited: Cell<u64>,
1941}
1942
1943struct GateCore {
1944    device: ash::Device,
1945    semaphore: vk::Semaphore,
1946    state: Mutex<GateState>,
1947}
1948
1949struct GateState {
1950    /// Cleared, under the lock, before the semaphore is destroyed.
1951    open: bool,
1952    raised: u64,
1953}
1954
1955/// The raising half of a [`VulkanHostGate`], for any thread.
1956#[derive(Clone)]
1957pub struct VulkanGateSignal {
1958    core: Arc<GateCore>,
1959}
1960
1961impl std::fmt::Debug for VulkanHostGate {
1962    fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
1963        formatter
1964            .debug_struct("VulkanHostGate")
1965            .field("awaited", &self.awaited.get())
1966            .finish_non_exhaustive()
1967    }
1968}
1969
1970impl std::fmt::Debug for VulkanGateSignal {
1971    fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
1972        formatter
1973            .debug_struct("VulkanGateSignal")
1974            .finish_non_exhaustive()
1975    }
1976}
1977
1978impl VulkanGateSignal {
1979    /// Raise the gate to `value`, releasing every submission waiting for it or less. A value at
1980    /// or below the gate's current one changes nothing, so raises may arrive in any order.
1981    /// Returns `false`, raising nothing, once the gate has been dropped.
1982    pub fn raise(&self, value: u64) -> Result<bool, BackendError> {
1983        // The state is two plain fields written together; a panic elsewhere cannot tear it.
1984        let mut state = self
1985            .core
1986            .state
1987            .lock()
1988            .unwrap_or_else(|poison| poison.into_inner());
1989        if !state.open {
1990            return Ok(false);
1991        }
1992        if value > state.raised {
1993            let info = vk::SemaphoreSignalInfo::default()
1994                .semaphore(self.core.semaphore)
1995                .value(value);
1996            // SAFETY: the gate is open, so its device and semaphore are live: the gate closes
1997            // under this lock before destroying either. The value exceeds the semaphore's current
1998            // one (`raised` tracks every signal, and only this path signals).
1999            unsafe { self.core.device.signal_semaphore(&info) }.map_err(backend_error)?;
2000            state.raised = value;
2001        }
2002        Ok(true)
2003    }
2004}
2005
2006impl VulkanHostGate {
2007    /// A raising handle for another thread.
2008    pub fn signal(&self) -> VulkanGateSignal {
2009        VulkanGateSignal {
2010            core: Arc::clone(&self.core),
2011        }
2012    }
2013}
2014
2015impl Drop for VulkanHostGate {
2016    fn drop(&mut self) {
2017        let mut state = self
2018            .core
2019            .state
2020            .lock()
2021            .unwrap_or_else(|poison| poison.into_inner());
2022        let awaited = self.awaited.get();
2023        if awaited > state.raised {
2024            let info = vk::SemaphoreSignalInfo::default()
2025                .semaphore(self.core.semaphore)
2026                .value(awaited);
2027            // SAFETY: the gate is still open and the value exceeds the current one. Releasing
2028            // the waiters is what lets the idle wait below return.
2029            if unsafe { self.core.device.signal_semaphore(&info) }.is_ok() {
2030                state.raised = awaited;
2031            }
2032        }
2033        state.open = false;
2034        drop(state);
2035        self.shared.wait_idle();
2036        // SAFETY: no submission still references the semaphore (the device is idle), no raise
2037        // can reach it (closed under the lock above), and it is destroyed exactly once.
2038        unsafe {
2039            self.shared
2040                .device
2041                .destroy_semaphore(self.core.semaphore, None)
2042        };
2043    }
2044}
2045
2046pub struct VulkanAccelerator {
2047    shared: Rc<Shared>,
2048    next_id: Cell<u64>,
2049}
2050
2051impl std::fmt::Debug for VulkanAccelerator {
2052    fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
2053        formatter
2054            .debug_struct("VulkanAccelerator")
2055            .field("device", &self.device_name())
2056            .finish_non_exhaustive()
2057    }
2058}
2059
2060impl VulkanAccelerator {
2061    /// Open the preferred Vulkan 1.3 compute device: discrete, integrated, virtual, then CPU.
2062    pub fn new() -> Result<Self, InitError> {
2063        let instance = Instance::create()?;
2064        let devices = enumerate(&instance)?;
2065        let physical = devices
2066            .into_iter()
2067            .min_by_key(PhysicalDeviceRecord::rank)
2068            .ok_or(InitError::DeviceUnavailable)?;
2069        Self::open(instance, physical, VulkanOptions::default())
2070    }
2071
2072    /// Open the device whose enumerated name (`available_devices`) equals `device`.
2073    pub fn with_device(device: &str) -> Result<Self, InitError> {
2074        Self::with_device_options(device, VulkanOptions::default())
2075    }
2076
2077    /// [`with_device`](Self::with_device) with explicit [`VulkanOptions`].
2078    pub fn with_device_options(device: &str, options: VulkanOptions) -> Result<Self, InitError> {
2079        let instance = Instance::create()?;
2080        let physical = enumerate(&instance)?
2081            .into_iter()
2082            .find(|record| record.name == device)
2083            .ok_or(InitError::DeviceUnavailable)?;
2084        Self::open(instance, physical, options)
2085    }
2086
2087    /// Enumerate the names of every suitable device visible through the loader.
2088    pub fn available_devices() -> Result<Vec<String>, InitError> {
2089        let instance = Instance::create()?;
2090        Ok(enumerate(&instance)?
2091            .into_iter()
2092            .map(|record| record.name)
2093            .collect())
2094    }
2095
2096    fn open(
2097        instance: Instance,
2098        physical: PhysicalDeviceRecord,
2099        options: VulkanOptions,
2100    ) -> Result<Self, InitError> {
2101        Ok(Self {
2102            shared: Shared::open(instance, physical, options)?,
2103            next_id: Cell::new(0),
2104        })
2105    }
2106
2107    /// The enumerated name of the device this instance executes on.
2108    pub fn device_name(&self) -> &str {
2109        &self.shared.physical.name
2110    }
2111
2112    /// Whether this instance observed device loss and refuses further work.
2113    pub fn is_poisoned(&self) -> bool {
2114        self.shared.poisoned.get()
2115    }
2116
2117    /// A new [`VulkanHostGate`] at value 0, or `Unsupported` without timeline semaphores.
2118    pub fn host_gate(&self) -> Result<VulkanHostGate, BackendError> {
2119        let shared = &self.shared;
2120        shared.check_live()?;
2121        if !shared.physical.timeline_semaphore {
2122            return Err(BackendError::Unsupported);
2123        }
2124        let mut timeline = vk::SemaphoreTypeCreateInfo::default()
2125            .semaphore_type(vk::SemaphoreType::TIMELINE)
2126            .initial_value(0);
2127        let info = vk::SemaphoreCreateInfo::default().push_next(&mut timeline);
2128        // SAFETY: the device is live with `timelineSemaphore` enabled; `info` is a live local.
2129        let semaphore = unsafe { shared.device.create_semaphore(&info, None) }
2130            .map_err(|result| shared.fail(result))?;
2131        Ok(VulkanHostGate {
2132            shared: Rc::clone(shared),
2133            core: Arc::new(GateCore {
2134                device: shared.device.clone(),
2135                semaphore,
2136                state: Mutex::new(GateState {
2137                    open: true,
2138                    raised: 0,
2139                }),
2140            }),
2141            awaited: Cell::new(0),
2142        })
2143    }
2144
2145    /// [`Accelerator::submit`], the work waiting on the device until `gate` reaches `value`
2146    /// (ADR 0013). The bindings are held from now, as `submit` holds them, with one difference
2147    /// the gate exists for: the bytes of the program's *input* bindings may still be written
2148    /// until the gate is raised to `value` (by host stores into an imported buffer, or by a
2149    /// device's DMA the host has seen complete), because raising the gate orders every host
2150    /// operation before it ahead of the device's wait. Outputs, and inputs after the raise,
2151    /// follow `submit`'s rules.
2152    // The result type is `Accelerator::submit`'s, whose event travels in the failure.
2153    #[allow(clippy::result_large_err)]
2154    pub fn submit_after(
2155        &self,
2156        queue: &VulkanQueue,
2157        program: &VulkanProgram,
2158        bindings: &[BindingRef<'_, VulkanBuffer>],
2159        gate: &VulkanHostGate,
2160        value: u64,
2161    ) -> Result<VulkanEvent, SubmitFailure<VulkanEvent>> {
2162        if !Rc::ptr_eq(&gate.shared, &self.shared) {
2163            return Err(SubmitFailure::Rejected(BackendError::InvalidArgument));
2164        }
2165        gate.awaited.set(gate.awaited.get().max(value));
2166        let wait = Some((gate.core.semaphore, value));
2167        self.submit_waiting(queue, program, bindings, Timeout::Infinite, wait)
2168    }
2169
2170    /// The alignment [`import_host_buffer`](Self::import_host_buffer) requires of a pointer and a
2171    /// length, or `None` when the device cannot import host memory.
2172    pub fn host_import_alignment(&self) -> Option<u64> {
2173        self.shared.physical.host_import_alignment
2174    }
2175
2176    /// A buffer over caller-owned host memory, imported rather than copied (ADR 0013): the device
2177    /// addresses `memory..memory + len` itself, so bytes placed there by the host, or by another
2178    /// device's DMA, are the buffer's contents with no `write_buffer`. Behaves as an allocated
2179    /// buffer of `desc` in every other respect: bound by submissions, gated while in flight, read
2180    /// and written by the explicit transfers, released by `free_buffer`. Its domain must be
2181    /// `Host` or `Shared`, the domains whose contract is host-visible memory.
2182    ///
2183    /// This is a host-side API of this backend, not a protocol feature: the protocol's external
2184    /// memory import remains deferred.
2185    ///
2186    /// # Safety
2187    ///
2188    /// `memory..memory + len` must be live host memory (anonymous or huge-page mappings; not a
2189    /// device mapping) that stays allocated and mapped, and is not remapped, until the buffer is
2190    /// released and no submission that bound it is still executing. While a submission that binds
2191    /// the buffer is in flight, the host must not write bytes that submission reads or read bytes
2192    /// it writes.
2193    pub unsafe fn import_host_buffer(
2194        &self,
2195        context: &VulkanContext,
2196        desc: BufferDesc,
2197        memory: NonNull<u8>,
2198        len: u64,
2199    ) -> Result<AllocatedBuffer<VulkanBuffer>, BackendError> {
2200        let shared = &self.shared;
2201        shared.info.validate_buffer_desc(desc)?;
2202        shared.check_live()?;
2203        let alignment = shared
2204            .physical
2205            .host_import_alignment
2206            .ok_or(BackendError::Unsupported)?;
2207        if !matches!(desc.domain, MemoryDomain::Host | MemoryDomain::Shared) {
2208            return Err(BackendError::Unsupported);
2209        }
2210        if memory.as_ptr() as u64 % alignment != 0 || len % alignment != 0 || desc.bytes() > len {
2211            return Err(BackendError::InvalidArgument);
2212        }
2213        if shared.counters.buffers.get()
2214            >= u64::from(MAX_BUFFERS_PER_CONTEXT) * u64::from(MAX_CONTEXTS)
2215        {
2216            return Err(BackendError::ResourceLimit);
2217        }
2218        // SAFETY: the caller vouches for the range; alignment was checked above.
2219        let raw = unsafe { RawAllocation::import(shared, memory, len) }?;
2220        if raw.measured_alignment < desc.alignment() {
2221            return Err(BackendError::ResourceLimit);
2222        }
2223        let mut properties = BufferProperties::DIRECT_BINDING | BufferProperties::HOST_VISIBLE;
2224        if raw
2225            .memory_flags
2226            .contains(vk::MemoryPropertyFlags::DEVICE_LOCAL)
2227        {
2228            properties |= BufferProperties::DEVICE_LOCAL;
2229        }
2230        let info = BufferInfo::new(
2231            desc,
2232            raw.allocation_bytes,
2233            raw.measured_alignment,
2234            properties,
2235        )?;
2236        let (buffer, device_memory, _) = raw.into_parts();
2237        increment(&shared.counters.buffers, 1);
2238        Ok(AllocatedBuffer::new(
2239            VulkanBuffer {
2240                context: Rc::clone(&context.inner),
2241                desc,
2242                buffer,
2243                memory: device_memory,
2244                mapped: Some(memory),
2245                imported: true,
2246                state: Rc::new(BufferState::default()),
2247            },
2248            info,
2249        ))
2250    }
2251
2252    /// Cumulative count of buffers admitted as direct bindings.
2253    pub fn direct_binding_admissions(&self) -> u64 {
2254        self.shared.counters.direct_binding_admissions.get()
2255    }
2256
2257    /// Cumulative bytes moved by explicit `write_buffer`/`read_buffer` transfers.
2258    pub fn explicit_transfer_bytes(&self) -> u64 {
2259        self.shared.counters.explicit_transfer_bytes.get()
2260    }
2261
2262    /// Provider handles currently alive for this instance.
2263    pub fn live_resources(&self) -> LiveResources {
2264        let counters = &self.shared.counters;
2265        LiveResources {
2266            contexts: counters.contexts.get(),
2267            buffers: counters.buffers.get(),
2268            programs: counters.programs.get(),
2269            queues: counters.queues.get(),
2270            events: counters.events.get(),
2271        }
2272    }
2273
2274    fn next_id(&self) -> Result<u64, BackendError> {
2275        let id = self.next_id.get();
2276        if id == u64::MAX {
2277            return Err(BackendError::ResourceLimit);
2278        }
2279        self.next_id.set(id + 1);
2280        Ok(id)
2281    }
2282
2283    fn checked_range(
2284        buffer: &VulkanBuffer,
2285        offset: u64,
2286        bytes: u64,
2287    ) -> Result<(usize, usize), BackendError> {
2288        if bytes == 0 {
2289            return Err(BackendError::InvalidArgument);
2290        }
2291        let end = offset
2292            .checked_add(bytes)
2293            .filter(|end| *end <= buffer.desc.bytes())
2294            .ok_or(BackendError::OutOfBounds)?;
2295        let start = usize::try_from(offset).map_err(|_| BackendError::OutOfBounds)?;
2296        let end = usize::try_from(end).map_err(|_| BackendError::OutOfBounds)?;
2297        Ok((start, end))
2298    }
2299
2300    fn lowering_error(error: LoweringError) -> BackendError {
2301        match error {
2302            LoweringError::Parse(_) | LoweringError::Analysis(_) => BackendError::InvalidArgument,
2303            LoweringError::UnsupportedTarget => BackendError::Incompatible,
2304            LoweringError::UnsupportedGraph
2305            | LoweringError::UnsupportedType(_)
2306            | LoweringError::UnsupportedOperator(_) => BackendError::Unsupported,
2307            LoweringError::ResourceLimit => BackendError::ResourceLimit,
2308        }
2309    }
2310
2311    /// Write into a device-local buffer through a bounded staging allocation.
2312    fn staged_write(
2313        &self,
2314        buffer: &VulkanBuffer,
2315        start: u64,
2316        data: &dyn ByteSource,
2317        len: u64,
2318    ) -> Result<(), BackendError> {
2319        let shared = &self.shared;
2320        let mut staging = Staging::new(shared, len.min(STAGING_BYTES))?;
2321        write_through_staging(
2322            &buffer.context,
2323            buffer.buffer,
2324            start,
2325            data,
2326            len,
2327            &mut staging,
2328        )?;
2329        increment(&shared.counters.explicit_transfer_bytes, len);
2330        Ok(())
2331    }
2332
2333    /// Read out of a device-local buffer through a bounded staging allocation.
2334    fn staged_read(
2335        &self,
2336        buffer: &VulkanBuffer,
2337        start: u64,
2338        data: &mut dyn ByteSink,
2339        len: u64,
2340    ) -> Result<(), BackendError> {
2341        let shared = &self.shared;
2342        let mut staging = Staging::new(shared, len.min(STAGING_BYTES))?;
2343        read_through_staging(
2344            &buffer.context,
2345            buffer.buffer,
2346            start,
2347            data,
2348            len,
2349            &mut staging,
2350        )?;
2351        increment(&shared.counters.explicit_transfer_bytes, len);
2352        Ok(())
2353    }
2354
2355    /// `Accelerator::submit`, the batch first waiting for `wait`'s semaphore to reach its value
2356    /// when one is given (ADR 0013).
2357    // The result type is `Accelerator::submit`'s, whose event travels in the failure.
2358    #[allow(clippy::result_large_err)]
2359    fn submit_waiting(
2360        &self,
2361        queue: &VulkanQueue,
2362        program: &VulkanProgram,
2363        bindings: &[BindingRef<'_, VulkanBuffer>],
2364        timeout: Timeout,
2365        wait: Option<(vk::Semaphore, u64)>,
2366    ) -> Result<VulkanEvent, SubmitFailure<VulkanEvent>> {
2367        let shared = &self.shared;
2368        let reject = SubmitFailure::Rejected;
2369        shared.check_live().map_err(reject)?;
2370        // Vulkan has no cancel primitive, so a finite deadline is refused before admission rather
2371        // than latched against retained resources (ADR 0006).
2372        if let Timeout::AfterNs(_) = timeout {
2373            return Err(reject(BackendError::DeadlineExpired));
2374        }
2375        if bindings.is_empty() || bindings.len() > shared.physical.tuning.max_bindings() as usize {
2376            return Err(reject(BackendError::ResourceLimit));
2377        }
2378        if !Rc::ptr_eq(&queue.context, &program.context) {
2379            return Err(reject(BackendError::InvalidArgument));
2380        }
2381        let context = &queue.context;
2382        let plan = &program.plan;
2383
2384        // Per-binding reasons (bounds, access, slot) are reported before the aggregate count
2385        // check so a host learns the most specific rejection first.
2386        let offset_alignment = shared
2387            .physical
2388            .limits
2389            .min_storage_buffer_offset_alignment
2390            .max(WORD_BYTES);
2391        let tuning = shared.physical.tuning;
2392        let mut descriptors = vec![vk::DescriptorBufferInfo::default(); tuning.buffers as usize];
2393        let mut seen = 0_u32;
2394        for binding in bindings {
2395            if !Rc::ptr_eq(&binding.buffer.context, context) {
2396                return Err(reject(BackendError::InvalidArgument));
2397            }
2398            if !binding.buffer.desc.allows_access(binding.access) {
2399                return Err(reject(BackendError::PermissionDenied));
2400            }
2401            let (start, _) =
2402                Self::checked_range(binding.buffer, binding.range.offset, binding.range.bytes())
2403                    .map_err(reject)?;
2404            let index = plan
2405                .slots
2406                .iter()
2407                .position(|slot| slot.slot == binding.slot)
2408                .ok_or(reject(BackendError::Incompatible))?;
2409            if seen & (1 << index) != 0 {
2410                return Err(reject(BackendError::InvalidArgument));
2411            }
2412            seen |= 1 << index;
2413            let slot_plan = &plan.slots[index];
2414            let expected_access = match slot_plan.role {
2415                SlotRole::Input => AccessMode::Read,
2416                SlotRole::Output => AccessMode::Write,
2417            };
2418            if binding.access != expected_access {
2419                return Err(reject(BackendError::Incompatible));
2420            }
2421            // The descriptor covers the range directly: exact tensor bytes, word- and
2422            // `minStorageBufferOffsetAlignment`-aligned start. Byte-storage tensors are
2423            // addressed by whole words, so their descriptor range extends to the containing
2424            // word; the allocation behind every buffer is word-sized so that word exists.
2425            if binding.range.bytes() != slot_plan.byte_len || (start as u64) % offset_alignment != 0
2426            {
2427                return Err(reject(BackendError::Incompatible));
2428            }
2429            descriptors[index] = vk::DescriptorBufferInfo {
2430                buffer: binding.buffer.buffer,
2431                offset: start as u64,
2432                range: binding.range.bytes().div_ceil(WORD_BYTES) * WORD_BYTES,
2433            };
2434        }
2435        if bindings.len() != plan.slots.len() {
2436            return Err(reject(BackendError::Incompatible));
2437        }
2438        // The arena follows the slots; every element the program never addresses is filled
2439        // with the first bound buffer so the whole array is valid.
2440        if let Some(arena) = &program.arena {
2441            descriptors[plan.arena_buffer_index() as usize] = vk::DescriptorBufferInfo {
2442                buffer: arena.buffer,
2443                offset: 0,
2444                range: arena.bytes,
2445            };
2446        }
2447        let filler = descriptors[0];
2448        for descriptor in &mut descriptors[plan.buffer_count() as usize..] {
2449            *descriptor = filler;
2450        }
2451        // A TOSA graph's inputs and outputs are distinct tensors, so one allocation may back
2452        // several read-only slots but never a written slot together with any other slot: that
2453        // aliasing is a program incompatibility, reported here rather than as a transient `Busy`
2454        // from the in-flight gates below.
2455        for (index, binding) in bindings.iter().enumerate() {
2456            let aliased = bindings[..index].iter().any(|prior| {
2457                Rc::ptr_eq(&prior.buffer.state, &binding.buffer.state)
2458                    && (prior.access != AccessMode::Read || binding.access != AccessMode::Read)
2459            });
2460            if aliased {
2461                return Err(reject(BackendError::Incompatible));
2462            }
2463        }
2464
2465        let slot_index = context
2466            .claim_slot()
2467            .ok_or(reject(BackendError::ResourceLimit))?;
2468        let mut guards: [Option<Guard>; MAX_BINDINGS_PER_SUBMISSION as usize] =
2469            [const { None }; MAX_BINDINGS_PER_SUBMISSION as usize];
2470        for (index, binding) in bindings.iter().enumerate() {
2471            let exclusive = binding.access != AccessMode::Read;
2472            // Aliasing with a written slot was rejected above, so a conflict here can only come
2473            // from another in-flight submission: a transient `Busy`.
2474            match Guard::acquire(&binding.buffer.state, exclusive) {
2475                Ok(guard) => guards[index] = Some(guard),
2476                Err(error) => {
2477                    drop(guards);
2478                    context.release_slot(slot_index);
2479                    return Err(reject(error));
2480                }
2481            }
2482        }
2483
2484        let slot = &context.slots[slot_index as usize];
2485        match self.record_and_submit(slot, program, &descriptors, wait) {
2486            Ok(()) => {}
2487            Err(vk::Result::ERROR_DEVICE_LOST) => {
2488                // Past the admission boundary with an ambiguous outcome: the event owns the slot
2489                // and latches the loss; the instance is poisoned (ADR 0006).
2490                shared.poisoned.set(true);
2491                program
2492                    .state
2493                    .in_flight
2494                    .set(program.state.in_flight.get() + 1);
2495                increment(&shared.counters.events, 1);
2496                let event = VulkanEvent {
2497                    context: Rc::clone(context),
2498                    slot: slot_index,
2499                    program: Rc::clone(&program.state),
2500                    guards: RefCell::new(guards),
2501                    latched: Cell::new(None),
2502                    released: Cell::new(false),
2503                };
2504                event.latch(EventState::Failed(BackendError::DeviceLost));
2505                return Err(SubmitFailure::Indeterminate {
2506                    error: BackendError::DeviceLost,
2507                    event,
2508                });
2509            }
2510            Err(result) => {
2511                // Recording and submission failures before the queue accepted the work leave
2512                // every resource untouched (Vulkan guarantees this for out-of-memory results).
2513                drop(guards);
2514                context.release_slot(slot_index);
2515                return Err(reject(shared.fail(result)));
2516            }
2517        }
2518        program
2519            .state
2520            .in_flight
2521            .set(program.state.in_flight.get() + 1);
2522        increment(
2523            &shared.counters.direct_binding_admissions,
2524            bindings.len() as u64,
2525        );
2526        increment(&shared.counters.events, 1);
2527        Ok(VulkanEvent {
2528            context: Rc::clone(context),
2529            slot: slot_index,
2530            program: Rc::clone(&program.state),
2531            guards: RefCell::new(guards),
2532            latched: Cell::new(None),
2533            released: Cell::new(false),
2534        })
2535    }
2536
2537    /// Record every dispatch of `program` for one claimed slot and submit it with the slot's
2538    /// fence.
2539    fn record_and_submit(
2540        &self,
2541        slot: &Slot,
2542        program: &VulkanProgram,
2543        descriptors: &[vk::DescriptorBufferInfo],
2544        wait: Option<(vk::Semaphore, u64)>,
2545    ) -> Result<(), vk::Result> {
2546        let shared = &self.shared;
2547        let device = &shared.device;
2548        // One write covers the whole descriptor array: bound slots, the arena, and valid
2549        // filler for elements this program never addresses.
2550        let write = vk::WriteDescriptorSet::default()
2551            .dst_set(slot.descriptor_set)
2552            .dst_binding(0)
2553            .dst_array_element(0)
2554            .descriptor_type(vk::DescriptorType::STORAGE_BUFFER)
2555            .buffer_info(descriptors);
2556        let begin = vk::CommandBufferBeginInfo::default()
2557            .flags(vk::CommandBufferUsageFlags::ONE_TIME_SUBMIT);
2558        // Between dependent dispatches. A `COMPUTE_SHADER → COMPUTE_SHADER` barrier is an
2559        // execution dependency on every prior compute command, which alone orders a later write
2560        // after earlier reads (WAR: an arena region reused after its last reader). The access
2561        // masks add the memory dependency the RAW and WAW cases need: prior storage writes made
2562        // available, then visible to the next dispatch's storage reads and writes. Read accesses
2563        // never appear in a source mask because a read leaves nothing to make available.
2564        let compute_barrier = [vk::MemoryBarrier2::default()
2565            .src_stage_mask(vk::PipelineStageFlags2::COMPUTE_SHADER)
2566            .src_access_mask(vk::AccessFlags2::SHADER_STORAGE_WRITE)
2567            .dst_stage_mask(vk::PipelineStageFlags2::COMPUTE_SHADER)
2568            .dst_access_mask(
2569                vk::AccessFlags2::SHADER_STORAGE_READ | vk::AccessFlags2::SHADER_STORAGE_WRITE,
2570            )];
2571        let compute_dependency = vk::DependencyInfo::default().memory_barriers(&compute_barrier);
2572        // After the last dispatch: make the shader's storage writes visible to host reads once
2573        // the fence signals, and to the staging copies a later `read_buffer`/`write_buffer` of a
2574        // device-local buffer submits (there is no implicit dependency between submissions).
2575        let host_barrier = [vk::MemoryBarrier2::default()
2576            .src_stage_mask(vk::PipelineStageFlags2::COMPUTE_SHADER)
2577            .src_access_mask(vk::AccessFlags2::SHADER_STORAGE_WRITE)
2578            .dst_stage_mask(vk::PipelineStageFlags2::HOST | vk::PipelineStageFlags2::COPY)
2579            .dst_access_mask(
2580                vk::AccessFlags2::HOST_READ
2581                    | vk::AccessFlags2::TRANSFER_READ
2582                    | vk::AccessFlags2::TRANSFER_WRITE,
2583            )];
2584        let host_dependency = vk::DependencyInfo::default().memory_barriers(&host_barrier);
2585        let submit_buffers =
2586            [vk::CommandBufferSubmitInfo::default().command_buffer(slot.command_buffer)];
2587        // A gated submission's first dispatch waits for the gate's value; everything the host
2588        // wrote before signalling it is then visible to the device (the semaphore signal operation
2589        // is a host-to-device memory dependency).
2590        let waits: Vec<vk::SemaphoreSubmitInfo> = wait
2591            .map(|(semaphore, value)| {
2592                vk::SemaphoreSubmitInfo::default()
2593                    .semaphore(semaphore)
2594                    .value(value)
2595                    .stage_mask(vk::PipelineStageFlags2::COMPUTE_SHADER)
2596            })
2597            .into_iter()
2598            .collect();
2599        let submits = [vk::SubmitInfo2::default()
2600            .wait_semaphore_infos(&waits)
2601            .command_buffer_infos(&submit_buffers)];
2602        // SAFETY: the slot is free (no submission references its command buffer, fence, or
2603        // descriptor set), the descriptor infos name live buffers whose ranges were validated,
2604        // every pipeline is live and in-flight-counted by the caller, and the pool flag lets
2605        // `begin_command_buffer` reset the buffer implicitly. Host writes made before this
2606        // submission are visible to the device by the implicit host-write ordering guarantee.
2607        unsafe {
2608            device.update_descriptor_sets(std::slice::from_ref(&write), &[]);
2609            device.reset_fences(&[slot.fence])?;
2610            device.begin_command_buffer(slot.command_buffer, &begin)?;
2611            device.cmd_bind_descriptor_sets(
2612                slot.command_buffer,
2613                vk::PipelineBindPoint::COMPUTE,
2614                shared.pipeline_layout,
2615                0,
2616                &[slot.descriptor_set],
2617                &[],
2618            );
2619            for dispatch in &program.dispatches {
2620                if dispatch.barrier_before {
2621                    device.cmd_pipeline_barrier2(slot.command_buffer, &compute_dependency);
2622                }
2623                device.cmd_bind_pipeline(
2624                    slot.command_buffer,
2625                    vk::PipelineBindPoint::COMPUTE,
2626                    dispatch.pipeline,
2627                );
2628                device.cmd_dispatch(
2629                    slot.command_buffer,
2630                    dispatch.workgroups[0],
2631                    dispatch.workgroups[1],
2632                    dispatch.workgroups[2],
2633                );
2634            }
2635            device.cmd_pipeline_barrier2(slot.command_buffer, &host_dependency);
2636            device.end_command_buffer(slot.command_buffer)?;
2637            device.queue_submit2(shared.queue, &submits, slot.fence)
2638        }
2639    }
2640
2641    /// Upload every constant of `plan` into `arena` through the context's staging path.
2642    fn upload_constants(
2643        &self,
2644        context: &ContextInner,
2645        arena: &Arena,
2646        plan: &ProgramPlan,
2647    ) -> Result<(), BackendError> {
2648        let shared = &self.shared;
2649        let largest = plan
2650            .constants
2651            .iter()
2652            .map(|constant| constant.bytes.len() as u64)
2653            .max()
2654            .unwrap_or(0);
2655        if largest == 0 {
2656            return Ok(());
2657        }
2658        if let Some(mapped) = arena.mapped {
2659            for constant in &plan.constants {
2660                let offset =
2661                    usize::try_from(constant.offset).map_err(|_| BackendError::OutOfBounds)?;
2662                // SAFETY: the arena was just created with `plan.arena_bytes` bytes, every
2663                // constant's region lies inside it (lowering placed them), the mapping covers
2664                // the whole allocation, and nothing else references the arena yet. Coherent
2665                // memory needs no flush.
2666                let target = unsafe {
2667                    std::slice::from_raw_parts_mut(
2668                        mapped.as_ptr().add(offset),
2669                        constant.bytes.len(),
2670                    )
2671                };
2672                target.copy_from_slice(&constant.bytes);
2673            }
2674            return Ok(());
2675        }
2676        let mut staging = Staging::new(shared, largest.min(STAGING_BYTES))?;
2677        for constant in &plan.constants {
2678            write_through_staging(
2679                context,
2680                arena.buffer,
2681                constant.offset,
2682                constant.bytes.as_slice(),
2683                constant.bytes.len() as u64,
2684                &mut staging,
2685            )?;
2686        }
2687        Ok(())
2688    }
2689}
2690
2691/// Destroys pipelines and the arena of a program whose creation fails midway.
2692struct PartialProgram<'a> {
2693    shared: &'a Shared,
2694    pipelines: Vec<vk::Pipeline>,
2695    arena: Option<Arena>,
2696}
2697
2698impl Drop for PartialProgram<'_> {
2699    fn drop(&mut self) {
2700        // SAFETY: every handle here was created by `load_program` and not yet handed to a
2701        // program; nothing references them.
2702        unsafe {
2703            for pipeline in self.pipelines.drain(..) {
2704                if pipeline != vk::Pipeline::null() {
2705                    self.shared.device.destroy_pipeline(pipeline, None);
2706                }
2707            }
2708            if let Some(arena) = self.arena.take() {
2709                self.shared.device.destroy_buffer(arena.buffer, None);
2710                self.shared.device.free_memory(arena.memory, None);
2711            }
2712        }
2713    }
2714}
2715
2716impl TosaCapabilityProvider for VulkanAccelerator {
2717    fn tosa_capabilities(&self) -> &'static [CapabilityDescriptor] {
2718        // Neither tier needs a device feature (ADR 0008, ADR 0009): every conversion is
2719        // crate-owned integer and binary32 code, so both are advertised on every device the
2720        // backend opens, with numerics identical everywhere.
2721        &[
2722            crate::VULKAN_TOSA_FP16_CAPABILITY,
2723            crate::VULKAN_TOSA_FP8_CAPABILITY,
2724        ]
2725    }
2726}
2727
2728impl Accelerator for VulkanAccelerator {
2729    type Context = VulkanContext;
2730    type Buffer = VulkanBuffer;
2731    type Program = VulkanProgram;
2732    type Queue = VulkanQueue;
2733    type Event = VulkanEvent;
2734
2735    fn device_info(&self) -> Result<DeviceInfo, BackendError> {
2736        Ok(self.shared.info)
2737    }
2738
2739    fn create_context(&self, desc: ContextDesc) -> Result<Self::Context, BackendError> {
2740        self.shared.info.validate_context_desc(desc)?;
2741        self.shared.check_live()?;
2742        if self.shared.counters.contexts.get() >= u64::from(MAX_CONTEXTS) {
2743            return Err(BackendError::ResourceLimit);
2744        }
2745        let inner = ContextInner::create(&self.shared, self.next_id()?)?;
2746        Ok(VulkanContext { inner })
2747    }
2748
2749    fn destroy_context(&self, context: Self::Context) -> Result<(), ReleaseFailure<Self::Context>> {
2750        if Rc::strong_count(&context.inner) > 1 {
2751            return Err(ReleaseFailure::Rejected {
2752                error: BackendError::Busy,
2753                resource: context,
2754            });
2755        }
2756        Ok(())
2757    }
2758
2759    fn allocate_buffer(
2760        &self,
2761        context: &Self::Context,
2762        desc: BufferDesc,
2763    ) -> Result<AllocatedBuffer<Self::Buffer>, BackendError> {
2764        let shared = &self.shared;
2765        shared.info.validate_buffer_desc(desc)?;
2766        shared.check_live()?;
2767        if shared.counters.buffers.get()
2768            >= u64::from(MAX_BUFFERS_PER_CONTEXT) * u64::from(MAX_CONTEXTS)
2769        {
2770            return Err(BackendError::ResourceLimit);
2771        }
2772        let memory_type = shared
2773            .memory_plan
2774            .for_domain(desc.domain)
2775            .ok_or(BackendError::Unsupported)?;
2776        // Mapped whenever the type allows: on a unified-memory device that includes `Device`.
2777        let map = shared
2778            .memory_plan
2779            .is_mapped(&shared.physical.memory, desc.domain);
2780        let raw = RawAllocation::create(shared, desc.bytes(), memory_type, map)?;
2781        if raw.measured_alignment < desc.alignment() {
2782            return Err(BackendError::ResourceLimit);
2783        }
2784        let mut properties = BufferProperties::DIRECT_BINDING;
2785        if raw.mapped.is_some() {
2786            properties |= BufferProperties::HOST_VISIBLE;
2787        }
2788        if raw
2789            .memory_flags
2790            .contains(vk::MemoryPropertyFlags::DEVICE_LOCAL)
2791        {
2792            properties |= BufferProperties::DEVICE_LOCAL;
2793        }
2794        let info = BufferInfo::new(
2795            desc,
2796            raw.allocation_bytes,
2797            raw.measured_alignment,
2798            properties,
2799        )?;
2800        let (buffer, memory, mapped) = raw.into_parts();
2801        increment(&shared.counters.buffers, 1);
2802        Ok(AllocatedBuffer::new(
2803            VulkanBuffer {
2804                context: Rc::clone(&context.inner),
2805                desc,
2806                buffer,
2807                memory,
2808                mapped,
2809                imported: false,
2810                state: Rc::new(BufferState::default()),
2811            },
2812            info,
2813        ))
2814    }
2815
2816    fn write_buffer(
2817        &self,
2818        buffer: &mut Self::Buffer,
2819        offset: u64,
2820        data: &dyn ByteSource,
2821    ) -> Result<(), BackendError> {
2822        if !buffer
2823            .desc
2824            .usage
2825            .contains(BufferUsage::TRANSFER_DESTINATION)
2826        {
2827            return Err(BackendError::PermissionDenied);
2828        }
2829        if buffer.in_flight() != 0 {
2830            return Err(BackendError::Busy);
2831        }
2832        self.shared.check_live()?;
2833        let (start, end) = Self::checked_range(buffer, offset, data.len())?;
2834        let len = end - start;
2835        let Some(target) = buffer.mapped_at(start) else {
2836            return self.staged_write(buffer, start as u64, data, len as u64);
2837        };
2838        // SAFETY: `target..target + len` is inside the persistent host-coherent mapping of a
2839        // buffer that is exclusively borrowed and not in flight; the source is a distinct
2840        // borrowed region. Coherent memory needs no flush.
2841        let target = unsafe { std::slice::from_raw_parts_mut(target, len) };
2842        match data.as_contiguous() {
2843            Some(source) if source.len() == len => target.copy_from_slice(source),
2844            Some(_) => return Err(BackendError::InvalidArgument),
2845            None => data.read_at(0, target)?,
2846        }
2847        increment(&self.shared.counters.explicit_transfer_bytes, len as u64);
2848        Ok(())
2849    }
2850
2851    fn read_buffer(
2852        &self,
2853        buffer: &Self::Buffer,
2854        offset: u64,
2855        data: &mut dyn ByteSink,
2856    ) -> Result<(), BackendError> {
2857        if !buffer.desc.usage.contains(BufferUsage::TRANSFER_SOURCE) {
2858            return Err(BackendError::PermissionDenied);
2859        }
2860        if buffer.in_flight() != 0 {
2861            return Err(BackendError::Busy);
2862        }
2863        self.shared.check_live()?;
2864        let (start, end) = Self::checked_range(buffer, offset, data.len())?;
2865        let len = end - start;
2866        let Some(source) = buffer.mapped_at(start) else {
2867            return self.staged_read(buffer, start as u64, data, len as u64);
2868        };
2869        // SAFETY: the range is inside the mapping and the in-flight gate proved no submission
2870        // still writes this buffer; every completed submission's writes were made host-visible
2871        // by its command buffer's barrier before its fence signaled.
2872        let source = unsafe { std::slice::from_raw_parts(source.cast_const(), len) };
2873        match data.as_contiguous_mut() {
2874            Some(target) if target.len() == len => target.copy_from_slice(source),
2875            Some(_) => return Err(BackendError::InvalidArgument),
2876            None => data.write_at(0, source)?,
2877        }
2878        increment(&self.shared.counters.explicit_transfer_bytes, len as u64);
2879        Ok(())
2880    }
2881
2882    fn free_buffer(&self, buffer: Self::Buffer) -> Result<(), ReleaseFailure<Self::Buffer>> {
2883        if buffer.in_flight() != 0 {
2884            return Err(ReleaseFailure::Rejected {
2885                error: BackendError::Busy,
2886                resource: buffer,
2887            });
2888        }
2889        Ok(())
2890    }
2891
2892    fn load_program(
2893        &self,
2894        context: &Self::Context,
2895        artifact: ArtifactRef<'_>,
2896    ) -> Result<Self::Program, BackendError> {
2897        let shared = &self.shared;
2898        if artifact.payload.len() > shared.info.limits.max_artifact_bytes {
2899            return Err(BackendError::ResourceLimit);
2900        }
2901        if artifact.resident_bytes != REQUIRED_RESIDENT_BYTES {
2902            return Err(BackendError::ResourceLimit);
2903        }
2904        shared.check_live()?;
2905        if shared.counters.programs.get()
2906            >= u64::from(MAX_PROGRAMS_PER_CONTEXT) * u64::from(MAX_CONTEXTS)
2907        {
2908            return Err(BackendError::ResourceLimit);
2909        }
2910        let mut owned = Vec::new();
2911        let bytes = match artifact.payload.as_contiguous() {
2912            Some(bytes) => bytes,
2913            None => {
2914                let len = usize::try_from(artifact.payload.len())
2915                    .map_err(|_| BackendError::ResourceLimit)?;
2916                owned
2917                    .try_reserve_exact(len)
2918                    .map_err(|_| BackendError::OutOfMemory)?;
2919                owned.resize(len, 0);
2920                artifact.payload.read_at(0, &mut owned)?;
2921                &owned
2922            }
2923        };
2924        let plan = if artifact.format == crate::nvfp4::VULKAN_NVFP4_FORMAT {
2925            lower_nvfp4(bytes).map_err(Self::lowering_error)?
2926        } else if artifact.format == virtio_accel_tosa::ARTIFACT_FORMAT {
2927            let target = virtio_accel_tosa::Target::from_identity(artifact.target)
2928                .map_err(|_| BackendError::Incompatible)?;
2929            lower_tosa(bytes, target).map_err(Self::lowering_error)?
2930        } else {
2931            return Err(BackendError::Unsupported);
2932        };
2933        let tuning = shared.physical.tuning;
2934        let limits = &shared.physical.limits;
2935        // The plan's slots and arena must fit the descriptor array, and the arena one storage
2936        // buffer descriptor.
2937        if plan.buffer_count() > tuning.buffers
2938            || plan.slots.len() > tuning.max_bindings() as usize
2939            || plan.arena_bytes > u64::from(limits.max_storage_buffer_range)
2940        {
2941            return Err(BackendError::ResourceLimit);
2942        }
2943        let mut workgroups = Vec::with_capacity(plan.dispatches.len());
2944        for dispatch in &plan.dispatches {
2945            workgroups.push(
2946                tuning
2947                    .workgroups(dispatch.work, dispatch.kernel, limits)
2948                    .ok_or(BackendError::ResourceLimit)?,
2949            );
2950        }
2951
2952        let mut partial = PartialProgram {
2953            shared,
2954            pipelines: Vec::new(),
2955            arena: None,
2956        };
2957        if plan.arena_bytes != 0 {
2958            // Device-local when the device has such memory: intermediates never leave the GPU.
2959            let (memory_type, map) = match shared.memory_plan.device {
2960                Some(device) => (
2961                    device,
2962                    shared
2963                        .memory_plan
2964                        .is_mapped(&shared.physical.memory, MemoryDomain::Device),
2965                ),
2966                None => (shared.memory_plan.host, true),
2967            };
2968            let raw = RawAllocation::create(shared, plan.arena_bytes, memory_type, map)?;
2969            let (buffer, memory, mapped) = raw.into_parts();
2970            partial.arena = Some(Arena {
2971                buffer,
2972                memory,
2973                bytes: plan.arena_bytes,
2974                mapped,
2975            });
2976            self.upload_constants(
2977                &context.inner,
2978                partial.arena.as_ref().expect("arena set above"),
2979                &plan,
2980            )?;
2981        }
2982
2983        // One shader module per distinct kernel variant, one pipeline per dispatch, created in
2984        // a single call against the instance's pipeline cache.
2985        let device = &shared.device;
2986        let mut modules: Vec<(KernelKey, vk::ShaderModule)> = Vec::new();
2987        let mut module_for = |key: KernelKey| -> Result<vk::ShaderModule, BackendError> {
2988            if let Some((_, module)) = modules.iter().find(|(existing, _)| *existing == key) {
2989                return Ok(*module);
2990            }
2991            let code = shared.module(key);
2992            let module_info = vk::ShaderModuleCreateInfo::default().code(&code);
2993            // SAFETY: `code` is the crate-assembled SPIR-V module, live for the call.
2994            let module = unsafe { device.create_shader_module(&module_info, None) }
2995                .map_err(|result| shared.fail(result))?;
2996            modules.push((key, module));
2997            Ok(module)
2998        };
2999        struct ModuleGuard<'a>(&'a ash::Device, Vec<vk::ShaderModule>);
3000        impl Drop for ModuleGuard<'_> {
3001            fn drop(&mut self) {
3002                for module in self.1.drain(..) {
3003                    // SAFETY: modules are no longer needed once pipeline creation returned (or
3004                    // failed); each is destroyed exactly once.
3005                    unsafe { self.0.destroy_shader_module(module, None) };
3006                }
3007            }
3008        }
3009        let mut stage_modules = Vec::with_capacity(plan.dispatches.len());
3010        for dispatch in &plan.dispatches {
3011            let key = tuning.key(dispatch.kernel);
3012            debug_assert_eq!(dispatch.spec.len() as u32, key.spec_constant_count());
3013            match module_for(key) {
3014                Ok(module) => stage_modules.push(module),
3015                Err(error) => {
3016                    drop(ModuleGuard(
3017                        device,
3018                        modules.into_iter().map(|(_, m)| m).collect(),
3019                    ));
3020                    return Err(error);
3021                }
3022            }
3023        }
3024        let module_guard = ModuleGuard(device, modules.into_iter().map(|(_, m)| m).collect());
3025        let entries: Vec<Vec<vk::SpecializationMapEntry>> = plan
3026            .dispatches
3027            .iter()
3028            .map(|dispatch| {
3029                (0..dispatch.spec.len() as u32)
3030                    .map(|id| vk::SpecializationMapEntry {
3031                        constant_id: id,
3032                        offset: id * 4,
3033                        size: 4,
3034                    })
3035                    .collect()
3036            })
3037            .collect();
3038        let spec_data: Vec<Vec<u8>> = plan
3039            .dispatches
3040            .iter()
3041            .map(|dispatch| {
3042                dispatch
3043                    .spec
3044                    .iter()
3045                    .flat_map(|word| word.to_ne_bytes())
3046                    .collect()
3047            })
3048            .collect();
3049        let specializations: Vec<vk::SpecializationInfo<'_>> = entries
3050            .iter()
3051            .zip(&spec_data)
3052            .map(|(entries, data)| {
3053                vk::SpecializationInfo::default()
3054                    .map_entries(entries)
3055                    .data(data)
3056            })
3057            .collect();
3058        let stages: Vec<vk::PipelineShaderStageCreateInfo<'_>> = stage_modules
3059            .iter()
3060            .zip(&specializations)
3061            .map(|(module, specialization)| {
3062                vk::PipelineShaderStageCreateInfo::default()
3063                    .stage(vk::ShaderStageFlags::COMPUTE)
3064                    .module(*module)
3065                    .name(c"main")
3066                    .specialization_info(specialization)
3067            })
3068            .collect();
3069        let pipeline_infos: Vec<vk::ComputePipelineCreateInfo<'_>> = stages
3070            .iter()
3071            .map(|stage| {
3072                vk::ComputePipelineCreateInfo::default()
3073                    .stage(*stage)
3074                    .layout(shared.pipeline_layout)
3075            })
3076            .collect();
3077        if !pipeline_infos.is_empty() {
3078            // SAFETY: modules, layout, and cache are live; every pointed-to structure outlives
3079            // the call. On failure ash returns the partially created array, whose non-null
3080            // entries are destroyed by the partial-program guard.
3081            let created = unsafe {
3082                device.create_compute_pipelines(shared.pipeline_cache, &pipeline_infos, None)
3083            };
3084            match created {
3085                Ok(pipelines) => partial.pipelines = pipelines,
3086                Err((pipelines, result)) => {
3087                    partial.pipelines = pipelines;
3088                    drop(module_guard);
3089                    return Err(shared.fail(result));
3090                }
3091            }
3092        }
3093        drop(module_guard);
3094
3095        let dispatches = partial
3096            .pipelines
3097            .drain(..)
3098            .zip(&plan.dispatches)
3099            .zip(workgroups)
3100            .map(|((pipeline, dispatch), workgroups)| Dispatch {
3101                pipeline,
3102                workgroups,
3103                barrier_before: dispatch.barrier_before,
3104            })
3105            .collect();
3106        let arena = partial.arena.take();
3107        increment(&shared.counters.programs, 1);
3108        Ok(VulkanProgram {
3109            context: Rc::clone(&context.inner),
3110            dispatches,
3111            arena,
3112            plan,
3113            state: Rc::new(ProgramState::default()),
3114        })
3115    }
3116
3117    fn unload_program(&self, program: Self::Program) -> Result<(), ReleaseFailure<Self::Program>> {
3118        if program.state.in_flight.get() != 0 {
3119            return Err(ReleaseFailure::Rejected {
3120                error: BackendError::Busy,
3121                resource: program,
3122            });
3123        }
3124        Ok(())
3125    }
3126
3127    fn create_queue(
3128        &self,
3129        context: &Self::Context,
3130        desc: QueueDesc,
3131    ) -> Result<Self::Queue, BackendError> {
3132        self.shared.info.validate_queue_desc(desc)?;
3133        self.shared.check_live()?;
3134        if self.shared.counters.queues.get()
3135            >= u64::from(MAX_QUEUES_PER_CONTEXT) * u64::from(MAX_CONTEXTS)
3136        {
3137            return Err(BackendError::ResourceLimit);
3138        }
3139        increment(&self.shared.counters.queues, 1);
3140        Ok(VulkanQueue {
3141            context: Rc::clone(&context.inner),
3142        })
3143    }
3144
3145    fn destroy_queue(&self, _queue: Self::Queue) -> Result<(), ReleaseFailure<Self::Queue>> {
3146        Ok(())
3147    }
3148
3149    fn submit(
3150        &self,
3151        queue: &Self::Queue,
3152        program: &Self::Program,
3153        bindings: &[BindingRef<'_, Self::Buffer>],
3154        timeout: Timeout,
3155    ) -> Result<Self::Event, SubmitFailure<Self::Event>> {
3156        self.submit_waiting(queue, program, bindings, timeout, None)
3157    }
3158
3159    fn poll_event(&self, event: &Self::Event) -> Result<EventState, BackendError> {
3160        event.poll()
3161    }
3162
3163    fn destroy_event(&self, event: Self::Event) -> Result<(), ReleaseFailure<Self::Event>> {
3164        match event.poll() {
3165            Ok(EventState::Pending) => Err(ReleaseFailure::Rejected {
3166                error: BackendError::Busy,
3167                resource: event,
3168            }),
3169            Ok(_) => {
3170                event.release();
3171                Ok(())
3172            }
3173            Err(error) => Err(ReleaseFailure::Rejected {
3174                error,
3175                resource: event,
3176            }),
3177        }
3178    }
3179}
3180
3181#[cfg(test)]
3182mod tests {
3183    use super::*;
3184
3185    /// The memory types a Radeon 860M (RADV, Mesa 26.1.8) reports. The ordinary types come first
3186    /// and `VK_AMD_device_coherent_memory` appends its own after them, which is what makes a
3187    /// last-wins tie-break select exactly the types that require an enabled feature.
3188    fn amd_device_coherent_memory() -> vk::PhysicalDeviceMemoryProperties {
3189        use vk::MemoryPropertyFlags as Flags;
3190        let amd = Flags::DEVICE_COHERENT_AMD | Flags::DEVICE_UNCACHED_AMD;
3191        let host_coherent = Flags::HOST_VISIBLE | Flags::HOST_COHERENT;
3192        let layout = [
3193            (Flags::DEVICE_LOCAL, 1),
3194            (Flags::DEVICE_LOCAL, 1),
3195            (host_coherent, 0),
3196            (Flags::DEVICE_LOCAL | host_coherent, 1),
3197            (Flags::DEVICE_LOCAL | host_coherent, 1),
3198            (host_coherent | Flags::HOST_CACHED, 0),
3199            (host_coherent | Flags::HOST_CACHED, 0),
3200            (Flags::DEVICE_LOCAL | amd, 1),
3201            (host_coherent | amd, 0),
3202            (Flags::DEVICE_LOCAL | host_coherent | amd, 1),
3203            (host_coherent | Flags::HOST_CACHED | amd, 0),
3204        ];
3205        let mut memory = vk::PhysicalDeviceMemoryProperties::default();
3206        for (slot, (flags, heap)) in layout.iter().enumerate() {
3207            memory.memory_types[slot] = vk::MemoryType::default()
3208                .property_flags(*flags)
3209                .heap_index(*heap);
3210        }
3211        memory.memory_type_count =
3212            u32::try_from(layout.len()).expect("the fixture declares eleven memory types");
3213        memory.memory_heap_count = 2;
3214        memory
3215    }
3216
3217    /// `VK_MEMORY_PROPERTY_DEVICE_COHERENT_BIT_AMD` may not be allocated from unless the
3218    /// `deviceCoherentMemory` feature is enabled, which this backend does not request, and the
3219    /// spec advises against that memory anyway: it is uncached, so repeated accesses to nearby
3220    /// locations — a tiled MATMUL — are slower. No CI device exposes these types, so the layout
3221    /// is a fixture rather than a live probe.
3222    #[test]
3223    fn never_selects_memory_that_requires_an_unrequested_feature() {
3224        let memory = amd_device_coherent_memory();
3225        let plan = MemoryPlan::select(&memory, u32::MAX, false)
3226            .expect("a host-visible coherent type is present in the fixture");
3227
3228        for (domain, selected) in [
3229            ("host", Some(plan.host)),
3230            ("device", plan.device),
3231            ("shared", plan.shared),
3232        ] {
3233            let Some(index) = selected else { continue };
3234            let flags = memory.memory_types[index as usize].property_flags;
3235            assert!(
3236                !flags.intersects(
3237                    vk::MemoryPropertyFlags::DEVICE_COHERENT_AMD
3238                        | vk::MemoryPropertyFlags::RDMA_CAPABLE_NV
3239                ),
3240                "{domain} domain selected memory type {index}, which requires a feature the \
3241                 backend never enables: property flags {:#x}",
3242                flags.as_raw(),
3243            );
3244        }
3245    }
3246
3247    /// Excluding those types must not cost a domain: the AMD extension adds its memory types
3248    /// alongside the ordinary ones rather than replacing them, so every domain stays reachable.
3249    #[test]
3250    fn excluding_them_strands_no_memory_domain() {
3251        let plan = MemoryPlan::select(&amd_device_coherent_memory(), u32::MAX, false)
3252            .expect("a host-visible coherent type is present in the fixture");
3253        assert!(plan.device.is_some(), "device-local domain lost");
3254        assert!(plan.shared.is_some(), "shared domain lost");
3255    }
3256}