Skip to main content

navi_core/
model_router.rs

1//! Benchmark-driven model routing contracts.
2
3use 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}