use std::any::Any;
use std::panic;
use std::sync::Arc;
use std::sync::Mutex;
use std::sync::OnceLock;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::mpsc;
use std::time::Duration;
pub trait PhaseEndTrigger: Send + Sync + 'static {
fn fire(&self, event: &PhaseEndEvent);
fn name(&self) -> &str {
"phase-end-trigger"
}
}
#[derive(Debug, Clone)]
pub struct PhaseEndEvent {
pub phase_name: String,
pub phase_labels: String,
pub outcome: PhaseOutcome,
pub duration_secs: f64,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum PhaseOutcome {
Completed,
Failed { error: String },
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub struct TriggerId(u64);
struct Entry {
id: TriggerId,
trigger: Arc<dyn PhaseEndTrigger>,
}
struct Registry {
next_id: u64,
triggers: Vec<Entry>,
dispatch: Option<mpsc::Sender<PhaseEndEvent>>,
}
static REGISTRY: OnceLock<Mutex<Registry>> = OnceLock::new();
static TRIGGER_COUNT: AtomicUsize = AtomicUsize::new(0);
fn registry() -> &'static Mutex<Registry> {
REGISTRY.get_or_init(|| {
Mutex::new(Registry {
next_id: 1,
triggers: Vec::new(),
dispatch: None,
})
})
}
pub fn register(trigger: Arc<dyn PhaseEndTrigger>) -> TriggerId {
let mut reg = registry()
.lock()
.expect("phase-end-triggers registry poisoned");
let id = TriggerId(reg.next_id);
reg.next_id += 1;
reg.triggers.push(Entry { id, trigger });
TRIGGER_COUNT.store(reg.triggers.len(), Ordering::Release);
if reg.dispatch.is_none() {
let (tx, rx) = mpsc::channel::<PhaseEndEvent>();
reg.dispatch = Some(tx);
std::thread::Builder::new()
.name("phase-end-trigger-worker".into())
.spawn(move || dispatch_loop(rx))
.expect("spawn phase-end-trigger worker");
}
id
}
pub fn unregister(id: TriggerId) {
let mut reg = registry()
.lock()
.expect("phase-end-triggers registry poisoned");
reg.triggers.retain(|e| e.id != id);
TRIGGER_COUNT.store(reg.triggers.len(), Ordering::Release);
}
#[cfg(test)]
pub fn reset_for_tests() {
let mut reg = registry()
.lock()
.expect("phase-end-triggers registry poisoned");
reg.triggers.clear();
TRIGGER_COUNT.store(0, Ordering::Release);
reg.next_id = 1;
}
pub fn fire_phase_completed(name: &str, labels: &str, duration_secs: f64) {
fire(PhaseEndEvent {
phase_name: name.to_string(),
phase_labels: labels.to_string(),
outcome: PhaseOutcome::Completed,
duration_secs,
});
}
pub fn fire_phase_failed(name: &str, labels: &str, error: &str) {
fire(PhaseEndEvent {
phase_name: name.to_string(),
phase_labels: labels.to_string(),
outcome: PhaseOutcome::Failed {
error: error.to_string(),
},
duration_secs: 0.0,
});
}
fn fire(event: PhaseEndEvent) {
if TRIGGER_COUNT.load(Ordering::Acquire) == 0 {
return;
}
let reg = registry()
.lock()
.expect("phase-end-triggers registry poisoned");
if reg.triggers.is_empty() {
return;
}
if let Some(tx) = reg.dispatch.as_ref() {
let _ = tx.send(event);
}
}
fn dispatch_loop(rx: mpsc::Receiver<PhaseEndEvent>) {
loop {
match rx.recv_timeout(Duration::from_secs(60)) {
Ok(event) => {
let snap: Vec<Arc<dyn PhaseEndTrigger>> = {
let reg = match registry().lock() {
Ok(r) => r,
Err(_) => return, };
reg.triggers.iter().map(|e| e.trigger.clone()).collect()
};
for trigger in snap {
let result = panic::catch_unwind(panic::AssertUnwindSafe(|| {
trigger.fire(&event);
}));
if let Err(payload) = result {
let msg = payload_to_message(payload);
crate::diag!(
crate::observer::LogLevel::Warn,
"phase-end trigger '{name}' panicked: {msg}",
name = trigger.name(),
);
}
}
}
Err(mpsc::RecvTimeoutError::Timeout) => {
}
Err(mpsc::RecvTimeoutError::Disconnected) => {
return;
}
}
}
}
fn payload_to_message(payload: Box<dyn Any + Send>) -> String {
if let Some(s) = payload.downcast_ref::<&'static str>() {
(*s).to_string()
} else if let Some(s) = payload.downcast_ref::<String>() {
s.clone()
} else {
"<non-string panic payload>".to_string()
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::Mutex as StdMutex;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::time::Instant;
struct CountingTrigger {
count: Arc<AtomicUsize>,
names: Arc<StdMutex<Vec<String>>>,
name: &'static str,
}
impl PhaseEndTrigger for CountingTrigger {
fn name(&self) -> &str {
self.name
}
fn fire(&self, event: &PhaseEndEvent) {
self.count.fetch_add(1, Ordering::Release);
self.names.lock().unwrap().push(event.phase_name.clone());
}
}
fn wait_for(count: &AtomicUsize, target: usize, timeout: Duration) -> bool {
let start = Instant::now();
while count.load(Ordering::Acquire) < target {
if start.elapsed() > timeout {
return false;
}
std::thread::sleep(Duration::from_millis(5));
}
true
}
static TEST_LOCK: std::sync::Mutex<()> = std::sync::Mutex::new(());
#[test]
fn registered_trigger_fires_for_completed_phase() {
let _g = TEST_LOCK.lock().unwrap();
reset_for_tests();
let count = Arc::new(AtomicUsize::new(0));
let names = Arc::new(StdMutex::new(Vec::new()));
let trig = Arc::new(CountingTrigger {
count: count.clone(),
names: names.clone(),
name: "test",
});
let _id = register(trig);
fire_phase_completed("setup", "", 1.5);
assert!(
wait_for(&count, 1, Duration::from_secs(2)),
"trigger did not fire within 2s"
);
assert_eq!(*names.lock().unwrap(), vec!["setup".to_string()]);
reset_for_tests();
}
#[test]
fn registered_trigger_fires_for_failed_phase() {
let _g = TEST_LOCK.lock().unwrap();
reset_for_tests();
let count = Arc::new(AtomicUsize::new(0));
let names = Arc::new(StdMutex::new(Vec::new()));
let trig = Arc::new(CountingTrigger {
count: count.clone(),
names: names.clone(),
name: "test",
});
let _id = register(trig);
fire_phase_failed("query", "k=10", "timeout");
assert!(wait_for(&count, 1, Duration::from_secs(2)));
assert_eq!(*names.lock().unwrap(), vec!["query".to_string()]);
reset_for_tests();
}
#[test]
fn unregister_stops_subsequent_dispatches() {
let _g = TEST_LOCK.lock().unwrap();
reset_for_tests();
let count = Arc::new(AtomicUsize::new(0));
let names = Arc::new(StdMutex::new(Vec::new()));
let trig = Arc::new(CountingTrigger {
count: count.clone(),
names: names.clone(),
name: "test",
});
let id = register(trig);
fire_phase_completed("a", "", 0.1);
assert!(wait_for(&count, 1, Duration::from_secs(2)));
unregister(id);
fire_phase_completed("b", "", 0.2);
std::thread::sleep(Duration::from_millis(100));
assert_eq!(
count.load(Ordering::Acquire),
1,
"trigger fired after unregister"
);
reset_for_tests();
}
#[test]
fn multiple_triggers_fire_in_registration_order() {
let _g = TEST_LOCK.lock().unwrap();
reset_for_tests();
let count_a = Arc::new(AtomicUsize::new(0));
let count_b = Arc::new(AtomicUsize::new(0));
let names = Arc::new(StdMutex::new(Vec::new()));
let _id_a = register(Arc::new(CountingTrigger {
count: count_a.clone(),
names: names.clone(),
name: "a",
}));
let _id_b = register(Arc::new(CountingTrigger {
count: count_b.clone(),
names: names.clone(),
name: "b",
}));
fire_phase_completed("phase1", "", 0.0);
assert!(wait_for(&count_a, 1, Duration::from_secs(2)));
assert!(wait_for(&count_b, 1, Duration::from_secs(2)));
assert_eq!(
*names.lock().unwrap(),
vec!["phase1".to_string(), "phase1".to_string()]
);
reset_for_tests();
}
#[test]
fn panic_in_one_trigger_does_not_stop_others() {
let _g = TEST_LOCK.lock().unwrap();
reset_for_tests();
struct PanickingTrigger;
impl PhaseEndTrigger for PanickingTrigger {
fn name(&self) -> &str {
"panicker"
}
fn fire(&self, _: &PhaseEndEvent) {
panic!("boom");
}
}
let count = Arc::new(AtomicUsize::new(0));
let names = Arc::new(StdMutex::new(Vec::new()));
let _a = register(Arc::new(PanickingTrigger));
let _b = register(Arc::new(CountingTrigger {
count: count.clone(),
names: names.clone(),
name: "after-panic",
}));
fire_phase_completed("phase", "", 0.0);
assert!(
wait_for(&count, 1, Duration::from_secs(2)),
"downstream trigger lost to upstream panic"
);
reset_for_tests();
}
#[test]
fn fire_with_no_triggers_is_noop() {
let _g = TEST_LOCK.lock().unwrap();
reset_for_tests();
fire_phase_completed("x", "", 1.0);
fire_phase_failed("y", "", "oops");
}
}