1use crate::backend::Backend;
35use rlx_driver::Device;
36use std::collections::HashMap;
37use std::sync::{OnceLock, RwLock};
38
39pub type BackendFactory = fn() -> Box<dyn Backend>;
45
46struct Registry {
47 factories: RwLock<HashMap<Device, BackendFactory>>,
48}
49
50fn registry() -> &'static Registry {
51 static REGISTRY: OnceLock<Registry> = OnceLock::new();
52 REGISTRY.get_or_init(|| {
53 let r = Registry {
54 factories: RwLock::new(HashMap::new()),
55 };
56 register_builtin(&r);
57 r
58 })
59}
60
61#[allow(unused_mut, unused_variables)]
65fn register_builtin(r: &Registry) {
66 let mut map = r.factories.write().expect("registry poisoned");
67
68 #[cfg(feature = "cpu")]
69 map.insert(Device::Cpu, || {
70 Box::new(crate::backend::cpu_backend::CpuBackend) as Box<dyn Backend>
71 });
72
73 #[cfg(all(feature = "metal", target_vendor = "apple", not(target_os = "watchos")))]
74 map.insert(Device::Metal, || {
75 Box::new(crate::backend::metal_backend::MetalBackend) as Box<dyn Backend>
76 });
77
78 #[cfg(all(feature = "mlx", rlx_mlx_host))]
79 map.insert(Device::Mlx, || {
80 Box::new(crate::backend::mlx_backend::MlxBackend) as Box<dyn Backend>
81 });
82
83 #[cfg(all(
84 feature = "coreml",
85 target_vendor = "apple",
86 not(target_os = "watchos")
87 ))]
88 map.insert(Device::Ane, || {
89 Box::new(crate::backend::coreml_backend::CoremlBackend) as Box<dyn Backend>
90 });
91
92 #[cfg(feature = "gpu")]
93 map.insert(Device::Gpu, || {
94 Box::new(crate::backend::wgpu_backend::WgpuBackend) as Box<dyn Backend>
95 });
96
97 #[cfg(feature = "webgpu")]
99 map.insert(Device::WebGpu, || {
100 Box::new(crate::backend::wgpu_backend::WgpuBackend) as Box<dyn Backend>
101 });
102
103 #[cfg(feature = "opengl")]
104 map.insert(Device::OpenGl, || {
105 Box::new(crate::backend::webgl_backend::WebglBackend) as Box<dyn Backend>
106 });
107
108 #[cfg(feature = "vulkan")]
109 map.insert(Device::Vulkan, || {
110 Box::new(crate::backend::vulkan_backend::VulkanBackend) as Box<dyn Backend>
111 });
112
113 #[cfg(feature = "cuda")]
114 map.insert(Device::Cuda, || {
115 Box::new(crate::backend::cuda_backend::CudaBackend) as Box<dyn Backend>
116 });
117
118 #[cfg(feature = "rocm")]
119 map.insert(Device::Rocm, || {
120 Box::new(crate::backend::rocm_backend::RocmBackend) as Box<dyn Backend>
121 });
122
123 #[cfg(feature = "oneapi")]
124 map.insert(Device::OneApi, || {
125 Box::new(crate::backend::oneapi_backend::OneApiBackend) as Box<dyn Backend>
126 });
127
128 #[cfg(feature = "tpu")]
129 map.insert(Device::Tpu, || {
130 Box::new(crate::backend::tpu_backend::TpuBackend) as Box<dyn Backend>
131 });
132
133 #[cfg(feature = "qnn")]
134 map.insert(Device::Hexagon, || {
135 Box::new(crate::backend::qnn_backend::QnnBackend) as Box<dyn Backend>
136 });
137}
138
139pub fn register_backend(device: Device, factory: BackendFactory) {
148 let r = registry();
149 let mut map = r.factories.write().expect("registry poisoned");
150 map.insert(device, factory);
151}
152
153pub fn backend_for(device: Device) -> Option<Box<dyn Backend>> {
156 let r = registry();
157 let map = r.factories.read().expect("registry poisoned");
158 map.get(&device).map(|f| f())
159}
160
161pub fn registered_devices() -> Vec<Device> {
163 let r = registry();
164 let map = r.factories.read().expect("registry poisoned");
165 let mut out: Vec<Device> = map.keys().copied().collect();
166 out.sort_by_key(|d| format!("{d:?}"));
167 out
168}