Skip to main content

llm/
parser.rs

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    /// Create a new parser with custom provider factories
34    pub fn new(factories: HashMap<String, CreateProviderFn>) -> Self {
35        Self { factories, provider_connections: ProviderConnectionOverrides::default() }
36    }
37}
38
39impl Default for ModelProviderParser {
40    /// Create a parser with all built-in providers registered
41    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    /// Create a provider from a typed `LlmModel`
117    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    /// Parse a model specification string and create a provider instance.
124    ///
125    /// Returns both the provider and an `LlmModel` describing the identity
126    /// of the first (or only) provider in the spec.
127    ///
128    /// # Format
129    ///
130    /// - `"provider:model"` - Single provider (e.g., "anthropic:claude-3.5-sonnet")
131    /// - `"provider1:model1,provider2:model2"` - Multiple providers create an `AlloyedModelProvider`
132    ///
133    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
192/// Factory function type for creating model providers
193///
194/// Takes a model name and returns a boxed future that resolves to a `StreamingModelProvider`
195pub 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}