aria_router_algorithm/
lib.rs1use 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}