use crate::tools::{ToolCallRequest, ToolCallResponse, UnifiedToolError};
use async_trait::async_trait;
use std::ops::Add;
use std::sync::Arc;
use std::sync::atomic::{AtomicU64, Ordering};
pub type ToolRequest = ToolCallRequest;
pub type ToolResponse = ToolCallResponse;
pub type MiddlewareResult<T> = Result<T, UnifiedToolError>;
#[async_trait]
pub trait Middleware: Send + Sync {
async fn before_execute(&self, _req: &ToolRequest) -> MiddlewareResult<()> {
Ok(())
}
async fn after_execute(&self, _req: &ToolRequest, _res: &ToolResponse) -> MiddlewareResult<()> {
Ok(())
}
async fn on_error(&self, _req: &ToolRequest, _err: &UnifiedToolError) -> MiddlewareResult<()> {
Ok(())
}
}
pub struct NoopMiddleware;
#[async_trait]
impl Middleware for NoopMiddleware {}
#[derive(Clone)]
pub struct MiddlewareChain {
middlewares: Vec<Arc<dyn Middleware>>,
}
impl MiddlewareChain {
pub fn new() -> Self {
Self { middlewares: Vec::new() }
}
pub fn push(mut self, mw: Arc<dyn Middleware>) -> Self {
self.middlewares.push(mw);
self
}
pub async fn before_execute_opt(&self, req: &ToolRequest, fail_open: bool) -> MiddlewareResult<()> {
if !fail_open {
return self.before_execute(req).await;
}
let mut last_err = None;
for mw in &self.middlewares {
if let Err(err) = mw.before_execute(req).await {
tracing::warn!(
tool = %req.tool_name,
error = %err,
"middleware before_execute failed; ignoring (fail-open)"
);
last_err = Some(err);
}
}
if let Some(err) = last_err { Err(err) } else { Ok(()) }
}
pub async fn before_execute(&self, req: &ToolRequest) -> MiddlewareResult<()> {
for mw in &self.middlewares {
mw.before_execute(req).await?;
}
Ok(())
}
pub async fn after_execute(&self, req: &ToolRequest, res: &ToolResponse) -> MiddlewareResult<()> {
for mw in self.middlewares.iter().rev() {
mw.after_execute(req, res).await?;
}
Ok(())
}
pub async fn on_error(&self, req: &ToolRequest, err: &UnifiedToolError) -> MiddlewareResult<()> {
for mw in self.middlewares.iter().rev() {
if let Err(handler_err) = mw.on_error(req, err).await {
tracing::warn!(
error = %handler_err,
"Middleware error handler itself failed"
);
}
}
Ok(())
}
}
impl Add<Arc<dyn Middleware>> for MiddlewareChain {
type Output = Self;
fn add(self, rhs: Arc<dyn Middleware>) -> Self::Output {
self.push(rhs)
}
}
impl Default for MiddlewareChain {
fn default() -> Self {
Self::new()
}
}
pub struct LoggingMiddleware {
name: String,
}
impl LoggingMiddleware {
pub fn new(name: impl Into<String>) -> Arc<Self> {
Arc::new(Self { name: name.into() })
}
}
#[async_trait]
impl Middleware for LoggingMiddleware {
async fn before_execute(&self, req: &ToolRequest) -> MiddlewareResult<()> {
tracing::info!(
middleware = %self.name,
tool = %req.tool_name,
"Executing tool"
);
Ok(())
}
async fn after_execute(&self, req: &ToolRequest, res: &ToolResponse) -> MiddlewareResult<()> {
let duration_ms = res.duration_ms.unwrap_or(0);
let cache_hit = res.cache_hit.unwrap_or(false);
tracing::info!(
middleware = %self.name,
tool = %req.tool_name,
duration_ms,
cache_hit,
"Completed tool"
);
Ok(())
}
async fn on_error(&self, req: &ToolRequest, err: &UnifiedToolError) -> MiddlewareResult<()> {
tracing::error!(
middleware = %self.name,
tool = %req.tool_name,
error = %err,
"Tool execution failed"
);
Ok(())
}
}
#[derive(Clone, Copy, Debug)]
pub struct MetricsSnapshot {
pub total_calls: u64,
pub successful_calls: u64,
pub failed_calls: u64,
pub total_duration_ms: u64,
pub cache_hits: u64,
}
pub struct MetricsMiddleware {
total_calls: Arc<AtomicU64>,
successful_calls: Arc<AtomicU64>,
failed_calls: Arc<AtomicU64>,
total_duration_ms: Arc<AtomicU64>,
cache_hits: Arc<AtomicU64>,
}
impl MetricsMiddleware {
fn new_inner() -> Self {
Self {
total_calls: Arc::new(AtomicU64::new(0)),
successful_calls: Arc::new(AtomicU64::new(0)),
failed_calls: Arc::new(AtomicU64::new(0)),
total_duration_ms: Arc::new(AtomicU64::new(0)),
cache_hits: Arc::new(AtomicU64::new(0)),
}
}
pub fn new() -> Arc<Self> {
Arc::new(Self::new_inner())
}
pub async fn snapshot(&self) -> MetricsSnapshot {
MetricsSnapshot {
total_calls: self.total_calls.load(Ordering::Relaxed),
successful_calls: self.successful_calls.load(Ordering::Relaxed),
failed_calls: self.failed_calls.load(Ordering::Relaxed),
total_duration_ms: self.total_duration_ms.load(Ordering::Relaxed),
cache_hits: self.cache_hits.load(Ordering::Relaxed),
}
}
}
#[async_trait]
impl Middleware for MetricsMiddleware {
async fn before_execute(&self, _: &ToolRequest) -> MiddlewareResult<()> {
self.total_calls.fetch_add(1, Ordering::Relaxed);
Ok(())
}
async fn after_execute(&self, _: &ToolRequest, res: &ToolResponse) -> MiddlewareResult<()> {
self.successful_calls.fetch_add(1, Ordering::Relaxed);
self.total_duration_ms
.fetch_add(res.duration_ms.unwrap_or(0), Ordering::Relaxed);
if res.cache_hit.unwrap_or(false) {
self.cache_hits.fetch_add(1, Ordering::Relaxed);
}
Ok(())
}
async fn on_error(&self, _: &ToolRequest, _: &UnifiedToolError) -> MiddlewareResult<()> {
self.failed_calls.fetch_add(1, Ordering::Relaxed);
Ok(())
}
}
impl Default for MetricsMiddleware {
fn default() -> Self {
Self::new_inner()
}
}
#[cfg(test)]
mod tests {
use super::*;
use serde_json::Value;
#[tokio::test]
async fn test_chain_execution() {
let chain = MiddlewareChain::new()
.push(LoggingMiddleware::new("test"))
.push(MetricsMiddleware::new());
let req = ToolRequest {
id: "req-1".to_string(),
tool_name: "test_tool".into(),
args: Value::Null,
metadata: Some(Default::default()),
};
chain.before_execute(&req).await.unwrap();
let res = ToolResponse {
id: "req-1".to_string(),
success: true,
result: Some(Value::Null),
error: None,
duration_ms: Some(100),
cache_hit: Some(false),
};
chain.after_execute(&req, &res).await.unwrap();
}
#[tokio::test]
async fn test_metrics_tracking() {
let metrics = MetricsMiddleware::new();
let chain = MiddlewareChain::new().push(metrics.clone());
let req = ToolRequest {
id: "req-2".to_string(),
tool_name: "test".into(),
args: Value::Null,
metadata: Some(Default::default()),
};
for i in 0..5 {
chain.before_execute(&req).await.unwrap();
let res = ToolResponse {
id: format!("req-2-{i}"),
success: true,
result: Some(Value::Null),
error: None,
duration_ms: Some(10),
cache_hit: Some(true),
};
chain.after_execute(&req, &res).await.unwrap();
}
let snapshot = metrics.snapshot().await;
assert_eq!(snapshot.total_calls, 5);
assert_eq!(snapshot.successful_calls, 5);
assert_eq!(snapshot.cache_hits, 5);
}
}