use crate::config::constants::execution;
use crate::tools::lru_cache::LruCache;
use crate::tools::pattern_detection::{PatternDetector, ToolEvent};
use crate::tools::tool_middleware::{MiddlewareChain, MiddlewareResult, ToolRequest, ToolResponse};
use crate::tools::{UnifiedErrorKind, UnifiedToolError};
use serde_json::Value;
use std::collections::hash_map::DefaultHasher;
use std::hash::Hasher;
use std::sync::Arc;
use std::sync::RwLock;
use std::sync::atomic::{AtomicU64, Ordering};
use std::time::Duration;
use tokio::time::timeout;
use tracing::warn;
static REQUEST_ID: AtomicU64 = AtomicU64::new(1);
fn next_request_id() -> String {
let id = REQUEST_ID.fetch_add(1, Ordering::Relaxed);
format!("req-{id}")
}
#[repr(align(64))]
struct PaddedAtomicU64(AtomicU64);
impl PaddedAtomicU64 {
fn new(value: u64) -> Self {
Self(AtomicU64::new(value))
}
#[inline]
fn fetch_add(&self, value: u64, ordering: Ordering) -> u64 {
self.0.fetch_add(value, ordering)
}
#[inline]
fn load(&self, ordering: Ordering) -> u64 {
self.0.load(ordering)
}
}
#[derive(Clone, Debug)]
pub struct ExecutorStats {
pub total_calls: u64,
pub successful_calls: u64,
pub failed_calls: u64,
pub cache_hits: u64,
pub cache_misses: u64,
pub avg_duration_ms: u64,
pub patterns_detected: usize,
}
struct AtomicStats {
total_calls: PaddedAtomicU64,
successful_calls: PaddedAtomicU64,
failed_calls: PaddedAtomicU64,
cache_hits: PaddedAtomicU64,
cache_misses: PaddedAtomicU64,
total_duration_ms: PaddedAtomicU64,
duration_count: PaddedAtomicU64,
}
impl AtomicStats {
fn new() -> Self {
Self {
total_calls: PaddedAtomicU64::new(0),
successful_calls: PaddedAtomicU64::new(0),
failed_calls: PaddedAtomicU64::new(0),
cache_hits: PaddedAtomicU64::new(0),
cache_misses: PaddedAtomicU64::new(0),
total_duration_ms: PaddedAtomicU64::new(0),
duration_count: PaddedAtomicU64::new(0),
}
}
#[inline]
fn record_success(&self, duration_ms: u64) {
self.successful_calls.fetch_add(1, Ordering::Relaxed);
self.total_duration_ms.fetch_add(duration_ms, Ordering::Relaxed);
self.duration_count.fetch_add(1, Ordering::Relaxed);
}
#[inline]
fn record_failure(&self, duration_ms: u64) {
self.failed_calls.fetch_add(1, Ordering::Relaxed);
self.total_duration_ms.fetch_add(duration_ms, Ordering::Relaxed);
self.duration_count.fetch_add(1, Ordering::Relaxed);
}
fn snapshot(&self) -> ExecutorStats {
let count = self.duration_count.load(Ordering::Relaxed);
let total = self.total_duration_ms.load(Ordering::Relaxed);
let avg = if count > 0 { total / count } else { 0 };
ExecutorStats {
total_calls: self.total_calls.load(Ordering::Relaxed),
successful_calls: self.successful_calls.load(Ordering::Relaxed),
failed_calls: self.failed_calls.load(Ordering::Relaxed),
cache_hits: self.cache_hits.load(Ordering::Relaxed),
cache_misses: self.cache_misses.load(Ordering::Relaxed),
avg_duration_ms: avg,
patterns_detected: 0, }
}
}
pub struct CachedToolExecutor {
cache: Arc<LruCache<Value>>,
middleware: MiddlewareChain,
patterns: Arc<RwLock<PatternDetector>>,
stats: Arc<AtomicStats>,
}
impl CachedToolExecutor {
pub fn new() -> Self {
Self::with_config(1000, Duration::from_secs(3600), 3)
}
pub fn with_config(cache_capacity: usize, cache_ttl: Duration, pattern_window: usize) -> Self {
let cache = Arc::new(LruCache::<Value>::new(cache_capacity, cache_ttl));
let middleware = MiddlewareChain::new();
let patterns = Arc::new(RwLock::new(PatternDetector::new(pattern_window)));
let stats = Arc::new(AtomicStats::new());
Self { cache, middleware, patterns, stats }
}
pub fn with_middleware(mut self, mw: Arc<dyn crate::tools::tool_middleware::Middleware>) -> Self {
self.middleware = self.middleware.push(mw);
self
}
pub async fn execute(&self, tool_name: &str, args: Value) -> MiddlewareResult<Value> {
let r = self.execute_shared_owned(tool_name, args).await?;
Ok((*r).clone())
}
pub async fn execute_shared(&self, tool_name: &str, args: Arc<Value>) -> MiddlewareResult<Arc<Value>> {
let start = std::time::Instant::now();
let cache_key = make_cache_key(tool_name, &args);
self.stats.total_calls.fetch_add(1, Ordering::Relaxed);
let owned_args = Arc::clone(&args);
let req = ToolRequest {
id: next_request_id(),
tool_name: tool_name.into(),
args: (*owned_args).clone(),
metadata: Some(Default::default()),
};
if let Err(err) = self.middleware.before_execute(&req).await {
self.record_error(tool_name, start.elapsed(), &req, &err).await;
return Err(err);
}
if let Some(result) = self.cache.get(&cache_key).await {
let duration_ms = start.elapsed().as_millis() as u64;
self.stats.record_success(duration_ms);
self.stats.cache_hits.fetch_add(1, Ordering::Relaxed);
let res = ToolResponse {
id: req.id.clone(),
success: true,
result: Some((*result).clone()),
error: None,
duration_ms: Some(duration_ms),
cache_hit: Some(true),
};
if let Err(err) = self.middleware.after_execute(&req, &res).await {
self.record_error(tool_name, start.elapsed(), &req, &err).await;
return Err(err);
}
self.record_pattern(tool_name, true, duration_ms).await;
return Ok(Arc::clone(&result));
}
self.stats.cache_misses.fetch_add(1, Ordering::Relaxed);
let timeout_secs = execution::DEFAULT_TIMEOUT_SECS;
let result = match timeout(
Duration::from_secs(timeout_secs),
self.execute_tool_internal(tool_name, &owned_args),
)
.await
{
Ok(result) => match result {
Ok(result) => result,
Err(err) => {
self.record_error(tool_name, start.elapsed(), &req, &err).await;
return Err(err);
}
},
Err(_) => {
let err = UnifiedToolError::new(
UnifiedErrorKind::Timeout,
format!("Tool '{tool_name}' timed out after {timeout_secs} seconds"),
)
.with_tool_name(tool_name);
self.record_error(tool_name, start.elapsed(), &req, &err).await;
return Err(err);
}
};
let duration_ms = start.elapsed().as_millis() as u64;
let arc_res = Arc::new(result);
self.cache.insert_arc(cache_key, Arc::clone(&arc_res)).await;
self.stats.record_success(duration_ms);
let res = ToolResponse {
id: req.id.clone(),
success: true,
result: Some((*arc_res).clone()),
error: None,
duration_ms: Some(duration_ms),
cache_hit: Some(false),
};
if let Err(err) = self.middleware.after_execute(&req, &res).await {
self.record_error(tool_name, start.elapsed(), &req, &err).await;
return Err(err);
}
self.record_pattern(tool_name, true, duration_ms).await;
Ok(arc_res)
}
pub async fn execute_shared_owned(&self, tool_name: &str, args: Value) -> MiddlewareResult<Arc<Value>> {
let arg = Arc::new(args);
self.execute_shared(tool_name, arg).await
}
async fn execute_tool_internal(&self, _tool_name: &str, _args: &Value) -> MiddlewareResult<Value> {
Ok(serde_json::json!({"status": "ok"}))
}
async fn record_error(&self, tool_name: &str, elapsed: Duration, req: &ToolRequest, err: &UnifiedToolError) {
self.stats.record_failure(elapsed.as_millis() as u64);
let _ = self.middleware.on_error(req, err).await;
self.record_pattern(tool_name, false, elapsed.as_millis() as u64).await;
}
async fn record_pattern(&self, tool_name: &str, success: bool, duration_ms: u64) {
let mut patterns = self.patterns.write().unwrap_or_else(|poisoned| {
warn!("pattern detector lock poisoned; recovering");
poisoned.into_inner()
});
patterns.record_event(ToolEvent {
tool_name: tool_name.to_string(),
success,
duration_ms,
timestamp: std::time::Instant::now(),
});
}
pub async fn stats(&self) -> ExecutorStats {
let mut stats = self.stats.snapshot();
let patterns = self.patterns.read().unwrap_or_else(|poisoned| {
warn!("pattern detector lock poisoned; recovering");
poisoned.into_inner()
});
stats.patterns_detected = patterns.patterns().len();
stats
}
pub async fn cache_stats(&self) -> crate::tools::lru_cache::CacheStats {
self.cache.stats().await
}
pub async fn patterns(&self) -> Vec<crate::tools::pattern_detection::DetectedPattern> {
let patterns = self.patterns.read().unwrap_or_else(|poisoned| {
warn!("pattern detector lock poisoned; recovering");
poisoned.into_inner()
});
patterns.patterns().to_vec()
}
pub async fn feature_vector(&self) -> Vec<f64> {
let patterns = self.patterns.read().unwrap_or_else(|poisoned| {
warn!("pattern detector lock poisoned; recovering");
poisoned.into_inner()
});
patterns.feature_vector()
}
pub async fn clear_cache(&self) {
self.cache.clear().await;
}
pub async fn clear_patterns(&self) {
let mut patterns = self.patterns.write().unwrap_or_else(|poisoned| {
warn!("pattern detector lock poisoned; recovering");
poisoned.into_inner()
});
patterns.reset();
}
pub async fn report(&self) {
let stats = self.stats().await;
let cache_stats = self.cache_stats().await;
let patterns = self.patterns().await;
println!("\n=== ToolExecutor Report ===\n");
println!("Execution Statistics:");
println!(" Total calls: {}", stats.total_calls);
println!(" Successful: {}", stats.successful_calls);
println!(" Failed: {}", stats.failed_calls);
println!(" Avg duration: {}ms", stats.avg_duration_ms);
println!("\nCache Performance:");
println!(" Hits: {}", cache_stats.hits);
println!(" Misses: {}", cache_stats.misses);
println!(" Hit rate: {:.1}%", cache_stats.hit_rate());
println!(" Evictions: {}", cache_stats.evictions);
println!(" Expirations: {}", cache_stats.expirations);
println!("\nWorkflow Patterns ({} detected):", patterns.len());
for (i, pattern) in patterns.iter().take(5).enumerate() {
println!(" {}. {:?}", i + 1, pattern.sequence);
println!(" Frequency: {}, Confidence: {:.1}%", pattern.frequency, pattern.confidence * 100.0);
}
println!("\n");
}
}
#[inline]
fn make_cache_key(tool_name: &str, args: &Value) -> String {
use std::hash::Hash;
let mut hasher = DefaultHasher::new();
tool_name.hash(&mut hasher);
if let Ok(bytes) = serde_json::to_vec(args) {
hasher.write(&bytes);
}
let h = hasher.finish();
format!("{tool_name}:{h:x}")
}
impl Default for CachedToolExecutor {
fn default() -> Self {
Self::new()
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::tools::tool_middleware::{LoggingMiddleware, MetricsMiddleware, MiddlewareResult};
use crate::tools::{UnifiedErrorKind, UnifiedToolError};
use async_trait::async_trait;
use std::sync::Arc;
struct FailingMiddleware;
#[async_trait]
impl crate::tools::tool_middleware::Middleware for FailingMiddleware {
async fn before_execute(&self, _req: &ToolRequest) -> MiddlewareResult<()> {
Err(UnifiedToolError::new(UnifiedErrorKind::ExecutionFailed, "middleware rejected request"))
}
}
#[tokio::test]
async fn test_executor_basic() -> anyhow::Result<()> {
let executor = CachedToolExecutor::new();
let result = executor.execute("test_tool", serde_json::json!({"arg": 1})).await?;
assert_eq!(result, serde_json::json!({"status": "ok"}));
let stats = executor.stats().await;
assert_eq!(stats.total_calls, 1);
assert_eq!(stats.cache_misses, 1);
Ok(())
}
#[tokio::test]
async fn test_executor_cache_hit() -> anyhow::Result<()> {
let executor = CachedToolExecutor::new();
executor.execute("test_tool", serde_json::json!({"arg": 1})).await?;
executor.execute("test_tool", serde_json::json!({"arg": 1})).await?;
let stats = executor.stats().await;
assert_eq!(stats.total_calls, 2);
assert_eq!(stats.successful_calls, 2);
assert_eq!(stats.cache_hits, 1);
assert_eq!(stats.cache_misses, 1);
Ok(())
}
#[tokio::test]
async fn cache_hit_after_repeat_call() {
let exec = CachedToolExecutor::with_config(10, Duration::from_secs(60), 3);
let args = serde_json::json!({"x": 1});
let _first = exec.execute("test_tool", args.clone()).await.unwrap();
let _second = exec.execute("test_tool", args.clone()).await.unwrap();
let stats = exec.stats().await;
assert!(stats.cache_hits >= 1);
assert!(stats.cache_misses >= 1);
}
#[tokio::test]
async fn test_executor_with_middleware() -> anyhow::Result<()> {
let executor = CachedToolExecutor::new().with_middleware(LoggingMiddleware::new("test"));
executor.execute("test_tool", serde_json::json!({})).await?;
let stats = executor.stats().await;
assert_eq!(stats.total_calls, 1);
Ok(())
}
#[tokio::test]
async fn test_executor_patterns() -> anyhow::Result<()> {
let executor = CachedToolExecutor::new();
for _ in 0..6 {
executor.execute("tool_a", serde_json::json!({})).await?;
executor.execute("tool_b", serde_json::json!({})).await?;
}
let patterns = executor.patterns().await;
assert!(!patterns.is_empty());
Ok(())
}
#[tokio::test]
async fn test_executor_clear() -> anyhow::Result<()> {
let executor = CachedToolExecutor::new();
executor.execute("test", serde_json::json!({})).await?;
let stats_before = executor.stats().await;
assert_eq!(stats_before.total_calls, 1);
executor.clear_cache().await;
executor.clear_patterns().await;
let cache_stats = executor.cache_stats().await;
assert_eq!(cache_stats.hits + cache_stats.misses, 0);
Ok(())
}
#[tokio::test]
async fn test_executor_reports_typed_error_to_middleware_metrics() {
let metrics = MetricsMiddleware::new();
let executor = CachedToolExecutor::new()
.with_middleware(metrics.clone())
.with_middleware(Arc::new(FailingMiddleware));
let err = executor.execute("test", serde_json::json!({})).await;
assert!(err.is_err());
let stats = executor.stats().await;
assert_eq!(stats.failed_calls, 1);
let snapshot = metrics.snapshot().await;
assert_eq!(snapshot.total_calls, 1);
assert_eq!(snapshot.failed_calls, 1);
}
}