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, Deserialize, PartialEq, Eq)]
8pub struct ChainLeg {
9    pub provider: String,
10    pub model: String,
11}
12
13#[derive(Debug, Clone, Deserialize)]
14struct RouteEntry {
15    legs: Vec<ChainLeg>,
16}
17
18#[derive(Debug, Clone, Deserialize)]
19struct RoutesFile {
20    routes: HashMap<String, RouteEntry>,
21}
22
23#[derive(Debug, Clone)]
24pub struct RouteTable {
25    routes: HashMap<String, Vec<ChainLeg>>,
26}
27
28impl RouteTable {
29    pub fn from_toml_str(s: &str) -> anyhow::Result<Self> {
30        toml::from_str::<RoutesFile>(s)?
31            .routes
32            .into_iter()
33            .map(|(name, entry)| (name, entry.legs))
34            .collect::<HashMap<_, _>>()
35            .pipe(|routes| Self { routes })
36            .pipe(Ok)
37    }
38
39    /// Ordered legs for a model alias, or `None` if the alias is unknown.
40    pub fn legs(&self, model: &str) -> Option<&[ChainLeg]> {
41        self.routes.get(model).map(Vec::as_slice)
42    }
43
44    /// All registered aliases (for `/v1/models`), sorted for stable output.
45    pub fn aliases(&self) -> Vec<String> {
46        let mut v: Vec<String> = self.routes.keys().cloned().collect();
47        v.sort();
48        v
49    }
50
51    /// Provider ids referenced by any leg (for fail-fast credential validation).
52    pub fn referenced_providers(&self) -> std::collections::HashSet<String> {
53        self.routes
54            .values()
55            .flatten()
56            .map(|l| l.provider.clone())
57            .collect()
58    }
59}
60
61#[cfg(test)]
62mod tests {
63    use super::*;
64
65    const SAMPLE: &str = r#"
66        [routes."gemini-pro"]
67        legs = [
68          { provider = "vertex", model = "gemini-3-pro" },
69          { provider = "qwen", model = "qwen-max" },
70        ]
71        [routes."fast"]
72        legs = [{ provider = "vertex", model = "gemini-3-flash" }]
73    "#;
74
75    #[test]
76    fn parses_and_resolves_legs_in_order() {
77        let t = RouteTable::from_toml_str(SAMPLE).unwrap();
78        let legs = t.legs("gemini-pro").unwrap();
79        assert_eq!(legs.len(), 2);
80        assert_eq!(
81            legs[0],
82            ChainLeg {
83                provider: "vertex".into(),
84                model: "gemini-3-pro".into()
85            }
86        );
87        assert_eq!(legs[1].provider, "qwen");
88        assert!(t.legs("nope").is_none());
89    }
90
91    #[test]
92    fn aliases_are_sorted() {
93        let t = RouteTable::from_toml_str(SAMPLE).unwrap();
94        assert_eq!(
95            t.aliases(),
96            vec!["fast".to_string(), "gemini-pro".to_string()]
97        );
98    }
99
100    #[test]
101    fn referenced_providers_collected() {
102        let t = RouteTable::from_toml_str(SAMPLE).unwrap();
103        let p = t.referenced_providers();
104        assert!(p.contains("vertex"));
105        assert!(p.contains("qwen"));
106    }
107}