use std::{
panic, sync::{
mpsc::{Receiver, Sender},
Arc, Mutex,
}, thread::{self, JoinHandle}
};
use tklog::{debugs, sync::Logger, warns, LEVEL};
type Job = Box<dyn FnOnce() + Send + 'static>;
pub struct HjThreadPool {
workers: Vec<Worker>,
job_sender: Option<Sender<Job>>,
logger: Arc<Mutex<Logger>>,
}
pub struct HjThreadPoolCfg {
pub num_workers: usize, pub log_level: HjThreadPoolLogLevel, }
pub enum HjThreadPoolLogLevel {
Debug, Info, Warn, Error, }
struct Worker {
worker_id: usize,
job_receiver: Arc<Mutex<Receiver<Job>>>,
join_handle: Option<JoinHandle<()>>,
started: bool,
start_lock: Mutex<()>,
logger: Arc<Mutex<Logger>>,
}
impl Worker {
fn new(
worker_id: usize,
job_receiver: Arc<Mutex<Receiver<Job>>>,
logger: Arc<Mutex<Logger>>,
) -> Self {
return Worker {
worker_id: worker_id,
job_receiver: job_receiver,
join_handle: None,
started: false,
start_lock: Mutex::new(()),
logger: logger,
};
}
fn start(&mut self) {
{
let _lock = self.start_lock.lock().unwrap();
if self.started {
panic!("Worker {} already started", self.worker_id);
}
self.started = true;
}
let job_receiver = Arc::clone(&self.job_receiver);
let worker_id = self.worker_id;
let mut logger = Arc::clone(&self.logger);
let join_handle = thread::spawn(move || loop {
let recv_job_result = job_receiver.lock().unwrap().recv();
if recv_job_result.is_err() {
debugs!(&mut logger, format!("Worker {} exiting", worker_id));
break;
}
let job = recv_job_result.unwrap();
let result = panic::catch_unwind(panic::AssertUnwindSafe(|| {
job();
}));
match result {
Ok(_) => {
}
Err(err) => {
let panic_info = {
if let Some(message) = err.downcast_ref::<String>() {
message.clone()
} else if let Some(message) = err.downcast_ref::<&str>() {
message.to_string()
} else {
"Cause panic with unknown reason".to_string()
}
};
warns!(&mut logger, format!("Worker {} got panic while executing job, but the panic is caught and the worker will continue to work, so don't worry. The panic info is: {}", worker_id, panic_info));
}
}
});
self.join_handle = Some(join_handle);
}
}
impl HjThreadPool {
pub fn new(cfg: HjThreadPoolCfg) -> Self {
let num_workers = cfg.num_workers;
let log_level = match cfg.log_level {
HjThreadPoolLogLevel::Debug => LEVEL::Debug,
HjThreadPoolLogLevel::Info => LEVEL::Info,
HjThreadPoolLogLevel::Warn => LEVEL::Warn,
HjThreadPoolLogLevel::Error => LEVEL::Error,
};
let mut logger = Arc::new(Mutex::new(Logger::new()));
logger.lock().unwrap().set_level(log_level);
debugs!(
&mut logger,
format!("Creating HjThreadPool with {} workers", num_workers)
);
let (job_sender, job_receiver) = std::sync::mpsc::channel();
let job_receiver = Arc::new(Mutex::new(job_receiver));
let mut workers = Vec::new();
for worker_id in 0..num_workers {
let mut worker = Worker::new(worker_id, Arc::clone(&job_receiver), Arc::clone(&logger));
worker.start();
workers.push(worker);
}
return HjThreadPool {
workers: workers,
job_sender: Some(job_sender),
logger: logger,
};
}
pub fn execute<F>(&self, f: F)
where
F: FnOnce() + Send + 'static,
{
let job = Box::new(f);
self.job_sender.as_ref().unwrap().send(job).unwrap();
}
}
impl Drop for HjThreadPool {
fn drop(&mut self) {
debugs!(&mut self.logger, "Dropping HjThreadPool");
drop(self.job_sender.take());
for worker in self.workers.iter_mut() {
worker.join_handle.take().unwrap().join().unwrap();
}
}
}
#[cfg(test)]
mod tests {
use std::time::Duration;
use super::*;
#[test]
fn new() {
let _pool = HjThreadPool::new(HjThreadPoolCfg {
num_workers: 2,
log_level: HjThreadPoolLogLevel::Debug,
});
}
#[test]
fn test_execute() {
let pool = HjThreadPool::new(HjThreadPoolCfg {
num_workers: 2,
log_level: HjThreadPoolLogLevel::Debug,
});
pool.execute(|| {
for i in 0..10 {
println!("Task 1: {}", i);
thread::sleep(Duration::from_secs(1));
}
});
pool.execute(|| {
for i in 0..10 {
println!("Task 2: {}", i);
thread::sleep(Duration::from_secs(1));
}
});
pool.execute(|| {
for i in 0..10 {
println!("Task 3: {}", i);
thread::sleep(Duration::from_secs(1));
}
});
}
#[test]
fn test_job_panic() {
let pool = HjThreadPool::new(HjThreadPoolCfg {
num_workers: 1,
log_level: HjThreadPoolLogLevel::Debug,
});
pool.execute(|| {
panic!("Panic in job");
});
pool.execute(|| {
for i in 0..10 {
println!("Task 2: {}", i);
thread::sleep(Duration::from_secs(1));
}
});
}
}