#![forbid(unsafe_code)]
mod atomic_counter;
use atomic_counter::AtomicCounter;
use core::fmt::{Debug, Display, Formatter};
use core::time::Duration;
use std::error::Error;
use std::io::ErrorKind;
use std::sync::mpsc::{Receiver, RecvTimeoutError, SyncSender, TrySendError};
use std::sync::{Arc, Mutex};
use std::time::Instant;
#[cfg(feature = "testing")]
#[doc(hidden)]
pub static INTERNAL_MAX_THREADS: core::sync::atomic::AtomicUsize =
core::sync::atomic::AtomicUsize::new(usize::MAX);
fn sleep_ms(ms: u64) {
std::thread::sleep(Duration::from_millis(ms));
}
fn err_eq(a: &std::io::Error, b: &std::io::Error) -> bool {
a.kind() == b.kind() && format!("{}", a) == format!("{}", b)
}
#[derive(Debug)]
pub enum StartThreadsError {
NoThreads(std::io::Error),
Respawn(std::io::Error),
}
impl Display for StartThreadsError {
fn fmt(&self, f: &mut Formatter<'_>) -> Result<(), std::fmt::Error> {
match self {
StartThreadsError::NoThreads(e) => write!(
f,
"ThreadPool workers all panicked, failed starting replacement threads: {}",
e
),
StartThreadsError::Respawn(e) => {
write!(
f,
"ThreadPool failed starting threads to replace panicked threads: {}",
e
)
}
}
}
}
impl Error for StartThreadsError {}
impl PartialEq for StartThreadsError {
fn eq(&self, other: &Self) -> bool {
match (self, other) {
(StartThreadsError::NoThreads(a), StartThreadsError::NoThreads(b))
| (StartThreadsError::Respawn(a), StartThreadsError::Respawn(b)) => err_eq(a, b),
_ => false,
}
}
}
impl Eq for StartThreadsError {}
#[derive(Debug)]
pub enum NewThreadPoolError {
Parameter(String),
Spawn(std::io::Error),
}
impl Display for NewThreadPoolError {
fn fmt(&self, f: &mut Formatter<'_>) -> Result<(), std::fmt::Error> {
match self {
NewThreadPoolError::Parameter(s) => write!(f, "{}", s),
NewThreadPoolError::Spawn(e) => {
write!(f, "ThreadPool failed starting threads: {}", e)
}
}
}
}
impl Error for NewThreadPoolError {}
impl PartialEq for NewThreadPoolError {
fn eq(&self, other: &Self) -> bool {
match (self, other) {
(NewThreadPoolError::Parameter(a), NewThreadPoolError::Parameter(b)) => a == b,
(NewThreadPoolError::Spawn(a), NewThreadPoolError::Spawn(b)) => err_eq(a, b),
_ => false,
}
}
}
impl Eq for NewThreadPoolError {}
impl From<StartThreadsError> for NewThreadPoolError {
fn from(err: StartThreadsError) -> Self {
match err {
StartThreadsError::NoThreads(e) | StartThreadsError::Respawn(e) => {
NewThreadPoolError::Spawn(e)
}
}
}
}
impl From<NewThreadPoolError> for std::io::Error {
fn from(new_thread_pool_error: NewThreadPoolError) -> Self {
match new_thread_pool_error {
NewThreadPoolError::Parameter(s) => std::io::Error::new(ErrorKind::InvalidInput, s),
NewThreadPoolError::Spawn(s) => {
std::io::Error::new(ErrorKind::Other, format!("failed to start threads: {}", s))
}
}
}
}
#[derive(Debug)]
pub enum TryScheduleError {
QueueFull,
NoThreads(std::io::Error),
Respawn(std::io::Error),
}
impl Display for TryScheduleError {
fn fmt(&self, f: &mut Formatter<'_>) -> Result<(), std::fmt::Error> {
match self {
TryScheduleError::QueueFull => write!(f, "ThreadPool queue is full"),
TryScheduleError::NoThreads(e) => write!(
f,
"ThreadPool workers all panicked, failed starting replacement threads: {}",
e
),
TryScheduleError::Respawn(e) => {
write!(
f,
"ThreadPool failed starting threads to replace panicked threads: {}",
e
)
}
}
}
}
impl Error for TryScheduleError {}
impl PartialEq for TryScheduleError {
fn eq(&self, other: &Self) -> bool {
match (self, other) {
(TryScheduleError::QueueFull, TryScheduleError::QueueFull) => true,
(TryScheduleError::NoThreads(a), TryScheduleError::NoThreads(b))
| (TryScheduleError::Respawn(a), TryScheduleError::Respawn(b)) => err_eq(a, b),
_ => false,
}
}
}
impl Eq for TryScheduleError {}
impl From<StartThreadsError> for TryScheduleError {
fn from(err: StartThreadsError) -> Self {
match err {
StartThreadsError::NoThreads(e) => TryScheduleError::NoThreads(e),
StartThreadsError::Respawn(e) => TryScheduleError::Respawn(e),
}
}
}
impl From<TryScheduleError> for std::io::Error {
fn from(try_schedule_error: TryScheduleError) -> Self {
match try_schedule_error {
TryScheduleError::QueueFull => {
std::io::Error::new(ErrorKind::WouldBlock, "TryScheduleError::QueueFull")
}
TryScheduleError::NoThreads(e) => std::io::Error::new(
e.kind(),
format!(
"ThreadPool workers all panicked, failed starting replacement threads: {}",
e
),
),
TryScheduleError::Respawn(e) => std::io::Error::new(
e.kind(),
format!(
"ThreadPool failed starting threads to replace panicked threads: {}",
e
),
),
}
}
}
struct Inner {
name: &'static str,
next_name_num: AtomicCounter,
size: usize,
receiver: Mutex<Receiver<Box<dyn FnOnce() + Send>>>,
}
impl Inner {
fn num_live_threads(self: &Arc<Self>) -> usize {
Arc::strong_count(self) - 1
}
fn work(self: &Arc<Self>) {
loop {
let recv_result = self
.receiver
.lock()
.unwrap()
.recv_timeout(Duration::from_millis(500));
match recv_result {
Ok(f) => {
let _ignored = self.start_threads();
f();
}
Err(RecvTimeoutError::Timeout) => {}
Err(RecvTimeoutError::Disconnected) => return,
};
let _ignored = self.start_threads();
}
}
#[allow(clippy::unused_self)]
#[allow(unused_variables)]
fn spawn_thread(
&self,
num_live_threads: usize,
name: String,
f: impl FnOnce() + Send + 'static,
) -> Result<(), std::io::Error> {
#[cfg(feature = "testing")]
if num_live_threads >= INTERNAL_MAX_THREADS.load(std::sync::atomic::Ordering::Acquire) {
return Err(std::io::Error::new(
std::io::ErrorKind::Other,
"err1".to_string(),
));
}
std::thread::Builder::new().name(name).spawn(f)?;
Ok(())
}
fn start_thread(self: &Arc<Self>) -> Result<(), StartThreadsError> {
let self_clone = self.clone();
let num_live_threads = self.num_live_threads() - 1;
if num_live_threads < self.size {
self.spawn_thread(
num_live_threads,
format!("{}{}", self.name, self.next_name_num.next()),
move || self_clone.work(),
)
.map_err(|e| {
if num_live_threads == 0 {
StartThreadsError::NoThreads(e)
} else {
StartThreadsError::Respawn(e)
}
})?;
}
Ok(())
}
fn start_threads(self: &Arc<Self>) -> Result<(), StartThreadsError> {
while self.num_live_threads() < self.size {
self.start_thread()?;
}
Ok(())
}
}
pub struct ThreadPool {
inner: Arc<Inner>,
sender: SyncSender<Box<dyn FnOnce() + Send>>,
}
impl ThreadPool {
pub fn new(name: &'static str, size: usize) -> Result<Self, NewThreadPoolError> {
if name.is_empty() {
return Err(NewThreadPoolError::Parameter(
"ThreadPool::new called with empty name".to_string(),
));
}
if size < 1 {
return Err(NewThreadPoolError::Parameter(format!(
"ThreadPool::new called with invalid size value: {:?}",
size
)));
}
let (sender, receiver) = std::sync::mpsc::sync_channel(size * 200);
let pool = ThreadPool {
inner: Arc::new(Inner {
name,
next_name_num: AtomicCounter::new(),
size,
receiver: Mutex::new(receiver),
}),
sender,
};
pool.inner.start_threads()?;
Ok(pool)
}
#[must_use]
pub fn size(&self) -> usize {
self.inner.size
}
#[must_use]
pub fn num_live_threads(&self) -> usize {
self.inner.num_live_threads()
}
#[cfg(feature = "testing")]
#[doc(hidden)]
#[must_use]
pub fn num_live_threads_fn(&self) -> Box<dyn Fn() -> usize> {
let inner_clone = self.inner.clone();
Box::new(move || inner_clone.num_live_threads())
}
#[allow(clippy::missing_panics_doc)]
pub fn schedule<F: FnOnce() + Send + 'static>(&self, f: F) {
let mut opt_box_f: Option<Box<dyn FnOnce() + Send + 'static>> = Some(Box::new(f));
loop {
match self.inner.start_threads() {
Ok(()) | Err(StartThreadsError::Respawn(_)) => {
}
Err(StartThreadsError::NoThreads(_)) => {
sleep_ms(10);
continue;
}
}
opt_box_f = match self.sender.try_send(opt_box_f.take().unwrap()) {
Ok(()) => return,
Err(TrySendError::Disconnected(_)) => unreachable!(),
Err(TrySendError::Full(box_f)) => Some(box_f),
};
sleep_ms(10);
}
}
#[allow(clippy::missing_panics_doc)]
pub fn try_schedule(&self, f: impl FnOnce() + Send + 'static) -> Result<(), TryScheduleError> {
match self.sender.try_send(Box::new(f)) {
Ok(_) => {}
Err(TrySendError::Disconnected(_)) => unreachable!(),
Err(TrySendError::Full(_)) => return Err(TryScheduleError::QueueFull),
};
self.inner.start_threads().map_err(std::convert::Into::into)
}
pub fn join(self) {
let inner = self.inner.clone();
drop(self);
while inner.num_live_threads() > 0 {
sleep_ms(10);
}
}
pub fn try_join(self, timeout: Duration) -> Result<(), String> {
let inner = self.inner.clone();
drop(self);
let deadline = Instant::now() + timeout;
loop {
if inner.num_live_threads() < 1 {
return Ok(());
}
if deadline < Instant::now() {
return Err("timed out waiting for ThreadPool workers to stop".to_string());
}
sleep_ms(10);
}
}
}
impl Debug for ThreadPool {
fn fmt(&self, f: &mut Formatter<'_>) -> Result<(), core::fmt::Error> {
write!(
f,
"ThreadPool{{{:?},size={:?}}}",
self.inner.name, self.inner.size
)
}
}