use std::collections::VecDeque;
use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering};
use std::sync::{Arc, Condvar, Mutex, Weak};
use std::thread::{self, JoinHandle};
type Job = Box<dyn FnOnce() + Send + 'static>;
struct Shared {
global: Mutex<VecDeque<Job>>,
has_work: Condvar,
workers: Mutex<Vec<Weak<Worker>>>,
shutdown: AtomicBool,
park_seq: AtomicUsize,
}
struct Worker {
id: usize,
shared: Arc<Shared>,
local: Mutex<VecDeque<Job>>,
idle_seq: AtomicUsize,
}
impl Worker {
fn run(&self) {
loop {
if self.shared.shutdown.load(Ordering::Acquire) {
return;
}
if let Some(job) = self.local.lock().unwrap().pop_back() {
run_job(job);
continue;
}
{
let mut g = self.shared.global.lock().unwrap();
if let Some(job) = g.pop_front() {
drop(g);
run_job(job);
continue;
}
}
if let Some(job) = self.steal() {
run_job(job);
continue;
}
self.idle_seq.store(
self.shared.park_seq.fetch_add(1, Ordering::Relaxed) + 1,
Ordering::Relaxed,
);
let mut g = self.shared.global.lock().unwrap();
while g.is_empty() && !self.shared.shutdown.load(Ordering::Acquire) {
g = self.shared.has_work.wait(g).unwrap();
}
}
}
fn steal(&self) -> Option<Job> {
let workers = {
let w = self.shared.workers.lock().unwrap();
if w.len() <= 1 {
return None;
}
w.iter()
.filter_map(|x| x.upgrade())
.filter(|x| x.id != self.id)
.collect::<Vec<Arc<Worker>>>()
};
if workers.is_empty() {
return None;
}
let victim = workers
.iter()
.min_by_key(|w| w.idle_seq.load(Ordering::Relaxed))
.unwrap();
let stolen = victim.local.lock().unwrap().pop_front();
stolen
}
}
fn run_job(job: Job) {
let result = std::panic::catch_unwind(std::panic::AssertUnwindSafe(job));
if let Err(p) = result {
let msg = p
.downcast_ref::<&str>()
.map(|s| (*s).to_string())
.or_else(|| p.downcast_ref::<String>().cloned())
.unwrap_or_else(|| "unknown panic".to_string());
eprintln!("courierust pool: job panicked: {msg}");
}
}
pub struct ThreadPool {
shared: Arc<Shared>,
workers: Arc<Vec<Arc<Worker>>>,
handles: Vec<JoinHandle<()>>,
}
impl Default for ThreadPool {
fn default() -> Self {
Self::new().unwrap_or_else(|_| Self::with_size(1).expect("pool"))
}
}
impl ThreadPool {
pub fn new() -> std::io::Result<Self> {
let count = thread::available_parallelism()
.map(|n| n.get())
.unwrap_or(4);
Self::with_size(count)
}
pub fn with_size(size: usize) -> std::io::Result<Self> {
let size = size.max(1);
let shared = Arc::new(Shared {
global: Mutex::new(VecDeque::new()),
has_work: Condvar::new(),
workers: Mutex::new(Vec::new()),
shutdown: AtomicBool::new(false),
park_seq: AtomicUsize::new(0),
});
let mut handles = Vec::with_capacity(size);
let mut worker_arcs = Vec::with_capacity(size);
for id in 0..size {
let worker = Arc::new(Worker {
id,
shared: shared.clone(),
local: Mutex::new(VecDeque::new()),
idle_seq: AtomicUsize::new(0),
});
worker_arcs.push(worker.clone());
shared.workers.lock().unwrap().push(Arc::downgrade(&worker));
let w2 = worker.clone();
handles.push(
thread::Builder::new()
.name(format!("courierust-worker-{id}"))
.spawn(move || w2.run())?,
);
}
Ok(Self {
shared,
workers: Arc::new(worker_arcs),
handles,
})
}
pub fn len(&self) -> usize {
self.workers.len()
}
pub fn is_empty(&self) -> bool {
self.workers.is_empty()
}
pub fn spawn<F>(&self, f: F)
where
F: FnOnce() + Send + 'static,
{
if self.shared.shutdown.load(Ordering::Acquire) {
return;
}
self.shared.global.lock().unwrap().push_back(Box::new(f));
self.shared.has_work.notify_one();
}
}
impl Drop for ThreadPool {
fn drop(&mut self) {
self.shared.shutdown.store(true, Ordering::Release);
self.shared.has_work.notify_all();
for h in self.handles.drain(..) {
let _ = h.join();
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::time::Duration;
#[test]
fn runs_jobs() {
let pool = ThreadPool::new().unwrap();
let counter = Arc::new(AtomicUsize::new(0));
let n = 128;
for _ in 0..n {
let c = counter.clone();
pool.spawn(move || {
c.fetch_add(1, Ordering::SeqCst);
});
}
let deadline = std::time::Instant::now() + Duration::from_secs(10);
while counter.load(Ordering::SeqCst) < n && std::time::Instant::now() < deadline {
thread::sleep(Duration::from_millis(2));
}
assert_eq!(counter.load(Ordering::SeqCst), n);
}
#[test]
fn jobs_can_spawn_jobs() {
let pool = Arc::new(ThreadPool::with_size(4).unwrap());
let counter = Arc::new(AtomicUsize::new(0));
let n = 32;
for _ in 0..n {
let c = counter.clone();
let p_inner = pool.clone();
pool.spawn(move || {
for _ in 0..4 {
let c = c.clone();
let p2 = p_inner.clone();
p2.spawn(move || {
c.fetch_add(1, Ordering::SeqCst);
});
}
});
}
let deadline = std::time::Instant::now() + Duration::from_secs(10);
while counter.load(Ordering::SeqCst) < n * 4 && std::time::Instant::now() < deadline {
thread::sleep(Duration::from_millis(2));
}
assert_eq!(counter.load(Ordering::SeqCst), n * 4);
}
}