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 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    /// Create a provider from a typed `LlmModel`
118    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    /// Parse a model specification string and create a provider instance.
125    ///
126    /// Returns both the provider and an `LlmModel` describing the identity
127    /// of the first (or only) provider in the spec.
128    ///
129    /// # Format
130    ///
131    /// - `"provider:model"` - Single provider (e.g., "anthropic:claude-3.5-sonnet")
132    /// - `"provider1:model1,provider2:model2"` - Multiple providers create an `AlloyedModelProvider`
133    ///
134    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
194/// Factory function type for creating model providers
195///
196/// Takes a model name and returns a boxed future that resolves to a `StreamingModelProvider`
197pub 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}