#![forbid(unsafe_code)]
use std::{
num::NonZero,
sync::{mpsc, Arc, Mutex},
thread::{self, JoinHandle},
};
#[derive(Debug)]
pub struct ThreadPool {
workers: Vec<ThreadWorker>,
producer: Option<mpsc::Sender<ThreadJob>>,
}
impl ThreadPool {
pub fn builder() -> ThreadPoolBuilder {
ThreadPoolBuilder::default()
}
#[inline]
pub fn execute<F>(&self, f: F)
where
F: FnOnce() + Send + 'static,
{
let job = Box::new(f);
self.producer
.as_ref()
.expect("err acquiring sender ref")
.send(ThreadJob::Run(job))
.expect("send error")
}
pub fn join(&mut self) {
(0..self.workers.len()).for_each(|_| {
self.producer
.as_ref()
.unwrap()
.send(ThreadJob::Stop)
.unwrap();
});
drop(self.producer.take());
self.workers.iter_mut().for_each(|worker| {
if let Some(thread) = worker.thread.take() {
thread.join().unwrap();
}
});
}
#[doc(alias = "available_parallelism")]
#[doc(alias = "available_concurrency")]
#[doc(alias = "available_workers")]
#[doc(alias = "available_threads")]
pub fn num_threads(&self) -> usize {
self.workers.len()
}
}
impl Drop for ThreadPool {
fn drop(&mut self) {
if self.producer.is_some() {
self.join();
}
}
}
#[derive(Debug)]
pub struct ThreadPoolBuilder {
num_threads: NonZero<usize>,
stack_size: Option<usize>,
name_prefix: Option<String>,
}
impl Default for ThreadPoolBuilder {
fn default() -> ThreadPoolBuilder {
ThreadPoolBuilder {
num_threads: thread::available_parallelism().unwrap(),
stack_size: Option::default(),
name_prefix: Option::default(),
}
}
}
impl ThreadPoolBuilder {
pub fn new(num_threads: usize, stack_size: usize, name_prefix: String) -> ThreadPoolBuilder {
assert!(num_threads > 0);
ThreadPoolBuilder {
num_threads: NonZero::new(num_threads).unwrap(),
stack_size: Some(stack_size),
name_prefix: Some(name_prefix),
}
}
pub fn build(&self) -> ThreadPool {
let (producer, consumer) = mpsc::channel();
let consumer = Arc::new(Mutex::new(consumer));
let mut workers = Vec::with_capacity(self.num_threads.into());
(0..self.num_threads.into()).for_each(|id| {
let consumer = Arc::clone(&consumer);
let mut builder = thread::Builder::new();
if let Some(stack_size) = self.stack_size {
builder = builder.stack_size(stack_size);
}
if let Some(prefix) = &self.name_prefix {
builder = builder.name(format!("{}-{}", prefix, id));
}
let worker = ThreadWorker::new(id, consumer, builder);
workers.push(worker);
});
ThreadPool {
workers,
producer: Some(producer),
}
}
pub fn num_threads(mut self, num_threads: usize) -> ThreadPoolBuilder {
assert!(num_threads > 0);
self.num_threads = NonZero::new(num_threads).unwrap();
self
}
pub fn stack_size(mut self, stack_size: usize) -> ThreadPoolBuilder {
self.stack_size = Some(stack_size);
self
}
pub fn name_prefix(mut self, name_prefix: String) -> ThreadPoolBuilder {
self.name_prefix = Some(name_prefix);
self
}
}
enum ThreadJob {
Stop,
Run(Box<dyn FnOnce() + Send + 'static>),
}
#[derive(Debug)]
struct ThreadWorker {
id: usize,
thread: Option<JoinHandle<()>>,
}
impl ThreadWorker {
fn new(
id: usize,
consumer: Arc<Mutex<mpsc::Receiver<ThreadJob>>>,
builder: thread::Builder,
) -> ThreadWorker {
let thread = builder
.spawn(move || loop {
let job = consumer.lock().unwrap().recv().unwrap();
match job {
ThreadJob::Run(job) => job(),
ThreadJob::Stop => break,
};
})
.unwrap();
ThreadWorker {
id,
thread: Some(thread),
}
}
}
impl std::fmt::Display for ThreadWorker {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(f, "[{}]", self.id)
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::{
sync::{
atomic::{AtomicBool, Ordering},
Arc, Mutex,
},
thread, time,
};
mod helpers {
use super::*;
const TARGET: usize = 500_000_000;
pub fn get_sequential_speed() -> time::Duration {
let mut value = 0;
let start = time::Instant::now();
(0..TARGET).for_each(|_| {
value += 1;
});
start.elapsed()
}
pub fn get_parallel_speed() -> time::Duration {
let mut pool = ThreadPoolBuilder::default().build();
let num_threads = pool.num_threads();
let value = Arc::new(Mutex::new(0));
let start = time::Instant::now();
assert!(num_threads > 0);
(0..num_threads).for_each(|_| {
let value = Arc::clone(&value);
let mut ir = 0;
pool.execute(move || {
(0..TARGET / num_threads).for_each(|_| {
ir += 1;
});
let mut value = value.lock().unwrap();
*value += ir;
});
});
pool.join();
start.elapsed()
}
}
#[test]
fn construct_pool() {
let mut pool = ThreadPoolBuilder::default().build();
let p = Arc::new(Mutex::new(5));
let v = Arc::clone(&p);
pool.execute(move || {
let mut lock = v.lock().unwrap();
*lock += 1;
thread::sleep(time::Duration::from_secs(5));
});
pool.execute(|| {
thread::sleep(time::Duration::from_secs(10));
});
pool.join();
assert_eq!(*p.lock().unwrap(), 6);
}
#[test]
fn test_sequential_vs_parallel_speed() {
let sequential = helpers::get_sequential_speed();
let parallel = helpers::get_parallel_speed();
println!("sequential speed: {sequential:#?}\nparallel speed: {parallel:#?}");
assert!(sequential > parallel);
assert!(sequential > parallel / 2);
}
#[test]
fn test_join_disposal() {
use std::sync::atomic::{AtomicBool, Ordering};
let mut pool = ThreadPoolBuilder::default().num_threads(2).build();
let task_completed = Arc::new(AtomicBool::new(false));
let task_completed_clone = Arc::clone(&task_completed);
pool.execute(move || {
thread::sleep(time::Duration::from_millis(2500));
task_completed_clone.store(true, Ordering::SeqCst);
});
pool.execute(|| {
thread::sleep(time::Duration::from_secs(1));
});
pool.join();
assert!(
task_completed.load(Ordering::SeqCst),
"task not completed before shutdown"
);
assert!(
pool.producer.is_none(),
"producer isn't none after pool join"
);
}
#[test]
fn test_setup_builder_default() {
let pool = ThreadPoolBuilder::default();
assert_eq!(pool.num_threads, thread::available_parallelism().unwrap());
assert_eq!(pool.name_prefix, None);
assert_eq!(pool.stack_size, None);
}
#[test]
fn test_setup_builder_new() {
let pool = ThreadPoolBuilder::new(1, 5 * 1024, "PrivatePool".to_string());
assert_eq!(pool.num_threads, NonZero::new(1).unwrap());
assert_eq!(pool.stack_size, Some(5 * 1024));
assert_eq!(pool.name_prefix, Some("PrivatePool".to_string()));
}
#[test]
fn test_setup_builder_num_threads() {
let pool = ThreadPoolBuilder::default().num_threads(4).build();
assert_eq!(pool.num_threads(), 4);
}
#[test]
fn test_setup_builder_prefix_name() {
let pool = ThreadPoolBuilder::default().name_prefix("DarkPrivatisedPool".to_string());
assert_eq!(pool.name_prefix, Some("DarkPrivatisedPool".to_string()));
}
#[test]
fn test_setup_builder_stack_size() {
let pool = ThreadPoolBuilder::default().stack_size(5 * 1024 * 1024);
assert_eq!(pool.stack_size, Some(5 * 1024 * 1024));
}
#[test]
#[should_panic(expected = "err acquiring sender ref")]
fn test_execute_after_join_panics() {
let mut pool = ThreadPoolBuilder::default().num_threads(2).build();
pool.join();
pool.execute(|| {
println!("shouldn't execute");
});
}
}