Skip to main content

synapse/routing/
embeddings.rs

1//! Embedding route table: alias → declared output dimension + ordered fallback legs.
2use crate::routing::table::ChainLeg;
3use serde::Deserialize;
4use std::collections::{HashMap, HashSet};
5
6#[derive(Debug, Clone, Deserialize)]
7struct EmbeddingEntry {
8    dimensions: u32,
9    legs: Vec<ChainLeg>,
10}
11
12#[derive(Debug, Clone, Deserialize)]
13struct EmbeddingsFile {
14    #[serde(default)]
15    embeddings: HashMap<String, EmbeddingEntry>,
16}
17
18#[derive(Debug, Clone, Default)]
19pub struct EmbeddingRouteTable {
20    routes: HashMap<String, EmbeddingEntry>,
21}
22
23impl EmbeddingRouteTable {
24    pub fn from_toml_str(s: &str) -> anyhow::Result<Self> {
25        let routes = toml::from_str::<EmbeddingsFile>(s)?.embeddings;
26        for (alias, e) in &routes {
27            anyhow::ensure!(
28                e.dimensions > 0,
29                "embedding alias '{alias}' has dimensions = 0"
30            );
31            anyhow::ensure!(!e.legs.is_empty(), "embedding alias '{alias}' has no legs");
32        }
33        Ok(Self { routes })
34    }
35
36    pub fn legs(&self, alias: &str) -> Option<&[ChainLeg]> {
37        self.routes.get(alias).map(|e| e.legs.as_slice())
38    }
39
40    pub fn dimensions(&self, alias: &str) -> Option<u32> {
41        self.routes.get(alias).map(|e| e.dimensions)
42    }
43
44    pub fn aliases(&self) -> Vec<String> {
45        let mut v: Vec<String> = self.routes.keys().cloned().collect();
46        v.sort();
47        v
48    }
49
50    /// This table with every leg belonging to `drop` removed, and any alias left
51    /// with no legs removed entirely. Mirrors `RouteTable::without_providers`.
52    pub fn without_providers(&self, drop: &HashSet<String>) -> Self {
53        let routes: HashMap<String, EmbeddingEntry> = self
54            .routes
55            .iter()
56            .map(|(alias, entry)| {
57                let legs: Vec<ChainLeg> = entry
58                    .legs
59                    .iter()
60                    .filter(|l| !drop.contains(&l.provider))
61                    .cloned()
62                    .collect();
63                (
64                    alias.clone(),
65                    EmbeddingEntry {
66                        dimensions: entry.dimensions,
67                        legs,
68                    },
69                )
70            })
71            .filter(|(_, entry)| !entry.legs.is_empty())
72            .collect();
73        Self { routes }
74    }
75
76    pub fn referenced_providers(&self) -> HashSet<String> {
77        self.routes
78            .values()
79            .flat_map(|e| &e.legs)
80            .map(|l| l.provider.clone())
81            .collect()
82    }
83}
84
85#[cfg(test)]
86mod tests {
87    #[test]
88    fn without_providers_drops_legs_then_empty_aliases() {
89        let table = EmbeddingRouteTable::from_toml_str(
90            r#"
91            [embeddings."multi"]
92            dimensions = 768
93            legs = [
94              { provider = "vertex", model = "text-embedding-004" },
95              { provider = "openai", model = "text-embedding-3-small" },
96            ]
97            [embeddings."vertex-only"]
98            dimensions = 768
99            legs = [{ provider = "vertex", model = "text-embedding-004" }]
100        "#,
101        )
102        .unwrap()
103        .without_providers(&["vertex".to_string()].into_iter().collect());
104
105        // multi survives on openai, keeping its declared dimensions.
106        assert_eq!(table.legs("multi").unwrap().len(), 1);
107        assert_eq!(table.legs("multi").unwrap()[0].provider, "openai");
108        assert_eq!(table.dimensions("multi"), Some(768));
109        // vertex-only has nothing left to serve it.
110        assert!(table.legs("vertex-only").is_none());
111        assert_eq!(table.aliases(), vec!["multi".to_string()]);
112    }
113
114    use super::*;
115
116    const SAMPLE: &str = r#"
117        [embeddings."default-embed"]
118        dimensions = 768
119        legs = [
120          { provider = "vertex", model = "text-embedding-004" },
121          { provider = "openai", model = "text-embedding-3-small" },
122        ]
123    "#;
124
125    #[test]
126    fn parses_dimensions_and_legs() {
127        let t = EmbeddingRouteTable::from_toml_str(SAMPLE).unwrap();
128        assert_eq!(t.dimensions("default-embed"), Some(768));
129        let legs = t.legs("default-embed").unwrap();
130        assert_eq!(legs.len(), 2);
131        assert_eq!(
132            legs[0],
133            ChainLeg {
134                provider: "vertex".into(),
135                model: "text-embedding-004".into(),
136                ..Default::default()
137            }
138        );
139        assert!(t.legs("nope").is_none());
140        assert!(t.referenced_providers().contains("openai"));
141    }
142
143    #[test]
144    fn missing_dimensions_is_rejected() {
145        let bad = r#"
146            [embeddings."x"]
147            legs = [{ provider = "vertex", model = "text-embedding-004" }]
148        "#;
149        assert!(EmbeddingRouteTable::from_toml_str(bad).is_err());
150    }
151
152    #[test]
153    fn empty_legs_is_rejected() {
154        let bad = r#"
155            [embeddings."x"]
156            dimensions = 768
157            legs = []
158        "#;
159        assert!(EmbeddingRouteTable::from_toml_str(bad).is_err());
160    }
161}