1use serde::Deserialize;
4use std::collections::HashMap;
5
6#[derive(Debug, Clone, Default, Deserialize, PartialEq, Eq)]
7pub struct ChainLeg {
8 pub provider: String,
9 pub model: String,
10 #[serde(default)]
15 pub region: Option<String>,
16}
17
18#[derive(Debug, Clone, Deserialize)]
19struct RouteEntry {
20 legs: Vec<ChainLeg>,
21 #[serde(default)]
22 policy: Option<String>,
23}
24
25#[derive(Debug, Clone, Deserialize)]
26struct RoutesFile {
27 routes: HashMap<String, RouteEntry>,
28}
29
30#[derive(Debug, Clone)]
31pub struct RouteTable {
32 routes: HashMap<String, Vec<ChainLeg>>,
33 policies: HashMap<String, String>,
34}
35
36impl RouteTable {
37 pub fn from_toml_str(s: &str) -> anyhow::Result<Self> {
38 let file = toml::from_str::<RoutesFile>(s)?;
39 let mut routes = HashMap::new();
40 let mut policies = HashMap::new();
41 for (name, entry) in file.routes {
42 if let Some(policy) = entry.policy {
43 policies.insert(name.clone(), policy);
44 }
45 routes.insert(name, entry.legs);
46 }
47 Ok(Self { routes, policies })
48 }
49
50 pub fn legs(&self, model: &str) -> Option<&[ChainLeg]> {
52 self.routes.get(model).map(Vec::as_slice)
53 }
54
55 pub fn policy_of(&self, model: &str) -> Option<&str> {
57 self.policies.get(model).map(String::as_str)
58 }
59
60 pub fn aliases(&self) -> Vec<String> {
62 let mut v: Vec<String> = self.routes.keys().cloned().collect();
63 v.sort();
64 v
65 }
66
67 pub fn referenced_providers(&self) -> std::collections::HashSet<String> {
69 self.routes
70 .values()
71 .flatten()
72 .map(|l| l.provider.clone())
73 .collect()
74 }
75
76 pub fn vertex_fallback_chain(&self, model: &str) -> VertexPassthroughChain {
84 let mut aliases: Vec<&String> = self.routes.keys().collect();
85 aliases.sort();
86
87 let from_first = aliases.iter().find_map(|alias| {
88 let legs = self.routes.get(*alias)?;
89 let first = legs.first()?;
90 if first.provider == "vertex" && first.model == model {
91 Some(((*alias).clone(), 0usize))
92 } else {
93 None
94 }
95 });
96
97 let resolved = from_first.or_else(|| {
98 aliases.iter().find_map(|alias| {
99 let legs = self.routes.get(*alias)?;
100 legs.iter()
101 .position(|l| l.provider == "vertex" && l.model == model)
102 .map(|idx| ((*alias).clone(), idx))
103 })
104 });
105
106 match resolved {
107 Some((route, start)) => {
108 let legs = self.routes.get(&route).expect("alias from routes keys");
109 let chain: Vec<ChainLeg> = legs[start..]
110 .iter()
111 .take_while(|l| l.provider == "vertex")
112 .cloned()
113 .collect();
114 VertexPassthroughChain {
115 route: Some(route),
116 legs: chain,
117 }
118 }
119 None => VertexPassthroughChain {
120 route: None,
121 legs: vec![ChainLeg {
122 provider: "vertex".into(),
123 model: model.to_string(),
124 region: None,
125 }],
126 },
127 }
128 }
129}
130
131#[derive(Debug, Clone, PartialEq, Eq)]
133pub struct VertexPassthroughChain {
134 pub route: Option<String>,
136 pub legs: Vec<ChainLeg>,
137}
138
139#[cfg(test)]
140mod tests {
141 use super::*;
142
143 const SAMPLE: &str = r#"
144 [routes."gemini-pro"]
145 legs = [
146 { provider = "vertex", model = "gemini-3-pro" },
147 { provider = "qwen", model = "qwen-max" },
148 ]
149 [routes."fast"]
150 legs = [{ provider = "vertex", model = "gemini-3-flash" }]
151 "#;
152
153 #[test]
154 fn parses_and_resolves_legs_in_order() {
155 let t = RouteTable::from_toml_str(SAMPLE).unwrap();
156 let legs = t.legs("gemini-pro").unwrap();
157 assert_eq!(legs.len(), 2);
158 assert_eq!(
159 legs[0],
160 ChainLeg {
161 provider: "vertex".into(),
162 model: "gemini-3-pro".into(),
163 ..Default::default()
164 }
165 );
166 assert_eq!(legs[1].provider, "qwen");
167 assert!(t.legs("nope").is_none());
168 }
169
170 #[test]
171 fn parses_optional_per_leg_region() {
172 let t = RouteTable::from_toml_str(
173 r#"
174 [routes."visual"]
175 legs = [
176 { provider = "vertex", model = "gemini-3.1-pro-preview", region = "global" },
177 { provider = "qwen", model = "qwen3-vl-plus" },
178 ]
179 "#,
180 )
181 .unwrap();
182 let legs = t.legs("visual").unwrap();
183 assert_eq!(legs[0].region.as_deref(), Some("global"));
184 assert_eq!(legs[1].region, None);
185 }
186
187 #[test]
188 fn aliases_are_sorted() {
189 let t = RouteTable::from_toml_str(SAMPLE).unwrap();
190 assert_eq!(
191 t.aliases(),
192 vec!["fast".to_string(), "gemini-pro".to_string()]
193 );
194 }
195
196 #[test]
197 fn referenced_providers_collected() {
198 let t = RouteTable::from_toml_str(SAMPLE).unwrap();
199 let p = t.referenced_providers();
200 assert!(p.contains("vertex"));
201 assert!(p.contains("qwen"));
202 }
203
204 #[test]
205 fn parses_optional_route_policy_and_defaults_to_none() {
206 let t = RouteTable::from_toml_str(
207 r#"
208 [routes."guarded"]
209 policy = "strict"
210 legs = [{ provider = "vertex", model = "gemini-3-pro" }]
211 [routes."plain"]
212 legs = [{ provider = "qwen", model = "qwen-max" }]
213 "#,
214 )
215 .unwrap();
216 assert_eq!(t.policy_of("guarded"), Some("strict"));
217 assert_eq!(t.policy_of("plain"), None);
218 assert_eq!(t.policy_of("missing"), None);
219 }
220
221 #[test]
222 fn vertex_fallback_chain_deep_like_and_stops_before_qwen() {
223 let t = RouteTable::from_toml_str(
224 r#"
225 [routes."conversation"]
226 legs = [
227 { provider = "vertex", model = "gemini-3.1-pro-preview", region = "global" },
228 { provider = "vertex", model = "gemini-2.5-pro", region = "us-central1" },
229 ]
230 [routes."visual"]
231 legs = [
232 { provider = "vertex", model = "gemini-flash-image", region = "global" },
233 { provider = "qwen", model = "qwen3-vl-plus" },
234 ]
235 "#,
236 )
237 .unwrap();
238 let c = t.vertex_fallback_chain("gemini-3.1-pro-preview");
239 assert_eq!(c.route.as_deref(), Some("conversation"));
240 assert_eq!(c.legs.len(), 2);
241 assert_eq!(c.legs[1].model, "gemini-2.5-pro");
242
243 let image = t.vertex_fallback_chain("gemini-flash-image");
244 assert_eq!(image.legs.len(), 1);
245 }
246
247 #[test]
248 fn vertex_fallback_chain_picks_lex_first_alias() {
249 let t = RouteTable::from_toml_str(
250 r#"
251 [routes."planning"]
252 legs = [
253 { provider = "vertex", model = "gemini-3.1-pro-preview" },
254 { provider = "vertex", model = "gemini-2.5-pro" },
255 ]
256 [routes."conversation"]
257 legs = [
258 { provider = "vertex", model = "gemini-3.1-pro-preview" },
259 { provider = "vertex", model = "gemini-2.5-pro" },
260 ]
261 "#,
262 )
263 .unwrap();
264 assert_eq!(
265 t.vertex_fallback_chain("gemini-3.1-pro-preview")
266 .route
267 .as_deref(),
268 Some("conversation")
269 );
270 }
271
272 #[test]
273 fn vertex_fallback_chain_unknown_is_synthetic_single() {
274 let t = RouteTable::from_toml_str(SAMPLE).unwrap();
275 let c = t.vertex_fallback_chain("totally-unknown");
276 assert_eq!(c.route, None);
277 assert_eq!(c.legs.len(), 1);
278 assert_eq!(c.legs[0].model, "totally-unknown");
279 }
280}