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