1use crate::tools::improvements_errors::ObservabilityContext;
7use crate::types::CompactStr;
8use serde_json::{Map, Value};
9use std::future::Future;
10use std::pin::Pin;
11use std::sync::Arc;
12use std::time::Instant;
13
14type AsyncContinuation<'a> =
16 Box<dyn Fn(ToolRequest) -> Pin<Box<dyn Future<Output = MiddlewareToolResult> + Send>> + Send + Sync + 'a>;
17
18type AsyncContinuationOwned =
20 Box<dyn Fn(ToolRequest) -> Pin<Box<dyn Future<Output = MiddlewareToolResult> + Send>> + Send + Sync>;
21
22#[async_trait::async_trait]
24pub trait AsyncMiddleware: Send + Sync {
25 fn name(&self) -> &str;
27
28 async fn execute<'a>(&'a self, request: ToolRequest, next: AsyncContinuation<'a>) -> MiddlewareToolResult;
30}
31
32#[derive(Clone, Debug)]
34pub struct ToolRequest {
35 pub tool_name: CompactStr,
36 pub arguments: String,
37 pub context: String,
38}
39
40#[derive(Clone, Debug)]
42pub struct MiddlewareToolResult {
43 pub success: bool,
44 pub output: Option<String>,
45 pub error: Option<String>,
46 pub duration_ms: u64,
47 pub from_cache: bool,
48}
49
50pub struct AsyncMiddlewareChain {
52 middlewares: Vec<Arc<dyn AsyncMiddleware>>,
53}
54
55impl AsyncMiddlewareChain {
56 pub fn new() -> Self {
57 Self { middlewares: Vec::new() }
58 }
59
60 pub fn with_middleware(mut self, middleware: Arc<dyn AsyncMiddleware>) -> Self {
61 self.middlewares.push(middleware);
62 self
63 }
64
65 pub async fn execute_simple<F>(&self, request: ToolRequest, executor: F) -> MiddlewareToolResult
67 where
68 F: Fn(ToolRequest) -> MiddlewareToolResult + Send + Sync + 'static,
69 {
70 if self.middlewares.is_empty() {
71 return executor(request);
72 }
73
74 let executor = Arc::new(executor);
75 let middlewares = self.middlewares.clone();
76
77 fn build_chain(
78 middlewares: &[Arc<dyn AsyncMiddleware>],
79 executor: Arc<dyn Fn(ToolRequest) -> MiddlewareToolResult + Send + Sync>,
80 ) -> AsyncContinuationOwned {
81 if middlewares.is_empty() {
82 Box::new(move |req: ToolRequest| {
83 let result = executor(req);
84 Box::pin(async move { result })
85 })
86 } else {
87 let current = middlewares[0].clone();
88 let rest = build_chain(&middlewares[1..], executor);
89 let rest = Arc::new(rest);
90 Box::new(move |req: ToolRequest| {
91 let current = current.clone();
92 let rest = rest.clone();
93 Box::pin(async move {
94 let next: AsyncContinuationOwned = Box::new(move |r: ToolRequest| {
95 let rest = rest.clone();
96 Box::pin(async move { rest(r).await })
97 });
98 current.execute(req, next).await
99 })
100 })
101 }
102 }
103
104 let chain = build_chain(&middlewares, executor);
105 chain(request).await
106 }
107}
108
109impl Default for AsyncMiddlewareChain {
110 fn default() -> Self {
111 Self::new()
112 }
113}
114
115fn normalize_context(context: &str) -> String {
116 let mut normalized = Map::new();
117 let parsed: Value = serde_json::from_str(context).unwrap_or_else(|_| Value::Object(Map::new()));
118
119 if let Some(session) = parsed.get("session_id").and_then(Value::as_str)
120 && !session.is_empty()
121 {
122 normalized.insert("session_id".into(), Value::String(session.to_string()));
123 }
124
125 if let Some(task) = parsed.get("task_id").and_then(Value::as_str)
126 && !task.is_empty()
127 {
128 normalized.insert("task_id".into(), Value::String(task.to_string()));
129 }
130
131 if let Some(version) = parsed.get("plan_version").and_then(Value::as_u64) {
132 normalized.insert("plan_version".into(), Value::Number(version.into()));
133 }
134
135 if let Some(plan) = parsed.get("plan_summary").and_then(Value::as_object) {
136 let mut summary = Map::new();
137 if let Some(status) = plan.get("status").and_then(Value::as_str) {
138 summary.insert("status".into(), Value::String(status.to_string()));
139 }
140 if let Some(total) = plan.get("total_steps").and_then(Value::as_u64) {
141 summary.insert("total_steps".into(), Value::Number(total.into()));
142 }
143 if let Some(completed) = plan.get("completed_steps").and_then(Value::as_u64) {
144 summary.insert("completed_steps".into(), Value::Number(completed.into()));
145 }
146 if !summary.is_empty() {
147 normalized.insert("plan_summary".into(), Value::Object(summary));
148 }
149 }
150
151 if let Some(phase) = parsed.get("plan_phase").and_then(|v| v.as_str()).filter(|p| !p.is_empty()) {
152 normalized.insert("plan_phase".into(), Value::String(phase.to_string()));
153 }
154
155 serde_json::to_string(&Value::Object(normalized)).unwrap_or_else(|_| "{}".to_string())
156}
157
158pub struct AsyncLoggingMiddleware {
160 obs_context: Arc<ObservabilityContext>,
161}
162
163impl AsyncLoggingMiddleware {
164 pub fn new(obs_context: Arc<ObservabilityContext>) -> Self {
165 Self { obs_context }
166 }
167}
168
169#[async_trait::async_trait]
170impl AsyncMiddleware for AsyncLoggingMiddleware {
171 fn name(&self) -> &str {
172 "async_logging"
173 }
174
175 async fn execute<'a>(
176 &'a self,
177 request: ToolRequest,
178 next: Box<
179 dyn Fn(ToolRequest) -> Pin<Box<dyn std::future::Future<Output = MiddlewareToolResult> + Send>>
180 + Send
181 + Sync
182 + 'a,
183 >,
184 ) -> MiddlewareToolResult {
185 let tool_name = request.tool_name.clone();
186 let normalized_context = normalize_context(&request.context);
187 let context_json: Option<Value> = serde_json::from_str(&normalized_context).ok();
188 let session_id = context_json
189 .as_ref()
190 .and_then(|v| v.get("session_id").and_then(|s| s.as_str()))
191 .unwrap_or("");
192 let task_id = context_json
193 .as_ref()
194 .and_then(|v| v.get("task_id").and_then(|s| s.as_str()))
195 .unwrap_or("");
196 let plan_summary = context_json.as_ref().and_then(|v| v.get("plan_summary"));
197 let plan_status = plan_summary
198 .and_then(|v| v.get("status").and_then(|s| s.as_str()))
199 .unwrap_or("");
200 let plan_phase = context_json
201 .as_ref()
202 .and_then(|v| v.get("plan_phase").and_then(|p| p.as_str()))
203 .unwrap_or("");
204 let plan_total_steps = plan_summary
205 .and_then(|v| v.get("total_steps").and_then(|n| n.as_u64()))
206 .unwrap_or(0);
207 let plan_completed_steps = plan_summary
208 .and_then(|v| v.get("completed_steps").and_then(|n| n.as_u64()))
209 .unwrap_or(0);
210 let plan_version = context_json
211 .as_ref()
212 .and_then(|v| v.get("plan_version").and_then(|n| n.as_u64()))
213 .unwrap_or(0);
214
215 tracing::debug!(
216 tool = %tool_name,
217 session_id = %session_id,
218 task_id = %task_id,
219 plan_version,
220 plan_status = %plan_status,
221 plan_phase = %plan_phase,
222 plan_total_steps,
223 plan_completed_steps,
224 "tool execution started"
225 );
226 tracing::trace!(
227 tool = %tool_name,
228 context = %normalized_context,
229 "tool execution context payload"
230 );
231
232 let start = Instant::now();
233 let mut result = next(request).await;
234 let duration = start.elapsed().as_millis().min(u64::MAX as u128) as u64;
235
236 result.duration_ms = duration;
237
238 if result.success {
239 tracing::debug!(
240 tool = %tool_name,
241 duration_ms = duration,
242 session_id = %session_id,
243 task_id = %task_id,
244 plan_version,
245 plan_status = %plan_status,
246 plan_phase = %plan_phase,
247 plan_total_steps,
248 plan_completed_steps,
249 from_cache = result.from_cache,
250 "tool execution completed"
251 );
252 self.obs_context.event(
253 crate::tools::EventType::ToolSelected,
254 "executor",
255 format!("executed {tool_name} in {duration}ms"),
256 Some(1.0),
257 );
258 } else {
259 tracing::error!(
260 tool = %tool_name,
261 error = ?result.error,
262 session_id = %context_json
263 .as_ref()
264 .and_then(|v| v.get("session_id").and_then(|s| s.as_str()))
265 .unwrap_or(""),
266 task_id = %context_json
267 .as_ref()
268 .and_then(|v| v.get("task_id").and_then(|s| s.as_str()))
269 .unwrap_or(""),
270 "tool execution failed"
271 );
272 }
273
274 result
275 }
276}
277
278pub struct AsyncCachingMiddleware {
280 cache: Arc<crate::cache::UnifiedCache<AsyncCacheKey, String>>,
281 obs_context: Arc<ObservabilityContext>,
282}
283
284#[derive(Debug, Clone, Hash, PartialEq, Eq)]
285struct AsyncCacheKey(String);
286
287impl crate::cache::CacheKey for AsyncCacheKey {
288 fn to_cache_key(&self) -> String {
289 self.0.clone()
290 }
291}
292
293impl AsyncCachingMiddleware {
294 pub fn new(max_entries: usize, ttl_seconds: u64, obs_context: Arc<ObservabilityContext>) -> Self {
295 let cache = crate::cache::UnifiedCache::new(
296 max_entries,
297 std::time::Duration::from_secs(ttl_seconds),
298 crate::cache::EvictionPolicy::Lru,
299 );
300
301 Self { cache: Arc::new(cache), obs_context }
302 }
303
304 fn cache_key(tool: &str, args: &str, context: &str) -> String {
305 use std::collections::hash_map::DefaultHasher;
307 use std::hash::Hasher;
308 let mut hasher = DefaultHasher::new();
309 hasher.write(args.as_bytes());
310 let normalized = normalize_context(context);
311 if !normalized.is_empty() {
312 hasher.write(normalized.as_bytes());
313 }
314 format!("{}::{}", tool, hasher.finish())
315 }
316}
317
318#[async_trait::async_trait]
319impl AsyncMiddleware for AsyncCachingMiddleware {
320 fn name(&self) -> &str {
321 "async_caching"
322 }
323
324 async fn execute<'a>(
325 &'a self,
326 request: ToolRequest,
327 next: Box<
328 dyn Fn(ToolRequest) -> Pin<Box<dyn std::future::Future<Output = MiddlewareToolResult> + Send>>
329 + Send
330 + Sync
331 + 'a,
332 >,
333 ) -> MiddlewareToolResult {
334 let key = AsyncCacheKey(Self::cache_key(&request.tool_name, &request.arguments, &request.context));
335
336 if let Some(cached) = self.cache.get_owned(&key) {
338 self.obs_context
339 .event(crate::tools::EventType::CacheHit, "cache", "returning cached result", Some(1.0));
340
341 return MiddlewareToolResult {
342 success: true,
343 output: Some(cached),
344 error: None,
345 duration_ms: 0,
346 from_cache: true,
347 };
348 }
349
350 let result = next(request).await;
352
353 if result.success
355 && let Some(ref output) = result.output
356 {
357 let size = output.len() as u64;
358 self.cache.insert(key, output.clone(), size);
359 }
360
361 result
362 }
363}
364
365pub struct AsyncRetryMiddleware {
367 max_attempts: u32,
368 initial_backoff_ms: u64,
369 max_backoff_ms: u64,
370 obs_context: Arc<ObservabilityContext>,
371}
372
373impl AsyncRetryMiddleware {
374 pub fn new(
375 max_attempts: u32,
376 initial_backoff_ms: u64,
377 max_backoff_ms: u64,
378 obs_context: Arc<ObservabilityContext>,
379 ) -> Self {
380 Self {
381 max_attempts,
382 initial_backoff_ms,
383 max_backoff_ms,
384 obs_context,
385 }
386 }
387
388 fn backoff_duration(&self, attempt: u32) -> std::time::Duration {
389 let backoff = self.initial_backoff_ms * 2_u64.pow(attempt);
390 std::time::Duration::from_millis(backoff.min(self.max_backoff_ms))
391 }
392}
393
394#[async_trait::async_trait]
395impl AsyncMiddleware for AsyncRetryMiddleware {
396 fn name(&self) -> &str {
397 "async_retry"
398 }
399
400 async fn execute<'a>(
401 &'a self,
402 request: ToolRequest,
403 next: Box<
404 dyn Fn(ToolRequest) -> Pin<Box<dyn std::future::Future<Output = MiddlewareToolResult> + Send>>
405 + Send
406 + Sync
407 + 'a,
408 >,
409 ) -> MiddlewareToolResult {
410 for attempt in 0..self.max_attempts {
411 if attempt > 0 {
412 let backoff = self.backoff_duration(attempt - 1);
413 tracing::debug!(attempt = attempt, backoff_ms = backoff.as_millis(), "retrying after backoff");
414 tokio::time::sleep(backoff).await;
415 }
416
417 let mut result = next(request.clone()).await;
418
419 if result.success {
420 if attempt > 0 {
421 self.obs_context.event(
422 crate::tools::EventType::FallbackSuccess,
423 "retry",
424 format!("succeeded on attempt {}", attempt + 1),
425 Some(1.0),
426 );
427 }
428 return result;
429 }
430
431 if let Some(error_msg) = result.error.clone() {
434 let category = vtcode_commons::classify_error_message(&error_msg);
435 let guidance = vtcode_commons::detect_misconfiguration(category, &error_msg);
436 if !category.is_retryable() || guidance.is_some() {
437 if let Some(guidance) = guidance {
438 result.error = Some(format!("{error_msg}: {}", guidance.user_message()));
439 }
440 tracing::debug!(
441 attempt = attempt,
442 category = ?category,
443 "non-retryable error, skipping remaining attempts"
444 );
445 return result;
446 }
447 }
448
449 self.obs_context.event(
450 crate::tools::EventType::FallbackAttempt,
451 "retry",
452 format!("attempt {} failed", attempt + 1),
453 None,
454 );
455 }
456
457 MiddlewareToolResult {
458 success: false,
459 output: None,
460 error: Some(format!("all {} attempts failed", self.max_attempts)),
461 duration_ms: 0,
462 from_cache: false,
463 }
464 }
465}
466
467#[cfg(test)]
468mod tests {
469 use super::*;
470 use std::future::Future;
471 use std::pin::Pin;
472 use std::sync::atomic::{AtomicUsize, Ordering};
473
474 type BoxedToolFuture = Pin<Box<dyn Future<Output = MiddlewareToolResult> + Send>>;
475 type BoxedExecutor = Box<dyn Fn(ToolRequest) -> BoxedToolFuture + Send + Sync>;
476
477 fn make_executor(output: &'static str) -> BoxedExecutor {
478 Box::new(move |_req: ToolRequest| {
479 Box::pin(async move {
480 MiddlewareToolResult {
481 success: true,
482 output: Some(output.to_string()),
483 error: None,
484 duration_ms: 0,
485 from_cache: false,
486 }
487 })
488 })
489 }
490
491 #[tokio::test]
492 async fn test_async_logging_middleware() {
493 let obs = Arc::new(ObservabilityContext::noop());
494 let middleware = AsyncLoggingMiddleware::new(obs);
495
496 let request = ToolRequest {
497 tool_name: "test_tool".into(),
498 arguments: "arg1".to_string(),
499 context: "ctx".to_string(),
500 };
501
502 let executor = make_executor("result");
503
504 let result = middleware.execute(request, executor).await;
505
506 assert!(result.success);
507 }
508
509 #[tokio::test]
510 async fn test_async_caching_middleware() {
511 let obs = Arc::new(ObservabilityContext::noop());
512 let cache = AsyncCachingMiddleware::new(10, 60, obs);
513
514 let request = ToolRequest {
515 tool_name: "cached_tool".into(),
516 arguments: "arg1".to_string(),
517 context: "ctx".to_string(),
518 };
519
520 let executor1 = make_executor("result1");
522
523 let result1 = cache.execute(request.clone(), executor1).await;
524 assert!(!result1.from_cache);
525
526 let executor2 = make_executor("result2");
528
529 let result2 = cache.execute(request, executor2).await;
530 assert!(result2.from_cache);
531 assert_eq!(result2.output, Some("result1".to_string())); }
533
534 #[tokio::test]
535 async fn async_retry_skips_non_retryable_errors() {
536 let obs = Arc::new(ObservabilityContext::noop());
537 let middleware = AsyncRetryMiddleware::new(3, 1, 2, obs);
538 let attempts = Arc::new(AtomicUsize::new(0));
539
540 let executor_attempts = attempts.clone();
541 let executor: BoxedExecutor = Box::new(move |_req: ToolRequest| {
542 let executor_attempts = executor_attempts.clone();
543 Box::pin(async move {
544 executor_attempts.fetch_add(1, Ordering::SeqCst);
545 MiddlewareToolResult {
546 success: false,
547 output: None,
548 error: Some("invalid api key".to_string()),
549 duration_ms: 0,
550 from_cache: false,
551 }
552 })
553 });
554
555 let result = middleware
556 .execute(
557 ToolRequest {
558 tool_name: "auth_tool".into(),
559 arguments: "{}".to_string(),
560 context: "{}".to_string(),
561 },
562 executor,
563 )
564 .await;
565
566 assert!(!result.success);
567 assert_eq!(attempts.load(Ordering::SeqCst), 1);
568 }
569
570 #[tokio::test]
571 async fn async_retry_skips_misconfiguration_with_network_category() {
572 let obs = Arc::new(ObservabilityContext::noop());
573 let middleware = AsyncRetryMiddleware::new(3, 1, 2, obs);
574 let attempts = Arc::new(AtomicUsize::new(0));
575
576 let executor_attempts = attempts.clone();
577 let executor: BoxedExecutor = Box::new(move |_req: ToolRequest| {
578 let executor_attempts = executor_attempts.clone();
579 Box::pin(async move {
580 executor_attempts.fetch_add(1, Ordering::SeqCst);
581 MiddlewareToolResult {
582 success: false,
583 output: None,
584 error: Some("network error: invalid endpoint in base_url".to_string()),
585 duration_ms: 0,
586 from_cache: false,
587 }
588 })
589 });
590
591 let result = middleware
592 .execute(
593 ToolRequest {
594 tool_name: "provider_tool".into(),
595 arguments: "{}".to_string(),
596 context: "{}".to_string(),
597 },
598 executor,
599 )
600 .await;
601
602 assert!(!result.success);
603 assert_eq!(attempts.load(Ordering::SeqCst), 1);
604 assert!(
605 result
606 .error
607 .as_deref()
608 .is_some_and(|error| error.contains("Check settings/config first"))
609 );
610 }
611
612 #[tokio::test]
613 async fn async_retry_retries_retryable_errors_until_success() {
614 let obs = Arc::new(ObservabilityContext::noop());
615 let middleware = AsyncRetryMiddleware::new(3, 1, 2, obs);
616 let attempts = Arc::new(AtomicUsize::new(0));
617
618 let executor_attempts = attempts.clone();
619 let executor: BoxedExecutor = Box::new(move |_req: ToolRequest| {
620 let executor_attempts = executor_attempts.clone();
621 Box::pin(async move {
622 let attempt = executor_attempts.fetch_add(1, Ordering::SeqCst);
623 if attempt < 2 {
624 MiddlewareToolResult {
625 success: false,
626 output: None,
627 error: Some("429 Too Many Requests".to_string()),
628 duration_ms: 0,
629 from_cache: false,
630 }
631 } else {
632 MiddlewareToolResult {
633 success: true,
634 output: Some("ok".to_string()),
635 error: None,
636 duration_ms: 0,
637 from_cache: false,
638 }
639 }
640 })
641 });
642
643 let result = middleware
644 .execute(
645 ToolRequest {
646 tool_name: "rate_limited_tool".into(),
647 arguments: "{}".to_string(),
648 context: "{}".to_string(),
649 },
650 executor,
651 )
652 .await;
653
654 assert!(result.success);
655 assert_eq!(result.output.as_deref(), Some("ok"));
656 assert_eq!(attempts.load(Ordering::SeqCst), 3);
657 }
658}