use super::Middleware;
use crate::error::MiddlewareError;
use std::sync::Arc;
use tracing::warn;
pub struct MiddlewareChain {
middlewares: Vec<Arc<dyn Middleware>>,
}
impl MiddlewareChain {
pub fn new() -> Self {
Self {
middlewares: Vec::new(),
}
}
pub fn add(&mut self, middleware: Arc<dyn Middleware>) {
let name = middleware.name().to_string();
self.middlewares.push(middleware);
tracing::debug!(middleware_name = %name, "Middleware added");
}
pub async fn before(&self, task_name: &str) -> Result<(), MiddlewareError> {
for middleware in &self.middlewares {
if let Err(e) = middleware.before(task_name).await {
warn!(
middleware_name = %middleware.name(),
task_name = %task_name,
error = %e,
"Middleware before failed, interrupting chain"
);
return Err(e);
}
}
Ok(())
}
pub async fn after(
&self,
task_name: &str,
result: &Result<(), Box<dyn std::error::Error + Send + Sync>>,
) -> Result<(), MiddlewareError> {
let mut has_error = false;
for middleware in &self.middlewares {
if let Err(e) = middleware.after(task_name, result).await {
warn!(
middleware_name = %middleware.name(),
task_name = %task_name,
error = %e,
"Middleware after failed"
);
has_error = true;
}
}
if has_error {
Err(MiddlewareError::ExecutionFailed {
name: "middleware-chain".to_string(),
reason: "One or more middlewares failed".to_string(),
})
} else {
Ok(())
}
}
pub fn middleware_count(&self) -> usize {
self.middlewares.len()
}
}
impl Default for MiddlewareChain {
fn default() -> Self {
Self::new()
}
}
#[cfg(test)]
mod tests {
use super::*;
struct TestMiddleware {
name: String,
}
impl Middleware for TestMiddleware {
fn name(&self) -> &str {
&self.name
}
}
#[tokio::test]
async fn test_middleware_chain_add() {
let mut chain = MiddlewareChain::new();
chain.add(Arc::new(TestMiddleware {
name: "test-middleware".to_string(),
}));
assert_eq!(chain.middleware_count(), 1);
}
#[tokio::test]
async fn test_middleware_chain_before() {
let mut chain = MiddlewareChain::new();
chain.add(Arc::new(TestMiddleware {
name: "test-middleware".to_string(),
}));
let result = chain.before("task-1").await;
assert!(result.is_ok());
}
#[tokio::test]
async fn test_middleware_chain_after() {
let mut chain = MiddlewareChain::new();
chain.add(Arc::new(TestMiddleware {
name: "test-middleware".to_string(),
}));
let task_result: Result<(), Box<dyn std::error::Error + Send + Sync>> = Ok(());
let result = chain.after("task-1", &task_result).await;
assert!(result.is_ok());
}
}