#![forbid(unsafe_code)]
use std::collections::HashMap;
use std::future::Future;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::{Arc, OnceLock};
use tokio::sync::{OwnedSemaphorePermit, Semaphore};
use tokio::task::{Id as TaskId, JoinError, JoinSet};
pub use crate::constants::MIN_CONCURRENCY;
pub use crate::constants::HARD_CAP;
pub use crate::constants::IO_OVERSUBSCRIBE;
pub use crate::constants::NON_LINUX_CPU_CAP;
pub use crate::constants::RAM_PER_TASK_BYTES;
const RAM_SAFETY_NUM: u64 = 1;
const RAM_SAFETY_DEN: u64 = 2;
static PROCESS_LIMIT: OnceLock<usize> = OnceLock::new();
static FAIL_FAST: std::sync::atomic::AtomicBool = std::sync::atomic::AtomicBool::new(false);
static SCP_FILE_CONCURRENCY: OnceLock<usize> = OnceLock::new();
static PEAK_IN_FLIGHT: AtomicUsize = AtomicUsize::new(0);
static CURRENT_IN_FLIGHT: AtomicUsize = AtomicUsize::new(0);
pub fn install_process_limit(limit: usize) {
let capped = limit.clamp(MIN_CONCURRENCY, HARD_CAP);
let _ = PROCESS_LIMIT.set(capped);
tracing::debug!(max_concurrency = capped, "installed process concurrency limit");
}
pub fn install_fail_fast(enabled: bool) {
FAIL_FAST.store(enabled, Ordering::Relaxed);
if enabled {
tracing::debug!("installed fail-fast multi-host policy");
}
}
#[must_use]
pub fn fail_fast_enabled() -> bool {
FAIL_FAST.load(Ordering::Relaxed)
}
pub fn install_scp_file_concurrency(n: usize) {
let capped = n.clamp(MIN_CONCURRENCY, HARD_CAP);
let _ = SCP_FILE_CONCURRENCY.set(capped);
tracing::debug!(scp_file_concurrency = capped, "installed scp file concurrency");
}
#[must_use]
pub fn scp_file_concurrency() -> usize {
SCP_FILE_CONCURRENCY.get().copied().unwrap_or(MIN_CONCURRENCY)
}
#[must_use]
pub fn effective_limit() -> usize {
if let Some(&n) = PROCESS_LIMIT.get() {
return n;
}
resolve_limit(None)
}
#[must_use]
pub fn resolve_limit(cli_override: Option<usize>) -> usize {
if let Some(n) = cli_override {
return n.clamp(MIN_CONCURRENCY, HARD_CAP);
}
auto_limit()
}
#[must_use]
pub fn auto_limit() -> usize {
let cpus = std::thread::available_parallelism()
.map(|n| n.get())
.unwrap_or(2);
let cpu_budget = cpus.saturating_mul(IO_OVERSUBSCRIBE).max(MIN_CONCURRENCY);
let ram_budget = match free_ram_bytes() {
Some(free) => {
let usable = free.saturating_mul(RAM_SAFETY_NUM) / RAM_SAFETY_DEN;
let tasks = usize::try_from(usable / RAM_PER_TASK_BYTES.max(1)).unwrap_or(usize::MAX);
tasks.max(MIN_CONCURRENCY)
}
None => cpu_budget.clamp(MIN_CONCURRENCY, NON_LINUX_CPU_CAP),
};
cpu_budget.min(ram_budget).clamp(MIN_CONCURRENCY, HARD_CAP)
}
#[must_use]
pub fn worker_threads() -> usize {
let cpus = std::thread::available_parallelism()
.map(|n| n.get())
.unwrap_or(2);
let budget = resolve_limit(None);
budget.min(cpus).clamp(2, 16)
}
#[must_use]
pub fn max_blocking_threads() -> usize {
resolve_limit(None).clamp(2, 32)
}
#[must_use]
pub fn free_ram_bytes() -> Option<u64> {
#[cfg(target_os = "linux")]
{
let text = std::fs::read_to_string("/proc/meminfo").ok()?;
for line in text.lines() {
if let Some(rest) = line.strip_prefix("MemAvailable:") {
let kb: u64 = rest
.split_whitespace()
.next()?
.parse()
.ok()?;
return Some(kb.saturating_mul(1024));
}
}
None
}
#[cfg(not(target_os = "linux"))]
{
None
}
}
#[must_use]
pub fn semaphore(limit: usize) -> Arc<Semaphore> {
Arc::new(Semaphore::new(limit.clamp(MIN_CONCURRENCY, HARD_CAP)))
}
pub async fn acquire_owned(sem: &Arc<Semaphore>) -> OwnedSemaphorePermit {
match Arc::clone(sem).acquire_owned().await {
Ok(p) => p,
Err(_) => {
tracing::error!(
"concurrency semaphore was closed unexpectedly; admitting via ephemeral permit (G-SEC-03)"
);
loop {
let emergency = Arc::new(Semaphore::new(1));
if let Ok(p) = emergency.acquire_owned().await {
return p;
}
tokio::task::yield_now().await;
}
}
}
}
#[must_use]
pub fn peak_in_flight() -> usize {
PEAK_IN_FLIGHT.load(Ordering::Relaxed)
}
pub fn reset_peak_counters() {
PEAK_IN_FLIGHT.store(0, Ordering::Relaxed);
CURRENT_IN_FLIGHT.store(0, Ordering::Relaxed);
}
fn track_enter() {
let cur = CURRENT_IN_FLIGHT.fetch_add(1, Ordering::Relaxed) + 1;
PEAK_IN_FLIGHT.fetch_max(cur, Ordering::Relaxed);
}
fn track_leave() {
CURRENT_IN_FLIGHT.fetch_sub(1, Ordering::Relaxed);
}
struct InFlightGuard;
impl Drop for InFlightGuard {
fn drop(&mut self) {
track_leave();
}
}
#[derive(Debug)]
pub struct IndexedResult<R> {
pub index: usize,
pub outcome: Result<R, JoinError>,
}
pub async fn map_bounded<T, R, F, Fut>(
items: Vec<T>,
limit: usize,
work: F,
) -> Vec<IndexedResult<R>>
where
T: Send + 'static,
R: Send + 'static,
F: Fn(T) -> Fut + Send + Sync + 'static,
Fut: Future<Output = R> + Send + 'static,
{
map_bounded_with(items, limit, work, |_r| false).await
}
pub async fn map_bounded_with<T, R, F, Fut, P>(
items: Vec<T>,
limit: usize,
work: F,
is_failure: P,
) -> Vec<IndexedResult<R>>
where
T: Send + 'static,
R: Send + 'static,
F: Fn(T) -> Fut + Send + Sync + 'static,
Fut: Future<Output = R> + Send + 'static,
P: Fn(&R) -> bool + Send + Sync + 'static,
{
let limit = limit.clamp(MIN_CONCURRENCY, HARD_CAP);
let sem = semaphore(limit);
let work = Arc::new(work);
let is_failure = Arc::new(is_failure);
let mut set: JoinSet<R> = JoinSet::new();
let mut task_index: HashMap<TaskId, usize> = HashMap::new();
let mut iter = items.into_iter().enumerate();
let mut results: Vec<IndexedResult<R>> = Vec::new();
let mut admit = true;
while set.len() < limit {
if crate::signals::should_stop() {
admit = false;
tracing::debug!("fan-out: stop admission (should_stop) during seed");
break;
}
let Some((index, item)) = iter.next() else {
break;
};
spawn_one(&mut set, &mut task_index, &sem, &work, index, item).await;
}
loop {
if set.is_empty() {
if !admit || crate::signals::should_stop() {
break;
}
if let Some((index, item)) = iter.next() {
spawn_one(&mut set, &mut task_index, &sem, &work, index, item).await;
continue;
}
break;
}
tokio::select! {
joined = set.join_next_with_id() => {
match joined {
Some(j) => {
let before = results.len();
push_joined(&mut results, &mut task_index, j);
if fail_fast_enabled() {
if let Some(last) = results.get(before..) {
for r in last {
if let Ok(ref val) = r.outcome {
if is_failure(val) && admit {
admit = false;
tracing::debug!(
index = r.index,
"fan-out: fail-fast stop admission"
);
}
}
}
}
}
}
None => break,
}
if crate::signals::is_force_exit() {
tracing::debug!(remaining = set.len(), "fan-out: force_exit abort_all");
set.abort_all();
while let Some(j) = set.join_next_with_id().await {
push_joined(&mut results, &mut task_index, j);
}
break;
}
if crate::signals::should_stop() {
if admit {
admit = false;
tracing::debug!("fan-out: stop admission (should_stop); draining");
}
continue;
}
if admit {
if let Some((index, item)) = iter.next() {
spawn_one(&mut set, &mut task_index, &sem, &work, index, item).await;
}
}
}
_ = tokio::time::sleep(std::time::Duration::from_millis(
crate::constants::FAN_OUT_SIGNAL_POLL_INTERVAL_MS,
)) => {
if crate::signals::is_force_exit() {
tracing::debug!(remaining = set.len(), "fan-out: force_exit abort_all (timer)");
set.abort_all();
while let Some(j) = set.join_next_with_id().await {
push_joined(&mut results, &mut task_index, j);
}
break;
}
if admit && crate::signals::should_stop() {
admit = false;
tracing::debug!("fan-out: stop admission (should_stop via timer)");
}
}
}
}
results.sort_by_key(|r| r.index);
results
}
fn push_joined<R>(
results: &mut Vec<IndexedResult<R>>,
task_index: &mut HashMap<TaskId, usize>,
joined: Result<(TaskId, R), JoinError>,
) {
match joined {
Ok((id, value)) => {
let index = task_index.remove(&id).unwrap_or(usize::MAX);
results.push(IndexedResult {
index,
outcome: Ok(value),
});
}
Err(e) => {
let index = task_index.remove(&e.id()).unwrap_or(usize::MAX);
results.push(IndexedResult {
index,
outcome: Err(e),
});
}
}
}
async fn spawn_one<T, R, F, Fut>(
set: &mut JoinSet<R>,
task_index: &mut HashMap<TaskId, usize>,
sem: &Arc<Semaphore>,
work: &Arc<F>,
index: usize,
item: T,
) where
T: Send + 'static,
R: Send + 'static,
F: Fn(T) -> Fut + Send + Sync + 'static,
Fut: Future<Output = R> + Send + 'static,
{
use tracing::Instrument;
let permit = acquire_owned(sem).await;
let available = sem.available_permits();
tracing::debug!(index, available_permits = available, "fan-out admit");
let span = tracing::info_span!("fan_out_unit", index, available_permits = available);
let work = Arc::clone(work);
let abort = set.spawn(
async move {
track_enter();
let _inflight = InFlightGuard;
let _permit = permit;
work(item).await
}
.instrument(span),
);
task_index.insert(abort.id(), index);
}
pub async fn map_bounded_ok<T, R, F, Fut>(items: Vec<T>, limit: usize, work: F) -> Vec<R>
where
T: Send + 'static,
R: Send + 'static,
F: Fn(T) -> Fut + Send + Sync + 'static,
Fut: Future<Output = R> + Send + 'static,
{
let mut out = Vec::with_capacity(items.len());
for r in map_bounded(items, limit, work).await {
match r.outcome {
Ok(v) => out.push(v),
Err(e) if e.is_panic() => std::panic::resume_unwind(e.into_panic()),
Err(_) => {
}
}
}
out
}
#[cfg(test)]
#[path = "concurrency_tests.rs"]
mod tests;