Skip to main content

rightkit_ort/
probe.rs

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