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::Other(format!("Unknown provider: {key}")))?;
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::Other("No models provided".to_string()));
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::Other(
152 "providers.bedrock.inferenceProfileArn cannot be used with multiple bedrock models in one alloy spec"
153 .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::Other(format!(
163 "providers.{provider_name}.requestModel cannot be used with multiple {provider_name} models in one alloy spec"
164 )));
165 }
166
167 let factory = self
168 .factories
169 .get(provider_name)
170 .ok_or_else(|| LlmError::Other(format!("Unknown provider: {provider_name}")))?;
171
172 let connection = self.provider_connections.config_for(provider_name);
173 providers.push(factory(model, connection).await?);
174
175 if first_identity.is_none() {
176 first_identity = Some(pair.parse::<LlmModel>().map_err(LlmError::Other)?);
177 }
178 }
179
180 let identity = first_identity.ok_or_else(|| LlmError::Other("No providers parsed".to_string()))?;
181
182 let provider: Box<dyn StreamingModelProvider> = if providers.len() == 1 {
183 providers.into_iter().next().ok_or_else(|| LlmError::Other("No providers available".to_string()))?
184 } else {
185 Box::new(AlloyedModelProvider::new(providers))
186 };
187
188 Ok((provider, identity))
189 }
190}
191
192pub type CreateProviderFn = Box<
196 dyn Fn(&str, ProviderConnectionConfig) -> BoxFuture<'static, Result<Box<dyn StreamingModelProvider>>> + Send + Sync,
197>;
198
199#[cfg(test)]
200mod tests {
201 use std::collections::BTreeMap;
202
203 use super::*;
204
205 #[tokio::test]
206 async fn test_parse_llamacpp() {
207 let parser = ModelProviderParser::default();
208 let result = parser.parse("llamacpp").await;
209 assert!(result.is_ok());
210 let (_, model) = result.unwrap();
211 assert_eq!(model, LlmModel::LlamaCpp(String::new()));
212 }
213
214 #[tokio::test]
215 async fn test_parse_anthropic() {
216 let parser = ModelProviderParser::default();
217 let result = parser.parse("anthropic:claude-sonnet-4-6").await;
218 match result {
219 Ok((_, model)) => {
220 assert_eq!(model, LlmModel::Anthropic(crate::catalog::AnthropicModel::ClaudeSonnet46));
221 }
222 Err(e) => {
223 let err = e.to_string();
224 assert!(
225 err.contains("API")
226 || err.contains("ANTHROPIC")
227 || err.contains("credentials")
228 || err.contains("JSON"),
229 "Should fail on API key or credentials, not parsing. Got: {err}"
230 );
231 }
232 }
233 }
234
235 #[tokio::test]
236 async fn test_parse_ollama() {
237 let parser = ModelProviderParser::default();
238 let result = parser.parse("ollama:llama3.2").await;
239 assert!(result.is_ok());
240 let (_, model) = result.unwrap();
241 assert_eq!(model, LlmModel::Ollama("llama3.2".to_string()));
242 }
243
244 #[tokio::test]
245 async fn test_parse_openai() {
246 let parser = ModelProviderParser::default();
247 let result = parser.parse("openai:gpt-4.1").await;
248 if let Err(e) = result {
249 let err = e.to_string();
250 assert!(err.contains("API") || err.contains("OPENAI"), "Should fail on API key, not parsing. Got: {err}");
251 }
252 }
253
254 #[tokio::test]
255 async fn test_parse_openrouter() {
256 let parser = ModelProviderParser::default();
257 let result = parser.parse("openrouter:google/gemini-2.5-flash").await;
258 if let Err(e) = result {
259 let err = e.to_string();
260 assert!(err.contains("API") || err.contains("OPENROUTER"), "Should fail on API key, not parsing");
261 }
262 }
263
264 #[tokio::test]
265 async fn test_parse_gemini() {
266 let parser = ModelProviderParser::default();
267 let result = parser.parse("gemini:gemini-2.5-flash").await;
268 if let Err(e) = result {
269 let err = e.to_string();
270 assert!(err.contains("API") || err.contains("GEMINI"), "Should fail on API key, not parsing");
271 }
272 }
273
274 #[tokio::test]
275 async fn test_parse_provider_without_model() {
276 let parser = ModelProviderParser::default();
277 let result = parser.parse("anthropic").await;
278 assert!(result.is_err());
279 }
280
281 #[cfg(feature = "bedrock")]
282 #[tokio::test]
283 async fn test_parse_rejects_bedrock_inference_profile_arn() {
284 let parser = ModelProviderParser::default();
285 let spec = "bedrock:arn:aws:bedrock:us-west-2:000000000000:inference-profile/us.anthropic.claude-opus-4-7";
286
287 let error = match parser.parse(spec).await {
288 Ok(_) => panic!("Bedrock ARN should be rejected"),
289 Err(error) => error.to_string(),
290 };
291
292 assert!(error.contains("providers.bedrock.inferenceProfileArn"), "{error}");
293 }
294
295 #[cfg(feature = "bedrock")]
296 #[tokio::test]
297 async fn test_parse_rejects_bedrock_application_inference_profile_arn() {
298 let parser = ModelProviderParser::default();
299 let spec = "bedrock:arn:aws:bedrock:us-west-2:000000000000:application-inference-profile/000000000000";
300
301 let error = match parser.parse(spec).await {
302 Ok(_) => panic!("Bedrock ARN should be rejected"),
303 Err(error) => error.to_string(),
304 };
305
306 assert!(error.contains("providers.bedrock.inferenceProfileArn"), "{error}");
307 }
308
309 #[tokio::test]
310 async fn test_parse_rejects_repeated_request_model_provider() {
311 let parser = ModelProviderParser::default().with_provider_connections(ProviderConnectionOverrides::new(
312 BTreeMap::from([(
313 "azure-foundry".to_string(),
314 crate::ProviderConnectionOverride {
315 base_url: Some("http://127.0.0.1:8787".to_string()),
316 auth_mode: Some(crate::ProviderAuthMode::None),
317 request_model: Some("production-coding".to_string()),
318 inference_profile_arn: None,
319 },
320 )]),
321 ));
322
323 let Err(error) = parser.parse("azure-foundry:gpt-5.5,azure-foundry:gpt-5.5").await else {
324 panic!("repeated request-model provider should be rejected");
325 };
326 assert!(error.to_string().contains("providers.azure-foundry.requestModel"));
327 }
328
329 #[tokio::test]
330 async fn test_parse_unknown_provider() {
331 let parser = ModelProviderParser::default();
332 let result = parser.parse("unknown:model").await;
333 assert!(result.is_err());
334 if let Err(e) = result {
335 assert!(e.to_string().contains("Unknown provider"));
336 }
337 }
338
339 #[tokio::test]
340 async fn test_with_custom_provider() {
341 let parser = ModelProviderParser::default().with_provider::<OllamaProvider>("custom");
342
343 let model = LlmModel::Ollama("test-model".to_string());
344 let result = parser.create_provider(&model).await;
345 assert!(result.is_ok());
346 }
347
348 #[tokio::test]
349 async fn test_parse_single_provider() {
350 let parser = ModelProviderParser::default();
351 let result = parser.parse("llamacpp").await;
352 assert!(result.is_ok());
353 }
354
355 #[tokio::test]
356 async fn test_parse_multiple_providers() {
357 let parser = ModelProviderParser::default();
358 let result = parser.parse("llamacpp,ollama:llama3.2").await;
359 assert!(result.is_ok());
360 let (_, model) = result.unwrap();
361 assert_eq!(model, LlmModel::LlamaCpp(String::new()));
362 }
363
364 #[tokio::test]
365 async fn test_parse_with_spaces() {
366 let parser = ModelProviderParser::default();
367 let result = parser.parse("llamacpp , ollama:llama3.2").await;
368 assert!(result.is_ok());
369 }
370
371 #[test]
372 fn test_parser_is_send_sync() {
373 fn assert_send_sync<T: Send + Sync>() {}
374 assert_send_sync::<ModelProviderParser>();
375 }
376}