use crate::{
task::{Task, TaskFn, TaskListeners},
worker::Worker,
TPError, TPResult,
};
use crossbeam_channel::{Receiver, Sender, TrySendError};
use std::{panic::UnwindSafe, sync::Arc, time::Duration};
pub enum RejectedTaskHandler {
Abort,
Discard,
CallerRuns,
}
pub struct ThreadPool {
pub(crate) sender: Option<Sender<Task>>,
pub(crate) reciver: Receiver<Task>,
pub(crate) core_workers: Vec<Worker>,
pub(crate) workers: Vec<Worker>,
pub(crate) core_pool_size: usize,
pub(crate) max_pool_size: usize,
pub(crate) keep_alive_time: Duration,
pub(crate) rejected_task_handler: RejectedTaskHandler,
pub(crate) next_task_id: usize,
pub(crate) task_lisenters: Arc<TaskListeners>,
}
impl ThreadPool {
pub fn execute<F>(&mut self, task_fn: F) -> Result<(), TPError>
where
F: FnOnce() + Send + UnwindSafe + 'static,
{
if self.sender.is_none() {
return Err(TPError::DisconnectedError);
}
let task = self.create_task(Box::new(task_fn));
if self.core_workers.len() < self.core_pool_size {
self.add_worker(task, true);
Ok(())
} else {
self.try_send_task(task)
}
}
pub fn close_channel(&mut self) {
self.sender.take();
}
pub fn wait(self) -> std::thread::Result<()> {
assert!(self.sender.is_none());
Self::wait_workers(self.core_workers)?;
Self::wait_workers(self.workers)
}
fn wait_workers(workers: Vec<Worker>) -> std::thread::Result<()> {
for worker in workers {
worker.join()?
}
Ok(())
}
fn create_task(&mut self, task_fn: TaskFn) -> Task {
let id = self.next_task_id;
self.next_task_id += 1;
Task::create(id, task_fn, self.task_lisenters.clone())
}
fn add_worker(&mut self, task: Task, is_core: bool) {
let new_worker = Worker::new(is_core, self.keep_alive_time, self.reciver.clone(), task);
if is_core {
self.core_workers.push(new_worker);
} else {
self.workers.push(new_worker);
}
}
fn try_send_task(&mut self, task: Task) -> TPResult<()> {
let sender = self.sender.as_ref();
if sender.is_none() {
return Err(TPError::DisconnectedError);
}
if let Err(err) = sender.unwrap().try_send(task) {
return match err {
TrySendError::Full(task) => self.process_task_if_channel_full(task),
TrySendError::Disconnected(_) => Err(TPError::DisconnectedError),
};
}
Ok(())
}
fn process_task_if_channel_full(&mut self, task: Task) -> TPResult<()> {
let idle_worker = self.workers.iter_mut().find(|worker| worker.is_finished());
if let Some(idle_worker) = idle_worker {
idle_worker.restart(task);
return Ok(());
}
if self.workers.len() < self.max_pool_size - self.core_pool_size {
self.add_worker(task, false);
Ok(())
} else {
self.reject(task)
}
}
fn reject(&self, task: Task) -> Result<(), TPError> {
match &self.rejected_task_handler {
RejectedTaskHandler::Abort => Err(TPError::AbortError),
RejectedTaskHandler::CallerRuns => {
task.run();
Ok(())
}
RejectedTaskHandler::Discard => Ok(()),
}
}
}