1use std::path::{Path, PathBuf};
4use std::sync::atomic::{AtomicBool, Ordering};
5use std::sync::{Mutex, OnceLock};
6
7use ort::execution_providers::ExecutionProvider;
8use ort::session::builder::GraphOptimizationLevel;
9use ort::session::Session;
10
11static EP_LOGGED: AtomicBool = AtomicBool::new(false);
14
15static ACTIVE_PROVIDER: OnceLock<&'static str> = OnceLock::new();
23static ORT_INITIALIZED: OnceLock<()> = OnceLock::new();
24static ORT_INIT_LOCK: Mutex<()> = Mutex::new(());
25
26pub fn active_execution_provider() -> &'static str {
29 ACTIVE_PROVIDER.get().copied().unwrap_or("CPU")
30}
31
32#[cfg(target_os = "windows")]
34const ORT_LIB_NAME: &str = "onnxruntime.dll";
35#[cfg(target_os = "macos")]
36const ORT_LIB_NAME: &str = "libonnxruntime.dylib";
37#[cfg(not(any(target_os = "windows", target_os = "macos")))]
38const ORT_LIB_NAME: &str = "libonnxruntime.so";
39
40pub struct OnnxRuntime;
51
52impl OnnxRuntime {
53 fn usable_runtime_file(path: &Path) -> bool {
54 std::fs::metadata(path)
55 .is_ok_and(|metadata| metadata.is_file() && metadata.len() >= 1024 * 1024)
56 }
57
58 fn runtime_file_name_matches(name: &str) -> bool {
59 #[cfg(target_os = "macos")]
60 {
61 name.starts_with("libonnxruntime") && name.ends_with(".dylib")
62 }
63 #[cfg(not(target_os = "macos"))]
64 {
65 name.starts_with(ORT_LIB_NAME)
66 }
67 }
68
69 fn find_runtime_in_dir(dir: &Path) -> Option<PathBuf> {
70 let direct = dir.join(ORT_LIB_NAME);
71 if Self::usable_runtime_file(&direct) {
72 return Some(direct);
73 }
74
75 let entries = std::fs::read_dir(dir).ok()?;
77 for entry in entries.flatten() {
78 let path = entry.path();
79 if let Some(name) = path.file_name().and_then(|n| n.to_str()) {
80 if Self::runtime_file_name_matches(name) && Self::usable_runtime_file(&path) {
81 return Some(path);
82 }
83 }
84 }
85
86 None
87 }
88
89 fn resolve_dylib_path() -> Option<PathBuf> {
94 if let Ok(path) = std::env::var("ORT_DYLIB_PATH") {
96 let p = PathBuf::from(&path);
97 if p.exists() {
98 tracing::info!("Using ONNX Runtime from ORT_DYLIB_PATH: {}", p.display());
99 return Some(p);
100 }
101 tracing::warn!("ORT_DYLIB_PATH set but file not found: {}", path);
102 }
103
104 if let Some(candidate) = crate::bootstrap::onnx_runtime_path() {
108 tracing::info!(
109 "Using ONNX Runtime from optional asset-pack path: {}",
110 candidate.display()
111 );
112 return Some(candidate);
113 }
114
115 let rel_dir = Path::new("libs").join("onnxruntime");
116
117 if let Ok(exe) = std::env::current_exe() {
122 if let Some(exe_dir) = exe.parent() {
123 for base in [exe_dir.to_path_buf(), exe_dir.join("..").join("..")] {
124 let candidate_dir = base.join(&rel_dir);
125 if let Some(candidate) = Self::find_runtime_in_dir(&candidate_dir) {
126 tracing::info!(
127 "Using ONNX Runtime from exe-relative path: {}",
128 candidate.display()
129 );
130 return Some(candidate);
131 }
132 }
133 }
134 }
135
136 if let Ok(cwd) = std::env::current_dir() {
140 let mut bases: Vec<PathBuf> = vec![cwd.clone()];
141 if let Some(parent) = cwd.parent() {
142 bases.push(parent.to_path_buf());
143 }
144 for base in bases {
145 let candidate_dir = base.join(&rel_dir);
146 if let Some(candidate) = Self::find_runtime_in_dir(&candidate_dir) {
147 tracing::info!(
148 "Using ONNX Runtime from cwd-relative path: {}",
149 candidate.display()
150 );
151 return Some(candidate);
152 }
153 }
154 }
155
156 None
157 }
158
159 pub fn init() -> ort::Result<Self> {
166 if ORT_INITIALIZED.get().is_some() {
167 return Ok(Self);
168 }
169
170 let _guard = ORT_INIT_LOCK.lock().map_err(|_| {
171 ort::Error::new("ONNX Runtime initialization lock is poisoned".to_string())
172 })?;
173 if ORT_INITIALIZED.get().is_some() {
174 return Ok(Self);
175 }
176
177 if let Some(dylib_path) = Self::resolve_dylib_path() {
178 ort::init_from(&dylib_path)?.commit();
179 let _ = ORT_INITIALIZED.set(());
180 tracing::info!(
181 "ONNX Runtime initialized (dynamic) from: {}",
182 dylib_path.display()
183 );
184 } else {
185 return Err(ort::Error::new(
186 format!(
187 "ONNX Runtime library not found. Set ORT_DYLIB_PATH or place {} (1.23.x) in libs/onnxruntime/",
188 ORT_LIB_NAME
189 ),
190 ));
191 }
192 Ok(Self)
193 }
194
195 pub fn load_model_with_threads<P: AsRef<Path>>(
209 &self,
210 path: P,
211 intra_threads: usize,
212 ) -> ort::Result<Session> {
213 let mut providers: Vec<ort::execution_providers::ExecutionProviderDispatch> = Vec::new();
217
218 #[cfg(target_os = "windows")]
220 {
221 providers.push(ort::execution_providers::DirectMLExecutionProvider::default().build());
222 }
223
224 #[cfg(target_os = "linux")]
225 {
226 providers.push(ort::execution_providers::CUDAExecutionProvider::default().build());
227
228 providers.push(
233 ort::execution_providers::OpenVINO::default()
234 .with_device_type("GPU")
235 .build(),
236 );
237 }
238
239 #[cfg(target_os = "macos")]
240 {
241 providers.push(ort::execution_providers::CoreMLExecutionProvider::default().build());
242 }
243
244 providers.push(
247 ort::execution_providers::OneDNN::default()
248 .with_arena_allocator(true)
249 .build(),
250 );
251 providers.push(ort::execution_providers::XNNPACK::default().build());
252
253 providers.push(ort::execution_providers::CPUExecutionProvider::default().build());
255
256 if !EP_LOGGED.swap(true, Ordering::Relaxed) {
257 Self::probe_and_log_providers();
258 }
259
260 let session = Session::builder()?
261 .with_optimization_level(GraphOptimizationLevel::Level3)?
262 .with_intra_threads(intra_threads)?
263 .with_execution_providers(providers)?
264 .commit_from_file(path)?;
265
266 Ok(session)
267 }
268
269 fn probe_and_log_providers() {
272 let mut available: Vec<&'static str> = Vec::new();
273
274 #[cfg(target_os = "windows")]
275 {
276 let ep = ort::execution_providers::DirectMLExecutionProvider::default();
277 if ep.is_available().unwrap_or(false) {
278 available.push("DirectML");
279 }
280 }
281
282 #[cfg(target_os = "linux")]
283 {
284 let ep = ort::execution_providers::CUDAExecutionProvider::default();
285 if ep.is_available().unwrap_or(false) {
286 available.push("CUDA");
287 }
288
289 let ep = ort::execution_providers::OpenVINO::default();
290 if ep.is_available().unwrap_or(false) {
291 available.push("OpenVINO");
292 }
293 }
294
295 #[cfg(target_os = "macos")]
296 {
297 let ep = ort::execution_providers::CoreMLExecutionProvider::default();
298 if ep.is_available().unwrap_or(false) {
299 available.push("CoreML");
300 }
301 }
302
303 {
304 let ep = ort::execution_providers::OneDNN::default();
305 if ep.is_available().unwrap_or(false) {
306 available.push("OneDNN");
307 }
308 }
309
310 {
311 let ep = ort::execution_providers::XNNPACK::default();
312 if ep.is_available().unwrap_or(false) {
313 available.push("XNNPACK");
314 }
315 }
316
317 available.push("CPU");
318
319 let chosen = *available.first().unwrap_or(&"CPU");
322 let _ = ACTIVE_PROVIDER.set(chosen);
323
324 if available.len() > 1 {
325 tracing::info!(
326 "Face inference will try execution providers in order: {}",
327 available.join(" -> ")
328 );
329 } else {
330 tracing::info!("Face inference will run on CPU (no GPU providers available)");
331 }
332 }
333}