Skip to main content

synapse/providers/
mod.rs

1//! Provider catalog: genai clients + circuit breakers, keyed by provider id.
2pub mod genai_provider;
3pub mod vertex_auth;
4
5use std::collections::HashMap;
6use std::sync::Arc;
7use std::time::Duration;
8
9use crate::config::vertex_project_from_env;
10use crate::providers::genai_provider::{
11    build_openai_compat_provider, build_vertex_provider, OpenAiCompatConfig, Provider,
12    VertexProviderConfig,
13};
14use crate::providers::vertex_auth::VertexAuth;
15
16/// Built provider clients keyed by provider id.
17#[derive(Debug)]
18pub struct Catalog {
19    providers: HashMap<String, Arc<Provider>>,
20}
21
22impl Catalog {
23    pub fn get(&self, id: &str) -> Option<&Arc<Provider>> {
24        self.providers.get(id)
25    }
26
27    /// Build every provider referenced by `referenced`, validating credentials
28    /// fail-fast. Recognised ids: `vertex`, `qwen`, `openai`, `oai_compat`.
29    pub fn build(
30        env: &HashMap<String, String>,
31        referenced: &std::collections::HashSet<String>,
32        request_timeout: Duration,
33    ) -> anyhow::Result<Self> {
34        let get = |k: &str| env.get(k).cloned().filter(|s| !s.trim().is_empty());
35        let mut providers: HashMap<String, Arc<Provider>> = HashMap::new();
36
37        for id in referenced {
38            let provider = match id.as_str() {
39                "vertex" => {
40                    let project = vertex_project_from_env(env).ok_or_else(|| {
41                        anyhow::anyhow!(
42                            "route references provider 'vertex' but VERTEX_PROJECT_ID and VERTEX_PROJECT are unset"
43                        )
44                    })?;
45                    build_vertex_provider(
46                        "vertex",
47                        VertexProviderConfig {
48                            project,
49                            region: "global".into(),
50                            request_timeout,
51                            endpoint_override: None,
52                        },
53                        Arc::new(VertexAuth::from_adc()),
54                    )?
55                }
56                "qwen" => build_openai_compat_provider(
57                    "qwen",
58                    OpenAiCompatConfig {
59                        base_url: get("DASHSCOPE_BASE_URL").unwrap_or_else(|| {
60                            "https://dashscope-intl.aliyuncs.com/compatible-mode/v1".into()
61                        }),
62                        api_key: get("DASHSCOPE_API_KEY").ok_or_else(|| {
63                            anyhow::anyhow!("route references provider 'qwen' but DASHSCOPE_API_KEY is unset")
64                        })?,
65                        request_timeout,
66                        endpoint_override: None,
67                    },
68                )?,
69                "openai" => build_openai_compat_provider(
70                    "openai",
71                    OpenAiCompatConfig {
72                        base_url: get("OPENAI_BASE_URL").unwrap_or_else(|| "https://api.openai.com/v1".into()),
73                        api_key: get("OPENAI_API_KEY").ok_or_else(|| {
74                            anyhow::anyhow!("route references provider 'openai' but OPENAI_API_KEY is unset")
75                        })?,
76                        request_timeout,
77                        endpoint_override: None,
78                    },
79                )?,
80                "oai_compat" => build_openai_compat_provider(
81                    "oai_compat",
82                    OpenAiCompatConfig {
83                        base_url: get("OAI_COMPAT_BASE_URL").ok_or_else(|| {
84                            anyhow::anyhow!("route references provider 'oai_compat' but OAI_COMPAT_BASE_URL is unset")
85                        })?,
86                        api_key: get("OAI_COMPAT_API_KEY").unwrap_or_else(|| "not-needed".into()),
87                        request_timeout,
88                        endpoint_override: None,
89                    },
90                )?,
91                other => anyhow::bail!("unknown provider id in route table: '{other}'"),
92            };
93            providers.insert(id.clone(), Arc::new(provider));
94        }
95        Ok(Self { providers })
96    }
97
98    /// Construct a catalog from pre-built providers (tests and embedders).
99    pub fn from_map(providers: HashMap<String, Arc<Provider>>) -> Self {
100        Self { providers }
101    }
102
103    /// Test-only catalog of OpenAI-compatible providers from (id, base_url) pairs.
104    #[cfg(test)]
105    pub fn for_test(pairs: Vec<(&'static str, String)>) -> Self {
106        let mut providers: HashMap<String, Arc<Provider>> = HashMap::new();
107        for (id, base_url) in pairs {
108            let p = build_openai_compat_provider(
109                id,
110                OpenAiCompatConfig {
111                    base_url,
112                    api_key: "k".into(),
113                    request_timeout: Duration::from_secs(5),
114                    endpoint_override: None,
115                },
116            )
117            .unwrap();
118            providers.insert(id.to_string(), Arc::new(p));
119        }
120        Self { providers }
121    }
122}
123
124#[cfg(test)]
125mod catalog_tests {
126    use super::*;
127
128    fn env(pairs: &[(&str, &str)]) -> HashMap<String, String> {
129        pairs
130            .iter()
131            .map(|(k, v)| (k.to_string(), v.to_string()))
132            .collect()
133    }
134    fn refs(ids: &[&str]) -> std::collections::HashSet<String> {
135        ids.iter().map(|s| s.to_string()).collect()
136    }
137
138    #[test]
139    fn builds_vertex_when_project_id_present() {
140        let cat = Catalog::build(
141            &env(&[("VERTEX_PROJECT_ID", "my-gcp-project")]),
142            &refs(&["vertex"]),
143            Duration::from_secs(5),
144        )
145        .unwrap();
146        assert!(cat.get("vertex").is_some());
147    }
148
149    #[test]
150    fn builds_vertex_when_legacy_project_present() {
151        let cat = Catalog::build(
152            &env(&[("VERTEX_PROJECT", "my-gcp-project")]),
153            &refs(&["vertex"]),
154            Duration::from_secs(5),
155        )
156        .unwrap();
157        assert!(cat.get("vertex").is_some());
158    }
159
160    #[test]
161    fn missing_dashscope_key_fails_fast_with_named_error() {
162        let err = Catalog::build(&env(&[]), &refs(&["qwen"]), Duration::from_secs(5)).unwrap_err();
163        let msg = err.to_string();
164        assert!(msg.contains("qwen"), "{msg}");
165        assert!(msg.contains("DASHSCOPE_API_KEY"), "{msg}");
166    }
167
168    #[test]
169    fn builds_qwen_when_key_present() {
170        let cat = Catalog::build(
171            &env(&[("DASHSCOPE_API_KEY", "sk-test")]),
172            &refs(&["qwen"]),
173            Duration::from_secs(5),
174        )
175        .unwrap();
176        assert!(cat.get("qwen").is_some());
177        assert!(cat.get("vertex").is_none());
178    }
179
180    #[test]
181    fn unknown_provider_id_errors() {
182        let err = Catalog::build(&env(&[]), &refs(&["bogus"]), Duration::from_secs(5)).unwrap_err();
183        assert!(err.to_string().contains("unknown provider id"));
184    }
185}