1use 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#[derive(Debug, Clone, PartialEq)]
23pub enum RoutingDecision {
24 Local {
26 score: f64,
28 worker: WorkerKind,
30 },
31 LocalWithFallback {
34 score: f64,
36 primary: WorkerKind,
38 fallback: WorkerKind,
40 },
41 Cloud {
43 score: f64,
45 worker: WorkerKind,
47 },
48}
49
50impl RoutingDecision {
51 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 pub fn is_local(&self) -> bool {
70 matches!(self, Self::Local { .. })
71 }
72
73 pub fn is_cloud(&self) -> bool {
79 matches!(self, Self::Cloud { .. })
80 }
81
82 pub fn is_fallback(&self) -> bool {
88 matches!(self, Self::LocalWithFallback { .. })
89 }
90}
91
92pub struct ModelRouter {
103 scorer: ComplexityScorer,
104 config: RoutingConfig,
105 cost_tracker: CostTracker,
106
107 effective_cloud_threshold: RwLock<f64>,
110
111 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 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 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 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 pub fn cost_tracker(&self) -> &CostTracker {
243 &self.cost_tracker
244 }
245
246 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 pub fn scorer(&self) -> &ComplexityScorer {
270 &self.scorer
271 }
272
273 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 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 let new = (*threshold + self.config.adaptive.step).min(1.0);
300 *threshold = new;
301 }
302 }
303 }
304}
305
306#[cfg(test)]
309mod tests {
310 use super::*;
311
312 fn default_router() -> ModelRouter {
313 ModelRouter::new(RoutingConfig::default())
314 }
315
316 #[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 let prompt = "Fix this:\n```\nfoo()\n```\n1. Step one\n2. Step two";
371 let decision = router.route(prompt);
372 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 #[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 #[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 #[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 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 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 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 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 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 #[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 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 #[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 #[test]
699 fn test_cost_savings_after_routing_decisions() {
700 let router = default_router();
701
702 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 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 assert_eq!(snap.local_tokens, 5000);
723 assert_eq!(snap.cloud_tokens, 1000);
724 assert!(snap.savings_usd > 0.0);
726 assert!(snap.savings_percent > 80.0);
728 }
729
730 #[test]
733 fn test_model_router_debug_does_not_panic() {
734 let router = default_router();
735 let _ = format!("{router:?}");
736 }
737}