1pub 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#[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 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 #[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}