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