1use crate::Result;
2use crate::catalog::LlmModel;
3#[cfg(feature = "bedrock")]
4use crate::providers::bedrock::BedrockProvider;
5#[cfg(feature = "codex")]
6use crate::providers::codex::CodexProvider;
7use crate::providers::{
8 anthropic::AnthropicProvider,
9 gemini::GeminiProvider,
10 local::{llama_cpp::LlamaCppProvider, ollama::OllamaProvider},
11 openai::OpenAiProvider,
12 openai_compatible::generic::{self, GenericOpenAiProvider},
13 openrouter::OpenRouterProvider,
14};
15use crate::{
16 LlmError, ProviderConnectionConfig, ProviderConnectionOverrides, ProviderFactory, StreamingModelProvider,
17 alloyed::AlloyedModelProvider,
18};
19#[cfg(feature = "codex")]
20use aether_auth::OAuthCredentialStorage;
21use futures::future::BoxFuture;
22use std::collections::{HashMap, HashSet};
23#[cfg(feature = "codex")]
24use std::sync::Arc;
25
26#[doc = include_str!("docs/parser.md")]
27pub struct ModelProviderParser {
28 factories: HashMap<String, CreateProviderFn>,
29 provider_connections: ProviderConnectionOverrides,
30}
31
32impl ModelProviderParser {
33 pub fn new(factories: HashMap<String, CreateProviderFn>) -> Self {
35 Self { factories, provider_connections: ProviderConnectionOverrides::default() }
36 }
37}
38
39impl Default for ModelProviderParser {
40 fn default() -> Self {
42 let parser = Self::new(HashMap::new())
43 .with_provider::<AnthropicProvider>("anthropic")
44 .with_provider::<GeminiProvider>("gemini")
45 .with_provider::<OpenRouterProvider>("openrouter")
46 .with_provider::<OllamaProvider>("ollama")
47 .with_provider::<LlamaCppProvider>("llamacpp")
48 .with_provider::<OpenAiProvider>("openai")
49 .with_openai_provider("deepseek", &generic::DEEPSEEK)
50 .with_openai_provider("moonshot", &generic::MOONSHOT)
51 .with_openai_provider("zai", &generic::ZAI)
52 .with_openai_provider("azure-foundry", &generic::AZURE_FOUNDRY)
53 .with_openai_provider("fireworks", &generic::FIREWORKS);
54
55 #[cfg(feature = "bedrock")]
56 let parser = parser.with_provider::<BedrockProvider>("bedrock");
57
58 parser
59 }
60}
61
62impl ModelProviderParser {
63 pub fn with_provider_connections(mut self, connections: ProviderConnectionOverrides) -> Self {
64 self.provider_connections = connections;
65 self
66 }
67
68 pub fn with_provider<P: ProviderFactory + StreamingModelProvider + 'static>(
69 mut self,
70 name: impl Into<String>,
71 ) -> Self {
72 self.factories.insert(
73 name.into(),
74 Box::new(|model: &str, connection: ProviderConnectionConfig| {
75 let model = model.to_string();
76 Box::pin(
77 async move { Ok(Box::new(P::from_env_with_connection(connection).await?.with_model(&model)) as _) },
78 )
79 }),
80 );
81 self
82 }
83
84 #[cfg(feature = "codex")]
85 pub fn with_codex_provider(mut self, store: Arc<dyn OAuthCredentialStorage>) -> Self {
86 self.factories.insert(
87 "codex".to_string(),
88 Box::new(move |model: &str, connection: ProviderConnectionConfig| {
89 let store = Arc::clone(&store);
90 let model = model.to_string();
91 Box::pin(async move {
92 Ok(Box::new(CodexProvider::new(store).with_connection(connection).with_model(&model)) as _)
93 })
94 }),
95 );
96 self
97 }
98
99 pub fn with_openai_provider(mut self, name: impl Into<String>, config: &'static generic::ProviderConfig) -> Self {
100 self.factories.insert(
101 name.into(),
102 Box::new(move |model: &str, connection: ProviderConnectionConfig| {
103 let model = model.to_string();
104 Box::pin(async move {
105 Ok(
106 Box::new(
107 GenericOpenAiProvider::from_env_with_connection(config, connection)?.with_model(&model),
108 ) as _,
109 )
110 })
111 }),
112 );
113 self
114 }
115
116 pub async fn create_provider(&self, model: &LlmModel) -> Result<Box<dyn StreamingModelProvider>> {
118 let key = model.provider();
119 let factory = self.factories.get(key).ok_or_else(|| LlmError::UnknownProvider { provider: key.to_string() })?;
120 factory(&model.model_id(), self.provider_connections.config_for(key)).await
121 }
122
123 pub async fn parse(&self, models_str: &str) -> Result<(Box<dyn StreamingModelProvider>, LlmModel)> {
134 let provider_model_pairs: Vec<&str> = models_str.split(',').map(str::trim).collect();
135 if provider_model_pairs.is_empty() {
136 return Err(LlmError::EmptyModelSpec);
137 }
138
139 let bedrock_has_inference_profile_arn =
140 self.provider_connections.config_for("bedrock").inference_profile_arn.is_some();
141 let mut seen_bedrock = false;
142 let mut seen_request_model_providers = HashSet::new();
143 let mut providers = Vec::new();
144 let mut first_identity: Option<LlmModel> = None;
145
146 for pair in provider_model_pairs {
147 let (provider_name, model) = pair.split_once(':').unwrap_or((pair, ""));
148
149 if provider_name == "bedrock" && bedrock_has_inference_profile_arn {
150 if seen_bedrock {
151 return Err(LlmError::DuplicateProvider {
152 provider: "bedrock".to_string(),
153 field: "inferenceProfileArn".to_string(),
154 });
155 }
156 seen_bedrock = true;
157 }
158
159 if self.provider_connections.config_for(provider_name).request_model.is_some()
160 && !seen_request_model_providers.insert(provider_name)
161 {
162 return Err(LlmError::DuplicateProvider {
163 provider: provider_name.to_string(),
164 field: "requestModel".to_string(),
165 });
166 }
167
168 let factory = self
169 .factories
170 .get(provider_name)
171 .ok_or_else(|| LlmError::UnknownProvider { provider: provider_name.to_string() })?;
172
173 let connection = self.provider_connections.config_for(provider_name);
174 providers.push(factory(model, connection).await?);
175
176 if first_identity.is_none() {
177 first_identity = Some(pair.parse::<LlmModel>().map_err(LlmError::InvalidModelSpec)?);
178 }
179 }
180
181 let identity = first_identity.ok_or(LlmError::EmptyModelSpec)?;
182
183 let provider: Box<dyn StreamingModelProvider> = if providers.len() == 1 {
184 providers.into_iter().next().ok_or(LlmError::EmptyModelSpec)?
185 } else {
186 Box::new(AlloyedModelProvider::new(providers))
187 };
188
189 Ok((provider, identity))
190 }
191}
192
193pub type CreateProviderFn = Box<
197 dyn Fn(&str, ProviderConnectionConfig) -> BoxFuture<'static, Result<Box<dyn StreamingModelProvider>>> + Send + Sync,
198>;
199
200#[cfg(test)]
201mod tests {
202 use std::collections::BTreeMap;
203
204 use super::*;
205
206 #[tokio::test]
207 async fn test_parse_llamacpp() {
208 let parser = ModelProviderParser::default();
209 let result = parser.parse("llamacpp").await;
210 assert!(result.is_ok());
211 let (_, model) = result.unwrap();
212 assert_eq!(model, LlmModel::LlamaCpp(String::new()));
213 }
214
215 #[tokio::test]
216 async fn test_parse_anthropic() {
217 let parser = ModelProviderParser::default();
218 let result = parser.parse("anthropic:claude-sonnet-4-6").await;
219 match result {
220 Ok((_, model)) => {
221 assert_eq!(model, LlmModel::Anthropic(crate::catalog::AnthropicModel::ClaudeSonnet46));
222 }
223 Err(e) => {
224 let err = e.to_string();
225 assert!(
226 err.contains("API")
227 || err.contains("ANTHROPIC")
228 || err.contains("credentials")
229 || err.contains("JSON"),
230 "Should fail on API key or credentials, not parsing. Got: {err}"
231 );
232 }
233 }
234 }
235
236 #[tokio::test]
237 async fn test_parse_ollama() {
238 let parser = ModelProviderParser::default();
239 let result = parser.parse("ollama:llama3.2").await;
240 assert!(result.is_ok());
241 let (_, model) = result.unwrap();
242 assert_eq!(model, LlmModel::Ollama("llama3.2".to_string()));
243 }
244
245 #[tokio::test]
246 async fn test_parse_openai() {
247 let parser = ModelProviderParser::default();
248 let result = parser.parse("openai:gpt-4.1").await;
249 if let Err(e) = result {
250 let err = e.to_string();
251 assert!(err.contains("API") || err.contains("OPENAI"), "Should fail on API key, not parsing. Got: {err}");
252 }
253 }
254
255 #[tokio::test]
256 async fn test_parse_openrouter() {
257 let parser = ModelProviderParser::default();
258 let result = parser.parse("openrouter:google/gemini-2.5-flash").await;
259 if let Err(e) = result {
260 let err = e.to_string();
261 assert!(err.contains("API") || err.contains("OPENROUTER"), "Should fail on API key, not parsing");
262 }
263 }
264
265 #[tokio::test]
266 async fn test_parse_gemini() {
267 let parser = ModelProviderParser::default();
268 let result = parser.parse("gemini:gemini-2.5-flash").await;
269 if let Err(e) = result {
270 let err = e.to_string();
271 assert!(err.contains("API") || err.contains("GEMINI"), "Should fail on API key, not parsing");
272 }
273 }
274
275 #[tokio::test]
276 async fn test_parse_provider_without_model() {
277 let parser = ModelProviderParser::default();
278 let result = parser.parse("anthropic").await;
279 assert!(result.is_err());
280 }
281
282 #[cfg(feature = "bedrock")]
283 #[tokio::test]
284 async fn test_parse_rejects_bedrock_inference_profile_arn() {
285 let parser = ModelProviderParser::default();
286 let spec = "bedrock:arn:aws:bedrock:us-west-2:000000000000:inference-profile/us.anthropic.claude-opus-4-7";
287
288 let error = match parser.parse(spec).await {
289 Ok(_) => panic!("Bedrock ARN should be rejected"),
290 Err(error) => error.to_string(),
291 };
292
293 assert!(error.contains("providers.bedrock.inferenceProfileArn"), "{error}");
294 }
295
296 #[cfg(feature = "bedrock")]
297 #[tokio::test]
298 async fn test_parse_rejects_bedrock_application_inference_profile_arn() {
299 let parser = ModelProviderParser::default();
300 let spec = "bedrock:arn:aws:bedrock:us-west-2:000000000000:application-inference-profile/000000000000";
301
302 let error = match parser.parse(spec).await {
303 Ok(_) => panic!("Bedrock ARN should be rejected"),
304 Err(error) => error.to_string(),
305 };
306
307 assert!(error.contains("providers.bedrock.inferenceProfileArn"), "{error}");
308 }
309
310 #[tokio::test]
311 async fn test_parse_rejects_repeated_request_model_provider() {
312 let parser = ModelProviderParser::default().with_provider_connections(ProviderConnectionOverrides::new(
313 BTreeMap::from([(
314 "azure-foundry".to_string(),
315 crate::ProviderConnectionOverride {
316 base_url: Some("http://127.0.0.1:8787".to_string()),
317 auth_mode: Some(crate::ProviderAuthMode::None),
318 request_model: Some("production-coding".to_string()),
319 inference_profile_arn: None,
320 },
321 )]),
322 ));
323
324 let Err(error) = parser.parse("azure-foundry:gpt-5.5,azure-foundry:gpt-5.5").await else {
325 panic!("repeated request-model provider should be rejected");
326 };
327 assert!(error.to_string().contains("providers.azure-foundry.requestModel"));
328 }
329
330 #[tokio::test]
331 async fn test_parse_unknown_provider() {
332 let parser = ModelProviderParser::default();
333 let result = parser.parse("unknown:model").await;
334 assert!(result.is_err());
335 if let Err(e) = result {
336 assert!(e.to_string().contains("Unknown provider"));
337 }
338 }
339
340 #[tokio::test]
341 async fn test_with_custom_provider() {
342 let parser = ModelProviderParser::default().with_provider::<OllamaProvider>("custom");
343
344 let model = LlmModel::Ollama("test-model".to_string());
345 let result = parser.create_provider(&model).await;
346 assert!(result.is_ok());
347 }
348
349 #[tokio::test]
350 async fn test_parse_single_provider() {
351 let parser = ModelProviderParser::default();
352 let result = parser.parse("llamacpp").await;
353 assert!(result.is_ok());
354 }
355
356 #[tokio::test]
357 async fn test_parse_multiple_providers() {
358 let parser = ModelProviderParser::default();
359 let result = parser.parse("llamacpp,ollama:llama3.2").await;
360 assert!(result.is_ok());
361 let (_, model) = result.unwrap();
362 assert_eq!(model, LlmModel::LlamaCpp(String::new()));
363 }
364
365 #[tokio::test]
366 async fn test_parse_with_spaces() {
367 let parser = ModelProviderParser::default();
368 let result = parser.parse("llamacpp , ollama:llama3.2").await;
369 assert!(result.is_ok());
370 }
371
372 #[test]
373 fn test_parser_is_send_sync() {
374 fn assert_send_sync<T: Send + Sync>() {}
375 assert_send_sync::<ModelProviderParser>();
376 }
377}