navi_core/
model_router.rs1use serde::{Deserialize, Serialize};
4use std::collections::BTreeMap;
5
6#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
7pub struct ModelScorecard {
8 pub provider_id: String,
9 pub model: String,
10 pub role: ModelRouteRole,
11 pub success_rate: f64,
12 pub verifier_pass_rate: f64,
13 pub tool_call_validity: f64,
14 pub cost_per_1k_tokens: f64,
15 pub latency_ms: f64,
16 pub retry_rate: f64,
17 pub unsafe_action_rate: f64,
18}
19
20#[derive(Debug, Clone, PartialEq, Eq, PartialOrd, Ord, Serialize, Deserialize)]
21#[serde(rename_all = "snake_case")]
22pub enum ModelRouteRole {
23 Planner,
24 Router,
25 Coder,
26 Reviewer,
27 VerifierJudge,
28 Summarizer,
29 MemoryMiner,
30}
31
32#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
33pub struct ModelRoute {
34 pub role: ModelRouteRole,
35 pub provider_id: String,
36 pub model: String,
37 pub score: f64,
38 pub fallback: Option<Box<ModelRoute>>,
39}
40
41#[derive(Debug, Clone, Default, Serialize, Deserialize)]
42pub struct ModelRouter {
43 scorecards: BTreeMap<ModelRouteRole, Vec<ModelScorecard>>,
44}
45
46impl ModelRouter {
47 pub fn add_scorecard(&mut self, scorecard: ModelScorecard) {
48 self.scorecards
49 .entry(scorecard.role.clone())
50 .or_default()
51 .push(scorecard);
52 }
53
54 pub fn route(&self, role: ModelRouteRole, high_risk: bool) -> Option<ModelRoute> {
55 let mut candidates = self.scorecards.get(&role)?.clone();
56 candidates.sort_by(|left, right| {
57 route_score(right, high_risk)
58 .partial_cmp(&route_score(left, high_risk))
59 .unwrap_or(std::cmp::Ordering::Equal)
60 });
61 let best = candidates.first()?;
62 let fallback = candidates.get(1).map(|candidate| {
63 Box::new(ModelRoute {
64 role: role.clone(),
65 provider_id: candidate.provider_id.clone(),
66 model: candidate.model.clone(),
67 score: route_score(candidate, high_risk),
68 fallback: None,
69 })
70 });
71 Some(ModelRoute {
72 role,
73 provider_id: best.provider_id.clone(),
74 model: best.model.clone(),
75 score: route_score(best, high_risk),
76 fallback,
77 })
78 }
79}
80
81fn route_score(card: &ModelScorecard, high_risk: bool) -> f64 {
82 let mut score =
83 card.success_rate * 35.0 + card.verifier_pass_rate * 30.0 + card.tool_call_validity * 15.0
84 - card.retry_rate * 10.0
85 - card.unsafe_action_rate * if high_risk { 40.0 } else { 15.0 }
86 - (card.cost_per_1k_tokens * 2.0)
87 - (card.latency_ms / 10_000.0).min(10.0);
88 if high_risk && card.verifier_pass_rate >= 0.95 {
89 score += 5.0;
90 }
91 score
92}
93
94#[cfg(test)]
95mod tests {
96 use super::*;
97
98 fn card(model: &str, unsafe_rate: f64, cost: f64) -> ModelScorecard {
99 ModelScorecard {
100 provider_id: "p".to_string(),
101 model: model.to_string(),
102 role: ModelRouteRole::Coder,
103 success_rate: 0.9,
104 verifier_pass_rate: 0.95,
105 tool_call_validity: 0.98,
106 cost_per_1k_tokens: cost,
107 latency_ms: 1000.0,
108 retry_rate: 0.1,
109 unsafe_action_rate: unsafe_rate,
110 }
111 }
112
113 #[test]
114 fn high_risk_routes_away_from_unsafe_model() {
115 let mut router = ModelRouter::default();
116 router.add_scorecard(card("cheap-risky", 0.2, 0.1));
117 router.add_scorecard(card("safe", 0.0, 1.0));
118
119 let route = router.route(ModelRouteRole::Coder, true).unwrap();
120
121 assert_eq!(route.model, "safe");
122 assert_eq!(route.fallback.unwrap().model, "cheap-risky");
123 }
124}