Skip to main content

tokio_prompt_orchestrator/routing/
router.rs

1//! Model routing logic.
2//!
3//! The [`ModelRouter`] combines a [`ComplexityScorer`]
4//! with a [`RoutingConfig`] and a
5//! [`CostTracker`] to decide which worker backend should
6//! serve each prompt and to adaptively adjust routing thresholds based on
7//! observed outcomes.
8
9use crate::config::WorkerKind;
10use std::sync::atomic::{AtomicU64, Ordering};
11use std::sync::RwLock;
12
13use super::config::RoutingConfig;
14use super::cost_tracker::CostTracker;
15use super::scorer::ComplexityScorer;
16
17/// The routing decision for a single prompt.
18///
19/// # Panics
20///
21/// This type never panics.
22#[derive(Debug, Clone, PartialEq)]
23pub enum RoutingDecision {
24    /// Route to the local model (complexity below `local_threshold`).
25    Local {
26        /// The complexity score that drove this decision.
27        score: f64,
28        /// The worker kind selected.
29        worker: WorkerKind,
30    },
31    /// Route to the local model first, fall back to cloud on failure
32    /// (complexity between `local_threshold` and `cloud_threshold`).
33    LocalWithFallback {
34        /// The complexity score that drove this decision.
35        score: f64,
36        /// Primary worker (local).
37        primary: WorkerKind,
38        /// Fallback worker (cloud).
39        fallback: WorkerKind,
40    },
41    /// Route directly to the cloud model (complexity above `cloud_threshold`).
42    Cloud {
43        /// The complexity score that drove this decision.
44        score: f64,
45        /// The worker kind selected.
46        worker: WorkerKind,
47    },
48}
49
50impl RoutingDecision {
51    /// Return the complexity score associated with this decision.
52    ///
53    /// # Panics
54    ///
55    /// This function never panics.
56    pub fn score(&self) -> f64 {
57        match self {
58            Self::Local { score, .. }
59            | Self::LocalWithFallback { score, .. }
60            | Self::Cloud { score, .. } => *score,
61        }
62    }
63
64    /// Return `true` if the decision routes to the local worker exclusively.
65    ///
66    /// # Panics
67    ///
68    /// This function never panics.
69    pub fn is_local(&self) -> bool {
70        matches!(self, Self::Local { .. })
71    }
72
73    /// Return `true` if the decision routes to the cloud worker exclusively.
74    ///
75    /// # Panics
76    ///
77    /// This function never panics.
78    pub fn is_cloud(&self) -> bool {
79        matches!(self, Self::Cloud { .. })
80    }
81
82    /// Return `true` if the decision uses local-with-fallback routing.
83    ///
84    /// # Panics
85    ///
86    /// This function never panics.
87    pub fn is_fallback(&self) -> bool {
88        matches!(self, Self::LocalWithFallback { .. })
89    }
90}
91
92/// Intelligent model router.
93///
94/// Combines prompt complexity scoring, threshold-based routing, adaptive
95/// threshold adjustment, and cost tracking into a single entry point.
96///
97/// Thread-safe: all mutable state uses atomics or interior `RwLock`.
98///
99/// # Panics
100///
101/// This type and its methods never panic.
102pub struct ModelRouter {
103    scorer: ComplexityScorer,
104    config: RoutingConfig,
105    cost_tracker: CostTracker,
106
107    /// Effective cloud threshold — may diverge from config when adaptive
108    /// adjustment is active.
109    effective_cloud_threshold: RwLock<f64>,
110
111    // Adaptive tracking: fallback-zone outcomes.
112    fallback_successes: AtomicU64,
113    fallback_failures: AtomicU64,
114}
115
116impl std::fmt::Debug for ModelRouter {
117    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
118        f.debug_struct("ModelRouter")
119            .field("config", &self.config)
120            .field("effective_cloud_threshold", &self.effective_cloud_threshold)
121            .finish()
122    }
123}
124
125impl ModelRouter {
126    /// Create a new router with the given configuration.
127    ///
128    /// # Arguments
129    ///
130    /// * `config` — Routing thresholds, cost rates, and adaptive settings.
131    ///
132    /// # Returns
133    ///
134    /// A new [`ModelRouter`] ready to route prompts.
135    ///
136    /// # Panics
137    ///
138    /// This function never panics.
139    pub fn new(config: RoutingConfig) -> Self {
140        let cost_tracker = CostTracker::new(
141            config.local_cost_per_1k_tokens,
142            config.cloud_cost_per_1k_tokens,
143        );
144        let effective_cloud_threshold = RwLock::new(config.cloud_threshold);
145
146        Self {
147            scorer: ComplexityScorer::new(),
148            config,
149            cost_tracker,
150            effective_cloud_threshold,
151            fallback_successes: AtomicU64::new(0),
152            fallback_failures: AtomicU64::new(0),
153        }
154    }
155
156    /// Route a prompt to the appropriate worker.
157    ///
158    /// # Arguments
159    ///
160    /// * `prompt` — The raw prompt text to analyse and route.
161    ///
162    /// # Returns
163    ///
164    /// A [`RoutingDecision`] indicating which worker(s) to use.
165    ///
166    /// # Panics
167    ///
168    /// This function never panics.
169    pub fn route(&self, prompt: &str) -> RoutingDecision {
170        let score = self.scorer.score(prompt);
171
172        let effective_threshold = self
173            .effective_cloud_threshold
174            .read()
175            .unwrap_or_else(|p| {
176                tracing::warn!("router: RwLock poisoned, using stale state");
177                p.into_inner()
178            })
179            .to_owned();
180
181        if score < self.config.local_threshold {
182            RoutingDecision::Local {
183                score,
184                worker: WorkerKind::LlamaCpp,
185            }
186        } else if score > effective_threshold {
187            RoutingDecision::Cloud {
188                score,
189                worker: WorkerKind::Anthropic,
190            }
191        } else {
192            RoutingDecision::LocalWithFallback {
193                score,
194                primary: WorkerKind::LlamaCpp,
195                fallback: WorkerKind::Anthropic,
196            }
197        }
198    }
199
200    /// Record the outcome of a routing decision.
201    ///
202    /// For [`RoutingDecision::LocalWithFallback`] outcomes, this feeds the
203    /// adaptive threshold algorithm. When enough samples accumulate and the
204    /// local failure rate exceeds `adaptive.failure_ceiling`, the effective
205    /// cloud threshold is lowered (sending more requests directly to cloud).
206    ///
207    /// # Arguments
208    ///
209    /// * `decision` — The routing decision that was made.
210    /// * `success` — Whether the primary worker succeeded.
211    /// * `tokens` — Number of tokens processed (for cost tracking).
212    ///
213    /// # Panics
214    ///
215    /// This function never panics.
216    pub fn record_outcome(&self, decision: &RoutingDecision, success: bool, tokens: u64) {
217        match decision {
218            RoutingDecision::Local { .. } => {
219                self.cost_tracker.record_local(tokens);
220            }
221            RoutingDecision::Cloud { .. } => {
222                self.cost_tracker.record_cloud(tokens);
223            }
224            RoutingDecision::LocalWithFallback { .. } => {
225                if success {
226                    self.fallback_successes.fetch_add(1, Ordering::Relaxed);
227                    self.cost_tracker.record_local(tokens);
228                } else {
229                    self.fallback_failures.fetch_add(1, Ordering::Relaxed);
230                    self.cost_tracker.record_fallback(tokens);
231                }
232                self.maybe_adapt();
233            }
234        }
235    }
236
237    /// Return a reference to the underlying cost tracker.
238    ///
239    /// # Panics
240    ///
241    /// This function never panics.
242    pub fn cost_tracker(&self) -> &CostTracker {
243        &self.cost_tracker
244    }
245
246    /// Return the current effective cloud threshold.
247    ///
248    /// This may differ from `config.cloud_threshold` if adaptive adjustment
249    /// has been triggered.
250    ///
251    /// # Panics
252    ///
253    /// This function never panics.
254    pub fn effective_cloud_threshold(&self) -> f64 {
255        self.effective_cloud_threshold
256            .read()
257            .unwrap_or_else(|p| {
258                tracing::warn!("router: RwLock poisoned, using stale state");
259                p.into_inner()
260            })
261            .to_owned()
262    }
263
264    /// Return a reference to the scorer for external breakdown queries.
265    ///
266    /// # Panics
267    ///
268    /// This function never panics.
269    pub fn scorer(&self) -> &ComplexityScorer {
270        &self.scorer
271    }
272
273    // ── Adaptive threshold adjustment ──────────────────────────────────
274
275    /// Check whether the adaptive threshold should be adjusted.
276    fn maybe_adapt(&self) {
277        if !self.config.adaptive.enabled {
278            return;
279        }
280
281        let successes = self.fallback_successes.load(Ordering::Relaxed);
282        let failures = self.fallback_failures.load(Ordering::Relaxed);
283        let total = successes + failures;
284
285        if total < self.config.adaptive.min_samples {
286            return;
287        }
288
289        let failure_rate = failures as f64 / total as f64;
290
291        if let Ok(mut threshold) = self.effective_cloud_threshold.write() {
292            if failure_rate > self.config.adaptive.failure_ceiling {
293                // Too many local failures in fallback zone → lower threshold
294                // so more requests go directly to cloud.
295                let new = (*threshold - self.config.adaptive.step).max(self.config.local_threshold);
296                *threshold = new;
297            } else if failure_rate < self.config.adaptive.failure_ceiling / 2.0 {
298                // Local is doing well → raise threshold (save more money).
299                let new = (*threshold + self.config.adaptive.step).min(1.0);
300                *threshold = new;
301            }
302        }
303    }
304}
305
306// ── Tests ──────────────────────────────────────────────────────────────
307
308#[cfg(test)]
309mod tests {
310    use super::*;
311
312    fn default_router() -> ModelRouter {
313        ModelRouter::new(RoutingConfig::default())
314    }
315
316    // -- routing decisions -----------------------------------------------
317
318    #[test]
319    fn test_route_simple_greeting_returns_local() {
320        let router = default_router();
321        let decision = router.route("Say hello");
322        assert!(
323            decision.is_local(),
324            "simple greeting should route local, got: {decision:?}"
325        );
326        assert_eq!(
327            match &decision {
328                RoutingDecision::Local { worker, .. } => worker,
329                _ => &WorkerKind::Echo,
330            },
331            &WorkerKind::LlamaCpp
332        );
333    }
334
335    #[test]
336    fn test_route_complex_rust_prompt_returns_cloud() {
337        let router = default_router();
338        let prompt = r#"Debug this Rust code that has a borrow checker error with lifetime issues:
339
340```rust
341fn process<'a>(data: &'a [u8]) -> &'a str {
342    let owned = String::from_utf8_lossy(data).to_string();
343    &owned
344}
345```
346
3471. Explain why the borrow checker rejects this
3482. Show the correct implementation with proper lifetime annotations
3493. Add unit tests for the fix
350
351Also, that function should handle it properly when those bytes are invalid UTF-8. Fix the thing so it works with tokio async fn properly."#;
352        let decision = router.route(prompt);
353        assert!(
354            decision.is_cloud(),
355            "complex Rust debugging should route to cloud, got: {decision:?}"
356        );
357        assert_eq!(
358            match &decision {
359                RoutingDecision::Cloud { worker, .. } => worker,
360                _ => &WorkerKind::Echo,
361            },
362            &WorkerKind::Anthropic
363        );
364    }
365
366    #[test]
367    fn test_route_medium_complexity_returns_fallback() {
368        let router = default_router();
369        // Code block alone gives 0.2, numbered list gives 0.2 → 0.4 (boundary)
370        let prompt = "Fix this:\n```\nfoo()\n```\n1. Step one\n2. Step two";
371        let decision = router.route(prompt);
372        // Score = 0.4, which is >= local_threshold(0.4) and <= cloud_threshold(0.7)
373        assert!(
374            decision.is_fallback(),
375            "medium complexity should route to fallback, got: {decision:?}"
376        );
377    }
378
379    #[test]
380    fn test_route_empty_prompt_returns_local() {
381        let router = default_router();
382        let decision = router.route("");
383        assert!(decision.is_local());
384    }
385
386    // -- score accessor --------------------------------------------------
387
388    #[test]
389    fn test_routing_decision_score_accessor() {
390        let router = default_router();
391        let decision = router.route("Hello");
392        assert!(decision.score() >= 0.0 && decision.score() <= 1.0);
393    }
394
395    // -- outcome recording -----------------------------------------------
396
397    #[test]
398    fn test_record_outcome_local_success_tracks_local_cost() {
399        let router = default_router();
400        let decision = RoutingDecision::Local {
401            score: 0.1,
402            worker: WorkerKind::LlamaCpp,
403        };
404        router.record_outcome(&decision, true, 100);
405        let snap = router.cost_tracker().snapshot();
406        assert_eq!(snap.local_tokens, 100);
407        assert_eq!(snap.local_requests, 1);
408    }
409
410    #[test]
411    fn test_record_outcome_cloud_tracks_cloud_cost() {
412        let router = default_router();
413        let decision = RoutingDecision::Cloud {
414            score: 0.9,
415            worker: WorkerKind::Anthropic,
416        };
417        router.record_outcome(&decision, true, 500);
418        let snap = router.cost_tracker().snapshot();
419        assert_eq!(snap.cloud_tokens, 500);
420        assert_eq!(snap.cloud_requests, 1);
421    }
422
423    #[test]
424    fn test_record_outcome_fallback_success_tracks_local() {
425        let router = default_router();
426        let decision = RoutingDecision::LocalWithFallback {
427            score: 0.5,
428            primary: WorkerKind::LlamaCpp,
429            fallback: WorkerKind::Anthropic,
430        };
431        router.record_outcome(&decision, true, 200);
432        let snap = router.cost_tracker().snapshot();
433        assert_eq!(snap.local_tokens, 200);
434    }
435
436    #[test]
437    fn test_record_outcome_fallback_failure_tracks_cloud() {
438        let router = default_router();
439        let decision = RoutingDecision::LocalWithFallback {
440            score: 0.5,
441            primary: WorkerKind::LlamaCpp,
442            fallback: WorkerKind::Anthropic,
443        };
444        router.record_outcome(&decision, false, 200);
445        let snap = router.cost_tracker().snapshot();
446        assert_eq!(snap.cloud_tokens, 200);
447        assert_eq!(snap.fallback_requests, 1);
448    }
449
450    // -- adaptive threshold -----------------------------------------------
451
452    #[test]
453    fn test_adaptive_lowers_threshold_on_high_failure_rate() {
454        let config = RoutingConfig {
455            adaptive: super::super::config::AdaptiveConfig {
456                enabled: true,
457                step: 0.05,
458                min_samples: 5,
459                failure_ceiling: 0.3,
460            },
461            ..RoutingConfig::default()
462        };
463        let router = ModelRouter::new(config);
464        let original_threshold = router.effective_cloud_threshold();
465
466        let decision = RoutingDecision::LocalWithFallback {
467            score: 0.5,
468            primary: WorkerKind::LlamaCpp,
469            fallback: WorkerKind::Anthropic,
470        };
471
472        // Record 5 failures out of 5 (100% failure rate, > 30% ceiling)
473        for _ in 0..5 {
474            router.record_outcome(&decision, false, 100);
475        }
476
477        let new_threshold = router.effective_cloud_threshold();
478        assert!(
479            new_threshold < original_threshold,
480            "threshold should decrease: was {original_threshold}, now {new_threshold}"
481        );
482    }
483
484    #[test]
485    fn test_adaptive_raises_threshold_on_low_failure_rate() {
486        let config = RoutingConfig {
487            adaptive: super::super::config::AdaptiveConfig {
488                enabled: true,
489                step: 0.05,
490                min_samples: 5,
491                failure_ceiling: 0.3,
492            },
493            ..RoutingConfig::default()
494        };
495        let router = ModelRouter::new(config);
496        let original_threshold = router.effective_cloud_threshold();
497
498        let decision = RoutingDecision::LocalWithFallback {
499            score: 0.5,
500            primary: WorkerKind::LlamaCpp,
501            fallback: WorkerKind::Anthropic,
502        };
503
504        // Record 5 successes out of 5 (0% failure rate, well below ceiling/2)
505        for _ in 0..5 {
506            router.record_outcome(&decision, true, 100);
507        }
508
509        let new_threshold = router.effective_cloud_threshold();
510        assert!(
511            new_threshold > original_threshold,
512            "threshold should increase: was {original_threshold}, now {new_threshold}"
513        );
514    }
515
516    #[test]
517    fn test_adaptive_no_change_below_min_samples() {
518        let config = RoutingConfig {
519            adaptive: super::super::config::AdaptiveConfig {
520                enabled: true,
521                step: 0.05,
522                min_samples: 100,
523                failure_ceiling: 0.3,
524            },
525            ..RoutingConfig::default()
526        };
527        let router = ModelRouter::new(config);
528        let original = router.effective_cloud_threshold();
529
530        let decision = RoutingDecision::LocalWithFallback {
531            score: 0.5,
532            primary: WorkerKind::LlamaCpp,
533            fallback: WorkerKind::Anthropic,
534        };
535
536        // Only 5 outcomes, below min_samples of 100
537        for _ in 0..5 {
538            router.record_outcome(&decision, false, 100);
539        }
540
541        assert!(
542            (router.effective_cloud_threshold() - original).abs() < f64::EPSILON,
543            "threshold should not change with insufficient samples"
544        );
545    }
546
547    #[test]
548    fn test_adaptive_disabled_no_threshold_change() {
549        let config = RoutingConfig {
550            adaptive: super::super::config::AdaptiveConfig {
551                enabled: false,
552                step: 0.05,
553                min_samples: 5,
554                failure_ceiling: 0.3,
555            },
556            ..RoutingConfig::default()
557        };
558        let router = ModelRouter::new(config);
559        let original = router.effective_cloud_threshold();
560
561        let decision = RoutingDecision::LocalWithFallback {
562            score: 0.5,
563            primary: WorkerKind::LlamaCpp,
564            fallback: WorkerKind::Anthropic,
565        };
566
567        for _ in 0..10 {
568            router.record_outcome(&decision, false, 100);
569        }
570
571        assert!(
572            (router.effective_cloud_threshold() - original).abs() < f64::EPSILON,
573            "threshold should not change when adaptive is disabled"
574        );
575    }
576
577    #[test]
578    fn test_adaptive_threshold_never_below_local_threshold() {
579        let config = RoutingConfig {
580            local_threshold: 0.4,
581            cloud_threshold: 0.45,
582            adaptive: super::super::config::AdaptiveConfig {
583                enabled: true,
584                step: 0.1,
585                min_samples: 5,
586                failure_ceiling: 0.3,
587            },
588            ..RoutingConfig::default()
589        };
590        let router = ModelRouter::new(config);
591
592        let decision = RoutingDecision::LocalWithFallback {
593            score: 0.5,
594            primary: WorkerKind::LlamaCpp,
595            fallback: WorkerKind::Anthropic,
596        };
597
598        // Drive many failures to push threshold down
599        for _ in 0..50 {
600            router.record_outcome(&decision, false, 100);
601        }
602
603        assert!(
604            router.effective_cloud_threshold() >= 0.4,
605            "threshold must not drop below local_threshold"
606        );
607    }
608
609    #[test]
610    fn test_adaptive_threshold_never_above_1_0() {
611        let config = RoutingConfig {
612            cloud_threshold: 0.95,
613            adaptive: super::super::config::AdaptiveConfig {
614                enabled: true,
615                step: 0.1,
616                min_samples: 5,
617                failure_ceiling: 0.3,
618            },
619            ..RoutingConfig::default()
620        };
621        let router = ModelRouter::new(config);
622
623        let decision = RoutingDecision::LocalWithFallback {
624            score: 0.5,
625            primary: WorkerKind::LlamaCpp,
626            fallback: WorkerKind::Anthropic,
627        };
628
629        // Drive many successes to push threshold up
630        for _ in 0..50 {
631            router.record_outcome(&decision, true, 100);
632        }
633
634        assert!(
635            router.effective_cloud_threshold() <= 1.0,
636            "threshold must not exceed 1.0"
637        );
638    }
639
640    // -- custom thresholds -----------------------------------------------
641
642    #[test]
643    fn test_route_with_custom_thresholds() {
644        let config = RoutingConfig {
645            local_threshold: 0.2,
646            cloud_threshold: 0.5,
647            ..RoutingConfig::default()
648        };
649        let router = ModelRouter::new(config);
650
651        // Code block alone = 0.2, which is >= 0.2 local threshold
652        let prompt = "```\ncode\n```";
653        let decision = router.route(prompt);
654        assert!(
655            decision.is_fallback() || decision.is_cloud(),
656            "score 0.2 with local_threshold 0.2 should not be Local"
657        );
658    }
659
660    // -- RoutingDecision predicates --------------------------------------
661
662    #[test]
663    fn test_routing_decision_is_local() {
664        let d = RoutingDecision::Local {
665            score: 0.1,
666            worker: WorkerKind::LlamaCpp,
667        };
668        assert!(d.is_local());
669        assert!(!d.is_cloud());
670        assert!(!d.is_fallback());
671    }
672
673    #[test]
674    fn test_routing_decision_is_cloud() {
675        let d = RoutingDecision::Cloud {
676            score: 0.9,
677            worker: WorkerKind::Anthropic,
678        };
679        assert!(d.is_cloud());
680        assert!(!d.is_local());
681        assert!(!d.is_fallback());
682    }
683
684    #[test]
685    fn test_routing_decision_is_fallback() {
686        let d = RoutingDecision::LocalWithFallback {
687            score: 0.5,
688            primary: WorkerKind::LlamaCpp,
689            fallback: WorkerKind::Anthropic,
690        };
691        assert!(d.is_fallback());
692        assert!(!d.is_local());
693        assert!(!d.is_cloud());
694    }
695
696    // -- cost tracking integration ---------------------------------------
697
698    #[test]
699    fn test_cost_savings_after_routing_decisions() {
700        let router = default_router();
701
702        // Route 10 simple prompts locally
703        for _ in 0..10 {
704            let d = RoutingDecision::Local {
705                score: 0.1,
706                worker: WorkerKind::LlamaCpp,
707            };
708            router.record_outcome(&d, true, 500);
709        }
710
711        // Route 2 complex prompts to cloud
712        for _ in 0..2 {
713            let d = RoutingDecision::Cloud {
714                score: 0.9,
715                worker: WorkerKind::Anthropic,
716            };
717            router.record_outcome(&d, true, 500);
718        }
719
720        let snap = router.cost_tracker().snapshot();
721        // 5000 local + 1000 cloud = 6000 total tokens
722        assert_eq!(snap.local_tokens, 5000);
723        assert_eq!(snap.cloud_tokens, 1000);
724        // Savings should be positive (local is free)
725        assert!(snap.savings_usd > 0.0);
726        // Savings should be ~83.3% (10/12 of tokens were free)
727        assert!(snap.savings_percent > 80.0);
728    }
729
730    // -- debug -----------------------------------------------------------
731
732    #[test]
733    fn test_model_router_debug_does_not_panic() {
734        let router = default_router();
735        let _ = format!("{router:?}");
736    }
737}