1use 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#[derive(Clone)]
42#[non_exhaustive]
43pub struct MiddlewareComposite {
44 pub layers: Vec<Arc<dyn Middleware>>,
46}
47
48impl MiddlewareComposite {
49 pub fn new(layer: Arc<dyn Middleware>) -> Self {
51 Self {
52 layers: vec![layer],
53 }
54 }
55
56 pub fn len(&self) -> usize {
58 self.layers.len()
59 }
60
61 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
73impl 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
83pub struct M;
85
86impl M {
87 pub fn log() -> MiddlewareComposite {
89 MiddlewareComposite::new(Arc::new(LogMiddleware::new()))
90 }
91
92 pub fn latency() -> MiddlewareComposite {
94 MiddlewareComposite::new(Arc::new(LatencyMiddleware::new()))
95 }
96
97 pub fn timeout(duration: Duration) -> MiddlewareComposite {
101 MiddlewareComposite::new(Arc::new(TimeoutMiddleware {
102 name: "timeout".to_string(),
103 duration,
104 }))
105 }
106
107 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 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 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 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 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 pub fn cost() -> MiddlewareComposite {
150 MiddlewareComposite::new(Arc::new(CostMiddleware {
151 tool_calls: std::sync::atomic::AtomicU64::new(0),
152 }))
153 }
154
155 pub fn rate_limit(rps: u32) -> MiddlewareComposite {
159 MiddlewareComposite::new(Arc::new(RateLimitMiddleware::new(rps)))
160 }
161
162 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 pub fn trace() -> MiddlewareComposite {
172 MiddlewareComposite::new(Arc::new(TraceMiddleware))
173 }
174
175 pub fn audit() -> MiddlewareComposite {
177 MiddlewareComposite::new(Arc::new(AuditMiddleware {
178 log: parking_lot::Mutex::new(Vec::new()),
179 }))
180 }
181
182 #[doc(hidden)]
188 pub fn scope(_names: &[&str], inner: MiddlewareComposite) -> MiddlewareComposite {
189 inner
190 }
191
192 pub fn structured_log() -> MiddlewareComposite {
194 MiddlewareComposite::new(Arc::new(StructuredLogMiddleware))
195 }
196
197 pub fn dispatch_log() -> MiddlewareComposite {
199 MiddlewareComposite::new(Arc::new(DispatchLogMiddleware))
200 }
201
202 pub fn topology_log() -> MiddlewareComposite {
204 MiddlewareComposite::new(Arc::new(TopologyLogMiddleware))
205 }
206
207 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 #[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 pub fn cache() -> MiddlewareComposite {
230 MiddlewareComposite::new(Arc::new(CacheMiddleware {
231 cache: parking_lot::Mutex::new(std::collections::HashMap::new()),
232 }))
233 }
234
235 pub fn dedup() -> MiddlewareComposite {
237 MiddlewareComposite::new(Arc::new(DedupMiddleware {
238 last_request_hash: parking_lot::Mutex::new(None),
239 }))
240 }
241
242 pub fn sample(rate: f64) -> MiddlewareComposite {
244 MiddlewareComposite::new(Arc::new(SampleMiddleware {
245 rate: rate.clamp(0.0, 1.0),
246 }))
247 }
248
249 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 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 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 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 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 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 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 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 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#[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
352struct 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
371struct 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
389struct 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
411struct 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
429pub struct CostMiddleware {
433 tool_calls: std::sync::atomic::AtomicU64,
434}
435
436impl CostMiddleware {
437 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#[allow(dead_code)]
463struct RateLimitMiddleware {
464 min_interval: Duration,
466 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 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
507struct 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
554struct 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
601pub struct AuditMiddleware {
605 log: parking_lot::Mutex<Vec<AuditEntry>>,
606}
607
608#[derive(Debug, Clone)]
610pub struct AuditEntry {
611 pub tool_name: String,
613 pub args: serde_json::Value,
615 pub success: Option<bool>,
617}
618
619impl AuditMiddleware {
620 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
666struct 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#[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 Ok(())
701 }
702}
703
704pub struct CacheMiddleware {
708 cache: parking_lot::Mutex<std::collections::HashMap<u64, gemini_adk_rs::llm::LlmResponse>>,
709}
710
711impl CacheMiddleware {
712 pub fn len(&self) -> usize {
714 self.cache.lock().len()
715 }
716
717 pub fn is_empty(&self) -> bool {
719 self.cache.lock().is_empty()
720 }
721
722 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) }
758}
759
760#[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 return Err(AgentError::Other(
786 "Duplicate consecutive request".to_string(),
787 ));
788 }
789 *last = Some(hash);
790 Ok(None)
791 }
792}
793
794#[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 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
825pub struct MetricsMiddleware {
829 request_count: std::sync::atomic::AtomicU64,
830 error_count: std::sync::atomic::AtomicU64,
831}
832
833impl MetricsMiddleware {
834 pub fn request_count(&self) -> u64 {
836 self.request_count.load(std::sync::atomic::Ordering::SeqCst)
837 }
838
839 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
867struct 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
889struct 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
911struct 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
933struct 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
963struct 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
983struct 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
1003struct 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
1023struct 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
1043struct 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 let _ = event;
1056 Ok(())
1057 }
1058}
1059
1060struct DispatchLogMiddleware;
1063
1064#[async_trait]
1065impl Middleware for DispatchLogMiddleware {
1066 fn name(&self) -> &str {
1067 "dispatch_log"
1068 }
1069}
1070
1071struct 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}