Skip to main content

tibba_hook/
lib.rs

1// Copyright 2026 Tree xie.
2//
3// Licensed under the Apache License, Version 2.0 (the "License");
4// you may not use this file except in compliance with the License.
5// You may obtain a copy of the License at
6//
7// http://www.apache.org/licenses/LICENSE-2.0
8//
9// Unless required by applicable law or agreed to in writing, software
10// distributed under the License is distributed on an "AS IS" BASIS,
11// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12// See the License for the specific language governing permissions and
13// limitations under the License.
14
15use 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
23/// 该 crate 所有日志事件的 tracing target。
24/// 可通过 `RUST_LOG=tibba:hook=info`(或 `debug`)进行过滤。
25const LOG_TARGET: &str = "tibba:hook";
26
27type Result<T> = std::result::Result<T, Error>;
28
29/// 装箱的异步 Future,用于 trait object 场景下的异步方法返回类型。
30pub type BoxFuture<'a, T> = Pin<Box<dyn Future<Output = T> + Send + 'a>>;
31
32/// 生命周期钩子 trait,用于在应用启动/关闭时执行自定义逻辑。
33///
34/// 由于实例以 `Arc<dyn Task>` 存入全局注册表并跨线程使用,要求实现类型必须
35/// 满足 `Send + Sync`——把约束直接写在 trait 上,存储类型就是简洁的
36/// `Arc<dyn Task>`,不必到处复述 `+ Send + Sync`。
37///
38/// - `before`:应用启动前执行(如初始化资源),按优先级从低到高顺序调用,**fail-fast**。
39/// - `after`:应用关闭后执行(如释放资源),按优先级从高到低顺序调用,**best-effort**
40///   (任一任务出错只记日志、继续执行后续清理,确保资源被尽可能释放)。
41/// - 返回 `true` 表示该钩子实际执行了操作,会记录耗时日志;返回 `false` 则静默跳过。
42pub trait Task: Send + Sync {
43    /// 应用启动前的钩子,默认不执行任何操作。
44    fn before(&self) -> BoxFuture<'_, Result<bool>> {
45        Box::pin(async { Ok(false) })
46    }
47    /// 应用关闭后的钩子,默认不执行任何操作。
48    fn after(&self) -> BoxFuture<'_, Result<bool>> {
49        Box::pin(async { Ok(false) })
50    }
51    /// 执行优先级,数值越小优先级越高(before 阶段),after 阶段反之。默认为 0。
52    fn priority(&self) -> u8 {
53        0
54    }
55}
56
57/// 全局任务注册表,键为任务名称,值为线程安全的任务实例。
58static TASKS: LazyLock<DashMap<String, Arc<dyn Task>>> = LazyLock::new(DashMap::new);
59
60/// 任务执行阶段:启动前(Before)或关闭后(After)。
61#[derive(Clone, Copy)]
62enum TaskType {
63    Before,
64    After,
65}
66
67impl TaskType {
68    /// 仅用于日志输出的短标签。
69    fn label(self) -> &'static str {
70        match self {
71            TaskType::Before => "before",
72            TaskType::After => "after",
73        }
74    }
75}
76
77/// 收集所有已注册任务并按当前阶段所需顺序排序。
78fn 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    // 用 i16 承载 priority 以便取负数实现降序,u8 无法直接配合 sort_by_key + Reverse
85    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
95/// 按优先级顺序执行所有已注册的钩子任务。
96/// - Before 阶段 fail-fast:首个错误立即返回,跳过剩余任务(避免半初始化的启动状态)
97/// - After 阶段 best-effort:每个错误记日志后继续,确保所有清理任务都被尝试,最终始终返回 Ok
98async 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                    // 启动期 fail-fast,避免应用以半初始化状态对外提供服务
127                    return Err(err);
128                }
129                // After 阶段继续执行后续清理任务
130            }
131        }
132    }
133    Ok(())
134}
135
136/// 注册一个具名钩子任务。同名任务重复注册时,新任务会覆盖旧任务。
137pub fn register_task(name: impl Into<String>, task: Arc<dyn Task>) {
138    TASKS.insert(name.into(), task);
139}
140
141/// 按优先级升序执行所有已注册的 `before` 钩子(应用启动前调用)。
142/// 任一任务返回错误时立即停止并向上传播。
143pub async fn run_before_tasks() -> Result<()> {
144    run_tasks(TaskType::Before).await
145}
146
147/// 按优先级降序执行所有已注册的 `after` 钩子(应用关闭后调用)。
148/// 任务错误仅记日志,所有任务都会被执行;本函数始终返回 `Ok(())`。
149pub async fn run_after_tasks() -> Result<()> {
150    run_tasks(TaskType::After).await
151}
152
153#[cfg(test)]
154// 测试通过 std::sync::Mutex 串行化共享的全局 TASKS 注册表,guard 跨 await
155// 持有;每个 #[tokio::test] 跑在独占的单线程 runtime 上,不存在跨任务的死锁
156// 风险,因此放行 await_holding_lock。
157#[allow(clippy::await_holding_lock)]
158mod tests {
159    use super::*;
160    use pretty_assertions::assert_eq;
161    use std::sync::Mutex;
162
163    /// 共享执行轨迹,用于断言任务执行顺序与是否被调用过。
164    type Trace = Arc<Mutex<Vec<&'static str>>>;
165
166    /// 通用测试任务:可配置名称、优先级以及 before/after 的返回值。
167    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    /// 清空全局注册表。测试通过 serial mutex 串行化以避免相互干扰。
200    fn reset() {
201        TASKS.clear();
202    }
203
204    /// 全局串行锁:注册表是单例,多个并发测试会污染彼此的注册项。
205    static SERIAL: Mutex<()> = Mutex::new(());
206
207    /// 取锁(PoisonError 时仍取回 guard),保证一次只跑一个测试。
208    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        // priority=50 先于 priority=10
277        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        // 首个失败后第二个任务不应被执行
309        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, // 先跑(after 降序)
322                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        // After 即便首个出错也应返回 Ok,并执行后续任务
339        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}