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
77#[cfg(test)]
78mod tests {
79 use super::*;
80
81 const SAMPLE: &str = r#"
82 [routes."gemini-pro"]
83 legs = [
84 { provider = "vertex", model = "gemini-3-pro" },
85 { provider = "qwen", model = "qwen-max" },
86 ]
87 [routes."fast"]
88 legs = [{ provider = "vertex", model = "gemini-3-flash" }]
89 "#;
90
91 #[test]
92 fn parses_and_resolves_legs_in_order() {
93 let t = RouteTable::from_toml_str(SAMPLE).unwrap();
94 let legs = t.legs("gemini-pro").unwrap();
95 assert_eq!(legs.len(), 2);
96 assert_eq!(
97 legs[0],
98 ChainLeg {
99 provider: "vertex".into(),
100 model: "gemini-3-pro".into(),
101 ..Default::default()
102 }
103 );
104 assert_eq!(legs[1].provider, "qwen");
105 assert!(t.legs("nope").is_none());
106 }
107
108 #[test]
109 fn parses_optional_per_leg_region() {
110 let t = RouteTable::from_toml_str(
111 r#"
112 [routes."visual"]
113 legs = [
114 { provider = "vertex", model = "gemini-3.1-pro-preview", region = "global" },
115 { provider = "qwen", model = "qwen3-vl-plus" },
116 ]
117 "#,
118 )
119 .unwrap();
120 let legs = t.legs("visual").unwrap();
121 assert_eq!(legs[0].region.as_deref(), Some("global"));
122 assert_eq!(legs[1].region, None);
123 }
124
125 #[test]
126 fn aliases_are_sorted() {
127 let t = RouteTable::from_toml_str(SAMPLE).unwrap();
128 assert_eq!(
129 t.aliases(),
130 vec!["fast".to_string(), "gemini-pro".to_string()]
131 );
132 }
133
134 #[test]
135 fn referenced_providers_collected() {
136 let t = RouteTable::from_toml_str(SAMPLE).unwrap();
137 let p = t.referenced_providers();
138 assert!(p.contains("vertex"));
139 assert!(p.contains("qwen"));
140 }
141
142 #[test]
143 fn parses_optional_route_policy_and_defaults_to_none() {
144 let t = RouteTable::from_toml_str(
145 r#"
146 [routes."guarded"]
147 policy = "strict"
148 legs = [{ provider = "vertex", model = "gemini-3-pro" }]
149 [routes."plain"]
150 legs = [{ provider = "qwen", model = "qwen-max" }]
151 "#,
152 )
153 .unwrap();
154 assert_eq!(t.policy_of("guarded"), Some("strict"));
155 assert_eq!(t.policy_of("plain"), None);
156 assert_eq!(t.policy_of("missing"), None);
157 }
158}