1use dashmap::DashMap;
16use std::future::Future;
17use std::pin::Pin;
18use std::sync::Arc;
19use std::sync::LazyLock;
20use tibba_error::Error;
21use tracing::{error, info};
22
23const LOG_TARGET: &str = "tibba:hook";
26
27type Result<T> = std::result::Result<T, Error>;
28
29pub type BoxFuture<'a, T> = Pin<Box<dyn Future<Output = T> + Send + 'a>>;
31
32pub trait Task: Send + Sync {
43 fn before(&self) -> BoxFuture<'_, Result<bool>> {
45 Box::pin(async { Ok(false) })
46 }
47 fn after(&self) -> BoxFuture<'_, Result<bool>> {
49 Box::pin(async { Ok(false) })
50 }
51 fn priority(&self) -> u8 {
53 0
54 }
55}
56
57static TASKS: LazyLock<DashMap<String, Arc<dyn Task>>> = LazyLock::new(DashMap::new);
59
60#[derive(Clone, Copy)]
62enum TaskType {
63 Before,
64 After,
65}
66
67impl TaskType {
68 fn label(self) -> &'static str {
70 match self {
71 TaskType::Before => "before",
72 TaskType::After => "after",
73 }
74 }
75}
76
77fn collect_sorted(task_type: TaskType) -> Vec<(String, Arc<dyn Task>)> {
79 let mut tasks: Vec<(String, Arc<dyn Task>)> = TASKS
80 .iter()
81 .map(|item| (item.key().clone(), item.value().clone()))
82 .collect();
83
84 tasks.sort_by_key(|(_, task)| {
86 let p = task.priority() as i16;
87 match task_type {
88 TaskType::Before => p,
89 TaskType::After => -p,
90 }
91 });
92 tasks
93}
94
95async fn run_tasks(task_type: TaskType) -> Result<()> {
99 for (name, task) in collect_sorted(task_type) {
100 let start = std::time::Instant::now();
101 let outcome = match task_type {
102 TaskType::Before => task.before().await,
103 TaskType::After => task.after().await,
104 };
105
106 match outcome {
107 Ok(executed) => {
108 if executed {
109 info!(
110 target: LOG_TARGET,
111 task_type = task_type.label(),
112 name,
113 elapsed = start.elapsed().as_millis(),
114 );
115 }
116 }
117 Err(err) => {
118 error!(
119 target: LOG_TARGET,
120 task_type = task_type.label(),
121 name,
122 elapsed = start.elapsed().as_millis(),
123 error = %err,
124 );
125 if matches!(task_type, TaskType::Before) {
126 return Err(err);
128 }
129 }
131 }
132 }
133 Ok(())
134}
135
136pub fn register_task(name: impl Into<String>, task: Arc<dyn Task>) {
138 TASKS.insert(name.into(), task);
139}
140
141pub async fn run_before_tasks() -> Result<()> {
144 run_tasks(TaskType::Before).await
145}
146
147pub async fn run_after_tasks() -> Result<()> {
150 run_tasks(TaskType::After).await
151}
152
153#[cfg(test)]
154#[allow(clippy::await_holding_lock)]
158mod tests {
159 use super::*;
160 use pretty_assertions::assert_eq;
161 use std::sync::Mutex;
162
163 type Trace = Arc<Mutex<Vec<&'static str>>>;
165
166 struct ProbeTask {
168 name: &'static str,
169 priority: u8,
170 before_result: fn() -> Result<bool>,
171 after_result: fn() -> Result<bool>,
172 trace: Trace,
173 }
174
175 impl Task for ProbeTask {
176 fn priority(&self) -> u8 {
177 self.priority
178 }
179 fn before(&self) -> BoxFuture<'_, Result<bool>> {
180 let name = self.name;
181 let trace = self.trace.clone();
182 let f = self.before_result;
183 Box::pin(async move {
184 trace.lock().unwrap().push(name);
185 f()
186 })
187 }
188 fn after(&self) -> BoxFuture<'_, Result<bool>> {
189 let name = self.name;
190 let trace = self.trace.clone();
191 let f = self.after_result;
192 Box::pin(async move {
193 trace.lock().unwrap().push(name);
194 f()
195 })
196 }
197 }
198
199 fn reset() {
201 TASKS.clear();
202 }
203
204 static SERIAL: Mutex<()> = Mutex::new(());
206
207 fn serial() -> std::sync::MutexGuard<'static, ()> {
209 SERIAL.lock().unwrap_or_else(|e| e.into_inner())
210 }
211
212 fn ok_true() -> Result<bool> {
213 Ok(true)
214 }
215 fn boom() -> Result<bool> {
216 Err(Error::new("boom"))
217 }
218
219 #[tokio::test]
220 async fn before_runs_in_ascending_priority_order() {
221 let _g = serial();
222 reset();
223 let trace: Trace = Arc::new(Mutex::new(Vec::new()));
224 register_task(
225 "high-prio",
226 Arc::new(ProbeTask {
227 name: "high",
228 priority: 1,
229 before_result: ok_true,
230 after_result: ok_true,
231 trace: trace.clone(),
232 }),
233 );
234 register_task(
235 "low-prio",
236 Arc::new(ProbeTask {
237 name: "low",
238 priority: 200,
239 before_result: ok_true,
240 after_result: ok_true,
241 trace: trace.clone(),
242 }),
243 );
244
245 run_before_tasks().await.unwrap();
246 assert_eq!(&*trace.lock().unwrap(), &["high", "low"]);
247 }
248
249 #[tokio::test]
250 async fn after_runs_in_descending_priority_order() {
251 let _g = serial();
252 reset();
253 let trace: Trace = Arc::new(Mutex::new(Vec::new()));
254 register_task(
255 "a",
256 Arc::new(ProbeTask {
257 name: "a",
258 priority: 10,
259 before_result: ok_true,
260 after_result: ok_true,
261 trace: trace.clone(),
262 }),
263 );
264 register_task(
265 "b",
266 Arc::new(ProbeTask {
267 name: "b",
268 priority: 50,
269 before_result: ok_true,
270 after_result: ok_true,
271 trace: trace.clone(),
272 }),
273 );
274
275 run_after_tasks().await.unwrap();
276 assert_eq!(&*trace.lock().unwrap(), &["b", "a"]);
278 }
279
280 #[tokio::test]
281 async fn before_is_fail_fast_on_first_error() {
282 let _g = serial();
283 reset();
284 let trace: Trace = Arc::new(Mutex::new(Vec::new()));
285 register_task(
286 "first",
287 Arc::new(ProbeTask {
288 name: "first",
289 priority: 0,
290 before_result: boom,
291 after_result: ok_true,
292 trace: trace.clone(),
293 }),
294 );
295 register_task(
296 "second",
297 Arc::new(ProbeTask {
298 name: "second",
299 priority: 10,
300 before_result: ok_true,
301 after_result: ok_true,
302 trace: trace.clone(),
303 }),
304 );
305
306 let err = run_before_tasks().await.unwrap_err();
307 assert!(err.to_string().contains("boom"));
308 assert_eq!(&*trace.lock().unwrap(), &["first"]);
310 }
311
312 #[tokio::test]
313 async fn after_is_best_effort_continues_past_errors() {
314 let _g = serial();
315 reset();
316 let trace: Trace = Arc::new(Mutex::new(Vec::new()));
317 register_task(
318 "first",
319 Arc::new(ProbeTask {
320 name: "first",
321 priority: 100, before_result: ok_true,
323 after_result: boom,
324 trace: trace.clone(),
325 }),
326 );
327 register_task(
328 "second",
329 Arc::new(ProbeTask {
330 name: "second",
331 priority: 10,
332 before_result: ok_true,
333 after_result: ok_true,
334 trace: trace.clone(),
335 }),
336 );
337
338 run_after_tasks().await.unwrap();
340 assert_eq!(&*trace.lock().unwrap(), &["first", "second"]);
341 }
342
343 #[tokio::test]
344 async fn register_task_overwrites_same_name() {
345 let _g = serial();
346 reset();
347 let trace: Trace = Arc::new(Mutex::new(Vec::new()));
348 register_task(
349 "dup",
350 Arc::new(ProbeTask {
351 name: "v1",
352 priority: 0,
353 before_result: ok_true,
354 after_result: ok_true,
355 trace: trace.clone(),
356 }),
357 );
358 register_task(
359 "dup",
360 Arc::new(ProbeTask {
361 name: "v2",
362 priority: 0,
363 before_result: ok_true,
364 after_result: ok_true,
365 trace: trace.clone(),
366 }),
367 );
368
369 run_before_tasks().await.unwrap();
370 assert_eq!(&*trace.lock().unwrap(), &["v2"]);
371 }
372
373 #[tokio::test]
374 async fn empty_registry_is_ok() {
375 let _g = serial();
376 reset();
377 run_before_tasks().await.unwrap();
378 run_after_tasks().await.unwrap();
379 }
380}