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