use std::sync::OnceLock;
use std::thread::JoinHandle;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum SchedPolicy {
Rr,
Fifo,
Nice,
None,
}
impl SchedPolicy {
pub fn parse(s: &str) -> Result<Self, String> {
match s.trim().to_ascii_lowercase().as_str() {
"rr" => Ok(Self::Rr),
"fifo" => Ok(Self::Fifo),
"nice" => Ok(Self::Nice),
"none" | "plain" | "" => Ok(Self::None),
other => Err(format!(
"unknown thread scheduling policy '{other}' (use rr|fifo|nice|none)"
)),
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum Affinity {
Auto,
Core(usize),
Off,
}
impl Affinity {
pub fn parse(s: &str) -> Result<Self, String> {
let t = s.trim().to_ascii_lowercase();
match t.as_str() {
"auto" => Ok(Self::Auto),
"off" | "none" | "" => Ok(Self::Off),
_ => t
.parse::<usize>()
.map(Self::Core)
.map_err(|_| format!("unknown thread affinity '{s}' (use auto|off|<core>)")),
}
}
}
#[derive(Debug, Clone, Copy)]
pub struct PoolSpec {
pub threads: usize,
pub sched: SchedPolicy,
pub pin: Affinity,
}
#[derive(Debug, Default, Clone)]
pub struct CliThreadOverrides {
pub timing: Option<String>,
pub io: Option<String>,
pub workers: Option<String>,
pub timing_sched: Option<String>,
pub timing_pin: Option<String>,
}
#[derive(Debug, Clone, Copy)]
pub struct ThreadPoolConfig {
pub timing: PoolSpec,
pub io: PoolSpec,
pub workers: usize,
}
impl ThreadPoolConfig {
pub fn defaults() -> Self {
let cores = std::thread::available_parallelism()
.map(|n| n.get())
.unwrap_or(1);
let reserved = if cores >= 3 { 1 } else { 0 };
let workers = cores.saturating_sub(reserved).max(1);
Self {
timing: PoolSpec {
threads: 1,
sched: SchedPolicy::Rr,
pin: Affinity::Auto,
},
io: PoolSpec {
threads: 2,
sched: SchedPolicy::None,
pin: Affinity::Off,
},
workers,
}
}
pub fn resolve() -> Result<Self, String> {
let mut cfg = Self::defaults();
let env_usize = |k: &str| -> Result<Option<usize>, String> {
match std::env::var(k) {
Ok(v) => v
.trim()
.parse::<usize>()
.map(Some)
.map_err(|_| format!("{k}='{v}' is not a thread count")),
Err(_) => Ok(None),
}
};
if let Some(n) = env_usize("NMBRS_THREADS_TIMING")? {
cfg.timing.threads = n;
}
if let Some(n) = env_usize("NMBRS_THREADS_IO")? {
cfg.io.threads = n;
}
if let Some(n) = env_usize("NMBRS_THREADS_WORKERS")? {
cfg.workers = n.max(1);
}
if let Ok(v) = std::env::var("NMBRS_THREADS_TIMING_SCHED") {
cfg.timing.sched = SchedPolicy::parse(&v)?;
}
if let Ok(v) = std::env::var("NMBRS_THREADS_TIMING_PIN") {
cfg.timing.pin = Affinity::parse(&v)?;
}
Ok(cfg)
}
pub fn resolve_with_cli(cli: &CliThreadOverrides) -> Result<Self, String> {
let mut cfg = Self::resolve()?;
let usize_flag = |v: &str, flag: &str| -> Result<usize, String> {
v.trim()
.parse::<usize>()
.map_err(|_| format!("{flag}='{v}' is not a thread count"))
};
if let Some(v) = &cli.timing {
cfg.timing.threads = usize_flag(v, "--threads.timing")?;
}
if let Some(v) = &cli.io {
cfg.io.threads = usize_flag(v, "--threads.io")?;
}
if let Some(v) = &cli.workers {
cfg.workers = usize_flag(v, "--threads.workers")?.max(1);
}
if let Some(v) = &cli.timing_sched {
cfg.timing.sched = SchedPolicy::parse(v)?;
}
if let Some(v) = &cli.timing_pin {
cfg.timing.pin = Affinity::parse(v)?;
}
Ok(cfg)
}
fn spec(&self, pool: &str) -> Result<PoolSpec, String> {
match pool {
"timing" => Ok(self.timing),
"io" => Ok(self.io),
other => Err(format!(
"unknown thread pool '{other}' (named pools: timing, io, workers)"
)),
}
}
}
pub struct ThreadPools {
config: ThreadPoolConfig,
top_core: usize,
}
static GLOBAL: OnceLock<ThreadPools> = OnceLock::new();
pub fn init(config: ThreadPoolConfig) {
let _ = GLOBAL.set(ThreadPools::new(config));
}
pub fn global() -> &'static ThreadPools {
GLOBAL.get_or_init(|| ThreadPools::new(ThreadPoolConfig::defaults()))
}
impl ThreadPools {
fn new(config: ThreadPoolConfig) -> Self {
let cores = std::thread::available_parallelism()
.map(|n| n.get())
.unwrap_or(1);
Self {
config,
top_core: cores.saturating_sub(1),
}
}
pub fn config(&self) -> &ThreadPoolConfig {
&self.config
}
fn resolve_pin(&self, spec: &PoolSpec) -> Option<usize> {
match spec.pin {
Affinity::Auto => Some(self.top_core),
Affinity::Core(c) => Some(c),
Affinity::Off => None,
}
}
pub fn spawn(
&self,
pool: &str,
name: &str,
f: impl FnOnce() + Send + 'static,
) -> Result<JoinHandle<()>, String> {
let spec = self.config.spec(pool)?;
let pin = self.resolve_pin(&spec);
let sched = spec.sched;
let pool_owned = pool.to_string();
let name_owned = name.to_string();
std::thread::Builder::new()
.name(format!("{pool}-{name}"))
.spawn(move || {
let achieved = apply_policy(sched, pin);
let line = format!("thread pool '{pool_owned}/{name_owned}': {achieved}");
if achieved.contains("denied") {
crate::diag::warn(&line);
} else {
crate::diag::info(&line);
}
f();
})
.map_err(|e| format!("spawn thread pool '{pool}/{name}': {e}"))
}
pub fn spawn_timing(
&self,
name: &str,
f: impl FnOnce() + Send + 'static,
) -> Result<JoinHandle<()>, String> {
self.spawn("timing", name, f)
}
}
#[cfg(target_os = "linux")]
fn apply_policy(sched: SchedPolicy, pin: Option<usize>) -> String {
let mut parts: Vec<String> = Vec::new();
match sched {
SchedPolicy::Rr | SchedPolicy::Fifo => {
let policy = if matches!(sched, SchedPolicy::Fifo) {
libc::SCHED_FIFO
} else {
libc::SCHED_RR
};
let prio = 10; let param = libc::sched_param {
sched_priority: prio,
};
let rc = unsafe { libc::sched_setscheduler(0, policy, ¶m) };
if rc == 0 {
let name = if policy == libc::SCHED_FIFO {
"fifo"
} else {
"rr"
};
parts.push(format!("sched={name}(prio {prio})"));
} else {
let rc2 = unsafe { libc::setpriority(libc::PRIO_PROCESS, 0, -5) };
if rc2 == 0 {
parts.push("sched=nice(-5) [realtime denied]".to_string());
} else {
parts.push("sched=plain [realtime+nice denied]".to_string());
}
}
}
SchedPolicy::Nice => {
let rc = unsafe { libc::setpriority(libc::PRIO_PROCESS, 0, -5) };
parts.push(if rc == 0 {
"sched=nice(-5)".to_string()
} else {
"sched=plain [nice denied]".to_string()
});
}
SchedPolicy::None => parts.push("sched=plain".to_string()),
}
if let Some(core) = pin {
let rc = unsafe {
let mut set: libc::cpu_set_t = std::mem::zeroed();
libc::CPU_ZERO(&mut set);
libc::CPU_SET(core, &mut set);
libc::sched_setaffinity(0, std::mem::size_of::<libc::cpu_set_t>(), &set)
};
parts.push(if rc == 0 {
format!("pin={core}")
} else {
format!("pin={core} [denied]")
});
}
parts.join(", ")
}
#[cfg(not(target_os = "linux"))]
fn apply_policy(_sched: SchedPolicy, _pin: Option<usize>) -> String {
"sched=plain [non-linux]".to_string()
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn defaults_reserve_a_core_when_affordable() {
let cfg = ThreadPoolConfig::defaults();
assert_eq!(cfg.timing.threads, 1);
assert_eq!(cfg.timing.sched, SchedPolicy::Rr);
assert!(matches!(cfg.timing.pin, Affinity::Auto));
assert!(cfg.workers >= 1);
}
#[test]
fn unknown_policy_and_pool_are_hard_errors() {
assert!(SchedPolicy::parse("bogus").is_err());
assert!(Affinity::parse("nonsense").is_err());
assert!(ThreadPoolConfig::defaults().spec("nope").is_err());
}
#[test]
fn policy_and_affinity_parse_roundtrip() {
assert_eq!(SchedPolicy::parse("RR").unwrap(), SchedPolicy::Rr);
assert_eq!(SchedPolicy::parse("none").unwrap(), SchedPolicy::None);
assert_eq!(Affinity::parse("auto").unwrap(), Affinity::Auto);
assert_eq!(Affinity::parse("3").unwrap(), Affinity::Core(3));
}
#[test]
fn cli_overrides_apply_and_win() {
let cli = CliThreadOverrides {
timing: Some("4".into()),
io: Some("3".into()),
workers: Some("7".into()),
timing_sched: Some("none".into()),
timing_pin: Some("off".into()),
..Default::default()
};
let cfg = ThreadPoolConfig::resolve_with_cli(&cli).unwrap();
assert_eq!(cfg.timing.threads, 4);
assert_eq!(cfg.io.threads, 3);
assert_eq!(cfg.workers, 7);
assert_eq!(cfg.timing.sched, SchedPolicy::None);
assert_eq!(cfg.timing.pin, Affinity::Off);
}
#[test]
fn cli_override_absent_leaves_resolved_value() {
let cfg = ThreadPoolConfig::resolve_with_cli(&CliThreadOverrides::default()).unwrap();
assert_eq!(cfg.timing.threads, 1);
assert_eq!(cfg.timing.sched, SchedPolicy::Rr);
}
#[test]
fn cli_malformed_value_is_a_hard_error() {
let bad_count = CliThreadOverrides {
timing: Some("abc".into()),
..Default::default()
};
let e = ThreadPoolConfig::resolve_with_cli(&bad_count).unwrap_err();
assert!(
e.contains("--threads.timing"),
"message names the flag: {e}"
);
assert!(
e.contains("not a thread count"),
"message explains why: {e}"
);
let bad_sched = CliThreadOverrides {
timing_sched: Some("bogus".into()),
..Default::default()
};
assert!(ThreadPoolConfig::resolve_with_cli(&bad_sched).is_err());
let bad_pin = CliThreadOverrides {
timing_pin: Some("nonsense".into()),
..Default::default()
};
assert!(ThreadPoolConfig::resolve_with_cli(&bad_pin).is_err());
}
#[test]
fn spawn_on_timing_runs_the_closure() {
let pools = ThreadPools::new(ThreadPoolConfig::defaults());
let (tx, rx) = std::sync::mpsc::channel();
let h = pools
.spawn_timing("unit", move || {
let _ = tx.send(42u8);
})
.unwrap();
assert_eq!(rx.recv().unwrap(), 42);
h.join().unwrap();
}
}