1use serde::Deserialize;
4use std::collections::HashMap;
5use tap::Pipe;
6
7#[derive(Debug, Clone, Default, Deserialize, PartialEq, Eq)]
8pub struct ChainLeg {
9 pub provider: String,
10 pub model: String,
11 #[serde(default)]
16 pub region: Option<String>,
17}
18
19#[derive(Debug, Clone, Deserialize)]
20struct RouteEntry {
21 legs: Vec<ChainLeg>,
22}
23
24#[derive(Debug, Clone, Deserialize)]
25struct RoutesFile {
26 routes: HashMap<String, RouteEntry>,
27}
28
29#[derive(Debug, Clone)]
30pub struct RouteTable {
31 routes: HashMap<String, Vec<ChainLeg>>,
32}
33
34impl RouteTable {
35 pub fn from_toml_str(s: &str) -> anyhow::Result<Self> {
36 toml::from_str::<RoutesFile>(s)?
37 .routes
38 .into_iter()
39 .map(|(name, entry)| (name, entry.legs))
40 .collect::<HashMap<_, _>>()
41 .pipe(|routes| Self { routes })
42 .pipe(Ok)
43 }
44
45 pub fn legs(&self, model: &str) -> Option<&[ChainLeg]> {
47 self.routes.get(model).map(Vec::as_slice)
48 }
49
50 pub fn aliases(&self) -> Vec<String> {
52 let mut v: Vec<String> = self.routes.keys().cloned().collect();
53 v.sort();
54 v
55 }
56
57 pub fn referenced_providers(&self) -> std::collections::HashSet<String> {
59 self.routes
60 .values()
61 .flatten()
62 .map(|l| l.provider.clone())
63 .collect()
64 }
65}
66
67#[cfg(test)]
68mod tests {
69 use super::*;
70
71 const SAMPLE: &str = r#"
72 [routes."gemini-pro"]
73 legs = [
74 { provider = "vertex", model = "gemini-3-pro" },
75 { provider = "qwen", model = "qwen-max" },
76 ]
77 [routes."fast"]
78 legs = [{ provider = "vertex", model = "gemini-3-flash" }]
79 "#;
80
81 #[test]
82 fn parses_and_resolves_legs_in_order() {
83 let t = RouteTable::from_toml_str(SAMPLE).unwrap();
84 let legs = t.legs("gemini-pro").unwrap();
85 assert_eq!(legs.len(), 2);
86 assert_eq!(
87 legs[0],
88 ChainLeg {
89 provider: "vertex".into(),
90 model: "gemini-3-pro".into(),
91 ..Default::default()
92 }
93 );
94 assert_eq!(legs[1].provider, "qwen");
95 assert!(t.legs("nope").is_none());
96 }
97
98 #[test]
99 fn parses_optional_per_leg_region() {
100 let t = RouteTable::from_toml_str(
101 r#"
102 [routes."visual"]
103 legs = [
104 { provider = "vertex", model = "gemini-3.1-pro-preview", region = "global" },
105 { provider = "qwen", model = "qwen3-vl-plus" },
106 ]
107 "#,
108 )
109 .unwrap();
110 let legs = t.legs("visual").unwrap();
111 assert_eq!(legs[0].region.as_deref(), Some("global"));
112 assert_eq!(legs[1].region, None);
113 }
114
115 #[test]
116 fn aliases_are_sorted() {
117 let t = RouteTable::from_toml_str(SAMPLE).unwrap();
118 assert_eq!(
119 t.aliases(),
120 vec!["fast".to_string(), "gemini-pro".to_string()]
121 );
122 }
123
124 #[test]
125 fn referenced_providers_collected() {
126 let t = RouteTable::from_toml_str(SAMPLE).unwrap();
127 let p = t.referenced_providers();
128 assert!(p.contains("vertex"));
129 assert!(p.contains("qwen"));
130 }
131}