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