Skip to main content

flare_core_runtime/middleware/
chain.rs

1//! 中间件链实现
2//!
3//! 管理所有中间件并按顺序执行
4
5use super::Middleware;
6use crate::error::MiddlewareError;
7use std::sync::Arc;
8use tracing::warn;
9
10/// 中间件链
11///
12/// 管理所有中间件并按顺序执行
13///
14/// # 示例
15///
16/// ```rust,ignore
17/// use flare_core_runtime::middleware::{MiddlewareChain, Middleware};
18///
19/// let mut chain = MiddlewareChain::new();
20/// chain.add(Arc::new(MyMiddleware));
21///
22/// // 执行中间件
23/// chain.before("task-1").await?;
24/// // ... 执行任务 ...
25/// chain.after("task-1", &result).await?;
26/// ```
27pub struct MiddlewareChain {
28    middlewares: Vec<Arc<dyn Middleware>>,
29}
30
31impl MiddlewareChain {
32    /// 创建新的中间件链
33    pub fn new() -> Self {
34        Self {
35            middlewares: Vec::new(),
36        }
37    }
38
39    /// 添加中间件
40    ///
41    /// # 参数
42    ///
43    /// * `middleware` - 要添加的中间件
44    pub fn add(&mut self, middleware: Arc<dyn Middleware>) {
45        let name = middleware.name().to_string();
46        self.middlewares.push(middleware);
47        tracing::debug!(middleware_name = %name, "Middleware added");
48    }
49
50    /// 执行所有中间件的 before 钩子
51    ///
52    /// 按注册顺序执行,如果某个中间件返回错误,则中断链
53    ///
54    /// # 参数
55    ///
56    /// * `task_name` - 任务名称
57    ///
58    /// # 返回
59    ///
60    /// - `Ok(())` - 所有中间件执行成功
61    /// - `Err(MiddlewareError)` - 某个中间件返回错误,链被中断
62    pub async fn before(&self, task_name: &str) -> Result<(), MiddlewareError> {
63        for middleware in &self.middlewares {
64            if let Err(e) = middleware.before(task_name).await {
65                warn!(
66                    middleware_name = %middleware.name(),
67                    task_name = %task_name,
68                    error = %e,
69                    "Middleware before failed, interrupting chain"
70                );
71                return Err(e);
72            }
73        }
74        Ok(())
75    }
76
77    /// 执行所有中间件的 after 钩子
78    ///
79    /// 按注册顺序执行,即使某个中间件失败也继续执行其他中间件
80    ///
81    /// # 参数
82    ///
83    /// * `task_name` - 任务名称
84    /// * `result` - 任务执行结果
85    ///
86    /// # 返回
87    ///
88    /// - `Ok(())` - 所有中间件执行成功
89    /// - `Err(MiddlewareError)` - 某个中间件返回错误(但所有中间件都会执行)
90    pub async fn after(
91        &self,
92        task_name: &str,
93        result: &Result<(), Box<dyn std::error::Error + Send + Sync>>,
94    ) -> Result<(), MiddlewareError> {
95        let mut has_error = false;
96
97        for middleware in &self.middlewares {
98            if let Err(e) = middleware.after(task_name, result).await {
99                warn!(
100                    middleware_name = %middleware.name(),
101                    task_name = %task_name,
102                    error = %e,
103                    "Middleware after failed"
104                );
105                has_error = true;
106            }
107        }
108
109        if has_error {
110            Err(MiddlewareError::ExecutionFailed {
111                name: "middleware-chain".to_string(),
112                reason: "One or more middlewares failed".to_string(),
113            })
114        } else {
115            Ok(())
116        }
117    }
118
119    /// 获取中间件数量
120    pub fn middleware_count(&self) -> usize {
121        self.middlewares.len()
122    }
123}
124
125impl Default for MiddlewareChain {
126    fn default() -> Self {
127        Self::new()
128    }
129}
130
131#[cfg(test)]
132mod tests {
133    use super::*;
134
135    struct TestMiddleware {
136        name: String,
137    }
138
139    impl Middleware for TestMiddleware {
140        fn name(&self) -> &str {
141            &self.name
142        }
143    }
144
145    #[tokio::test]
146    async fn test_middleware_chain_add() {
147        let mut chain = MiddlewareChain::new();
148        chain.add(Arc::new(TestMiddleware {
149            name: "test-middleware".to_string(),
150        }));
151
152        assert_eq!(chain.middleware_count(), 1);
153    }
154
155    #[tokio::test]
156    async fn test_middleware_chain_before() {
157        let mut chain = MiddlewareChain::new();
158        chain.add(Arc::new(TestMiddleware {
159            name: "test-middleware".to_string(),
160        }));
161
162        let result = chain.before("task-1").await;
163        assert!(result.is_ok());
164    }
165
166    #[tokio::test]
167    async fn test_middleware_chain_after() {
168        let mut chain = MiddlewareChain::new();
169        chain.add(Arc::new(TestMiddleware {
170            name: "test-middleware".to_string(),
171        }));
172
173        let task_result: Result<(), Box<dyn std::error::Error + Send + Sync>> = Ok(());
174        let result = chain.after("task-1", &task_result).await;
175        assert!(result.is_ok());
176    }
177}