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 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 ..Default::default()
84 }
85 );
86 assert!(t.legs("nope").is_none());
87 assert!(t.referenced_providers().contains("openai"));
88 }
89
90 #[test]
91 fn missing_dimensions_is_rejected() {
92 let bad = r#"
93 [embeddings."x"]
94 legs = [{ provider = "vertex", model = "text-embedding-004" }]
95 "#;
96 assert!(EmbeddingRouteTable::from_toml_str(bad).is_err());
97 }
98
99 #[test]
100 fn empty_legs_is_rejected() {
101 let bad = r#"
102 [embeddings."x"]
103 dimensions = 768
104 legs = []
105 "#;
106 assert!(EmbeddingRouteTable::from_toml_str(bad).is_err());
107 }
108}