use crate::context::RequestContext;
use crate::observability::MetricsCollector;
use pulseengine_auth::AuthenticationManager;
use pulseengine_mcp_protocol::*;
use pulseengine_mcp_security::SecurityMiddleware;
use async_trait::async_trait;
use std::sync::Arc;
use tracing::debug;
#[async_trait]
pub trait Middleware: Send + Sync {
async fn process_request(
&self,
request: Request,
context: &RequestContext,
) -> std::result::Result<Request, Error>;
async fn process_response(
&self,
response: Response,
context: &RequestContext,
) -> std::result::Result<Response, Error>;
}
#[derive(Clone)]
pub struct MiddlewareStack {
security: Option<SecurityMiddleware>,
auth: Option<Arc<AuthenticationManager>>,
monitoring: Option<Arc<MetricsCollector>>,
}
impl MiddlewareStack {
pub fn new() -> Self {
Self {
security: None,
auth: None,
monitoring: None,
}
}
pub fn with_security(mut self, security: SecurityMiddleware) -> Self {
self.security = Some(security);
self
}
pub fn with_auth(mut self, auth: Arc<AuthenticationManager>) -> Self {
self.auth = Some(auth);
self
}
pub fn with_monitoring(mut self, monitoring: Arc<MetricsCollector>) -> Self {
self.monitoring = Some(monitoring);
self
}
pub async fn process_request(
&self,
mut request: Request,
context: &RequestContext,
) -> std::result::Result<Request, crate::handler::HandlerError> {
debug!("Processing request through middleware stack");
if let Some(security) = &self.security {
let sec_context = pulseengine_mcp_security::middleware::RequestContext {
request_id: context.request_id,
};
request = security.process_request(request, &sec_context)?;
}
if let Some(monitoring) = &self.monitoring {
let mon_context = crate::observability::collector::RequestContext {
request_id: context.request_id,
};
request = monitoring.process_request(request, &mon_context)?;
}
Ok(request)
}
pub async fn process_response(
&self,
mut response: Response,
context: &RequestContext,
) -> std::result::Result<Response, crate::handler::HandlerError> {
debug!("Processing response through middleware stack");
if let Some(monitoring) = &self.monitoring {
let mon_context = crate::observability::collector::RequestContext {
request_id: context.request_id,
};
response = monitoring.process_response(response, &mon_context)?;
}
if let Some(security) = &self.security {
let sec_context = pulseengine_mcp_security::middleware::RequestContext {
request_id: context.request_id,
};
response = security.process_response(response, &sec_context)?;
}
Ok(response)
}
}
impl Default for MiddlewareStack {
fn default() -> Self {
Self::new()
}
}