Skip to main content

synapse/routing/
table.rs

1//! Route table: client-facing model alias → ordered fallback legs.
2
3use 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    /// Optional per-leg region override for the native Vertex lane. When unset,
11    /// the lane falls back to the provider's configured region (env
12    /// `VERTEX_LOCATION`). Lets a route pin a model to the region that serves it
13    /// (e.g. `global` for Gemini 3 previews) without a process-wide env change.
14    #[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    /// Ordered legs for a model alias, or `None` if the alias is unknown.
51    pub fn legs(&self, model: &str) -> Option<&[ChainLeg]> {
52        self.routes.get(model).map(Vec::as_slice)
53    }
54
55    /// Policy name selected by a route alias, or `None` when unset.
56    pub fn policy_of(&self, model: &str) -> Option<&str> {
57        self.policies.get(model).map(String::as_str)
58    }
59
60    /// All registered aliases (for `/v1/models`), sorted for stable output.
61    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    /// Provider ids referenced by any leg (for fail-fast credential validation).
68    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    /// Consecutive Vertex legs for Gemini-native passthrough fallback.
77    ///
78    /// Prefer a route whose first leg is `vertex` + `model` (lexicographically
79    /// first alias if several match). Otherwise start at the first matching
80    /// vertex leg in the lex-first route that contains it. Stop before the first
81    /// non-vertex leg. If nothing matches, return a single synthetic leg so
82    /// callers always have ≥1 attempt (same as today's single forward).
83    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/// Vertex-only fallback chain used by Gemini passthrough.
132#[derive(Debug, Clone, PartialEq, Eq)]
133pub struct VertexPassthroughChain {
134    /// Matched route alias, if any.
135    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}