1pub 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;
15use crate::telemetry::GatewayMetrics;
16
17#[derive(Debug)]
19pub struct Catalog {
20 providers: HashMap<String, Arc<Provider>>,
21}
22
23#[derive(Debug, Clone, PartialEq, Eq)]
25pub enum Unsatisfiable {
26 MissingCredential {
28 provider: String,
29 required: &'static str,
30 },
31 UnknownProvider { provider: String },
34}
35
36impl Unsatisfiable {
37 pub fn provider(&self) -> &str {
38 match self {
39 Self::MissingCredential { provider, .. } | Self::UnknownProvider { provider } => {
40 provider
41 }
42 }
43 }
44}
45
46impl std::fmt::Display for Unsatisfiable {
47 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
48 match self {
49 Self::MissingCredential { provider, required } => {
50 write!(f, "provider '{provider}' needs {required}, which is unset")
51 }
52 Self::UnknownProvider { provider } => {
53 write!(f, "provider '{provider}' is not recognised by this build")
54 }
55 }
56 }
57}
58
59pub fn unsatisfiable_providers(
62 env: &HashMap<String, String>,
63 referenced: &std::collections::HashSet<String>,
64) -> Vec<Unsatisfiable> {
65 let present = |k: &str| env.get(k).is_some_and(|v| !v.trim().is_empty());
66 let mut out: Vec<Unsatisfiable> = referenced
67 .iter()
68 .filter_map(|id| match id.as_str() {
69 "vertex" => {
71 vertex_project_from_env(env)
72 .is_none()
73 .then(|| Unsatisfiable::MissingCredential {
74 provider: id.clone(),
75 required: "VERTEX_PROJECT_ID",
76 })
77 }
78 "qwen" => (!present("DASHSCOPE_API_KEY")).then(|| Unsatisfiable::MissingCredential {
79 provider: id.clone(),
80 required: "DASHSCOPE_API_KEY",
81 }),
82 "openai" => (!present("OPENAI_API_KEY")).then(|| Unsatisfiable::MissingCredential {
83 provider: id.clone(),
84 required: "OPENAI_API_KEY",
85 }),
86 "typesafe" => {
87 (!present("TYPESAFE_API_KEY")).then(|| Unsatisfiable::MissingCredential {
88 provider: id.clone(),
89 required: "TYPESAFE_API_KEY",
90 })
91 }
92 "oai_compat" => {
93 (!present("OAI_COMPAT_BASE_URL")).then(|| Unsatisfiable::MissingCredential {
94 provider: id.clone(),
95 required: "OAI_COMPAT_BASE_URL",
96 })
97 }
98 _ => Some(Unsatisfiable::UnknownProvider {
99 provider: id.clone(),
100 }),
101 })
102 .collect();
103 out.sort_by(|a, b| a.provider().cmp(b.provider()));
104 out
105}
106
107impl Catalog {
108 pub fn get(&self, id: &str) -> Option<&Arc<Provider>> {
109 self.providers.get(id)
110 }
111
112 pub fn attach_metrics(&self, metrics: &Arc<GatewayMetrics>) {
114 self.providers
115 .values()
116 .for_each(|p| p.breaker.attach_metrics(metrics.clone()));
117 }
118
119 pub fn build(
123 env: &HashMap<String, String>,
124 referenced: &std::collections::HashSet<String>,
125 request_timeout: Duration,
126 ) -> anyhow::Result<Self> {
127 let get = |k: &str| env.get(k).cloned().filter(|s| !s.trim().is_empty());
128 let mut providers: HashMap<String, Arc<Provider>> = HashMap::new();
129
130 for id in referenced {
131 let provider = match id.as_str() {
132 "vertex" => {
133 let project = vertex_project_from_env(env).ok_or_else(|| {
134 anyhow::anyhow!(
135 "route references provider 'vertex' but VERTEX_PROJECT_ID and VERTEX_PROJECT are unset"
136 )
137 })?;
138 build_vertex_provider(
139 "vertex",
140 VertexProviderConfig {
141 project,
142 region: "global".into(),
143 request_timeout,
144 endpoint_override: None,
145 },
146 Arc::new(VertexAuth::from_adc()),
147 )?
148 }
149 "qwen" => build_openai_compat_provider(
150 "qwen",
151 OpenAiCompatConfig {
152 base_url: get("DASHSCOPE_BASE_URL").unwrap_or_else(|| {
153 "https://dashscope-intl.aliyuncs.com/compatible-mode/v1".into()
154 }),
155 api_key: get("DASHSCOPE_API_KEY").ok_or_else(|| {
156 anyhow::anyhow!("route references provider 'qwen' but DASHSCOPE_API_KEY is unset")
157 })?,
158 request_timeout,
159 endpoint_override: None,
160 },
161 )?,
162 "openai" => build_openai_compat_provider(
163 "openai",
164 OpenAiCompatConfig {
165 base_url: get("OPENAI_BASE_URL").unwrap_or_else(|| "https://api.openai.com/v1".into()),
166 api_key: get("OPENAI_API_KEY").ok_or_else(|| {
167 anyhow::anyhow!("route references provider 'openai' but OPENAI_API_KEY is unset")
168 })?,
169 request_timeout,
170 endpoint_override: None,
171 },
172 )?,
173 "typesafe" => {
179 if get("TYPESAFE_API_KEY").is_none() {
180 anyhow::bail!(
181 "route references provider 'typesafe' but TYPESAFE_API_KEY is unset"
182 );
183 }
184 continue;
185 }
186 "oai_compat" => build_openai_compat_provider(
187 "oai_compat",
188 OpenAiCompatConfig {
189 base_url: get("OAI_COMPAT_BASE_URL").ok_or_else(|| {
190 anyhow::anyhow!("route references provider 'oai_compat' but OAI_COMPAT_BASE_URL is unset")
191 })?,
192 api_key: get("OAI_COMPAT_API_KEY").unwrap_or_else(|| "not-needed".into()),
193 request_timeout,
194 endpoint_override: None,
195 },
196 )?,
197 other => anyhow::bail!("unknown provider id in route table: '{other}'"),
198 };
199 providers.insert(id.clone(), Arc::new(provider));
200 }
201 Ok(Self { providers })
202 }
203
204 pub fn from_map(providers: HashMap<String, Arc<Provider>>) -> Self {
206 Self { providers }
207 }
208
209 #[cfg(test)]
211 pub fn for_test(pairs: Vec<(&'static str, String)>) -> Self {
212 let mut providers: HashMap<String, Arc<Provider>> = HashMap::new();
213 for (id, base_url) in pairs {
214 let p = build_openai_compat_provider(
215 id,
216 OpenAiCompatConfig {
217 base_url,
218 api_key: "k".into(),
219 request_timeout: Duration::from_secs(5),
220 endpoint_override: None,
221 },
222 )
223 .unwrap();
224 providers.insert(id.to_string(), Arc::new(p));
225 }
226 Self { providers }
227 }
228}
229
230#[cfg(test)]
231mod catalog_tests {
232 use super::*;
233
234 fn env(pairs: &[(&str, &str)]) -> HashMap<String, String> {
235 pairs
236 .iter()
237 .map(|(k, v)| (k.to_string(), v.to_string()))
238 .collect()
239 }
240 fn refs(ids: &[&str]) -> std::collections::HashSet<String> {
241 ids.iter().map(|s| s.to_string()).collect()
242 }
243
244 #[test]
245 fn unsatisfiable_reports_the_env_var_each_provider_needs() {
246 let missing = unsatisfiable_providers(
247 &env(&[]),
248 &refs(&["vertex", "qwen", "openai", "typesafe", "oai_compat"]),
249 );
250 let pairs: Vec<(&str, &str)> = missing
251 .iter()
252 .map(|u| match u {
253 Unsatisfiable::MissingCredential { provider, required } => {
254 (provider.as_str(), *required)
255 }
256 Unsatisfiable::UnknownProvider { provider } => (provider.as_str(), "unknown"),
257 })
258 .collect();
259 assert_eq!(
260 pairs,
261 vec![
262 ("oai_compat", "OAI_COMPAT_BASE_URL"),
263 ("openai", "OPENAI_API_KEY"),
264 ("qwen", "DASHSCOPE_API_KEY"),
265 ("typesafe", "TYPESAFE_API_KEY"),
266 ("vertex", "VERTEX_PROJECT_ID"),
267 ]
268 );
269 }
270
271 #[test]
272 fn unsatisfiable_is_empty_when_every_credential_is_present() {
273 let satisfied = env(&[
274 ("VERTEX_PROJECT", "legacy-spelling-counts"),
275 ("DASHSCOPE_API_KEY", "k"),
276 ("TYPESAFE_API_KEY", "k"),
277 ]);
278 assert!(
279 unsatisfiable_providers(&satisfied, &refs(&["vertex", "qwen", "typesafe"])).is_empty()
280 );
281 }
282
283 #[test]
284 fn unsatisfiable_flags_an_id_this_build_does_not_know() {
285 assert_eq!(
287 unsatisfiable_providers(&env(&[]), &refs(&["from-the-future"])),
288 vec![Unsatisfiable::UnknownProvider {
289 provider: "from-the-future".into()
290 }]
291 );
292 }
293
294 #[test]
295 fn whitespace_only_credential_counts_as_unset() {
296 let blank = env(&[("TYPESAFE_API_KEY", " ")]);
297 assert_eq!(
298 unsatisfiable_providers(&blank, &refs(&["typesafe"])),
299 vec![Unsatisfiable::MissingCredential {
300 provider: "typesafe".into(),
301 required: "TYPESAFE_API_KEY"
302 }]
303 );
304 }
305
306 #[test]
307 fn builds_vertex_when_project_id_present() {
308 let cat = Catalog::build(
309 &env(&[("VERTEX_PROJECT_ID", "my-gcp-project")]),
310 &refs(&["vertex"]),
311 Duration::from_secs(5),
312 )
313 .unwrap();
314 assert!(cat.get("vertex").is_some());
315 }
316
317 #[test]
318 fn builds_vertex_when_legacy_project_present() {
319 let cat = Catalog::build(
320 &env(&[("VERTEX_PROJECT", "my-gcp-project")]),
321 &refs(&["vertex"]),
322 Duration::from_secs(5),
323 )
324 .unwrap();
325 assert!(cat.get("vertex").is_some());
326 }
327
328 #[test]
329 fn missing_dashscope_key_fails_fast_with_named_error() {
330 let err = Catalog::build(&env(&[]), &refs(&["qwen"]), Duration::from_secs(5)).unwrap_err();
331 let msg = err.to_string();
332 assert!(msg.contains("qwen"), "{msg}");
333 assert!(msg.contains("DASHSCOPE_API_KEY"), "{msg}");
334 }
335
336 #[test]
337 fn builds_qwen_when_key_present() {
338 let cat = Catalog::build(
339 &env(&[("DASHSCOPE_API_KEY", "sk-test")]),
340 &refs(&["qwen"]),
341 Duration::from_secs(5),
342 )
343 .unwrap();
344 assert!(cat.get("qwen").is_some());
345 assert!(cat.get("vertex").is_none());
346 }
347
348 #[test]
349 fn typesafe_without_key_fails_fast_with_named_error() {
350 let err =
351 Catalog::build(&env(&[]), &refs(&["typesafe"]), Duration::from_secs(5)).unwrap_err();
352 let msg = err.to_string();
353 assert!(msg.contains("typesafe"), "{msg}");
354 assert!(msg.contains("TYPESAFE_API_KEY"), "{msg}");
355 }
356
357 #[test]
358 fn typesafe_with_key_is_validated_but_not_built() {
359 let cat = Catalog::build(
360 &env(&[("TYPESAFE_API_KEY", "sk-test")]),
361 &refs(&["typesafe"]),
362 Duration::from_secs(5),
363 )
364 .unwrap();
365 assert!(cat.get("typesafe").is_none());
367 }
368
369 #[test]
370 fn unknown_provider_id_errors() {
371 let err = Catalog::build(&env(&[]), &refs(&["bogus"]), Duration::from_secs(5)).unwrap_err();
372 assert!(err.to_string().contains("unknown provider id"));
373 }
374
375 #[cfg(feature = "server")]
376 #[test]
377 fn attach_metrics_reaches_every_provider_breaker() {
378 let env = HashMap::from([("DASHSCOPE_API_KEY".to_string(), "sk".to_string())]);
379 let referenced = std::collections::HashSet::from(["qwen".to_string()]);
380 let catalog = Catalog::build(&env, &referenced, Duration::from_secs(5)).unwrap();
381 let (m, exporter) = crate::telemetry::test_metrics();
382 catalog.attach_metrics(&m);
383 let breaker = &catalog.get("qwen").unwrap().breaker;
384 (0..5).for_each(|_| breaker.record(&Err::<(), _>("boom")));
385 assert!(crate::telemetry::scrape(&exporter).contains(
386 r#"synapse_resilience_breaker_transitions_total{name="qwen",transition="open"} 1"#
387 ));
388 }
389}