Skip to main content

ironaccelerator_levelzero/
drv.rs

1//! Minimal `ze_loader` driver: `zeInit`, driver walk, device walk,
2//! `zeDeviceGetProperties`. Enough for the planner; kernel launch lives
3//! in higher layers.
4//!
5//! `drop(sym)` on a `libloading::Symbol` is intentional — it releases the
6//! borrow on the Library so the next `lib.get(...)` can proceed. Symbol is
7//! `Copy` and doesn't impl `Drop`, but the borrow lives in its lifetime
8//! parameter, which `drop` does end. We silence the spurious lint module-wide.
9//!
10//! Everything is loaded via `libloading` so the crate compiles on hosts
11//! without Level Zero installed; the backend simply reports unavailable.
12
13#![allow(clippy::drop_non_drop)] // see module docs above
14
15use core::ffi::c_void;
16use libloading::{Library, Symbol};
17use once_cell::sync::OnceCell;
18
19// ── Raw FFI types ──────────────────────────────────────────────────────────
20
21pub type ZeResult = u32;
22pub const ZE_RESULT_SUCCESS: ZeResult = 0;
23
24pub type ZeDriverHandle = *mut c_void;
25pub type ZeDeviceHandle = *mut c_void;
26
27pub const ZE_STRUCTURE_TYPE_DEVICE_PROPERTIES: u32 = 0x1;
28
29pub const ZE_DEVICE_TYPE_GPU: u32 = 1;
30pub const ZE_DEVICE_TYPE_CPU: u32 = 2;
31pub const ZE_DEVICE_TYPE_FPGA: u32 = 3;
32pub const ZE_DEVICE_TYPE_MCA: u32 = 4;
33pub const ZE_DEVICE_TYPE_VPU: u32 = 5;
34
35pub const ZE_MAX_DEVICE_NAME: usize = 256;
36
37#[repr(C)]
38#[derive(Clone, Copy)]
39pub struct ZeDeviceUuid {
40    pub id: [u8; 16],
41}
42
43#[repr(C)]
44#[derive(Clone, Copy)]
45pub struct ZeDeviceProperties {
46    pub stype: u32,
47    pub p_next: *mut c_void,
48    pub type_: u32,
49    pub vendor_id: u32,
50    pub device_id: u32,
51    pub subdevice_id: u32,
52    pub core_clock_rate: u32,
53    pub max_mem_alloc_size: u64,
54    pub max_hardware_contexts: u32,
55    pub max_command_queue_priority: u32,
56    pub num_threads_per_eu: u32,
57    pub physical_eu_simd_width: u32,
58    pub num_eus_per_subslice: u32,
59    pub num_subslices_per_slice: u32,
60    pub num_slices: u32,
61    pub timer_resolution: u64,
62    pub timestamp_valid_bits: u32,
63    pub kernel_timestamp_valid_bits: u32,
64    pub uuid: ZeDeviceUuid,
65    pub name: [core::ffi::c_char; ZE_MAX_DEVICE_NAME],
66}
67
68pub type ZeContextHandle = *mut c_void;
69pub type ZeCommandQueueHandle = *mut c_void;
70pub type ZeCommandListHandle = *mut c_void;
71
72pub const ZE_STRUCTURE_TYPE_CONTEXT_DESC: u32 = 0x2;
73pub const ZE_STRUCTURE_TYPE_COMMAND_QUEUE_DESC: u32 = 0x3;
74pub const ZE_STRUCTURE_TYPE_COMMAND_LIST_DESC: u32 = 0x4;
75pub const ZE_STRUCTURE_TYPE_DEVICE_MEM_ALLOC_DESC: u32 = 0xc;
76pub const ZE_STRUCTURE_TYPE_HOST_MEM_ALLOC_DESC: u32 = 0xd;
77pub const ZE_STRUCTURE_TYPE_MODULE_DESC: u32 = 0xf;
78pub const ZE_STRUCTURE_TYPE_KERNEL_DESC: u32 = 0x10;
79
80pub const ZE_MODULE_FORMAT_IL_SPIRV: u32 = 0;
81pub const ZE_MODULE_FORMAT_NATIVE: u32 = 1;
82
83pub type ZeModuleHandle = *mut c_void;
84pub type ZeKernelHandle = *mut c_void;
85
86#[repr(C)]
87pub struct ZeDeviceMemAllocDesc {
88    pub stype: u32,
89    pub p_next: *const c_void,
90    pub flags: u32,
91    pub ordinal: u32,
92}
93
94#[repr(C)]
95pub struct ZeHostMemAllocDesc {
96    pub stype: u32,
97    pub p_next: *const c_void,
98    pub flags: u32,
99}
100
101#[repr(C)]
102pub struct ZeModuleDesc {
103    pub stype: u32,
104    pub p_next: *const c_void,
105    pub format: u32,
106    pub input_size: usize,
107    pub p_input_module: *const u8,
108    pub p_build_flags: *const core::ffi::c_char,
109    pub p_constants: *const c_void,
110}
111
112#[repr(C)]
113pub struct ZeKernelDesc {
114    pub stype: u32,
115    pub p_next: *const c_void,
116    pub flags: u32,
117    pub p_kernel_name: *const core::ffi::c_char,
118}
119
120#[repr(C)]
121#[derive(Clone, Copy)]
122pub struct ZeGroupCount {
123    pub group_count_x: u32,
124    pub group_count_y: u32,
125    pub group_count_z: u32,
126}
127
128pub const ZE_COMMAND_QUEUE_MODE_DEFAULT: u32 = 0;
129pub const ZE_COMMAND_QUEUE_PRIORITY_NORMAL: u32 = 0;
130
131#[repr(C)]
132pub struct ZeContextDesc {
133    pub stype: u32,
134    pub p_next: *const c_void,
135    pub flags: u32,
136}
137
138#[repr(C)]
139pub struct ZeCommandQueueDesc {
140    pub stype: u32,
141    pub p_next: *const c_void,
142    pub ordinal: u32,
143    pub index: u32,
144    pub flags: u32,
145    pub mode: u32,
146    pub priority: u32,
147}
148
149#[repr(C)]
150pub struct ZeCommandListDesc {
151    pub stype: u32,
152    pub p_next: *const c_void,
153    pub command_queue_group_ordinal: u32,
154    pub flags: u32,
155}
156
157type ZeInitFn = unsafe extern "C" fn(flags: u32) -> ZeResult;
158pub type ZeContextCreateFn = unsafe extern "C" fn(
159    h_driver: ZeDriverHandle,
160    desc: *const ZeContextDesc,
161    ph_context: *mut ZeContextHandle,
162) -> ZeResult;
163pub type ZeContextDestroyFn = unsafe extern "C" fn(h_context: ZeContextHandle) -> ZeResult;
164pub type ZeCommandQueueCreateFn = unsafe extern "C" fn(
165    h_context: ZeContextHandle,
166    h_device: ZeDeviceHandle,
167    desc: *const ZeCommandQueueDesc,
168    ph_queue: *mut ZeCommandQueueHandle,
169) -> ZeResult;
170pub type ZeCommandQueueDestroyFn = unsafe extern "C" fn(h_queue: ZeCommandQueueHandle) -> ZeResult;
171pub type ZeCommandListCreateFn = unsafe extern "C" fn(
172    h_context: ZeContextHandle,
173    h_device: ZeDeviceHandle,
174    desc: *const ZeCommandListDesc,
175    ph_list: *mut ZeCommandListHandle,
176) -> ZeResult;
177pub type ZeCommandListDestroyFn = unsafe extern "C" fn(h_list: ZeCommandListHandle) -> ZeResult;
178type ZeDriverGetFn =
179    unsafe extern "C" fn(p_count: *mut u32, ph_drivers: *mut ZeDriverHandle) -> ZeResult;
180type ZeDeviceGetFn = unsafe extern "C" fn(
181    h_driver: ZeDriverHandle,
182    p_count: *mut u32,
183    ph_devices: *mut ZeDeviceHandle,
184) -> ZeResult;
185type ZeDeviceGetPropertiesFn = unsafe extern "C" fn(
186    h_device: ZeDeviceHandle,
187    p_properties: *mut ZeDeviceProperties,
188) -> ZeResult;
189
190pub type ZeMemAllocDeviceFn = unsafe extern "C" fn(
191    h_context: ZeContextHandle,
192    device_desc: *const ZeDeviceMemAllocDesc,
193    size: usize,
194    alignment: usize,
195    h_device: ZeDeviceHandle,
196    pptr: *mut *mut c_void,
197) -> ZeResult;
198pub type ZeMemAllocSharedFn = unsafe extern "C" fn(
199    h_context: ZeContextHandle,
200    device_desc: *const ZeDeviceMemAllocDesc,
201    host_desc: *const ZeHostMemAllocDesc,
202    size: usize,
203    alignment: usize,
204    h_device: ZeDeviceHandle,
205    pptr: *mut *mut c_void,
206) -> ZeResult;
207pub type ZeMemFreeFn =
208    unsafe extern "C" fn(h_context: ZeContextHandle, ptr: *mut c_void) -> ZeResult;
209pub type ZeModuleCreateFn = unsafe extern "C" fn(
210    h_context: ZeContextHandle,
211    h_device: ZeDeviceHandle,
212    desc: *const ZeModuleDesc,
213    ph_module: *mut ZeModuleHandle,
214    ph_build_log: *mut *mut c_void,
215) -> ZeResult;
216pub type ZeModuleDestroyFn = unsafe extern "C" fn(h_module: ZeModuleHandle) -> ZeResult;
217pub type ZeKernelCreateFn = unsafe extern "C" fn(
218    h_module: ZeModuleHandle,
219    desc: *const ZeKernelDesc,
220    ph_kernel: *mut ZeKernelHandle,
221) -> ZeResult;
222pub type ZeKernelDestroyFn = unsafe extern "C" fn(h_kernel: ZeKernelHandle) -> ZeResult;
223pub type ZeKernelSetGroupSizeFn =
224    unsafe extern "C" fn(h_kernel: ZeKernelHandle, gx: u32, gy: u32, gz: u32) -> ZeResult;
225pub type ZeKernelSetArgumentValueFn = unsafe extern "C" fn(
226    h_kernel: ZeKernelHandle,
227    arg_index: u32,
228    arg_size: usize,
229    p_arg_value: *const c_void,
230) -> ZeResult;
231pub type ZeCommandListAppendLaunchKernelFn = unsafe extern "C" fn(
232    h_list: ZeCommandListHandle,
233    h_kernel: ZeKernelHandle,
234    p_launch_args: *const ZeGroupCount,
235    h_signal_event: *mut c_void,
236    num_wait_events: u32,
237    ph_wait_events: *mut *mut c_void,
238) -> ZeResult;
239pub type ZeCommandListAppendMemoryCopyFn = unsafe extern "C" fn(
240    h_list: ZeCommandListHandle,
241    dstptr: *mut c_void,
242    srcptr: *const c_void,
243    size: usize,
244    h_signal_event: *mut c_void,
245    num_wait_events: u32,
246    ph_wait_events: *mut *mut c_void,
247) -> ZeResult;
248pub type ZeCommandListCloseFn = unsafe extern "C" fn(h_list: ZeCommandListHandle) -> ZeResult;
249pub type ZeCommandListResetFn = unsafe extern "C" fn(h_list: ZeCommandListHandle) -> ZeResult;
250pub type ZeCommandQueueExecuteCommandListsFn = unsafe extern "C" fn(
251    h_queue: ZeCommandQueueHandle,
252    num_lists: u32,
253    ph_lists: *const ZeCommandListHandle,
254    h_fence: *mut c_void,
255) -> ZeResult;
256pub type ZeCommandQueueSynchronizeFn =
257    unsafe extern "C" fn(h_queue: ZeCommandQueueHandle, timeout: u64) -> ZeResult;
258
259// ── Loader ─────────────────────────────────────────────────────────────────
260
261static LIB: OnceCell<Option<Loaded>> = OnceCell::new();
262
263pub struct Loaded {
264    _lib: Library,
265    pub ze_driver_get: ZeDriverGetFn,
266    pub ze_device_get: ZeDeviceGetFn,
267    pub ze_device_get_properties: ZeDeviceGetPropertiesFn,
268    pub ze_context_create: ZeContextCreateFn,
269    pub ze_context_destroy: ZeContextDestroyFn,
270    pub ze_command_queue_create: ZeCommandQueueCreateFn,
271    pub ze_command_queue_destroy: ZeCommandQueueDestroyFn,
272    pub ze_command_list_create: ZeCommandListCreateFn,
273    pub ze_command_list_destroy: ZeCommandListDestroyFn,
274    pub ze_mem_alloc_device: ZeMemAllocDeviceFn,
275    pub ze_mem_alloc_shared: ZeMemAllocSharedFn,
276    pub ze_mem_free: ZeMemFreeFn,
277    pub ze_module_create: ZeModuleCreateFn,
278    pub ze_module_destroy: ZeModuleDestroyFn,
279    pub ze_kernel_create: ZeKernelCreateFn,
280    pub ze_kernel_destroy: ZeKernelDestroyFn,
281    pub ze_kernel_set_group_size: ZeKernelSetGroupSizeFn,
282    pub ze_kernel_set_argument_value: ZeKernelSetArgumentValueFn,
283    pub ze_command_list_append_launch_kernel: ZeCommandListAppendLaunchKernelFn,
284    pub ze_command_list_append_memory_copy: ZeCommandListAppendMemoryCopyFn,
285    pub ze_command_list_close: ZeCommandListCloseFn,
286    pub ze_command_list_reset: ZeCommandListResetFn,
287    pub ze_command_queue_execute_command_lists: ZeCommandQueueExecuteCommandListsFn,
288    pub ze_command_queue_synchronize: ZeCommandQueueSynchronizeFn,
289}
290
291unsafe impl Send for Loaded {}
292unsafe impl Sync for Loaded {}
293
294#[cfg(target_os = "windows")]
295const LIB_CANDIDATES: &[&str] = &["ze_loader.dll"];
296#[cfg(any(target_os = "linux", target_os = "android"))]
297const LIB_CANDIDATES: &[&str] = &["libze_loader.so.1", "libze_loader.so"];
298#[cfg(not(any(target_os = "windows", target_os = "linux", target_os = "android")))]
299const LIB_CANDIDATES: &[&str] = &[];
300
301pub(crate) fn loaded() -> Option<&'static Loaded> {
302    LIB.get_or_init(load).as_ref()
303}
304
305fn load() -> Option<Loaded> {
306    for name in LIB_CANDIDATES {
307        let lib = match unsafe { Library::new(*name) } {
308            Ok(l) => l,
309            Err(_) => continue,
310        };
311        unsafe {
312            let init: Symbol<ZeInitFn> = match lib.get(b"zeInit\0") {
313                Ok(s) => s,
314                Err(_) => continue,
315            };
316            if init(0) != ZE_RESULT_SUCCESS {
317                continue;
318            }
319            let dg: Symbol<ZeDriverGetFn> = match lib.get(b"zeDriverGet\0") {
320                Ok(s) => s,
321                Err(_) => continue,
322            };
323            let devg: Symbol<ZeDeviceGetFn> = match lib.get(b"zeDeviceGet\0") {
324                Ok(s) => s,
325                Err(_) => continue,
326            };
327            let devp: Symbol<ZeDeviceGetPropertiesFn> = match lib.get(b"zeDeviceGetProperties\0") {
328                Ok(s) => s,
329                Err(_) => continue,
330            };
331            let ze_driver_get = *dg;
332            let ze_device_get = *devg;
333            let ze_device_get_properties = *devp;
334            drop(dg);
335            drop(devg);
336            drop(devp);
337            drop(init);
338
339            let ctx_create: Symbol<ZeContextCreateFn> = match lib.get(b"zeContextCreate\0") {
340                Ok(s) => s,
341                Err(_) => continue,
342            };
343            let ctx_destroy: Symbol<ZeContextDestroyFn> = match lib.get(b"zeContextDestroy\0") {
344                Ok(s) => s,
345                Err(_) => continue,
346            };
347            let q_create: Symbol<ZeCommandQueueCreateFn> = match lib.get(b"zeCommandQueueCreate\0")
348            {
349                Ok(s) => s,
350                Err(_) => continue,
351            };
352            let q_destroy: Symbol<ZeCommandQueueDestroyFn> =
353                match lib.get(b"zeCommandQueueDestroy\0") {
354                    Ok(s) => s,
355                    Err(_) => continue,
356                };
357            let l_create: Symbol<ZeCommandListCreateFn> = match lib.get(b"zeCommandListCreate\0") {
358                Ok(s) => s,
359                Err(_) => continue,
360            };
361            let l_destroy: Symbol<ZeCommandListDestroyFn> = match lib.get(b"zeCommandListDestroy\0")
362            {
363                Ok(s) => s,
364                Err(_) => continue,
365            };
366
367            let ze_context_create = *ctx_create;
368            let ze_context_destroy = *ctx_destroy;
369            let ze_command_queue_create = *q_create;
370            let ze_command_queue_destroy = *q_destroy;
371            let ze_command_list_create = *l_create;
372            let ze_command_list_destroy = *l_destroy;
373            drop(ctx_create);
374            drop(ctx_destroy);
375            drop(q_create);
376            drop(q_destroy);
377            drop(l_create);
378            drop(l_destroy);
379
380            macro_rules! sym {
381                ($ty:ty, $name:literal) => {{
382                    let s: Symbol<$ty> = match lib.get($name) {
383                        Ok(s) => s,
384                        Err(_) => continue,
385                    };
386                    let f = *s;
387                    drop(s);
388                    f
389                }};
390            }
391            let ze_mem_alloc_device = sym!(ZeMemAllocDeviceFn, b"zeMemAllocDevice\0");
392            let ze_mem_alloc_shared = sym!(ZeMemAllocSharedFn, b"zeMemAllocShared\0");
393            let ze_mem_free = sym!(ZeMemFreeFn, b"zeMemFree\0");
394            let ze_module_create = sym!(ZeModuleCreateFn, b"zeModuleCreate\0");
395            let ze_module_destroy = sym!(ZeModuleDestroyFn, b"zeModuleDestroy\0");
396            let ze_kernel_create = sym!(ZeKernelCreateFn, b"zeKernelCreate\0");
397            let ze_kernel_destroy = sym!(ZeKernelDestroyFn, b"zeKernelDestroy\0");
398            let ze_kernel_set_group_size = sym!(ZeKernelSetGroupSizeFn, b"zeKernelSetGroupSize\0");
399            let ze_kernel_set_argument_value =
400                sym!(ZeKernelSetArgumentValueFn, b"zeKernelSetArgumentValue\0");
401            let ze_command_list_append_launch_kernel = sym!(
402                ZeCommandListAppendLaunchKernelFn,
403                b"zeCommandListAppendLaunchKernel\0"
404            );
405            let ze_command_list_append_memory_copy = sym!(
406                ZeCommandListAppendMemoryCopyFn,
407                b"zeCommandListAppendMemoryCopy\0"
408            );
409            let ze_command_list_close = sym!(ZeCommandListCloseFn, b"zeCommandListClose\0");
410            let ze_command_list_reset = sym!(ZeCommandListResetFn, b"zeCommandListReset\0");
411            let ze_command_queue_execute_command_lists = sym!(
412                ZeCommandQueueExecuteCommandListsFn,
413                b"zeCommandQueueExecuteCommandLists\0"
414            );
415            let ze_command_queue_synchronize =
416                sym!(ZeCommandQueueSynchronizeFn, b"zeCommandQueueSynchronize\0");
417
418            return Some(Loaded {
419                _lib: lib,
420                ze_driver_get,
421                ze_device_get,
422                ze_device_get_properties,
423                ze_context_create,
424                ze_context_destroy,
425                ze_command_queue_create,
426                ze_command_queue_destroy,
427                ze_command_list_create,
428                ze_command_list_destroy,
429                ze_mem_alloc_device,
430                ze_mem_alloc_shared,
431                ze_mem_free,
432                ze_module_create,
433                ze_module_destroy,
434                ze_kernel_create,
435                ze_kernel_destroy,
436                ze_kernel_set_group_size,
437                ze_kernel_set_argument_value,
438                ze_command_list_append_launch_kernel,
439                ze_command_list_append_memory_copy,
440                ze_command_list_close,
441                ze_command_list_reset,
442                ze_command_queue_execute_command_lists,
443                ze_command_queue_synchronize,
444            });
445        }
446    }
447    None
448}
449
450// ── Enumeration ────────────────────────────────────────────────────────────
451
452#[derive(Debug, Clone)]
453pub struct EnumeratedDevice {
454    pub ordinal: u32,
455    pub type_: u32,
456    pub vendor_id: u32,
457    pub device_id: u32,
458    pub name: String,
459    pub core_clock_khz: u32,
460    pub max_mem_alloc_size: u64,
461    pub num_slices: u32,
462    pub num_subslices_per_slice: u32,
463    pub num_eus_per_subslice: u32,
464}
465
466pub fn is_available() -> bool {
467    loaded().is_some()
468}
469
470pub fn enumerate() -> Vec<EnumeratedDevice> {
471    let Some(l) = loaded() else { return Vec::new() };
472    let mut out = Vec::new();
473    let mut ordinal = 0u32;
474    unsafe {
475        let mut driver_count: u32 = 0;
476        if (l.ze_driver_get)(&mut driver_count, core::ptr::null_mut()) != ZE_RESULT_SUCCESS
477            || driver_count == 0
478        {
479            return out;
480        }
481        let mut drivers = vec![core::ptr::null_mut::<c_void>(); driver_count as usize];
482        if (l.ze_driver_get)(&mut driver_count, drivers.as_mut_ptr()) != ZE_RESULT_SUCCESS {
483            return out;
484        }
485        for driver in drivers.into_iter().take(driver_count as usize) {
486            let mut dev_count: u32 = 0;
487            if (l.ze_device_get)(driver, &mut dev_count, core::ptr::null_mut()) != ZE_RESULT_SUCCESS
488                || dev_count == 0
489            {
490                continue;
491            }
492            let mut devs = vec![core::ptr::null_mut::<c_void>(); dev_count as usize];
493            if (l.ze_device_get)(driver, &mut dev_count, devs.as_mut_ptr()) != ZE_RESULT_SUCCESS {
494                continue;
495            }
496            for dev in devs.into_iter().take(dev_count as usize) {
497                let mut props: ZeDeviceProperties = core::mem::zeroed();
498                props.stype = ZE_STRUCTURE_TYPE_DEVICE_PROPERTIES;
499                if (l.ze_device_get_properties)(dev, &mut props) != ZE_RESULT_SUCCESS {
500                    continue;
501                }
502                out.push(EnumeratedDevice {
503                    ordinal,
504                    type_: props.type_,
505                    vendor_id: props.vendor_id,
506                    device_id: props.device_id,
507                    name: c_name_to_string(&props.name),
508                    core_clock_khz: props.core_clock_rate.saturating_mul(1000),
509                    max_mem_alloc_size: props.max_mem_alloc_size,
510                    num_slices: props.num_slices,
511                    num_subslices_per_slice: props.num_subslices_per_slice,
512                    num_eus_per_subslice: props.num_eus_per_subslice,
513                });
514                ordinal += 1;
515            }
516        }
517    }
518    out
519}
520
521fn c_name_to_string(raw: &[core::ffi::c_char]) -> String {
522    let bytes: Vec<u8> = raw
523        .iter()
524        .take_while(|&&c| c != 0)
525        .map(|&c| c as u8)
526        .collect();
527    String::from_utf8_lossy(&bytes).into_owned()
528}
529
530#[cfg(test)]
531mod tests {
532    use super::*;
533
534    #[test]
535    fn probe_does_not_panic() {
536        let _ = is_available();
537        let _ = enumerate();
538    }
539}