use crate::HOOK_LOG_TARGET;
use dashmap::DashMap;
use std::future::Future;
use std::pin::Pin;
use std::sync::Arc;
use std::sync::LazyLock;
use tibba_error::Error;
use tracing::{error, info};
type Result<T> = std::result::Result<T, Error>;
pub type BoxFuture<'a, T> = Pin<Box<dyn Future<Output = T> + Send + 'a>>;
pub trait Task: Send + Sync {
fn before(&self) -> BoxFuture<'_, Result<bool>> {
Box::pin(async { Ok(false) })
}
fn after(&self) -> BoxFuture<'_, Result<bool>> {
Box::pin(async { Ok(false) })
}
fn priority(&self) -> u8 {
0
}
}
static TASKS: LazyLock<DashMap<String, Arc<dyn Task>>> = LazyLock::new(DashMap::new);
#[derive(Clone, Copy)]
enum TaskType {
Before,
After,
}
impl TaskType {
fn label(self) -> &'static str {
match self {
TaskType::Before => "before",
TaskType::After => "after",
}
}
}
fn collect_sorted(task_type: TaskType) -> Vec<(String, Arc<dyn Task>)> {
let mut tasks: Vec<(String, Arc<dyn Task>)> = TASKS
.iter()
.map(|item| (item.key().clone(), item.value().clone()))
.collect();
tasks.sort_by_key(|(_, task)| {
let p = task.priority() as i16;
match task_type {
TaskType::Before => p,
TaskType::After => -p,
}
});
tasks
}
async fn run_tasks(task_type: TaskType) -> Result<()> {
for (name, task) in collect_sorted(task_type) {
let start = std::time::Instant::now();
let outcome = match task_type {
TaskType::Before => task.before().await,
TaskType::After => task.after().await,
};
match outcome {
Ok(executed) => {
if executed {
info!(
target: HOOK_LOG_TARGET,
task_type = task_type.label(),
name,
elapsed = start.elapsed().as_millis(),
);
}
}
Err(err) => {
error!(
target: HOOK_LOG_TARGET,
task_type = task_type.label(),
name,
elapsed = start.elapsed().as_millis(),
error = %err,
);
if matches!(task_type, TaskType::Before) {
return Err(err);
}
}
}
}
Ok(())
}
pub fn register_task(name: impl Into<String>, task: Arc<dyn Task>) {
TASKS.insert(name.into(), task);
}
pub async fn run_before_tasks() -> Result<()> {
run_tasks(TaskType::Before).await
}
pub async fn run_after_tasks() -> Result<()> {
run_tasks(TaskType::After).await
}
#[cfg(test)]
#[allow(clippy::await_holding_lock)]
mod tests {
use super::*;
use pretty_assertions::assert_eq;
use std::sync::Mutex;
type Trace = Arc<Mutex<Vec<&'static str>>>;
struct ProbeTask {
name: &'static str,
priority: u8,
before_result: fn() -> Result<bool>,
after_result: fn() -> Result<bool>,
trace: Trace,
}
impl Task for ProbeTask {
fn priority(&self) -> u8 {
self.priority
}
fn before(&self) -> BoxFuture<'_, Result<bool>> {
let name = self.name;
let trace = self.trace.clone();
let f = self.before_result;
Box::pin(async move {
trace.lock().unwrap().push(name);
f()
})
}
fn after(&self) -> BoxFuture<'_, Result<bool>> {
let name = self.name;
let trace = self.trace.clone();
let f = self.after_result;
Box::pin(async move {
trace.lock().unwrap().push(name);
f()
})
}
}
fn reset() {
TASKS.clear();
}
static SERIAL: Mutex<()> = Mutex::new(());
fn serial() -> std::sync::MutexGuard<'static, ()> {
SERIAL.lock().unwrap_or_else(|e| e.into_inner())
}
fn ok_true() -> Result<bool> {
Ok(true)
}
fn boom() -> Result<bool> {
Err(Error::new("boom"))
}
#[tokio::test]
async fn before_runs_in_ascending_priority_order() {
let _g = serial();
reset();
let trace: Trace = Arc::new(Mutex::new(Vec::new()));
register_task(
"high-prio",
Arc::new(ProbeTask {
name: "high",
priority: 1,
before_result: ok_true,
after_result: ok_true,
trace: trace.clone(),
}),
);
register_task(
"low-prio",
Arc::new(ProbeTask {
name: "low",
priority: 200,
before_result: ok_true,
after_result: ok_true,
trace: trace.clone(),
}),
);
run_before_tasks().await.unwrap();
assert_eq!(&*trace.lock().unwrap(), &["high", "low"]);
}
#[tokio::test]
async fn after_runs_in_descending_priority_order() {
let _g = serial();
reset();
let trace: Trace = Arc::new(Mutex::new(Vec::new()));
register_task(
"a",
Arc::new(ProbeTask {
name: "a",
priority: 10,
before_result: ok_true,
after_result: ok_true,
trace: trace.clone(),
}),
);
register_task(
"b",
Arc::new(ProbeTask {
name: "b",
priority: 50,
before_result: ok_true,
after_result: ok_true,
trace: trace.clone(),
}),
);
run_after_tasks().await.unwrap();
assert_eq!(&*trace.lock().unwrap(), &["b", "a"]);
}
#[tokio::test]
async fn before_is_fail_fast_on_first_error() {
let _g = serial();
reset();
let trace: Trace = Arc::new(Mutex::new(Vec::new()));
register_task(
"first",
Arc::new(ProbeTask {
name: "first",
priority: 0,
before_result: boom,
after_result: ok_true,
trace: trace.clone(),
}),
);
register_task(
"second",
Arc::new(ProbeTask {
name: "second",
priority: 10,
before_result: ok_true,
after_result: ok_true,
trace: trace.clone(),
}),
);
let err = run_before_tasks().await.unwrap_err();
assert!(err.to_string().contains("boom"));
assert_eq!(&*trace.lock().unwrap(), &["first"]);
}
#[tokio::test]
async fn after_is_best_effort_continues_past_errors() {
let _g = serial();
reset();
let trace: Trace = Arc::new(Mutex::new(Vec::new()));
register_task(
"first",
Arc::new(ProbeTask {
name: "first",
priority: 100, before_result: ok_true,
after_result: boom,
trace: trace.clone(),
}),
);
register_task(
"second",
Arc::new(ProbeTask {
name: "second",
priority: 10,
before_result: ok_true,
after_result: ok_true,
trace: trace.clone(),
}),
);
run_after_tasks().await.unwrap();
assert_eq!(&*trace.lock().unwrap(), &["first", "second"]);
}
#[tokio::test]
async fn register_task_overwrites_same_name() {
let _g = serial();
reset();
let trace: Trace = Arc::new(Mutex::new(Vec::new()));
register_task(
"dup",
Arc::new(ProbeTask {
name: "v1",
priority: 0,
before_result: ok_true,
after_result: ok_true,
trace: trace.clone(),
}),
);
register_task(
"dup",
Arc::new(ProbeTask {
name: "v2",
priority: 0,
before_result: ok_true,
after_result: ok_true,
trace: trace.clone(),
}),
);
run_before_tasks().await.unwrap();
assert_eq!(&*trace.lock().unwrap(), &["v2"]);
}
#[tokio::test]
async fn empty_registry_is_ok() {
let _g = serial();
reset();
run_before_tasks().await.unwrap();
run_after_tasks().await.unwrap();
}
}