1use std::sync::OnceLock;
21use std::thread::JoinHandle;
22
23#[derive(Debug, Clone, Copy, PartialEq, Eq)]
25pub enum SchedPolicy {
26 Rr,
29 Fifo,
31 Nice,
33 None,
35}
36
37impl SchedPolicy {
38 pub fn parse(s: &str) -> Result<Self, String> {
41 match s.trim().to_ascii_lowercase().as_str() {
42 "rr" => Ok(Self::Rr),
43 "fifo" => Ok(Self::Fifo),
44 "nice" => Ok(Self::Nice),
45 "none" | "plain" | "" => Ok(Self::None),
46 other => Err(format!(
47 "unknown thread scheduling policy '{other}' (use rr|fifo|nice|none)"
48 )),
49 }
50 }
51}
52
53#[derive(Debug, Clone, Copy, PartialEq, Eq)]
55pub enum Affinity {
56 Auto,
58 Core(usize),
60 Off,
62}
63
64impl Affinity {
65 pub fn parse(s: &str) -> Result<Self, String> {
66 let t = s.trim().to_ascii_lowercase();
67 match t.as_str() {
68 "auto" => Ok(Self::Auto),
69 "off" | "none" | "" => Ok(Self::Off),
70 _ => t
71 .parse::<usize>()
72 .map(Self::Core)
73 .map_err(|_| format!("unknown thread affinity '{s}' (use auto|off|<core>)")),
74 }
75 }
76}
77
78#[derive(Debug, Clone, Copy)]
80pub struct PoolSpec {
81 pub threads: usize,
82 pub sched: SchedPolicy,
83 pub pin: Affinity,
84}
85
86#[derive(Debug, Default, Clone)]
95pub struct CliThreadOverrides {
96 pub timing: Option<String>,
98 pub io: Option<String>,
100 pub workers: Option<String>,
102 pub timing_sched: Option<String>,
104 pub timing_pin: Option<String>,
106}
107
108#[derive(Debug, Clone, Copy)]
110pub struct ThreadPoolConfig {
111 pub timing: PoolSpec,
113 pub io: PoolSpec,
115 pub workers: usize,
118}
119
120impl ThreadPoolConfig {
121 pub fn defaults() -> Self {
126 let cores = std::thread::available_parallelism()
127 .map(|n| n.get())
128 .unwrap_or(1);
129 let reserved = if cores >= 3 { 1 } else { 0 };
131 let workers = cores.saturating_sub(reserved).max(1);
132 Self {
133 timing: PoolSpec {
134 threads: 1,
135 sched: SchedPolicy::Rr,
136 pin: Affinity::Auto,
137 },
138 io: PoolSpec {
139 threads: 2,
140 sched: SchedPolicy::None,
141 pin: Affinity::Off,
142 },
143 workers,
144 }
145 }
146
147 pub fn resolve() -> Result<Self, String> {
151 let mut cfg = Self::defaults();
152 let env_usize = |k: &str| -> Result<Option<usize>, String> {
153 match std::env::var(k) {
154 Ok(v) => v
155 .trim()
156 .parse::<usize>()
157 .map(Some)
158 .map_err(|_| format!("{k}='{v}' is not a thread count")),
159 Err(_) => Ok(None),
160 }
161 };
162 if let Some(n) = env_usize("NMBRS_THREADS_TIMING")? {
163 cfg.timing.threads = n;
164 }
165 if let Some(n) = env_usize("NMBRS_THREADS_IO")? {
166 cfg.io.threads = n;
167 }
168 if let Some(n) = env_usize("NMBRS_THREADS_WORKERS")? {
169 cfg.workers = n.max(1);
170 }
171 if let Ok(v) = std::env::var("NMBRS_THREADS_TIMING_SCHED") {
172 cfg.timing.sched = SchedPolicy::parse(&v)?;
173 }
174 if let Ok(v) = std::env::var("NMBRS_THREADS_TIMING_PIN") {
175 cfg.timing.pin = Affinity::parse(&v)?;
176 }
177 Ok(cfg)
178 }
179
180 pub fn resolve_with_cli(cli: &CliThreadOverrides) -> Result<Self, String> {
186 let mut cfg = Self::resolve()?;
187 let usize_flag = |v: &str, flag: &str| -> Result<usize, String> {
188 v.trim()
189 .parse::<usize>()
190 .map_err(|_| format!("{flag}='{v}' is not a thread count"))
191 };
192 if let Some(v) = &cli.timing {
193 cfg.timing.threads = usize_flag(v, "--threads.timing")?;
194 }
195 if let Some(v) = &cli.io {
196 cfg.io.threads = usize_flag(v, "--threads.io")?;
197 }
198 if let Some(v) = &cli.workers {
199 cfg.workers = usize_flag(v, "--threads.workers")?.max(1);
200 }
201 if let Some(v) = &cli.timing_sched {
202 cfg.timing.sched = SchedPolicy::parse(v)?;
203 }
204 if let Some(v) = &cli.timing_pin {
205 cfg.timing.pin = Affinity::parse(v)?;
206 }
207 Ok(cfg)
208 }
209
210 fn spec(&self, pool: &str) -> Result<PoolSpec, String> {
211 match pool {
212 "timing" => Ok(self.timing),
213 "io" => Ok(self.io),
214 other => Err(format!(
215 "unknown thread pool '{other}' (named pools: timing, io, workers)"
216 )),
217 }
218 }
219}
220
221pub struct ThreadPools {
224 config: ThreadPoolConfig,
225 top_core: usize,
228}
229
230static GLOBAL: OnceLock<ThreadPools> = OnceLock::new();
231
232pub fn init(config: ThreadPoolConfig) {
235 let _ = GLOBAL.set(ThreadPools::new(config));
236}
237
238pub fn global() -> &'static ThreadPools {
241 GLOBAL.get_or_init(|| ThreadPools::new(ThreadPoolConfig::defaults()))
242}
243
244impl ThreadPools {
245 fn new(config: ThreadPoolConfig) -> Self {
246 let cores = std::thread::available_parallelism()
247 .map(|n| n.get())
248 .unwrap_or(1);
249 Self {
250 config,
251 top_core: cores.saturating_sub(1),
252 }
253 }
254
255 pub fn config(&self) -> &ThreadPoolConfig {
256 &self.config
257 }
258
259 fn resolve_pin(&self, spec: &PoolSpec) -> Option<usize> {
262 match spec.pin {
263 Affinity::Auto => Some(self.top_core),
264 Affinity::Core(c) => Some(c),
265 Affinity::Off => None,
266 }
267 }
268
269 pub fn spawn(
274 &self,
275 pool: &str,
276 name: &str,
277 f: impl FnOnce() + Send + 'static,
278 ) -> Result<JoinHandle<()>, String> {
279 let spec = self.config.spec(pool)?;
280 let pin = self.resolve_pin(&spec);
281 let sched = spec.sched;
282 let pool_owned = pool.to_string();
283 let name_owned = name.to_string();
284 std::thread::Builder::new()
285 .name(format!("{pool}-{name}"))
286 .spawn(move || {
287 let achieved = apply_policy(sched, pin);
288 let line = format!("thread pool '{pool_owned}/{name_owned}': {achieved}");
289 if achieved.contains("denied") {
290 crate::diag::warn(&line);
291 } else {
292 crate::diag::info(&line);
293 }
294 f();
295 })
296 .map_err(|e| format!("spawn thread pool '{pool}/{name}': {e}"))
297 }
298
299 pub fn spawn_timing(
301 &self,
302 name: &str,
303 f: impl FnOnce() + Send + 'static,
304 ) -> Result<JoinHandle<()>, String> {
305 self.spawn("timing", name, f)
306 }
307}
308
309#[cfg(target_os = "linux")]
313fn apply_policy(sched: SchedPolicy, pin: Option<usize>) -> String {
314 let mut parts: Vec<String> = Vec::new();
315 match sched {
316 SchedPolicy::Rr | SchedPolicy::Fifo => {
317 let policy = if matches!(sched, SchedPolicy::Fifo) {
318 libc::SCHED_FIFO
319 } else {
320 libc::SCHED_RR
321 };
322 let prio = 10; let param = libc::sched_param {
324 sched_priority: prio,
325 };
326 let rc = unsafe { libc::sched_setscheduler(0, policy, ¶m) };
328 if rc == 0 {
329 let name = if policy == libc::SCHED_FIFO {
330 "fifo"
331 } else {
332 "rr"
333 };
334 parts.push(format!("sched={name}(prio {prio})"));
335 } else {
336 let rc2 = unsafe { libc::setpriority(libc::PRIO_PROCESS, 0, -5) };
338 if rc2 == 0 {
339 parts.push("sched=nice(-5) [realtime denied]".to_string());
340 } else {
341 parts.push("sched=plain [realtime+nice denied]".to_string());
342 }
343 }
344 }
345 SchedPolicy::Nice => {
346 let rc = unsafe { libc::setpriority(libc::PRIO_PROCESS, 0, -5) };
347 parts.push(if rc == 0 {
348 "sched=nice(-5)".to_string()
349 } else {
350 "sched=plain [nice denied]".to_string()
351 });
352 }
353 SchedPolicy::None => parts.push("sched=plain".to_string()),
354 }
355 if let Some(core) = pin {
356 let rc = unsafe {
357 let mut set: libc::cpu_set_t = std::mem::zeroed();
358 libc::CPU_ZERO(&mut set);
359 libc::CPU_SET(core, &mut set);
360 libc::sched_setaffinity(0, std::mem::size_of::<libc::cpu_set_t>(), &set)
361 };
362 parts.push(if rc == 0 {
363 format!("pin={core}")
364 } else {
365 format!("pin={core} [denied]")
366 });
367 }
368 parts.join(", ")
369}
370
371#[cfg(not(target_os = "linux"))]
372fn apply_policy(_sched: SchedPolicy, _pin: Option<usize>) -> String {
373 "sched=plain [non-linux]".to_string()
375}
376
377#[cfg(test)]
378mod tests {
379 use super::*;
380
381 #[test]
382 fn defaults_reserve_a_core_when_affordable() {
383 let cfg = ThreadPoolConfig::defaults();
384 assert_eq!(cfg.timing.threads, 1);
385 assert_eq!(cfg.timing.sched, SchedPolicy::Rr);
386 assert!(matches!(cfg.timing.pin, Affinity::Auto));
387 assert!(cfg.workers >= 1);
388 }
389
390 #[test]
391 fn unknown_policy_and_pool_are_hard_errors() {
392 assert!(SchedPolicy::parse("bogus").is_err());
393 assert!(Affinity::parse("nonsense").is_err());
394 assert!(ThreadPoolConfig::defaults().spec("nope").is_err());
395 }
396
397 #[test]
398 fn policy_and_affinity_parse_roundtrip() {
399 assert_eq!(SchedPolicy::parse("RR").unwrap(), SchedPolicy::Rr);
400 assert_eq!(SchedPolicy::parse("none").unwrap(), SchedPolicy::None);
401 assert_eq!(Affinity::parse("auto").unwrap(), Affinity::Auto);
402 assert_eq!(Affinity::parse("3").unwrap(), Affinity::Core(3));
403 }
404
405 #[test]
406 fn cli_overrides_apply_and_win() {
407 let cli = CliThreadOverrides {
409 timing: Some("4".into()),
410 io: Some("3".into()),
411 workers: Some("7".into()),
412 timing_sched: Some("none".into()),
413 timing_pin: Some("off".into()),
414 ..Default::default()
415 };
416 let cfg = ThreadPoolConfig::resolve_with_cli(&cli).unwrap();
417 assert_eq!(cfg.timing.threads, 4);
418 assert_eq!(cfg.io.threads, 3);
419 assert_eq!(cfg.workers, 7);
420 assert_eq!(cfg.timing.sched, SchedPolicy::None);
421 assert_eq!(cfg.timing.pin, Affinity::Off);
422 }
423
424 #[test]
425 fn cli_override_absent_leaves_resolved_value() {
426 let cfg = ThreadPoolConfig::resolve_with_cli(&CliThreadOverrides::default()).unwrap();
429 assert_eq!(cfg.timing.threads, 1);
430 assert_eq!(cfg.timing.sched, SchedPolicy::Rr);
431 }
432
433 #[test]
434 fn cli_malformed_value_is_a_hard_error() {
435 let bad_count = CliThreadOverrides {
436 timing: Some("abc".into()),
437 ..Default::default()
438 };
439 let e = ThreadPoolConfig::resolve_with_cli(&bad_count).unwrap_err();
440 assert!(
441 e.contains("--threads.timing"),
442 "message names the flag: {e}"
443 );
444 assert!(
445 e.contains("not a thread count"),
446 "message explains why: {e}"
447 );
448
449 let bad_sched = CliThreadOverrides {
450 timing_sched: Some("bogus".into()),
451 ..Default::default()
452 };
453 assert!(ThreadPoolConfig::resolve_with_cli(&bad_sched).is_err());
454
455 let bad_pin = CliThreadOverrides {
456 timing_pin: Some("nonsense".into()),
457 ..Default::default()
458 };
459 assert!(ThreadPoolConfig::resolve_with_cli(&bad_pin).is_err());
460 }
461
462 #[test]
463 fn spawn_on_timing_runs_the_closure() {
464 let pools = ThreadPools::new(ThreadPoolConfig::defaults());
465 let (tx, rx) = std::sync::mpsc::channel();
466 let h = pools
467 .spawn_timing("unit", move || {
468 let _ = tx.send(42u8);
469 })
470 .unwrap();
471 assert_eq!(rx.recv().unwrap(), 42);
472 h.join().unwrap();
473 }
474}