synapse/routing/
embeddings.rs1use 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 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 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 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}