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
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}