flare_core_runtime/middleware/
chain.rs1use super::Middleware;
6use crate::error::MiddlewareError;
7use std::sync::Arc;
8use tracing::warn;
9
10pub struct MiddlewareChain {
28 middlewares: Vec<Arc<dyn Middleware>>,
29}
30
31impl MiddlewareChain {
32 pub fn new() -> Self {
34 Self {
35 middlewares: Vec::new(),
36 }
37 }
38
39 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 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 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 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}