use crossbeam_channel::{bounded, Receiver, Sender};
use rayon::ThreadPool as RayonPool;
use std::sync::{Arc, Condvar, Mutex};
type JobFn = Box<dyn FnOnce() + Send + 'static>;
struct PoolState {
pending: usize, }
pub struct TPool {
pool: Arc<RayonPool>,
slot_tx: Sender<()>,
slot_rx: Receiver<()>,
state: Arc<(Mutex<PoolState>, Condvar)>,
}
impl TPool {
pub fn new(nb_threads: usize, queue_size: usize) -> Option<Self> {
if nb_threads < 1 || queue_size < 1 {
return None;
}
let pool = rayon::ThreadPoolBuilder::new()
.num_threads(nb_threads)
.build()
.ok()?;
let capacity = queue_size + nb_threads;
let (slot_tx, slot_rx) = bounded(capacity);
for _ in 0..capacity {
slot_tx.send(()).ok()?;
}
let state = Arc::new((Mutex::new(PoolState { pending: 0 }), Condvar::new()));
Some(TPool {
pool: Arc::new(pool),
slot_tx,
slot_rx,
state,
})
}
pub fn submit_job(&self, job: JobFn) {
self.slot_rx.recv().expect("threadpool slot channel closed");
{
let (lock, _cvar) = &*self.state;
let mut s = lock.lock().unwrap();
s.pending += 1;
}
let state = Arc::clone(&self.state);
let slot_tx = self.slot_tx.clone();
self.pool.spawn(move || {
job();
let (lock, cvar) = &*state;
let mut s = lock.lock().unwrap();
s.pending -= 1;
if s.pending == 0 {
cvar.notify_all();
}
let _ = slot_tx.send(());
});
}
pub fn jobs_completed(&self) {
let (lock, cvar) = &*self.state;
let mut s = lock.lock().unwrap();
while s.pending > 0 {
s = cvar.wait(s).unwrap();
}
}
}
impl Drop for TPool {
fn drop(&mut self) {
self.jobs_completed();
}
}