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::UnknownProvider { provider: key.to_string() })?;
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::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
193/// Factory function type for creating model providers
194///
195/// Takes a model name and returns a boxed future that resolves to a `StreamingModelProvider`
196pub 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}