Skip to main content

rightkit_ort/
environment.rs

1//! Process-wide ORT environment initialisation (merged from HeardRight
2//! `heardright-onnx-asr/environment.rs`, product experiment switches removed).
3//!
4//! Call [`init_environment`] once, before any session is built or any worker
5//! thread starts. It loads the selected runtime explicitly, disables telemetry,
6//! optionally requests one shared global CPU thread pool, and materialises the
7//! C environment immediately (committing only stores Rust-side options).
8//! If a shared pool cannot be had, sessions fall back to private pools and the
9//! report says so; nothing is silently reconfigured.
10
11use std::sync::{
12    atomic::{AtomicBool, Ordering},
13    OnceLock,
14};
15
16use crate::RuntimeSelection;
17
18static REPORT: OnceLock<Result<EnvironmentReport, String>> = OnceLock::new();
19static PRIVATE_POOLS: AtomicBool = AtomicBool::new(false);
20
21/// Shared CPU pool request.
22#[derive(Debug, Clone, Copy, PartialEq, Eq)]
23pub struct GlobalPool {
24    pub intra_threads: usize,
25    pub inter_threads: usize,
26    /// Busy-wait in idle workers. Off by default: HeardRight measured spinning
27    /// as pure CPU cost for bursty dictation workloads.
28    pub spin: bool,
29}
30
31impl GlobalPool {
32    /// Budget derived from physical cores via [`cpu_thread_budget`], one inter
33    /// thread, no spinning.
34    pub fn from_budget() -> Self {
35        Self {
36            intra_threads: cpu_thread_budget(physical_cores().unwrap_or(1))
37                .min(std::thread::available_parallelism().map_or(1, usize::from)),
38            inter_threads: 1,
39            spin: false,
40        }
41    }
42}
43
44#[derive(Debug, Clone, Default)]
45pub struct EnvironmentOptions {
46    /// `None` keeps per-session private pools.
47    pub global_pool: Option<GlobalPool>,
48}
49
50#[derive(Debug, Clone, PartialEq, Eq)]
51pub struct EnvironmentReport {
52    pub runtime_path: std::path::PathBuf,
53    /// ORT build info string reported by the loaded library.
54    pub runtime_info: String,
55    pub shared_pool_requested: bool,
56    pub shared_pool_active: bool,
57}
58
59/// HeardRight's measured budget: leave headroom on small machines, cap at four.
60pub fn cpu_thread_budget(physical: usize) -> usize {
61    if physical <= 4 {
62        physical.saturating_sub(1).max(1)
63    } else {
64        4
65    }
66}
67
68#[cfg(target_os = "windows")]
69fn physical_cores() -> Option<usize> {
70    use windows_sys::Win32::System::SystemInformation::{
71        GetLogicalProcessorInformationEx, RelationProcessorCore,
72    };
73    let mut bytes = 0;
74    unsafe {
75        GetLogicalProcessorInformationEx(RelationProcessorCore, std::ptr::null_mut(), &mut bytes);
76    }
77    if bytes < 8 {
78        return None;
79    }
80    // u64 storage supplies required alignment; record sizes are variable.
81    let mut storage = vec![0_u64; (bytes as usize).div_ceil(8)];
82    if unsafe {
83        GetLogicalProcessorInformationEx(
84            RelationProcessorCore,
85            storage.as_mut_ptr().cast(),
86            &mut bytes,
87        )
88    } == 0
89    {
90        return None;
91    }
92    let data = unsafe { std::slice::from_raw_parts(storage.as_ptr().cast::<u8>(), bytes as usize) };
93    count_core_records(data)
94}
95
96#[cfg(not(target_os = "windows"))]
97fn physical_cores() -> Option<usize> {
98    std::thread::available_parallelism().ok().map(usize::from)
99}
100
101/// Count `RelationProcessorCore` records in a `GetLogicalProcessorInformationEx`
102/// buffer; rejects truncated or foreign records. Public so Windows hosts and
103/// tests can validate topology buffers.
104pub fn count_core_records(data: &[u8]) -> Option<usize> {
105    let mut offset = 0;
106    let mut cores = 0;
107    while offset < data.len() {
108        let header = data.get(offset..offset + 8)?;
109        let relation = u32::from_ne_bytes(header[..4].try_into().ok()?);
110        let size = u32::from_ne_bytes(header[4..].try_into().ok()?) as usize;
111        if relation != 0 || size < 8 || size > data.len() - offset {
112            return None;
113        }
114        cores += 1;
115        offset += size;
116    }
117    (cores > 0).then_some(cores)
118}
119
120/// True once [`init_environment`] succeeded with a shared pool and no session
121/// has since fallen back to private pools.
122pub fn shared_pool_active() -> bool {
123    matches!(REPORT.get(), Some(Ok(r)) if r.shared_pool_active)
124        && !PRIVATE_POOLS.load(Ordering::Acquire)
125}
126
127/// Initialise ORT exactly once for the process. Later calls return the first
128/// outcome unchanged (options of later calls are ignored).
129pub fn init_environment(
130    selection: &RuntimeSelection,
131    options: &EnvironmentOptions,
132) -> Result<EnvironmentReport, String> {
133    REPORT
134        .get_or_init(|| init_inner(selection, options))
135        .clone()
136}
137
138fn init_inner(
139    selection: &RuntimeSelection,
140    options: &EnvironmentOptions,
141) -> Result<EnvironmentReport, String> {
142    let mut builder = ort::init_from(&selection.path)
143        .map_err(|e| format!("load {}: {e}", selection.path.display()))?
144        .with_telemetry(false);
145    let mut shared_requested = false;
146    if let Some(pool) = options.global_pool {
147        shared_requested = true;
148        let pool_options = ort::environment::GlobalThreadPoolOptions::default()
149            .with_intra_threads(pool.intra_threads)
150            .map_err(|e| e.to_string())?
151            .with_inter_threads(pool.inter_threads)
152            .map_err(|e| e.to_string())?
153            .with_spin_control(pool.spin)
154            .map_err(|e| e.to_string())?;
155        builder = builder.with_global_thread_pool(pool_options);
156    }
157    if !builder.commit() {
158        return Err("ORT was configured before rightkit-ort::init_environment".into());
159    }
160    // Materialise the C environment now so failures surface at startup.
161    ort::environment::Environment::current().map_err(|e| format!("environment creation: {e}"))?;
162    Ok(EnvironmentReport {
163        runtime_path: selection.path.clone(),
164        runtime_info: ort::info().to_owned(),
165        shared_pool_requested: shared_requested,
166        shared_pool_active: shared_requested,
167    })
168}