use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering};
use std::sync::{Once, OnceLock};
use serde::{Deserialize, Serialize};
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
#[serde(default, deny_unknown_fields)]
pub struct ConcurrencyConfig {
pub max_threads: Option<usize>,
}
static POOL_INIT: Once = Once::new();
static ACTIVE_THREAD_BUDGET: AtomicUsize = AtomicUsize::new(0);
const DEFAULT_THREAD_CAP: usize = 8;
static DEFAULT_CAP_WARNED: AtomicBool = AtomicBool::new(false);
fn warn_default_thread_cap_once(already_warned: &AtomicBool, host_cpus: usize) {
if already_warned
.compare_exchange(false, true, Ordering::SeqCst, Ordering::SeqCst)
.is_ok()
{
tracing::warn!(
host_cpus,
thread_cap = DEFAULT_THREAD_CAP,
"detected {host_cpus} CPU cores but no `max_threads` is configured and no cgroup CPU \
quota was found; capping the thread budget at {DEFAULT_THREAD_CAP} \
(min(cpu_cores, {DEFAULT_THREAD_CAP})). Set `ConcurrencyConfig::max_threads` above \
{DEFAULT_THREAD_CAP} to use the remaining cores."
);
}
}
static CGROUP_QUOTA_CORES: OnceLock<Option<usize>> = OnceLock::new();
fn cgroup_cpu_quota_cores() -> Option<usize> {
*CGROUP_QUOTA_CORES.get_or_init(read_cgroup_cpu_quota_cores)
}
#[cfg(target_os = "linux")]
fn read_cgroup_cpu_quota_cores() -> Option<usize> {
cgroup_v2_quota_cores().or_else(cgroup_v1_quota_cores)
}
#[cfg(not(target_os = "linux"))]
fn read_cgroup_cpu_quota_cores() -> Option<usize> {
None
}
#[cfg(target_os = "linux")]
fn cgroup_v2_quota_cores() -> Option<usize> {
let contents = std::fs::read_to_string("/sys/fs/cgroup/cpu.max").ok()?;
let mut fields = contents.split_whitespace();
let quota_field = fields.next()?;
let period_field = fields.next()?;
if quota_field == "max" {
return None;
}
quota_period_to_cores(quota_field.parse().ok()?, period_field.parse().ok()?)
}
#[cfg(target_os = "linux")]
fn cgroup_v1_quota_cores() -> Option<usize> {
let quota: f64 = std::fs::read_to_string("/sys/fs/cgroup/cpu/cpu.cfs_quota_us")
.ok()?
.trim()
.parse()
.ok()?;
let period: f64 = std::fs::read_to_string("/sys/fs/cgroup/cpu/cpu.cfs_period_us")
.ok()?
.trim()
.parse()
.ok()?;
quota_period_to_cores(quota, period)
}
#[cfg(target_os = "linux")]
fn quota_period_to_cores(quota: f64, period: f64) -> Option<usize> {
if quota <= 0.0 || period <= 0.0 {
return None;
}
Some((quota / period).ceil().max(1.0) as usize)
}
pub(crate) fn resolve_thread_budget(config: Option<&ConcurrencyConfig>) -> usize {
resolve_thread_budget_inner(config, num_cpus::get(), cgroup_cpu_quota_cores())
}
#[cfg(feature = "captioning")]
pub(crate) fn resolve_llm_concurrency(
llm_config: &crate::core::config::LlmConfig,
concurrency: Option<&ConcurrencyConfig>,
) -> usize {
llm_config
.max_concurrency
.unwrap_or_else(|| resolve_thread_budget(concurrency))
.max(1)
}
fn resolve_thread_budget_inner(
config: Option<&ConcurrencyConfig>,
host_cpus: usize,
quota_cores: Option<usize>,
) -> usize {
resolve_thread_budget_with_guard(config, host_cpus, quota_cores, &DEFAULT_CAP_WARNED)
}
fn resolve_thread_budget_with_guard(
config: Option<&ConcurrencyConfig>,
host_cpus: usize,
quota_cores: Option<usize>,
already_warned: &AtomicBool,
) -> usize {
if let Some(n) = config.and_then(|c| c.max_threads) {
return n.max(1);
}
match quota_cores {
Some(quota_cores) => host_cpus.clamp(1, quota_cores.max(1)),
None => {
if host_cpus > DEFAULT_THREAD_CAP {
warn_default_thread_cap_once(already_warned, host_cpus);
}
host_cpus.clamp(1, DEFAULT_THREAD_CAP)
}
}
}
#[cfg(all(
not(target_arch = "wasm32"),
any(
test,
feature = "tokio-runtime",
feature = "late-interaction",
feature = "reranker",
feature = "sparse-embeddings"
)
))]
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(crate) struct BatchExecutionPlan {
pub workers: usize,
pub thread_budget: usize,
}
#[cfg(all(
not(target_arch = "wasm32"),
any(
test,
feature = "tokio-runtime",
feature = "late-interaction",
feature = "reranker",
feature = "sparse-embeddings"
)
))]
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(crate) enum LayoutBatchWorkload {
None,
#[cfg(layout_detection)]
Mixed,
#[cfg(layout_detection)]
All,
}
#[cfg(all(
not(target_arch = "wasm32"),
any(
test,
feature = "tokio-runtime",
feature = "late-interaction",
feature = "reranker",
feature = "sparse-embeddings"
)
))]
pub(crate) fn resolve_batch_execution_plan(
config: Option<&ConcurrencyConfig>,
layout_workload: LayoutBatchWorkload,
input_count: usize,
max_concurrent: Option<usize>,
) -> BatchExecutionPlan {
#[cfg(layout_detection)]
const MAX_NATIVE_LAYOUT_BATCH_WORKERS: usize = 1;
#[cfg(layout_detection)]
const MAX_MIXED_LAYOUT_BATCH_WORKERS: usize = 2;
let total_budget = resolve_thread_budget(config);
let available_inputs = input_count.max(1);
let worker_ceiling = max_concurrent
.unwrap_or(total_budget)
.max(1)
.min(total_budget)
.min(available_inputs);
let workers = match layout_workload {
LayoutBatchWorkload::None => worker_ceiling,
#[cfg(layout_detection)]
LayoutBatchWorkload::Mixed => worker_ceiling.min(MAX_MIXED_LAYOUT_BATCH_WORKERS),
#[cfg(layout_detection)]
LayoutBatchWorkload::All => worker_ceiling.min(MAX_NATIVE_LAYOUT_BATCH_WORKERS),
}
.max(1);
let thread_budget = (total_budget / workers).max(1);
debug_assert!(workers * thread_budget <= total_budget);
BatchExecutionPlan { workers, thread_budget }
}
#[cfg(all(
not(target_arch = "wasm32"),
any(feature = "late-interaction", feature = "reranker", feature = "sparse-embeddings")
))]
pub(crate) fn resolve_batch_concurrency(config: Option<&ConcurrencyConfig>, model_threads_active: bool) -> usize {
let budget = resolve_thread_budget(config);
if !model_threads_active {
return budget;
}
let cores = num_cpus::get().max(1);
(cores / budget).max(1).min(budget)
}
pub(crate) fn init_thread_pools(budget: usize) {
POOL_INIT.call_once(|| {
ACTIVE_THREAD_BUDGET.store(budget.max(1), Ordering::Relaxed);
#[cfg(not(target_arch = "wasm32"))]
if let Err(_err) = rayon::ThreadPoolBuilder::new().num_threads(budget).build_global() {
tracing::debug!(
budget,
"global rayon pool already initialized; reusing the existing pool \
(xberg thread budget not applied)"
);
}
#[cfg(target_arch = "wasm32")]
let _ = budget;
});
}
#[cfg(sceptre_ocr)]
pub(crate) fn active_thread_budget() -> usize {
match ACTIVE_THREAD_BUDGET.load(Ordering::Relaxed) {
0 => resolve_thread_budget(None),
budget => budget,
}
}
#[cfg(all(feature = "tokio-runtime", not(target_arch = "wasm32")))]
pub(crate) fn init_batch_thread_pool(config: Option<&ConcurrencyConfig>) -> usize {
let total_budget = resolve_thread_budget(config);
init_thread_pools(total_budget);
total_budget
}
#[cfg(test)]
mod tests {
use std::sync::{Arc, Mutex};
use tracing_subscriber::layer::SubscriberExt as _;
use tracing_subscriber::{EnvFilter, Layer};
use super::*;
#[cfg(feature = "captioning")]
#[test]
fn llm_concurrency_overrides_general_thread_budget() {
let llm = crate::core::config::LlmConfig {
max_concurrency: Some(3),
..Default::default()
};
let general = ConcurrencyConfig { max_threads: Some(12) };
assert_eq!(resolve_llm_concurrency(&llm, Some(&general)), 3);
}
#[cfg(feature = "captioning")]
#[test]
fn llm_concurrency_falls_back_to_general_thread_budget() {
let llm = crate::core::config::LlmConfig::default();
let general = ConcurrencyConfig { max_threads: Some(5) };
assert_eq!(resolve_llm_concurrency(&llm, Some(&general)), 5);
}
#[derive(Clone, Default)]
struct EventCapture {
levels: Arc<Mutex<Vec<tracing::Level>>>,
}
impl<S> Layer<S> for EventCapture
where
S: tracing::Subscriber,
{
fn on_event(&self, event: &tracing::Event<'_>, _ctx: tracing_subscriber::layer::Context<'_, S>) {
self.levels.lock().unwrap().push(*event.metadata().level());
}
}
fn warn_event_count(capture: &EventCapture) -> usize {
capture
.levels
.lock()
.unwrap()
.iter()
.filter(|level| **level == tracing::Level::WARN)
.count()
}
#[test]
fn test_resolve_thread_budget_none() {
assert_eq!(resolve_thread_budget_inner(None, 16, None), 8);
let budget = resolve_thread_budget(None);
assert!(budget >= 1, "the host always gets at least one thread");
}
#[test]
fn test_inner_pins_default_cap_when_no_quota_and_no_max_threads() {
assert_eq!(resolve_thread_budget_inner(None, 1, None), 1);
assert_eq!(resolve_thread_budget_inner(None, 4, None), 4);
assert_eq!(resolve_thread_budget_inner(None, 8, None), 8);
assert_eq!(resolve_thread_budget_inner(None, 16, None), 8);
assert_eq!(resolve_thread_budget_inner(None, 64, None), 8);
}
#[test]
fn test_inner_explicit_max_threads_wins_over_host_cpus_and_quota() {
let config = ConcurrencyConfig { max_threads: Some(20) };
assert_eq!(resolve_thread_budget_inner(Some(&config), 4, Some(2)), 20);
assert_eq!(resolve_thread_budget_inner(Some(&config), 64, None), 20);
}
#[test]
fn test_inner_explicit_max_threads_of_zero_clamps_to_one() {
let config = ConcurrencyConfig { max_threads: Some(0) };
assert_eq!(resolve_thread_budget_inner(Some(&config), 16, None), 1);
}
#[test]
fn test_inner_cgroup_quota_above_default_cap_is_not_clamped_to_eight() {
assert_eq!(resolve_thread_budget_inner(None, 64, Some(24)), 24);
assert_eq!(resolve_thread_budget_inner(None, 64, Some(9)), 9);
}
#[test]
fn test_inner_cgroup_quota_below_default_cap_is_used_as_is() {
assert_eq!(resolve_thread_budget_inner(None, 64, Some(3)), 3);
}
#[test]
fn test_inner_cgroup_quota_never_exceeds_host_cpus() {
assert_eq!(resolve_thread_budget_inner(None, 4, Some(16)), 4);
}
#[test]
#[serial_test::serial]
fn test_default_cap_warning_fires_exactly_once_when_cores_exceed_cap_and_unset() {
let already_warned = AtomicBool::new(false);
let capture = EventCapture::default();
let subscriber = tracing_subscriber::registry()
.with(EnvFilter::new("warn"))
.with(capture.clone());
tracing::subscriber::with_default(subscriber, || {
for _ in 0..5 {
assert_eq!(resolve_thread_budget_with_guard(None, 16, None, &already_warned), 8);
}
});
assert_eq!(
warn_event_count(&capture),
1,
"expected exactly one WARN event across repeated calls, got {:?}",
capture.levels.lock().unwrap()
);
}
#[test]
#[serial_test::serial]
fn test_default_cap_warning_does_not_fire_when_max_threads_is_set() {
let already_warned = AtomicBool::new(false);
let capture = EventCapture::default();
let subscriber = tracing_subscriber::registry()
.with(EnvFilter::new("warn"))
.with(capture.clone());
tracing::subscriber::with_default(subscriber, || {
let config = ConcurrencyConfig { max_threads: Some(4) };
resolve_thread_budget_with_guard(Some(&config), 16, None, &already_warned)
});
assert_eq!(warn_event_count(&capture), 0);
}
#[test]
#[serial_test::serial]
fn test_default_cap_warning_does_not_fire_when_cores_at_or_below_default_cap() {
let already_warned = AtomicBool::new(false);
let capture = EventCapture::default();
let subscriber = tracing_subscriber::registry()
.with(EnvFilter::new("warn"))
.with(capture.clone());
tracing::subscriber::with_default(subscriber, || {
resolve_thread_budget_with_guard(None, 8, None, &already_warned);
resolve_thread_budget_with_guard(None, 1, None, &already_warned);
});
assert_eq!(warn_event_count(&capture), 0);
}
#[test]
#[serial_test::serial]
fn test_default_cap_warning_does_not_fire_when_cgroup_quota_present() {
let already_warned = AtomicBool::new(false);
let capture = EventCapture::default();
let subscriber = tracing_subscriber::registry()
.with(EnvFilter::new("warn"))
.with(capture.clone());
tracing::subscriber::with_default(subscriber, || {
resolve_thread_budget_with_guard(None, 16, Some(12), &already_warned);
});
assert_eq!(warn_event_count(&capture), 0);
}
#[test]
fn test_resolve_thread_budget_with_config() {
let config = ConcurrencyConfig { max_threads: Some(4) };
assert_eq!(resolve_thread_budget(Some(&config)), 4);
}
#[test]
fn test_resolve_thread_budget_clamps_to_one() {
let config = ConcurrencyConfig { max_threads: Some(0) };
assert_eq!(resolve_thread_budget(Some(&config)), 1);
}
#[test]
fn test_resolve_thread_budget_no_max() {
let config = ConcurrencyConfig { max_threads: None };
assert_eq!(resolve_thread_budget_inner(Some(&config), 16, None), 8);
let budget = resolve_thread_budget(Some(&config));
assert!(budget >= 1, "the host always gets at least one thread");
}
#[test]
#[cfg(not(target_arch = "wasm32"))]
fn test_batch_plan_without_layout_uses_available_budget() {
let budget = resolve_thread_budget(None);
assert_eq!(
resolve_batch_execution_plan(None, LayoutBatchWorkload::None, budget, None),
BatchExecutionPlan {
workers: budget,
thread_budget: 1,
}
);
}
#[test]
#[cfg(all(not(target_arch = "wasm32"), layout_detection))]
fn test_layout_batch_plan_table() {
for budget in [1, 2, 4, 8] {
let config = ConcurrencyConfig {
max_threads: Some(budget),
};
assert_eq!(
resolve_batch_execution_plan(Some(&config), LayoutBatchWorkload::All, 16, None),
BatchExecutionPlan {
workers: 1,
thread_budget: budget,
}
);
}
}
#[test]
#[cfg(all(not(target_arch = "wasm32"), layout_detection))]
fn test_mixed_layout_batch_preserves_two_worker_cap() {
for (budget, workers, thread_budget) in [(1, 1, 1), (2, 2, 1), (4, 2, 2), (8, 2, 4)] {
let config = ConcurrencyConfig {
max_threads: Some(budget),
};
assert_eq!(
resolve_batch_execution_plan(Some(&config), LayoutBatchWorkload::Mixed, 16, None),
BatchExecutionPlan { workers, thread_budget }
);
}
}
#[test]
#[cfg(all(not(target_arch = "wasm32"), layout_detection))]
fn test_layout_batch_plan_respects_input_and_explicit_limits() {
let config = ConcurrencyConfig { max_threads: Some(8) };
assert_eq!(
resolve_batch_execution_plan(Some(&config), LayoutBatchWorkload::All, 1, Some(8)),
BatchExecutionPlan {
workers: 1,
thread_budget: 8,
}
);
assert_eq!(
resolve_batch_execution_plan(Some(&config), LayoutBatchWorkload::All, 8, Some(1)),
BatchExecutionPlan {
workers: 1,
thread_budget: 8,
}
);
}
#[test]
#[cfg(not(target_arch = "wasm32"))]
fn test_non_layout_batch_plan_divides_budget_at_explicit_worker_limit() {
let config = ConcurrencyConfig { max_threads: Some(8) };
let plan = resolve_batch_execution_plan(Some(&config), LayoutBatchWorkload::None, 16, Some(2));
assert_eq!(plan.workers, 2);
assert_eq!(plan.thread_budget, 4);
}
#[test]
#[cfg(not(target_arch = "wasm32"))]
fn test_non_layout_batch_plan_clamps_explicit_limit_to_total_budget() {
let config = ConcurrencyConfig { max_threads: Some(2) };
let plan = resolve_batch_execution_plan(Some(&config), LayoutBatchWorkload::None, 8, Some(6));
assert_eq!(plan.workers, 2);
assert_eq!(plan.thread_budget, 1);
}
#[test]
#[cfg(not(target_arch = "wasm32"))]
fn test_non_layout_batch_plan_gives_single_input_full_inner_budget() {
let config = ConcurrencyConfig { max_threads: Some(8) };
let plan = resolve_batch_execution_plan(Some(&config), LayoutBatchWorkload::None, 1, None);
assert_eq!(plan.workers, 1);
assert_eq!(plan.thread_budget, 8);
}
#[test]
#[cfg(not(target_arch = "wasm32"))]
fn test_batch_plan_never_exceeds_total_budget() {
for total_budget in 1..=8 {
let config = ConcurrencyConfig {
max_threads: Some(total_budget),
};
for input_count in 0..=12 {
for max_concurrent in [None, Some(0), Some(1), Some(3), Some(16)] {
#[cfg(layout_detection)]
let layout_workloads = [
LayoutBatchWorkload::None,
LayoutBatchWorkload::Mixed,
LayoutBatchWorkload::All,
];
#[cfg(not(layout_detection))]
let layout_workloads = [LayoutBatchWorkload::None];
for layout_workload in layout_workloads {
let plan =
resolve_batch_execution_plan(Some(&config), layout_workload, input_count, max_concurrent);
assert!(plan.workers * plan.thread_budget <= total_budget);
assert!(plan.workers <= total_budget);
assert!(plan.workers <= input_count.max(1));
if let Some(explicit) = max_concurrent {
assert!(plan.workers <= explicit.max(1));
}
}
}
}
}
}
#[test]
fn test_init_thread_pools_idempotent() {
init_thread_pools(2);
init_thread_pools(4);
}
#[test]
#[cfg(all(feature = "tokio-runtime", not(target_arch = "wasm32")))]
fn test_batch_thread_pool_uses_total_configured_budget() {
let config = ConcurrencyConfig { max_threads: Some(7) };
assert_eq!(init_batch_thread_pool(Some(&config)), 7);
}
#[test]
fn test_default() {
let config = ConcurrencyConfig::default();
assert!(config.max_threads.is_none());
}
#[test]
fn test_serde_roundtrip() {
let json = r#"{"max_threads": 2}"#;
let config: ConcurrencyConfig = serde_json::from_str(json).unwrap();
assert_eq!(config.max_threads, Some(2));
let serialized = serde_json::to_string(&config).unwrap();
let roundtripped: ConcurrencyConfig = serde_json::from_str(&serialized).unwrap();
assert_eq!(roundtripped.max_threads, Some(2));
}
#[test]
fn test_serde_empty() {
let json = r#"{}"#;
let config: ConcurrencyConfig = serde_json::from_str(json).unwrap();
assert!(config.max_threads.is_none());
}
}