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    pub fn referenced_providers(&self) -> HashSet<String> {
51        self.routes
52            .values()
53            .flat_map(|e| &e.legs)
54            .map(|l| l.provider.clone())
55            .collect()
56    }
57}
58
59#[cfg(test)]
60mod tests {
61    use super::*;
62
63    const SAMPLE: &str = r#"
64        [embeddings."default-embed"]
65        dimensions = 768
66        legs = [
67          { provider = "vertex", model = "text-embedding-004" },
68          { provider = "openai", model = "text-embedding-3-small" },
69        ]
70    "#;
71
72    #[test]
73    fn parses_dimensions_and_legs() {
74        let t = EmbeddingRouteTable::from_toml_str(SAMPLE).unwrap();
75        assert_eq!(t.dimensions("default-embed"), Some(768));
76        let legs = t.legs("default-embed").unwrap();
77        assert_eq!(legs.len(), 2);
78        assert_eq!(
79            legs[0],
80            ChainLeg {
81                provider: "vertex".into(),
82                model: "text-embedding-004".into(),
83            }
84        );
85        assert!(t.legs("nope").is_none());
86        assert!(t.referenced_providers().contains("openai"));
87    }
88
89    #[test]
90    fn missing_dimensions_is_rejected() {
91        let bad = r#"
92            [embeddings."x"]
93            legs = [{ provider = "vertex", model = "text-embedding-004" }]
94        "#;
95        assert!(EmbeddingRouteTable::from_toml_str(bad).is_err());
96    }
97
98    #[test]
99    fn empty_legs_is_rejected() {
100        let bad = r#"
101            [embeddings."x"]
102            dimensions = 768
103            legs = []
104        "#;
105        assert!(EmbeddingRouteTable::from_toml_str(bad).is_err());
106    }
107}