Skip to main content

rightkit_ort/
probe.rs

1//! Execution-provider availability probe against the loaded runtime.
2
3use ort::ep::ExecutionProvider as _;
4
5#[derive(Debug, Clone, Copy, PartialEq, Eq)]
6pub enum ProviderKind {
7    DirectMl,
8    CoreMl,
9    Cpu,
10}
11
12#[derive(Debug, Clone, PartialEq, Eq)]
13pub struct ProviderDiagnostic {
14    pub provider: ProviderKind,
15    pub available: bool,
16    pub detail: String,
17}
18
19/// Ask the loaded runtime which execution providers it was built with and
20/// whether this OS can host them. Does not create a session. A provider is
21/// `available` only when the platform supports it AND the library was compiled
22/// with it ("compiled in" is not "usable for this model"; sessions use
23/// `error_on_failure`, so a failed registration is reported, never hidden).
24///
25/// The runtime must already be configured (`configure_runtime` or
26/// `init_environment`); a missing library is reported instead of panicking.
27pub fn probe_providers() -> Vec<ProviderDiagnostic> {
28    let mut out = vec![ProviderDiagnostic {
29        provider: ProviderKind::Cpu,
30        available: true,
31        detail: "always compiled in".into(),
32    }];
33    let coreml = ort::ep::CoreML::default();
34    let directml = ort::ep::DirectML::default();
35    for (kind, supported, available) in [
36        (
37            ProviderKind::CoreMl,
38            coreml.supported_by_platform(),
39            guarded(|| coreml.is_available()),
40        ),
41        (
42            ProviderKind::DirectMl,
43            directml.supported_by_platform(),
44            guarded(|| directml.is_available()),
45        ),
46    ] {
47        let (available, detail) = match (supported, available) {
48            (false, _) => (false, "not supported on this OS".to_string()),
49            (true, Ok(true)) => (true, "compiled into the loaded runtime".to_string()),
50            (true, Ok(false)) => (false, "not compiled into the loaded runtime".to_string()),
51            (true, Err(e)) => (false, format!("probe failed: {e}")),
52        };
53        out.push(ProviderDiagnostic {
54            provider: kind,
55            available,
56            detail,
57        });
58    }
59    out
60}
61
62fn guarded(f: impl FnOnce() -> ort::Result<bool> + std::panic::UnwindSafe) -> Result<bool, String> {
63    match std::panic::catch_unwind(f) {
64        Ok(result) => result.map_err(|e| e.to_string()),
65        Err(_) => Err("ONNX Runtime library could not be loaded".into()),
66    }
67}