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