use std::sync::atomic::AtomicBool;
use std::sync::atomic::Ordering;
use std::sync::Arc;
use crossbeam_channel::unbounded;
use crossbeam_channel::Receiver;
use crossbeam_channel::Select;
use crossbeam_channel::Sender;
use humthreads::ErrorKind as HumthreadsErrorKind;
use humthreads::MapThread;
use humthreads::Thread;
use signal_hook::SigId;
use slog::debug;
use slog::o;
use slog::warn;
use slog::Discard;
use slog::Logger;
use replicante_util_failure::capture_fail;
use replicante_util_failure::failure_info;
pub struct Upkeep {
callbacks: Vec<Box<dyn Fn()>>,
logger: Logger,
registered_signals: Vec<SigId>,
signal_flag: Arc<AtomicBool>,
signal_receiver: Receiver<()>,
signal_sender: Option<Sender<()>>,
threads: Vec<ThreadMeta>,
}
impl Upkeep {
pub fn new() -> Upkeep {
let (signal_sender, signal_receiver) = unbounded();
let signal_sender = Some(signal_sender);
Upkeep {
callbacks: Vec::new(),
logger: Logger::root(Discard, o!()),
registered_signals: Vec::new(),
signal_flag: Arc::new(AtomicBool::new(false)),
signal_receiver,
signal_sender,
threads: Vec::new(),
}
}
pub fn keepalive(&mut self) -> bool {
let mut clean_exit = true;
loop {
let mut set = self.select_set();
let index = set.ready();
match index {
0 => {
warn!(self.logger, "Shutdown: signal received");
break;
}
n => {
let thread = &self.threads[n - 1];
let paniced = match thread.handle.join() {
Ok(()) => false,
Err(error) => match error.kind() {
HumthreadsErrorKind::Join(_) => {
capture_fail!(
&error,
self.logger,
"Thread paniced";
failure_info(&error),
);
clean_exit = false;
true
}
_ => false,
},
};
if paniced {
warn!(self.logger, "Shutdown: thread paniced");
break;
}
if thread.required {
warn!(self.logger, "Shutdown: thread exited");
break;
}
}
};
drop(set);
self.threads.remove(index - 1);
}
self.shutdown();
self.join_threads() && clean_exit
}
pub fn on_shutdown<F>(&mut self, callback: F)
where
F: Fn() + 'static,
{
self.callbacks.push(Box::new(callback))
}
pub fn register_signal(&mut self) -> Result<(), ::std::io::Error> {
let sender = match self.signal_sender.take() {
Some(sender) => sender,
None => return Ok(()),
};
let signals = vec![signal_hook::SIGINT, signal_hook::SIGTERM];
for signal in signals.into_iter() {
let signal_flag = Arc::clone(&self.signal_flag);
let signal_sender = sender.clone();
let callback = move || {
if signal_flag.load(Ordering::Relaxed) {
::std::process::exit(1);
}
signal_flag.store(true, Ordering::Relaxed);
let _ = signal_sender.send(());
};
let signal_id = unsafe { signal_hook::register(signal, callback) }?;
self.registered_signals.push(signal_id);
}
Ok(())
}
pub fn register_thread<T: Send + 'static>(&mut self, thread: Thread<T>) {
let thread = ThreadMeta {
handle: thread.map(|_| ()),
required: true,
};
self.threads.push(thread);
}
pub fn register_thread_optional<T: Send + 'static>(&mut self, thread: Thread<T>) {
let thread = ThreadMeta {
handle: thread.map(|_| ()),
required: false,
};
self.threads.push(thread);
}
pub fn set_logger(&mut self, logger: Logger) {
self.logger = logger;
}
fn join_threads(&mut self) -> bool {
debug!(self.logger, "Joining with registered threads");
let mut clean_exit = true;
for thread in self.threads.drain(..) {
if let Err(error) = thread.handle.join() {
if let HumthreadsErrorKind::JoinedAlready = error.kind() {
debug!(self.logger, "Joined thread twice");
continue;
}
capture_fail!(&error, self.logger, "Thread paniced"; failure_info(&error));
clean_exit = false;
}
}
clean_exit
}
fn select_set<'a, 'b: 'a>(&'b self) -> Select<'a> {
let mut set = Select::new();
set.recv(&self.signal_receiver);
for thread in &self.threads {
thread.handle.select_add(&mut set);
}
set
}
fn shutdown(&mut self) {
debug!(self.logger, "Requesting shutdowns for registered threads");
for thread in &self.threads {
thread.handle.request_shutdown();
}
debug!(self.logger, "Executing shutdown callbacks");
for callback in &self.callbacks {
callback();
}
}
}
impl Default for Upkeep {
fn default() -> Upkeep {
Upkeep::new()
}
}
impl Drop for Upkeep {
fn drop(&mut self) {
for signal in self.registered_signals.drain(..) {
signal_hook::unregister(signal);
}
}
}
struct ThreadMeta {
handle: MapThread<()>,
required: bool,
}
#[cfg(test)]
mod tests {
use std::sync::atomic::AtomicBool;
use std::sync::atomic::AtomicUsize;
use std::sync::atomic::Ordering;
use std::sync::Arc;
use std::time::Duration;
use humthreads::Builder;
use super::Upkeep;
#[test]
fn callback() {
let flag = Arc::new(AtomicBool::new(false));
let mut up = Upkeep::new();
let inner_flag = Arc::clone(&flag);
up.on_shutdown(move || inner_flag.store(true, Ordering::Relaxed));
up.shutdown();
assert_eq!(true, flag.load(Ordering::Relaxed));
}
#[test]
fn thread_optional() {
let count = Arc::new(AtomicUsize::new(0));
let inner_count = Arc::clone(&count);
let mut up = Upkeep::new();
let optional = Builder::new("thread_optional_two")
.spawn(|_| ::std::thread::sleep(Duration::from_millis(10)))
.expect("to spawn test thread");
up.register_thread_optional(optional);
let thread = Builder::new("thread_optional_one")
.spawn(move |scope| {
for _ in 0..5 {
::std::thread::sleep(Duration::from_millis(10));
if scope.should_shutdown() {
break;
}
inner_count.fetch_add(1, Ordering::Relaxed);
}
})
.expect("to spawn test thread");
up.register_thread(thread);
let clean = up.keepalive();
assert_eq!(true, clean);
assert_eq!(5, count.load(Ordering::Relaxed));
}
#[test]
fn thread_panics() {
let flag = Arc::new(AtomicBool::new(false));
let inner_flag = Arc::clone(&flag);
let mut up = Upkeep::new();
let thread = Builder::new("thread_panics")
.spawn(move |_| {
inner_flag.store(true, Ordering::Relaxed);
panic!("this panic is expected");
})
.expect("to spawn test thread");
up.register_thread(thread);
let clean = up.keepalive();
assert_eq!(true, flag.load(Ordering::Relaxed));
assert_eq!(false, clean);
}
#[test]
fn thread_shuts_down() {
let flag = Arc::new(AtomicBool::new(false));
let inner_flag = Arc::clone(&flag);
let thread = Builder::new("thread_shuts_down")
.spawn(move |scope| {
loop {
::std::thread::sleep(Duration::from_millis(10));
if scope.should_shutdown() {
break;
}
}
inner_flag.store(true, Ordering::Relaxed);
})
.expect("to spawn test thread");
let mut up = Upkeep::new();
up.register_thread(thread);
up.shutdown();
let clean = up.keepalive();
assert_eq!(true, flag.load(Ordering::Relaxed));
assert_eq!(true, clean);
}
}