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;
5use tap::Pipe;
6
7#[derive(Debug, Clone, Default, Deserialize, PartialEq, Eq)]
8pub struct ChainLeg {
9    pub provider: String,
10    pub model: String,
11    /// Optional per-leg region override for the native Vertex lane. When unset,
12    /// the lane falls back to the provider's configured region (env
13    /// `VERTEX_LOCATION`). Lets a route pin a model to the region that serves it
14    /// (e.g. `global` for Gemini 3 previews) without a process-wide env change.
15    #[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    /// Ordered legs for a model alias, or `None` if the alias is unknown.
46    pub fn legs(&self, model: &str) -> Option<&[ChainLeg]> {
47        self.routes.get(model).map(Vec::as_slice)
48    }
49
50    /// All registered aliases (for `/v1/models`), sorted for stable output.
51    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    /// Provider ids referenced by any leg (for fail-fast credential validation).
58    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}