Skip to main content

tibba_runtime/
hook.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
15//! 启动 / 关闭钩子:全局注册表 + 按优先级驱动执行。
16
17use 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
28/// 装箱的异步 Future,用于 trait object 场景下的异步方法返回类型。
29pub type BoxFuture<'a, T> = Pin<Box<dyn Future<Output = T> + Send + 'a>>;
30
31/// 生命周期钩子 trait,用于在应用启动/关闭时执行自定义逻辑。
32///
33/// 由于实例以 `Arc<dyn Task>` 存入全局注册表并跨线程使用,要求实现类型必须
34/// 满足 `Send + Sync`——把约束直接写在 trait 上,存储类型就是简洁的
35/// `Arc<dyn Task>`,不必到处复述 `+ Send + Sync`。
36///
37/// - `before`:应用启动前执行(如初始化资源),按优先级从低到高顺序调用,**fail-fast**。
38/// - `after`:应用关闭后执行(如释放资源),按优先级从高到低顺序调用,**best-effort**
39///   (任一任务出错只记日志、继续执行后续清理,确保资源被尽可能释放)。
40/// - 返回 `true` 表示该钩子实际执行了操作,会记录耗时日志;返回 `false` 则静默跳过。
41pub trait Task: Send + Sync {
42    /// 应用启动前的钩子,默认不执行任何操作。
43    fn before(&self) -> BoxFuture<'_, Result<bool>> {
44        Box::pin(async { Ok(false) })
45    }
46    /// 应用关闭后的钩子,默认不执行任何操作。
47    fn after(&self) -> BoxFuture<'_, Result<bool>> {
48        Box::pin(async { Ok(false) })
49    }
50    /// 执行优先级,数值越小优先级越高(before 阶段),after 阶段反之。默认为 0。
51    fn priority(&self) -> u8 {
52        0
53    }
54}
55
56/// 全局任务注册表,键为任务名称,值为线程安全的任务实例。
57static TASKS: LazyLock<DashMap<String, Arc<dyn Task>>> = LazyLock::new(DashMap::new);
58
59/// 任务执行阶段:启动前(Before)或关闭后(After)。
60#[derive(Clone, Copy)]
61enum TaskType {
62    Before,
63    After,
64}
65
66impl TaskType {
67    /// 仅用于日志输出的短标签。
68    fn label(self) -> &'static str {
69        match self {
70            TaskType::Before => "before",
71            TaskType::After => "after",
72        }
73    }
74}
75
76/// 收集所有已注册任务并按当前阶段所需顺序排序。
77fn 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    // 用 i16 承载 priority 以便取负数实现降序,u8 无法直接配合 sort_by_key + Reverse
84    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
94/// 按优先级顺序执行所有已注册的钩子任务。
95/// - Before 阶段 fail-fast:首个错误立即返回,跳过剩余任务(避免半初始化的启动状态)
96/// - After 阶段 best-effort:每个错误记日志后继续,确保所有清理任务都被尝试,最终始终返回 Ok
97async 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                    // 启动期 fail-fast,避免应用以半初始化状态对外提供服务
126                    return Err(err);
127                }
128                // After 阶段继续执行后续清理任务
129            }
130        }
131    }
132    Ok(())
133}
134
135/// 注册一个具名钩子任务。同名任务重复注册时,新任务会覆盖旧任务。
136pub fn register_task(name: impl Into<String>, task: Arc<dyn Task>) {
137    TASKS.insert(name.into(), task);
138}
139
140/// 按优先级升序执行所有已注册的 `before` 钩子(应用启动前调用)。
141/// 任一任务返回错误时立即停止并向上传播。
142pub async fn run_before_tasks() -> Result<()> {
143    run_tasks(TaskType::Before).await
144}
145
146/// 按优先级降序执行所有已注册的 `after` 钩子(应用关闭后调用)。
147/// 任务错误仅记日志,所有任务都会被执行;本函数始终返回 `Ok(())`。
148pub async fn run_after_tasks() -> Result<()> {
149    run_tasks(TaskType::After).await
150}
151
152#[cfg(test)]
153// 测试通过 std::sync::Mutex 串行化共享的全局 TASKS 注册表,guard 跨 await
154// 持有;每个 #[tokio::test] 跑在独占的单线程 runtime 上,不存在跨任务的死锁
155// 风险,因此放行 await_holding_lock。
156#[allow(clippy::await_holding_lock)]
157mod tests {
158    use super::*;
159    use pretty_assertions::assert_eq;
160    use std::sync::Mutex;
161
162    /// 共享执行轨迹,用于断言任务执行顺序与是否被调用过。
163    type Trace = Arc<Mutex<Vec<&'static str>>>;
164
165    /// 通用测试任务:可配置名称、优先级以及 before/after 的返回值。
166    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    /// 清空全局注册表。测试通过 serial mutex 串行化以避免相互干扰。
199    fn reset() {
200        TASKS.clear();
201    }
202
203    /// 全局串行锁:注册表是单例,多个并发测试会污染彼此的注册项。
204    static SERIAL: Mutex<()> = Mutex::new(());
205
206    /// 取锁(PoisonError 时仍取回 guard),保证一次只跑一个测试。
207    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        // priority=50 先于 priority=10
276        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        // 首个失败后第二个任务不应被执行
308        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, // 先跑(after 降序)
321                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        // After 即便首个出错也应返回 Ok,并执行后续任务
338        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}