Skip to main content

aria_router_algorithm/
lib.rs

1//! Selection algorithms (static / latency-aware / multi-factor).
2
3use aria_router_config::{DecisionCfg, RouterDocument};
4use aria_router_core::{RouterError, ModelCard};
5
6#[derive(Debug, Clone, Default)]
7pub struct RuntimeStats {
8    pub latency_ms: std::collections::HashMap<String, f32>,
9    pub load: std::collections::HashMap<String, f32>,
10    pub cost: std::collections::HashMap<String, f32>,
11}
12
13pub fn select(
14    _doc: &RouterDocument,
15    decision: &DecisionCfg,
16    eligible: &[ModelCard],
17    stats: &RuntimeStats,
18) -> Result<String, RouterError> {
19    if eligible.is_empty() {
20        return Err(RouterError::FailClosed("no eligible models".into()));
21    }
22    let algo = decision.algorithm.as_deref().unwrap_or("static");
23    if RouterDocument::unimplemented_algorithm(algo) {
24        return Err(RouterError::Unsupported(format!("algorithm {algo} not implemented")));
25    }
26    let names: Vec<String> = if decision.model_refs.is_empty() {
27        eligible.iter().map(|m| m.name.clone()).collect()
28    } else {
29        decision
30            .model_refs
31            .iter()
32            .map(|r| r.model.clone())
33            .filter(|n| eligible.iter().any(|e| e.name == *n))
34            .collect()
35    };
36    if names.is_empty() {
37        return Err(RouterError::FailClosed(
38            "decision modelRefs not in eligible pool".into(),
39        ));
40    }
41    match algo {
42        "static" => Ok(names[0].clone()),
43        "latency-aware" | "latency_aware" => {
44            let best = names
45                .iter()
46                .min_by(|a, b| {
47                    let la = stats.latency_ms.get(*a).copied().unwrap_or(1000.0);
48                    let lb = stats.latency_ms.get(*b).copied().unwrap_or(1000.0);
49                    la.partial_cmp(&lb).unwrap_or(std::cmp::Ordering::Equal)
50                })
51                .cloned()
52                .unwrap();
53            Ok(best)
54        }
55        "multi-factor" | "multi_factor" => {
56            let best = names
57                .iter()
58                .min_by(|a, b| {
59                    let sa = score(a, stats);
60                    let sb = score(b, stats);
61                    sa.partial_cmp(&sb).unwrap_or(std::cmp::Ordering::Equal)
62                })
63                .cloned()
64                .unwrap();
65            Ok(best)
66        }
67        other => Err(RouterError::Unsupported(format!("algorithm {other}"))),
68    }
69}
70
71fn score(name: &str, stats: &RuntimeStats) -> f32 {
72    let lat = stats.latency_ms.get(name).copied().unwrap_or(100.0);
73    let load = stats.load.get(name).copied().unwrap_or(0.0);
74    let cost = stats.cost.get(name).copied().unwrap_or(1.0);
75    lat * 0.5 + load * 20.0 + cost * 10.0
76}
77
78pub fn hard_filter(
79    doc: &RouterDocument,
80    names: &[String],
81    require_locality: Option<&str>,
82    require_modality: Option<&str>,
83) -> Vec<ModelCard> {
84    names
85        .iter()
86        .filter_map(|n| doc.provider(n))
87        .filter(|p| {
88            if let Some(loc) = require_locality {
89                if p.locality != loc {
90                    return false;
91                }
92            }
93            if let Some(mod_) = require_modality {
94                if p.modality != mod_ && p.modality != "any" {
95                    return false;
96                }
97            }
98            true
99        })
100        .map(|p| ModelCard {
101            name: p.name.clone(),
102            locality: p.locality.clone(),
103            modality: p.modality.clone(),
104            capabilities: p.capabilities.clone(),
105            provider_model_id: p.provider_model_id.clone(),
106        })
107        .collect()
108}
109
110#[cfg(test)]
111mod tests {
112    use super::*;
113
114    fn cards() -> Vec<ModelCard> {
115        vec![
116            ModelCard {
117                name: "a".into(),
118                locality: "local".into(),
119                modality: "text".into(),
120                capabilities: vec!["chat".into()],
121                provider_model_id: "a".into(),
122            },
123            ModelCard {
124                name: "b".into(),
125                locality: "local".into(),
126                modality: "text".into(),
127                capabilities: vec!["chat".into()],
128                provider_model_id: "b".into(),
129            },
130        ]
131    }
132
133    fn decision(algo: &str) -> DecisionCfg {
134        DecisionCfg {
135            name: "d".into(),
136            description: None,
137            priority: 1,
138            rules: Default::default(),
139            model_refs: vec![
140                aria_router_config::ModelRef { model: "a".into() },
141                aria_router_config::ModelRef { model: "b".into() },
142            ],
143            algorithm: Some(algo.into()),
144            plugins: vec![],
145            locality: None,
146        }
147    }
148
149    fn doc() -> RouterDocument {
150        RouterDocument::from_yaml_str(
151            r#"
152version: v0.3
153providers:
154  models:
155    - name: a
156      locality: local
157      backend_refs: [{name: p, endpoint: 127.0.0.1:1}]
158    - name: b
159      locality: local
160      backend_refs: [{name: p, endpoint: 127.0.0.1:2}]
161entrypoints:
162  - model_names: [auto]
163    router: semantic
164    recipe: r
165recipes:
166  - name: r
167    router: semantic
168    routing:
169      decisions:
170        - name: d
171          rules: { operator: AND, conditions: [] }
172          modelRefs: [{model: a}]
173"#,
174        )
175        .unwrap()
176    }
177
178    #[test]
179    fn static_first() {
180        let d = doc();
181        let got = select(&d, &decision("static"), &cards(), &RuntimeStats::default()).unwrap();
182        assert_eq!(got, "a");
183    }
184
185    #[test]
186    fn latency_aware_picks_faster() {
187        let d = doc();
188        let mut stats = RuntimeStats::default();
189        stats.latency_ms.insert("a".into(), 200.0);
190        stats.latency_ms.insert("b".into(), 10.0);
191        let got = select(&d, &decision("latency-aware"), &cards(), &stats).unwrap();
192        assert_eq!(got, "b");
193    }
194
195    #[test]
196    fn multi_factor_picks_cheaper() {
197        let d = doc();
198        let mut stats = RuntimeStats::default();
199        stats.cost.insert("a".into(), 9.0);
200        stats.cost.insert("b".into(), 1.0);
201        let got = select(&d, &decision("multi-factor"), &cards(), &stats).unwrap();
202        assert_eq!(got, "b");
203    }
204
205    #[test]
206    fn unimplemented_algorithm() {
207        let d = doc();
208        let err = select(&d, &decision("knn"), &cards(), &RuntimeStats::default()).unwrap_err();
209        assert!(matches!(err, RouterError::Unsupported(_)));
210    }
211}