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;
15
16#[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 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 pub fn from_map(providers: HashMap<String, Arc<Provider>>) -> Self {
100 Self { providers }
101 }
102
103 #[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}