1use std::sync::{Arc, Mutex, OnceLock, mpsc};
2use std::thread::JoinHandle;
3
4type Job = Box<dyn FnOnce() + Send + 'static>;
5
6pub struct WorkerPool {
8 sender: Option<mpsc::Sender<Job>>,
9 workers: Vec<JoinHandle<()>>,
10}
11
12impl WorkerPool {
13 pub fn new(size: usize) -> Self {
15 let (sender, receiver) = mpsc::channel::<Job>();
16 let receiver = Arc::new(Mutex::new(receiver));
17 let mut workers = Vec::with_capacity(size.max(1));
18 for _ in 0..size.max(1) {
19 let rx = receiver.clone();
20 workers.push(std::thread::spawn(move || {
21 loop {
22 let job = match rx.lock() {
23 Ok(receiver) => receiver.recv(),
24 Err(_) => break,
25 };
26 match job {
27 Ok(job) => job(),
28 Err(_) => break,
29 }
30 }
31 }));
32 }
33 Self {
34 sender: Some(sender),
35 workers,
36 }
37 }
38
39 pub fn execute<F: FnOnce() + Send + 'static>(&self, f: F) {
41 if let Some(sender) = &self.sender {
42 let _ = sender.send(Box::new(f));
43 }
44 }
45}
46
47impl Drop for WorkerPool {
48 fn drop(&mut self) {
49 self.sender.take();
50 for worker in self.workers.drain(..) {
51 let _ = worker.join();
52 }
53 }
54}
55
56pub fn default_worker_pool() -> &'static WorkerPool {
58 static POOL: OnceLock<WorkerPool> = OnceLock::new();
59 POOL.get_or_init(|| {
60 let size = std::thread::available_parallelism()
61 .map(|parallelism| parallelism.get())
62 .unwrap_or(1);
63 WorkerPool::new(size)
64 })
65}
66
67#[cfg(test)]
68mod tests {
69 use std::sync::mpsc;
70
71 use super::WorkerPool;
72
73 #[test]
74 fn worker_pool_executes_jobs_and_joins_on_drop() {
75 let (tx, rx) = mpsc::channel();
76 {
77 let pool = WorkerPool::new(2);
78 pool.execute(move || {
79 let _ = tx.send("done");
80 });
81 assert_eq!(rx.recv().unwrap(), "done");
82 }
83 }
84}