gemini_adk_fluent_rs/compose/
middleware.rs

1//! M — Middleware composition.
2//!
3//! Compose middleware in any order with `|`.
4//!
5//! ## Wiring status
6//!
7//! **TextAgent pipelines** (via `AgentBuilder::middleware` + `AgentBuilder::build`) —
8//! **fully wired**.  Every factory in this module produces a `MiddlewareComposite`
9//! whose layers are installed into the `LlmTextAgent` middleware chain at compile
10//! time.  Hooks fire in this order per `run()` call:
11//!
12//! 1. `before_model` (forward order) — may short-circuit with a cached response.
13//! 2. LLM call (skipped if `before_model` returned `Some`).
14//! 3. `after_model` (reverse order) — may replace the LLM response.
15//! 4. `before_tool` (forward) / `after_tool` (reverse) / `on_tool_error` (forward)
16//!    — called for each tool dispatch round.
17//! 5. `on_error` (forward) — called once if `run()` returns an error.
18//!
19//! **Live sessions** (via `Live::middleware`) — the **tool-lifecycle hooks**
20//! are wired: `before_tool` (a returned error vetoes the call), `after_tool`,
21//! and `on_tool_error` fire around every tool dispatch in the control lane,
22//! including background tools. Model-level hooks (`before_model`/`after_model`)
23//! do **not** apply to Live — a Live session streams over the wire and has no
24//! discrete `LlmRequest`/`LlmResponse` to intercept; use them on TextAgent
25//! pipelines instead.
26
27use std::sync::Arc;
28use std::time::Duration;
29
30use async_trait::async_trait;
31use gemini_genai_rs::prelude::FunctionCall;
32
33use gemini_adk_rs::context::AgentEvent;
34use gemini_adk_rs::error::{AgentError, ToolError};
35use gemini_adk_rs::middleware::{LatencyMiddleware, LogMiddleware, Middleware};
36
37/// A middleware composite — one or more middleware layers (`M::a() | M::b()`).
38///
39/// A single `Arc<dyn Middleware>` converts into a one-layer composite, so
40/// `.middleware(Arc::new(MyLayer))` works without the namespace.
41#[derive(Clone)]
42#[non_exhaustive]
43pub struct MiddlewareComposite {
44    /// The ordered list of middleware layers.
45    pub layers: Vec<Arc<dyn Middleware>>,
46}
47
48impl MiddlewareComposite {
49    /// Create a composite containing a single middleware layer.
50    pub fn new(layer: Arc<dyn Middleware>) -> Self {
51        Self {
52            layers: vec![layer],
53        }
54    }
55
56    /// Number of layers.
57    pub fn len(&self) -> usize {
58        self.layers.len()
59    }
60
61    /// Whether empty.
62    pub fn is_empty(&self) -> bool {
63        self.layers.is_empty()
64    }
65}
66
67impl From<Arc<dyn Middleware>> for MiddlewareComposite {
68    fn from(layer: Arc<dyn Middleware>) -> Self {
69        Self::new(layer)
70    }
71}
72
73/// Compose two middleware composites with `|`.
74impl std::ops::BitOr for MiddlewareComposite {
75    type Output = MiddlewareComposite;
76
77    fn bitor(mut self, rhs: MiddlewareComposite) -> Self::Output {
78        self.layers.extend(rhs.layers);
79        self
80    }
81}
82
83/// The `M` namespace — static factory methods for middleware.
84pub struct M;
85
86impl M {
87    /// Add logging middleware.
88    pub fn log() -> MiddlewareComposite {
89        MiddlewareComposite::new(Arc::new(LogMiddleware::new()))
90    }
91
92    /// Add latency tracking middleware.
93    pub fn latency() -> MiddlewareComposite {
94        MiddlewareComposite::new(Arc::new(LatencyMiddleware::new()))
95    }
96
97    /// Bound the agent run to `duration`. The text agent enforces the tightest
98    /// timeout across its middleware chain by wrapping the whole run; on elapse
99    /// it emits `AgentEvent::Timeout` and returns an error.
100    pub fn timeout(duration: Duration) -> MiddlewareComposite {
101        MiddlewareComposite::new(Arc::new(TimeoutMiddleware {
102            name: "timeout".to_string(),
103            duration,
104        }))
105    }
106
107    /// Add retry middleware — tracks errors and advises on retry.
108    pub fn retry(max_retries: u32) -> MiddlewareComposite {
109        MiddlewareComposite::new(Arc::new(gemini_adk_rs::middleware::RetryMiddleware::new(
110            max_retries,
111        )))
112    }
113
114    /// Add a custom event observer — called on every agent event.
115    pub fn tap(f: impl Fn(&AgentEvent) + Send + Sync + 'static) -> MiddlewareComposite {
116        MiddlewareComposite::new(Arc::new(TapMiddleware {
117            handler: Arc::new(f),
118        }))
119    }
120
121    /// Add a custom before-tool filter — called before every tool invocation.
122    pub fn before_tool(
123        f: impl Fn(&FunctionCall) -> Result<(), String> + Send + Sync + 'static,
124    ) -> MiddlewareComposite {
125        MiddlewareComposite::new(Arc::new(BeforeToolMiddleware {
126            handler: Arc::new(f),
127        }))
128    }
129
130    /// Add a custom after-tool hook — called after every successful tool invocation.
131    pub fn after_tool(
132        f: impl Fn(&FunctionCall, &serde_json::Value) -> Result<(), String> + Send + Sync + 'static,
133    ) -> MiddlewareComposite {
134        MiddlewareComposite::new(Arc::new(AfterToolMiddleware {
135            handler: Arc::new(f),
136        }))
137    }
138
139    /// Add a custom error observer — called when an agent-level error occurs.
140    pub fn on_error(
141        f: impl Fn(&AgentError) -> Result<(), String> + Send + Sync + 'static,
142    ) -> MiddlewareComposite {
143        MiddlewareComposite::new(Arc::new(OnErrorMiddleware {
144            handler: Arc::new(f),
145        }))
146    }
147
148    /// Add cost tracking middleware — records token usage estimates.
149    pub fn cost() -> MiddlewareComposite {
150        MiddlewareComposite::new(Arc::new(CostMiddleware {
151            tool_calls: std::sync::atomic::AtomicU64::new(0),
152        }))
153    }
154
155    /// Add rate-limiting middleware — spaces tool calls to at most `rps` per
156    /// second by delaying `before_tool` (concurrent calls queue rather than
157    /// burst).
158    pub fn rate_limit(rps: u32) -> MiddlewareComposite {
159        MiddlewareComposite::new(Arc::new(RateLimitMiddleware::new(rps)))
160    }
161
162    /// Add circuit breaker middleware — opens after consecutive failures.
163    pub fn circuit_breaker(threshold: u32) -> MiddlewareComposite {
164        MiddlewareComposite::new(Arc::new(CircuitBreakerMiddleware {
165            threshold,
166            consecutive_failures: std::sync::atomic::AtomicU32::new(0),
167        }))
168    }
169
170    /// Add tracing span middleware — creates spans for distributed tracing.
171    pub fn trace() -> MiddlewareComposite {
172        MiddlewareComposite::new(Arc::new(TraceMiddleware))
173    }
174
175    /// Add audit middleware — records all tool calls for review.
176    pub fn audit() -> MiddlewareComposite {
177        MiddlewareComposite::new(Arc::new(AuditMiddleware {
178            log: parking_lot::Mutex::new(Vec::new()),
179        }))
180    }
181
182    /// Scope middleware to specific agent names.
183    ///
184    /// Not yet enforced — agent-name routing requires dispatch-time filtering
185    /// the middleware chain doesn't expose, so this currently returns `inner`
186    /// unchanged. Hidden until real scoping lands to avoid implying behavior.
187    #[doc(hidden)]
188    pub fn scope(_names: &[&str], inner: MiddlewareComposite) -> MiddlewareComposite {
189        inner
190    }
191
192    /// Structured logging middleware — logs agent events as structured JSON.
193    pub fn structured_log() -> MiddlewareComposite {
194        MiddlewareComposite::new(Arc::new(StructuredLogMiddleware))
195    }
196
197    /// Dispatch logging middleware — logs dispatch/join events.
198    pub fn dispatch_log() -> MiddlewareComposite {
199        MiddlewareComposite::new(Arc::new(DispatchLogMiddleware))
200    }
201
202    /// Topology logging middleware — logs agent topology events.
203    pub fn topology_log() -> MiddlewareComposite {
204        MiddlewareComposite::new(Arc::new(TopologyLogMiddleware))
205    }
206
207    /// Add a tool input validator middleware.
208    pub fn validate(
209        f: impl Fn(&FunctionCall) -> Result<(), String> + Send + Sync + 'static,
210    ) -> MiddlewareComposite {
211        MiddlewareComposite::new(Arc::new(ValidateMiddleware {
212            validator: Arc::new(f),
213        }))
214    }
215
216    /// Fallback to an alternative model on error.
217    ///
218    /// Not yet enforced — swapping the model and retrying requires re-issuing
219    /// the LLM call, which the current `after_model`/`on_error` hooks can't do.
220    /// Hidden until real fallback lands to avoid implying behavior.
221    #[doc(hidden)]
222    pub fn fallback_model(model: &str) -> MiddlewareComposite {
223        MiddlewareComposite::new(Arc::new(FallbackModelMiddleware {
224            model: model.to_string(),
225        }))
226    }
227
228    /// Response caching middleware — caches model responses to avoid redundant calls.
229    pub fn cache() -> MiddlewareComposite {
230        MiddlewareComposite::new(Arc::new(CacheMiddleware {
231            cache: parking_lot::Mutex::new(std::collections::HashMap::new()),
232        }))
233    }
234
235    /// Deduplicate consecutive identical requests.
236    pub fn dedup() -> MiddlewareComposite {
237        MiddlewareComposite::new(Arc::new(DedupMiddleware {
238            last_request_hash: parking_lot::Mutex::new(None),
239        }))
240    }
241
242    /// Sample/pass-through a fraction of requests (0.0–1.0).
243    pub fn sample(rate: f64) -> MiddlewareComposite {
244        MiddlewareComposite::new(Arc::new(SampleMiddleware {
245            rate: rate.clamp(0.0, 1.0),
246        }))
247    }
248
249    /// Metrics collection middleware — tracks request counts, error counts, and latencies.
250    pub fn metrics() -> MiddlewareComposite {
251        MiddlewareComposite::new(Arc::new(MetricsMiddleware {
252            request_count: std::sync::atomic::AtomicU64::new(0),
253            error_count: std::sync::atomic::AtomicU64::new(0),
254        }))
255    }
256
257    /// Shortcut for a before-agent hook.
258    pub fn before_agent(
259        f: impl Fn(&gemini_adk_rs::context::InvocationContext) -> Result<(), String>
260        + Send
261        + Sync
262        + 'static,
263    ) -> MiddlewareComposite {
264        MiddlewareComposite::new(Arc::new(BeforeAgentMiddleware {
265            handler: Arc::new(f),
266        }))
267    }
268
269    /// Shortcut for an after-agent hook.
270    pub fn after_agent(
271        f: impl Fn(&gemini_adk_rs::context::InvocationContext) -> Result<(), String>
272        + Send
273        + Sync
274        + 'static,
275    ) -> MiddlewareComposite {
276        MiddlewareComposite::new(Arc::new(AfterAgentMiddleware {
277            handler: Arc::new(f),
278        }))
279    }
280
281    /// Shortcut for a before-model hook.
282    pub fn before_model(
283        f: impl Fn(&gemini_adk_rs::llm::LlmRequest) -> Result<(), String> + Send + Sync + 'static,
284    ) -> MiddlewareComposite {
285        MiddlewareComposite::new(Arc::new(BeforeModelMiddleware {
286            handler: Arc::new(f),
287        }))
288    }
289
290    /// Shortcut for an after-model hook.
291    pub fn after_model(
292        f: impl Fn(
293            &gemini_adk_rs::llm::LlmRequest,
294            &gemini_adk_rs::llm::LlmResponse,
295        ) -> Result<(), String>
296        + Send
297        + Sync
298        + 'static,
299    ) -> MiddlewareComposite {
300        MiddlewareComposite::new(Arc::new(AfterModelMiddleware {
301            handler: Arc::new(f),
302        }))
303    }
304
305    /// Loop iteration event hook — called on each iteration of a loop agent.
306    pub fn on_loop(f: impl Fn(u32) + Send + Sync + 'static) -> MiddlewareComposite {
307        MiddlewareComposite::new(Arc::new(OnLoopMiddleware {
308            handler: Arc::new(f),
309        }))
310    }
311
312    /// Timeout event hook — called when an agent times out.
313    pub fn on_timeout(f: impl Fn() + Send + Sync + 'static) -> MiddlewareComposite {
314        MiddlewareComposite::new(Arc::new(OnTimeoutMiddleware {
315            handler: Arc::new(f),
316        }))
317    }
318
319    /// Route decision event hook — called when a route agent selects a branch.
320    pub fn on_route(f: impl Fn(&str) + Send + Sync + 'static) -> MiddlewareComposite {
321        MiddlewareComposite::new(Arc::new(OnRouteMiddleware {
322            handler: Arc::new(f),
323        }))
324    }
325
326    /// Fallback event hook — called when a fallback agent activates.
327    pub fn on_fallback(f: impl Fn(&str) + Send + Sync + 'static) -> MiddlewareComposite {
328        MiddlewareComposite::new(Arc::new(OnFallbackMiddleware {
329            handler: Arc::new(f),
330        }))
331    }
332}
333
334/// Timeout middleware — stores the configured duration for runtime enforcement.
335#[allow(dead_code)]
336struct TimeoutMiddleware {
337    name: String,
338    duration: Duration,
339}
340
341#[async_trait::async_trait]
342impl Middleware for TimeoutMiddleware {
343    fn name(&self) -> &str {
344        &self.name
345    }
346
347    fn timeout(&self) -> Option<Duration> {
348        Some(self.duration)
349    }
350}
351
352// ── Tap Middleware ──────────────────────────────────────────────────────────
353
354struct TapMiddleware {
355    #[allow(clippy::type_complexity)]
356    handler: Arc<dyn Fn(&AgentEvent) + Send + Sync>,
357}
358
359#[async_trait]
360impl Middleware for TapMiddleware {
361    fn name(&self) -> &str {
362        "tap"
363    }
364
365    async fn on_event(&self, event: &AgentEvent) -> Result<(), AgentError> {
366        (self.handler)(event);
367        Ok(())
368    }
369}
370
371// ── BeforeTool Middleware ───────────────────────────────────────────────────
372
373struct BeforeToolMiddleware {
374    #[allow(clippy::type_complexity)]
375    handler: Arc<dyn Fn(&FunctionCall) -> Result<(), String> + Send + Sync>,
376}
377
378#[async_trait]
379impl Middleware for BeforeToolMiddleware {
380    fn name(&self) -> &str {
381        "before_tool"
382    }
383
384    async fn before_tool(&self, call: &FunctionCall) -> Result<(), AgentError> {
385        (self.handler)(call).map_err(AgentError::Other)
386    }
387}
388
389// ── AfterTool Middleware ────────────────────────────────────────────────────
390
391struct AfterToolMiddleware {
392    #[allow(clippy::type_complexity)]
393    handler: Arc<dyn Fn(&FunctionCall, &serde_json::Value) -> Result<(), String> + Send + Sync>,
394}
395
396#[async_trait]
397impl Middleware for AfterToolMiddleware {
398    fn name(&self) -> &str {
399        "after_tool"
400    }
401
402    async fn after_tool(
403        &self,
404        call: &FunctionCall,
405        result: &serde_json::Value,
406    ) -> Result<(), AgentError> {
407        (self.handler)(call, result).map_err(AgentError::Other)
408    }
409}
410
411// ── OnError Middleware ──────────────────────────────────────────────────────
412
413struct OnErrorMiddleware {
414    #[allow(clippy::type_complexity)]
415    handler: Arc<dyn Fn(&AgentError) -> Result<(), String> + Send + Sync>,
416}
417
418#[async_trait]
419impl Middleware for OnErrorMiddleware {
420    fn name(&self) -> &str {
421        "on_error"
422    }
423
424    async fn on_error(&self, err: &AgentError) -> Result<(), AgentError> {
425        (self.handler)(err).map_err(AgentError::Other)
426    }
427}
428
429// ── Cost Middleware ────────────────────────────────────────────────────────
430
431/// Tracks the number of tool calls as a proxy for cost.
432pub struct CostMiddleware {
433    tool_calls: std::sync::atomic::AtomicU64,
434}
435
436impl CostMiddleware {
437    /// Returns the total number of tool calls recorded.
438    pub fn tool_call_count(&self) -> u64 {
439        self.tool_calls.load(std::sync::atomic::Ordering::SeqCst)
440    }
441}
442
443#[async_trait]
444impl Middleware for CostMiddleware {
445    fn name(&self) -> &str {
446        "cost"
447    }
448
449    async fn after_tool(
450        &self,
451        _call: &FunctionCall,
452        _result: &serde_json::Value,
453    ) -> Result<(), AgentError> {
454        self.tool_calls
455            .fetch_add(1, std::sync::atomic::Ordering::SeqCst);
456        Ok(())
457    }
458}
459
460// ── RateLimit Middleware ───────────────────────────────────────────────────
461
462#[allow(dead_code)]
463struct RateLimitMiddleware {
464    /// Minimum spacing between successive tool calls.
465    min_interval: Duration,
466    /// Reserved start time of the most recent call (advances as calls queue).
467    last: parking_lot::Mutex<Option<std::time::Instant>>,
468}
469
470impl RateLimitMiddleware {
471    fn new(rps: u32) -> Self {
472        let rps = rps.max(1);
473        Self {
474            min_interval: Duration::from_secs_f64(1.0 / rps as f64),
475            last: parking_lot::Mutex::new(None),
476        }
477    }
478}
479
480#[async_trait]
481impl Middleware for RateLimitMiddleware {
482    fn name(&self) -> &str {
483        "rate_limit"
484    }
485
486    async fn before_tool(&self, _call: &FunctionCall) -> Result<(), AgentError> {
487        // Compute the wait needed to honor min_interval, reserving this call's
488        // slot so concurrent callers queue instead of bursting. The lock is
489        // released before the await (never held across it).
490        let wait = {
491            let mut last = self.last.lock();
492            let now = std::time::Instant::now();
493            let scheduled = match *last {
494                Some(prev) if prev + self.min_interval > now => prev + self.min_interval,
495                _ => now,
496            };
497            *last = Some(scheduled);
498            scheduled.saturating_duration_since(now)
499        };
500        if !wait.is_zero() {
501            tokio::time::sleep(wait).await;
502        }
503        Ok(())
504    }
505}
506
507// ── CircuitBreaker Middleware ──────────────────────────────────────────────
508
509struct CircuitBreakerMiddleware {
510    threshold: u32,
511    consecutive_failures: std::sync::atomic::AtomicU32,
512}
513
514#[async_trait]
515impl Middleware for CircuitBreakerMiddleware {
516    fn name(&self) -> &str {
517        "circuit_breaker"
518    }
519
520    async fn before_tool(&self, _call: &FunctionCall) -> Result<(), AgentError> {
521        let failures = self
522            .consecutive_failures
523            .load(std::sync::atomic::Ordering::SeqCst);
524        if failures >= self.threshold {
525            return Err(AgentError::Other(format!(
526                "Circuit breaker open: {} consecutive failures (threshold: {})",
527                failures, self.threshold
528            )));
529        }
530        Ok(())
531    }
532
533    async fn after_tool(
534        &self,
535        _call: &FunctionCall,
536        _result: &serde_json::Value,
537    ) -> Result<(), AgentError> {
538        self.consecutive_failures
539            .store(0, std::sync::atomic::Ordering::SeqCst);
540        Ok(())
541    }
542
543    async fn on_tool_error(
544        &self,
545        _call: &FunctionCall,
546        _err: &ToolError,
547    ) -> Result<(), AgentError> {
548        self.consecutive_failures
549            .fetch_add(1, std::sync::atomic::Ordering::SeqCst);
550        Ok(())
551    }
552}
553
554// ── Trace Middleware ──────────────────────────────────────────────────────
555
556/// Middleware that creates tracing spans for agent and tool lifecycle events.
557/// With an OTel exporter feature enabled, these spans are picked up by
558/// `tracing-opentelemetry` and exported as OTel spans.
559struct TraceMiddleware;
560
561#[async_trait]
562impl Middleware for TraceMiddleware {
563    fn name(&self) -> &str {
564        "trace"
565    }
566
567    async fn before_agent(
568        &self,
569        ctx: &gemini_adk_rs::context::InvocationContext,
570    ) -> Result<(), AgentError> {
571        let sid = ctx.session_id.as_deref().unwrap_or("unknown");
572        gemini_adk_rs::telemetry::logging::log_agent_started(sid, 0);
573        Ok(())
574    }
575
576    async fn before_tool(&self, call: &FunctionCall) -> Result<(), AgentError> {
577        gemini_adk_rs::telemetry::logging::log_tool_dispatch("fluent", &call.name, "function");
578        Ok(())
579    }
580
581    async fn after_tool(
582        &self,
583        call: &FunctionCall,
584        _result: &serde_json::Value,
585    ) -> Result<(), AgentError> {
586        gemini_adk_rs::telemetry::logging::log_tool_result("fluent", &call.name, true, 0.0);
587        Ok(())
588    }
589
590    async fn on_tool_error(&self, call: &FunctionCall, _err: &ToolError) -> Result<(), AgentError> {
591        gemini_adk_rs::telemetry::logging::log_tool_result("fluent", &call.name, false, 0.0);
592        Ok(())
593    }
594
595    async fn on_error(&self, err: &AgentError) -> Result<(), AgentError> {
596        gemini_adk_rs::telemetry::logging::log_agent_error("fluent", &err.to_string());
597        Ok(())
598    }
599}
600
601// ── Audit Middleware ─────────────────────────────────────────────────────
602
603/// Records all tool calls for audit review.
604pub struct AuditMiddleware {
605    log: parking_lot::Mutex<Vec<AuditEntry>>,
606}
607
608/// An audit log entry.
609#[derive(Debug, Clone)]
610pub struct AuditEntry {
611    /// Tool name.
612    pub tool_name: String,
613    /// Tool arguments.
614    pub args: serde_json::Value,
615    /// Whether the call succeeded.
616    pub success: Option<bool>,
617}
618
619impl AuditMiddleware {
620    /// Returns a snapshot of the audit log.
621    pub fn entries(&self) -> Vec<AuditEntry> {
622        self.log.lock().clone()
623    }
624}
625
626#[async_trait]
627impl Middleware for AuditMiddleware {
628    fn name(&self) -> &str {
629        "audit"
630    }
631
632    async fn before_tool(&self, call: &FunctionCall) -> Result<(), AgentError> {
633        let mut log = self.log.lock();
634        if log.len() >= 10_000 {
635            log.drain(..1_000);
636        }
637        log.push(AuditEntry {
638            tool_name: call.name.clone(),
639            args: call.args.clone(),
640            success: None,
641        });
642        Ok(())
643    }
644
645    async fn after_tool(
646        &self,
647        call: &FunctionCall,
648        _result: &serde_json::Value,
649    ) -> Result<(), AgentError> {
650        let mut log = self.log.lock();
651        if let Some(entry) = log.iter_mut().rev().find(|e| e.tool_name == call.name) {
652            entry.success = Some(true);
653        }
654        Ok(())
655    }
656
657    async fn on_tool_error(&self, call: &FunctionCall, _err: &ToolError) -> Result<(), AgentError> {
658        let mut log = self.log.lock();
659        if let Some(entry) = log.iter_mut().rev().find(|e| e.tool_name == call.name) {
660            entry.success = Some(false);
661        }
662        Ok(())
663    }
664}
665
666// ── Validate Middleware ──────────────────────────────────────────────────
667
668struct ValidateMiddleware {
669    #[allow(clippy::type_complexity)]
670    validator: Arc<dyn Fn(&FunctionCall) -> Result<(), String> + Send + Sync>,
671}
672
673#[async_trait]
674impl Middleware for ValidateMiddleware {
675    fn name(&self) -> &str {
676        "validate"
677    }
678
679    async fn before_tool(&self, call: &FunctionCall) -> Result<(), AgentError> {
680        (self.validator)(call).map_err(|e| AgentError::Tool(ToolError::InvalidArgs(e)))
681    }
682}
683
684// ── FallbackModel Middleware ──────────────────────────────────────────
685
686/// Middleware that falls back to an alternative model on error.
687#[allow(dead_code)]
688struct FallbackModelMiddleware {
689    model: String,
690}
691
692#[async_trait]
693impl Middleware for FallbackModelMiddleware {
694    fn name(&self) -> &str {
695        "fallback_model"
696    }
697
698    async fn on_error(&self, _err: &AgentError) -> Result<(), AgentError> {
699        // Runtime inspects the `model` field and retries with the fallback model.
700        Ok(())
701    }
702}
703
704// ── Cache Middleware ──────────────────────────────────────────────────
705
706/// Caches model responses keyed by request hash to avoid redundant LLM calls.
707pub struct CacheMiddleware {
708    cache: parking_lot::Mutex<std::collections::HashMap<u64, gemini_adk_rs::llm::LlmResponse>>,
709}
710
711impl CacheMiddleware {
712    /// Returns the number of cached entries.
713    pub fn len(&self) -> usize {
714        self.cache.lock().len()
715    }
716
717    /// Whether the cache is empty.
718    pub fn is_empty(&self) -> bool {
719        self.cache.lock().is_empty()
720    }
721
722    /// Clear all cached entries.
723    pub fn clear(&self) {
724        self.cache.lock().clear();
725    }
726}
727
728#[async_trait]
729impl Middleware for CacheMiddleware {
730    fn name(&self) -> &str {
731        "cache"
732    }
733
734    async fn before_model(
735        &self,
736        request: &gemini_adk_rs::llm::LlmRequest,
737    ) -> Result<Option<gemini_adk_rs::llm::LlmResponse>, AgentError> {
738        use std::hash::{Hash, Hasher};
739        let mut hasher = std::collections::hash_map::DefaultHasher::new();
740        format!("{request:?}").hash(&mut hasher);
741        let key = hasher.finish();
742        let cache = self.cache.lock();
743        Ok(cache.get(&key).cloned())
744    }
745
746    async fn after_model(
747        &self,
748        request: &gemini_adk_rs::llm::LlmRequest,
749        response: &gemini_adk_rs::llm::LlmResponse,
750    ) -> Result<Option<gemini_adk_rs::llm::LlmResponse>, AgentError> {
751        use std::hash::{Hash, Hasher};
752        let mut hasher = std::collections::hash_map::DefaultHasher::new();
753        format!("{request:?}").hash(&mut hasher);
754        let key = hasher.finish();
755        self.cache.lock().insert(key, response.clone());
756        Ok(None) // don't replace the response
757    }
758}
759
760// ── Dedup Middleware ─────────────────────────────────────────────────
761
762/// Deduplicates consecutive identical requests by hashing.
763#[allow(dead_code)]
764struct DedupMiddleware {
765    last_request_hash: parking_lot::Mutex<Option<u64>>,
766}
767
768#[async_trait]
769impl Middleware for DedupMiddleware {
770    fn name(&self) -> &str {
771        "dedup"
772    }
773
774    async fn before_model(
775        &self,
776        request: &gemini_adk_rs::llm::LlmRequest,
777    ) -> Result<Option<gemini_adk_rs::llm::LlmResponse>, AgentError> {
778        use std::hash::{Hash, Hasher};
779        let mut hasher = std::collections::hash_map::DefaultHasher::new();
780        format!("{request:?}").hash(&mut hasher);
781        let hash = hasher.finish();
782        let mut last = self.last_request_hash.lock();
783        if *last == Some(hash) {
784            // Duplicate consecutive request — signal skip by returning empty response.
785            return Err(AgentError::Other(
786                "Duplicate consecutive request".to_string(),
787            ));
788        }
789        *last = Some(hash);
790        Ok(None)
791    }
792}
793
794// ── Sample Middleware ────────────────────────────────────────────────
795
796/// Passes through only a fraction of requests, dropping the rest.
797#[allow(dead_code)]
798struct SampleMiddleware {
799    rate: f64,
800}
801
802#[async_trait]
803impl Middleware for SampleMiddleware {
804    fn name(&self) -> &str {
805        "sample"
806    }
807
808    async fn before_model(
809        &self,
810        _request: &gemini_adk_rs::llm::LlmRequest,
811    ) -> Result<Option<gemini_adk_rs::llm::LlmResponse>, AgentError> {
812        use std::hash::{Hash, Hasher};
813        // Use a fast pseudo-random check based on time.
814        let mut hasher = std::collections::hash_map::DefaultHasher::new();
815        std::time::Instant::now().hash(&mut hasher);
816        let hash = hasher.finish();
817        let normalized = (hash as f64) / (u64::MAX as f64);
818        if normalized > self.rate {
819            return Err(AgentError::Other("Sampled out".to_string()));
820        }
821        Ok(None)
822    }
823}
824
825// ── Metrics Middleware ──────────────────────────────────────────────
826
827/// Collects request and error counts.
828pub struct MetricsMiddleware {
829    request_count: std::sync::atomic::AtomicU64,
830    error_count: std::sync::atomic::AtomicU64,
831}
832
833impl MetricsMiddleware {
834    /// Returns the total number of requests observed.
835    pub fn request_count(&self) -> u64 {
836        self.request_count.load(std::sync::atomic::Ordering::SeqCst)
837    }
838
839    /// Returns the total number of errors observed.
840    pub fn error_count(&self) -> u64 {
841        self.error_count.load(std::sync::atomic::Ordering::SeqCst)
842    }
843}
844
845#[async_trait]
846impl Middleware for MetricsMiddleware {
847    fn name(&self) -> &str {
848        "metrics"
849    }
850
851    async fn before_agent(
852        &self,
853        _ctx: &gemini_adk_rs::context::InvocationContext,
854    ) -> Result<(), AgentError> {
855        self.request_count
856            .fetch_add(1, std::sync::atomic::Ordering::SeqCst);
857        Ok(())
858    }
859
860    async fn on_error(&self, _err: &AgentError) -> Result<(), AgentError> {
861        self.error_count
862            .fetch_add(1, std::sync::atomic::Ordering::SeqCst);
863        Ok(())
864    }
865}
866
867// ── BeforeAgent Middleware ───────────────────────────────────────────
868
869struct BeforeAgentMiddleware {
870    #[allow(clippy::type_complexity)]
871    handler:
872        Arc<dyn Fn(&gemini_adk_rs::context::InvocationContext) -> Result<(), String> + Send + Sync>,
873}
874
875#[async_trait]
876impl Middleware for BeforeAgentMiddleware {
877    fn name(&self) -> &str {
878        "before_agent"
879    }
880
881    async fn before_agent(
882        &self,
883        ctx: &gemini_adk_rs::context::InvocationContext,
884    ) -> Result<(), AgentError> {
885        (self.handler)(ctx).map_err(AgentError::Other)
886    }
887}
888
889// ── AfterAgent Middleware ───────────────────────────────────────────
890
891struct AfterAgentMiddleware {
892    #[allow(clippy::type_complexity)]
893    handler:
894        Arc<dyn Fn(&gemini_adk_rs::context::InvocationContext) -> Result<(), String> + Send + Sync>,
895}
896
897#[async_trait]
898impl Middleware for AfterAgentMiddleware {
899    fn name(&self) -> &str {
900        "after_agent"
901    }
902
903    async fn after_agent(
904        &self,
905        ctx: &gemini_adk_rs::context::InvocationContext,
906    ) -> Result<(), AgentError> {
907        (self.handler)(ctx).map_err(AgentError::Other)
908    }
909}
910
911// ── BeforeModel Middleware ──────────────────────────────────────────
912
913struct BeforeModelMiddleware {
914    #[allow(clippy::type_complexity)]
915    handler: Arc<dyn Fn(&gemini_adk_rs::llm::LlmRequest) -> Result<(), String> + Send + Sync>,
916}
917
918#[async_trait]
919impl Middleware for BeforeModelMiddleware {
920    fn name(&self) -> &str {
921        "before_model"
922    }
923
924    async fn before_model(
925        &self,
926        request: &gemini_adk_rs::llm::LlmRequest,
927    ) -> Result<Option<gemini_adk_rs::llm::LlmResponse>, AgentError> {
928        (self.handler)(request).map_err(AgentError::Other)?;
929        Ok(None)
930    }
931}
932
933// ── AfterModel Middleware ──────────────────────────────────────────
934
935struct AfterModelMiddleware {
936    #[allow(clippy::type_complexity)]
937    handler: Arc<
938        dyn Fn(
939                &gemini_adk_rs::llm::LlmRequest,
940                &gemini_adk_rs::llm::LlmResponse,
941            ) -> Result<(), String>
942            + Send
943            + Sync,
944    >,
945}
946
947#[async_trait]
948impl Middleware for AfterModelMiddleware {
949    fn name(&self) -> &str {
950        "after_model"
951    }
952
953    async fn after_model(
954        &self,
955        request: &gemini_adk_rs::llm::LlmRequest,
956        response: &gemini_adk_rs::llm::LlmResponse,
957    ) -> Result<Option<gemini_adk_rs::llm::LlmResponse>, AgentError> {
958        (self.handler)(request, response).map_err(AgentError::Other)?;
959        Ok(None)
960    }
961}
962
963// ── OnLoop Middleware ───────────────────────────────────────────────
964
965struct OnLoopMiddleware {
966    handler: Arc<dyn Fn(u32) + Send + Sync>,
967}
968
969#[async_trait]
970impl Middleware for OnLoopMiddleware {
971    fn name(&self) -> &str {
972        "on_loop"
973    }
974
975    async fn on_event(&self, event: &AgentEvent) -> Result<(), AgentError> {
976        if let AgentEvent::LoopIteration { iteration } = event {
977            (self.handler)(*iteration);
978        }
979        Ok(())
980    }
981}
982
983// ── OnTimeout Middleware ────────────────────────────────────────────
984
985struct OnTimeoutMiddleware {
986    handler: Arc<dyn Fn() + Send + Sync>,
987}
988
989#[async_trait]
990impl Middleware for OnTimeoutMiddleware {
991    fn name(&self) -> &str {
992        "on_timeout"
993    }
994
995    async fn on_event(&self, event: &AgentEvent) -> Result<(), AgentError> {
996        if let AgentEvent::Timeout = event {
997            (self.handler)();
998        }
999        Ok(())
1000    }
1001}
1002
1003// ── OnRoute Middleware ──────────────────────────────────────────────
1004
1005struct OnRouteMiddleware {
1006    handler: Arc<dyn Fn(&str) + Send + Sync>,
1007}
1008
1009#[async_trait]
1010impl Middleware for OnRouteMiddleware {
1011    fn name(&self) -> &str {
1012        "on_route"
1013    }
1014
1015    async fn on_event(&self, event: &AgentEvent) -> Result<(), AgentError> {
1016        if let AgentEvent::RouteSelected { agent_name } = event {
1017            (self.handler)(agent_name);
1018        }
1019        Ok(())
1020    }
1021}
1022
1023// ── OnFallback Middleware ───────────────────────────────────────────
1024
1025struct OnFallbackMiddleware {
1026    handler: Arc<dyn Fn(&str) + Send + Sync>,
1027}
1028
1029#[async_trait]
1030impl Middleware for OnFallbackMiddleware {
1031    fn name(&self) -> &str {
1032        "on_fallback"
1033    }
1034
1035    async fn on_event(&self, event: &AgentEvent) -> Result<(), AgentError> {
1036        if let AgentEvent::FallbackActivated { agent_name } = event {
1037            (self.handler)(agent_name);
1038        }
1039        Ok(())
1040    }
1041}
1042
1043// ── Structured Log Middleware ────────────────────────────────────────
1044
1045struct StructuredLogMiddleware;
1046
1047#[async_trait]
1048impl Middleware for StructuredLogMiddleware {
1049    fn name(&self) -> &str {
1050        "structured_log"
1051    }
1052
1053    async fn on_event(&self, event: &AgentEvent) -> Result<(), AgentError> {
1054        // Log events as structured format (uses tracing in production).
1055        let _ = event;
1056        Ok(())
1057    }
1058}
1059
1060// ── Dispatch Log Middleware ──────────────────────────────────────────
1061
1062struct DispatchLogMiddleware;
1063
1064#[async_trait]
1065impl Middleware for DispatchLogMiddleware {
1066    fn name(&self) -> &str {
1067        "dispatch_log"
1068    }
1069}
1070
1071// ── Topology Log Middleware ──────────────────────────────────────────
1072
1073struct TopologyLogMiddleware;
1074
1075#[async_trait]
1076impl Middleware for TopologyLogMiddleware {
1077    fn name(&self) -> &str {
1078        "topology_log"
1079    }
1080}
1081
1082#[cfg(test)]
1083mod tests {
1084    use super::*;
1085
1086    #[test]
1087    fn log_creates_composite() {
1088        let m = M::log();
1089        assert_eq!(m.len(), 1);
1090    }
1091
1092    #[test]
1093    fn latency_creates_composite() {
1094        let m = M::latency();
1095        assert_eq!(m.len(), 1);
1096    }
1097
1098    #[test]
1099    fn timeout_creates_composite() {
1100        let m = M::timeout(Duration::from_secs(30));
1101        assert_eq!(m.len(), 1);
1102    }
1103
1104    #[test]
1105    fn compose_with_bitor() {
1106        let m = M::log() | M::latency() | M::timeout(Duration::from_secs(5));
1107        assert_eq!(m.len(), 3);
1108    }
1109
1110    #[test]
1111    fn retry_creates_composite() {
1112        let m = M::retry(3);
1113        assert_eq!(m.len(), 1);
1114    }
1115
1116    #[test]
1117    fn tap_creates_composite() {
1118        let m = M::tap(|_event| {});
1119        assert_eq!(m.len(), 1);
1120    }
1121
1122    #[test]
1123    fn before_tool_creates_composite() {
1124        let m = M::before_tool(|_call| Ok(()));
1125        assert_eq!(m.len(), 1);
1126    }
1127
1128    #[test]
1129    fn cost_creates_composite() {
1130        let m = M::cost();
1131        assert_eq!(m.len(), 1);
1132    }
1133
1134    #[test]
1135    fn rate_limit_creates_composite() {
1136        let m = M::rate_limit(10);
1137        assert_eq!(m.len(), 1);
1138    }
1139
1140    #[test]
1141    fn circuit_breaker_creates_composite() {
1142        let m = M::circuit_breaker(5);
1143        assert_eq!(m.len(), 1);
1144    }
1145
1146    #[test]
1147    fn trace_creates_composite() {
1148        let m = M::trace();
1149        assert_eq!(m.len(), 1);
1150    }
1151
1152    #[test]
1153    fn audit_creates_composite() {
1154        let m = M::audit();
1155        assert_eq!(m.len(), 1);
1156    }
1157
1158    #[test]
1159    fn validate_creates_composite() {
1160        let m = M::validate(|_call| Ok(()));
1161        assert_eq!(m.len(), 1);
1162    }
1163
1164    #[test]
1165    fn fallback_model_creates_composite() {
1166        let m = M::fallback_model("gemini-1.5-flash");
1167        assert_eq!(m.len(), 1);
1168    }
1169
1170    #[test]
1171    fn cache_creates_composite() {
1172        let m = M::cache();
1173        assert_eq!(m.len(), 1);
1174    }
1175
1176    #[test]
1177    fn dedup_creates_composite() {
1178        let m = M::dedup();
1179        assert_eq!(m.len(), 1);
1180    }
1181
1182    #[test]
1183    fn sample_creates_composite() {
1184        let m = M::sample(0.5);
1185        assert_eq!(m.len(), 1);
1186    }
1187
1188    #[test]
1189    fn sample_clamps_rate() {
1190        let m = M::sample(2.0);
1191        assert_eq!(m.len(), 1);
1192        let m = M::sample(-1.0);
1193        assert_eq!(m.len(), 1);
1194    }
1195
1196    #[test]
1197    fn metrics_creates_composite() {
1198        let m = M::metrics();
1199        assert_eq!(m.len(), 1);
1200    }
1201
1202    #[test]
1203    fn before_agent_creates_composite() {
1204        let m = M::before_agent(|_ctx| Ok(()));
1205        assert_eq!(m.len(), 1);
1206    }
1207
1208    #[test]
1209    fn after_agent_creates_composite() {
1210        let m = M::after_agent(|_ctx| Ok(()));
1211        assert_eq!(m.len(), 1);
1212    }
1213
1214    #[test]
1215    fn before_model_creates_composite() {
1216        let m = M::before_model(|_req| Ok(()));
1217        assert_eq!(m.len(), 1);
1218    }
1219
1220    #[test]
1221    fn after_model_creates_composite() {
1222        let m = M::after_model(|_req, _resp| Ok(()));
1223        assert_eq!(m.len(), 1);
1224    }
1225
1226    #[test]
1227    fn on_loop_creates_composite() {
1228        let m = M::on_loop(|_iteration| {});
1229        assert_eq!(m.len(), 1);
1230    }
1231
1232    #[test]
1233    fn on_timeout_creates_composite() {
1234        let m = M::on_timeout(|| {});
1235        assert_eq!(m.len(), 1);
1236    }
1237
1238    #[test]
1239    fn on_route_creates_composite() {
1240        let m = M::on_route(|_name| {});
1241        assert_eq!(m.len(), 1);
1242    }
1243
1244    #[test]
1245    fn on_fallback_creates_composite() {
1246        let m = M::on_fallback(|_name| {});
1247        assert_eq!(m.len(), 1);
1248    }
1249
1250    #[test]
1251    fn compose_all_middleware() {
1252        let m = M::log()
1253            | M::latency()
1254            | M::timeout(Duration::from_secs(30))
1255            | M::retry(3)
1256            | M::cost()
1257            | M::rate_limit(10)
1258            | M::circuit_breaker(5)
1259            | M::trace()
1260            | M::audit()
1261            | M::validate(|_| Ok(()))
1262            | M::fallback_model("gemini-1.5-flash")
1263            | M::cache()
1264            | M::dedup()
1265            | M::sample(0.5)
1266            | M::metrics()
1267            | M::before_agent(|_| Ok(()))
1268            | M::after_agent(|_| Ok(()))
1269            | M::before_model(|_| Ok(()))
1270            | M::after_model(|_, _| Ok(()))
1271            | M::on_loop(|_| {})
1272            | M::on_timeout(|| {})
1273            | M::on_route(|_| {})
1274            | M::on_fallback(|_| {});
1275        assert_eq!(m.len(), 23);
1276    }
1277}