Skip to main content

vtcode_core/tools/
async_middleware.rs

1//! Async middleware for LLM-compatible tool execution
2//!
3//! Proper composition pattern with async/await support.
4//! Suitable for tokio-based systems handling LLM operations.
5
6use 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
14/// Type alias for the async continuation function
15type AsyncContinuation<'a> =
16    Box<dyn Fn(ToolRequest) -> Pin<Box<dyn Future<Output = MiddlewareToolResult> + Send>> + Send + Sync + 'a>;
17
18/// Type alias for the owned async continuation function
19type AsyncContinuationOwned =
20    Box<dyn Fn(ToolRequest) -> Pin<Box<dyn Future<Output = MiddlewareToolResult> + Send>> + Send + Sync>;
21
22/// Async middleware trait
23#[async_trait::async_trait]
24pub trait AsyncMiddleware: Send + Sync {
25    /// Middleware name
26    fn name(&self) -> &str;
27
28    /// Execute middleware
29    async fn execute<'a>(&'a self, request: ToolRequest, next: AsyncContinuation<'a>) -> MiddlewareToolResult;
30}
31
32/// Tool request
33#[derive(Clone, Debug)]
34pub struct ToolRequest {
35    pub tool_name: CompactStr,
36    pub arguments: String,
37    pub context: String,
38}
39
40/// Tool result
41#[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
50/// Async middleware chain executor
51pub 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    /// Execute request through chain (simplified)
66    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
158/// Async logging middleware
159pub 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
278/// Async caching middleware with UnifiedCache (migrated from LruCache)
279pub 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 a hashed key to avoid creating large string cache keys while still uniquely identifying args
306        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        // Check cache (migrated to UnifiedCache)
337        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        // Execute
351        let result = next(request).await;
352
353        // Cache successful result (migrated to UnifiedCache)
354        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
365/// Async retry middleware with exponential backoff
366pub 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            // Skip retry for non-retryable errors (auth failures, policy
432            // violations, invalid parameters) to fail fast.
433            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        // First call
521        let executor1 = make_executor("result1");
522
523        let result1 = cache.execute(request.clone(), executor1).await;
524        assert!(!result1.from_cache);
525
526        // Second call (should be cached)
527        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())); // Returns cached value
532    }
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}