rightkit_ort/
environment.rs1use 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#[derive(Debug, Clone, Copy, PartialEq, Eq)]
23pub struct GlobalPool {
24 pub intra_threads: usize,
25 pub inter_threads: usize,
26 pub spin: bool,
29}
30
31impl GlobalPool {
32 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 pub global_pool: Option<GlobalPool>,
48}
49
50#[derive(Debug, Clone, PartialEq, Eq)]
51pub struct EnvironmentReport {
52 pub runtime_path: std::path::PathBuf,
53 pub runtime_info: String,
55 pub shared_pool_requested: bool,
56 pub shared_pool_active: bool,
57}
58
59pub 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 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
101pub 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
120pub 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
127pub 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 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}