use std::collections::VecDeque;
use std::sync::{Arc, Condvar, Mutex};
use std::thread::JoinHandle;
use super::facility::{recover, run_facility_loop, run_isolated};
use crate::runtime::task::{StackSizeClass, ThreadPriority, enter_ioc_thread};
pub type Callback = Box<dyn FnOnce() + Send + 'static>;
pub const NUM_CALLBACK_PRIORITIES: usize = 3;
pub const DEFAULT_QUEUE_SIZE: usize = 2000;
pub const DEFAULT_THREADS_PER_PRIORITY: usize = 1;
#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub enum CallbackPriority {
Low,
Medium,
High,
}
impl CallbackPriority {
pub const ALL: [CallbackPriority; NUM_CALLBACK_PRIORITIES] = [
CallbackPriority::Low,
CallbackPriority::Medium,
CallbackPriority::High,
];
pub fn index(self) -> usize {
match self {
CallbackPriority::Low => 0,
CallbackPriority::Medium => 1,
CallbackPriority::High => 2,
}
}
pub fn name_prefix(self) -> &'static str {
match self {
CallbackPriority::Low => "cbLow",
CallbackPriority::Medium => "cbMedium",
CallbackPriority::High => "cbHigh",
}
}
pub fn os_priority(self) -> ThreadPriority {
let scan_low = ThreadPriority::ScanLow.value(); let scan_high = ThreadPriority::ScanHigh.value(); match self {
CallbackPriority::Low => ThreadPriority::Custom(scan_low - 1),
CallbackPriority::Medium => ThreadPriority::Custom(scan_low + 4),
CallbackPriority::High => ThreadPriority::Custom(scan_high + 1),
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum CallbackError {
QueueFull,
}
struct QueueState {
queue: VecDeque<Callback>,
overflow: bool,
overflows: u64,
shutdown: bool,
}
struct PriorityQueue {
capacity: usize,
state: Mutex<QueueState>,
wake: Condvar,
}
impl PriorityQueue {
fn new(capacity: usize) -> Self {
PriorityQueue {
capacity,
state: Mutex::new(QueueState {
queue: VecDeque::with_capacity(capacity.min(1024)),
overflow: false,
overflows: 0,
shutdown: false,
}),
wake: Condvar::new(),
}
}
fn request(&self, name: &str, cb: Callback) -> Result<(), CallbackError> {
let mut st = recover(FACILITY, self.state.lock());
if st.shutdown {
drop(st);
tracing::trace!(
target: "epics_base_rs::runtime::callback",
band = name,
"callbackRequest after shutdown dropped"
);
return Ok(());
}
if st.overflow {
return Err(CallbackError::QueueFull);
}
if st.queue.len() >= self.capacity {
st.overflow = true;
st.overflows += 1;
tracing::error!(
target: "epics_base_rs::runtime::callback",
band = name,
"callbackRequest: ERROR {} ring buffer full",
name
);
return Err(CallbackError::QueueFull);
}
st.queue.push_back(cb);
drop(st);
self.wake.notify_one();
Ok(())
}
}
const FACILITY: &str = "callback band";
fn worker_loop(pq: &PriorityQueue) {
loop {
let mut st = recover(FACILITY, pq.state.lock());
while st.queue.is_empty() && !st.shutdown {
st = recover(FACILITY, pq.wake.wait(st));
}
if st.queue.is_empty() {
return;
}
let cb = st.queue.pop_front().unwrap();
st.overflow = false;
drop(st);
run_isolated(FACILITY, cb);
}
}
#[derive(Clone)]
pub struct CallbackHandle {
queues: [Arc<PriorityQueue>; NUM_CALLBACK_PRIORITIES],
}
impl CallbackHandle {
pub fn request(&self, priority: CallbackPriority, cb: Callback) -> Result<(), CallbackError> {
let pq = &self.queues[priority.index()];
pq.request(priority.name_prefix(), cb)
}
pub fn overflow_count(&self, priority: CallbackPriority) -> u64 {
recover(FACILITY, self.queues[priority.index()].state.lock()).overflows
}
}
pub struct CallbackPool {
queues: [Arc<PriorityQueue>; NUM_CALLBACK_PRIORITIES],
workers: Vec<JoinHandle<()>>,
}
impl CallbackPool {
pub fn new() -> Self {
Self::with_config(DEFAULT_QUEUE_SIZE, DEFAULT_THREADS_PER_PRIORITY)
}
pub fn with_config(queue_size: usize, threads_per_priority: usize) -> Self {
let capacity = queue_size.max(1);
let threads = threads_per_priority.max(1);
let queues: [Arc<PriorityQueue>; NUM_CALLBACK_PRIORITIES] = [
Arc::new(PriorityQueue::new(capacity)),
Arc::new(PriorityQueue::new(capacity)),
Arc::new(PriorityQueue::new(capacity)),
];
let mut workers = Vec::with_capacity(NUM_CALLBACK_PRIORITIES * threads);
for prio in CallbackPriority::ALL {
let pq = &queues[prio.index()];
for j in 0..threads {
let name = if threads > 1 {
format!("{}-{}", prio.name_prefix(), j)
} else {
prio.name_prefix().to_string()
};
let pq = Arc::clone(pq);
let builder = std::thread::Builder::new()
.name(name)
.stack_size(StackSizeClass::Big.bytes());
let handle = builder
.spawn(move || {
let _ = enter_ioc_thread(prio.os_priority());
run_facility_loop(
FACILITY,
|| worker_loop(&pq),
|| recover(FACILITY, pq.state.lock()).shutdown = true,
);
})
.expect("failed to spawn callback worker thread");
workers.push(handle);
}
}
CallbackPool { queues, workers }
}
pub fn handle(&self) -> CallbackHandle {
CallbackHandle {
queues: self.queues.clone(),
}
}
pub fn request(&self, priority: CallbackPriority, cb: Callback) -> Result<(), CallbackError> {
self.queues[priority.index()].request(priority.name_prefix(), cb)
}
pub fn overflow_count(&self, priority: CallbackPriority) -> u64 {
recover(FACILITY, self.queues[priority.index()].state.lock()).overflows
}
pub fn shutdown(&mut self) {
for pq in &self.queues {
recover(FACILITY, pq.state.lock()).shutdown = true;
pq.wake.notify_all();
}
for w in self.workers.drain(..) {
let _ = w.join();
}
}
}
impl Default for CallbackPool {
fn default() -> Self {
Self::new()
}
}
impl Drop for CallbackPool {
fn drop(&mut self) {
self.shutdown();
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::mpsc;
use std::time::Duration;
const T: Duration = Duration::from_secs(5);
#[test]
fn a_panicking_callback_does_not_stop_the_band() {
let pool = CallbackPool::new();
pool.request(
CallbackPriority::Medium,
Box::new(|| panic!("a callback panicked on its band")),
)
.expect("enqueue the panicking callback");
let (tx, rx) = mpsc::channel();
pool.request(
CallbackPriority::Medium,
Box::new(move || tx.send(42u32).unwrap()),
)
.expect("enqueue the next callback");
assert_eq!(
rx.recv_timeout(T).unwrap(),
42,
"the callback after a panicking one never ran: the band worker died with it"
);
}
#[test]
fn enqueued_callback_runs() {
let pool = CallbackPool::new();
let (tx, rx) = mpsc::channel();
pool.request(
CallbackPriority::Medium,
Box::new(move || tx.send(42u32).unwrap()),
)
.unwrap();
assert_eq!(rx.recv_timeout(T).unwrap(), 42);
}
#[test]
fn priority_bands_are_independent() {
let pool = CallbackPool::new();
let (started_tx, started_rx) = mpsc::channel();
let (gate_tx, gate_rx) = mpsc::channel::<()>();
pool.request(
CallbackPriority::Low,
Box::new(move || {
started_tx.send(()).unwrap();
gate_rx.recv().unwrap();
}),
)
.unwrap();
started_rx.recv_timeout(T).unwrap();
let (high_tx, high_rx) = mpsc::channel();
pool.request(
CallbackPriority::High,
Box::new(move || high_tx.send(()).unwrap()),
)
.unwrap();
high_rx
.recv_timeout(T)
.expect("High band stalled behind a blocked Low worker");
gate_tx.send(()).unwrap(); }
#[test]
fn full_ring_latches_overflow_then_recovers() {
let mut pool = CallbackPool::with_config(1, 1);
let (started_tx, started_rx) = mpsc::channel();
let (gate_tx, gate_rx) = mpsc::channel::<()>();
pool.request(
CallbackPriority::Low,
Box::new(move || {
started_tx.send(()).unwrap();
gate_rx.recv().unwrap();
}),
)
.unwrap();
started_rx.recv_timeout(T).unwrap();
pool.request(CallbackPriority::Low, Box::new(|| {}))
.unwrap();
assert_eq!(
pool.request(CallbackPriority::Low, Box::new(|| {})),
Err(CallbackError::QueueFull)
);
assert_eq!(
pool.request(CallbackPriority::Low, Box::new(|| {})),
Err(CallbackError::QueueFull)
);
assert_eq!(pool.overflow_count(CallbackPriority::Low), 1);
gate_tx.send(()).unwrap(); pool.shutdown();
}
#[test]
fn request_after_shutdown_is_silent_noop() {
let pool = CallbackPool::new();
let h = pool.handle();
drop(pool);
let ran = Arc::new(AtomicBool::new(false));
let r = Arc::clone(&ran);
let res = h.request(
CallbackPriority::High,
Box::new(move || r.store(true, Ordering::SeqCst)),
);
assert_eq!(res, Ok(())); assert!(
!ran.load(Ordering::SeqCst),
"callback ran after shutdown; it must be dropped, not invoked"
);
}
}