Skip to main content

llm_unified/
factory.rs

1//! Provider factory.
2//!
3//! Creates `LlmProvider` instances based on `LlmConfig` and `ModelRegistry`.
4
5use std::sync::Arc;
6
7use llm_trait::{LlmConfig, LlmError, LlmProvider, Protocol};
8
9use crate::generic::{GenericProvider, ProfiledProvider};
10use crate::model_registry::{ModelProfile, ModelRegistry};
11use crate::protocol::anthropic::AnthropicProtocol;
12use crate::protocol::openai::OpenAiProtocol;
13
14/// Global registry (lazy-initialized).
15static MODEL_REGISTRY: std::sync::LazyLock<ModelRegistry> =
16    std::sync::LazyLock::new(ModelRegistry::builtin);
17
18/// Create a provider from configuration.
19///
20/// Routing:
21/// 1. Query ModelRegistry for profile (protocol + capabilities + reasoning mode)
22/// 2. Build protocol implementation with profile injected
23/// 3. Wrap in ProfiledProvider for correct info() and capabilities()
24pub fn create_provider(config: &LlmConfig) -> Result<Arc<dyn LlmProvider>, LlmError> {
25    if config.api_key.is_empty() {
26        return Err(LlmError::config("LLM_API_KEY is required"));
27    }
28    if config.model.is_empty() {
29        return Err(LlmError::config("LLM_MODEL is required"));
30    }
31    if config.base_url.is_empty() {
32        return Err(LlmError::config("LLM_BASE_URL is required"));
33    }
34
35    // 1. Validate explicit protocol
36    if let Some(protocol) = config.protocol
37        && !matches!(protocol, Protocol::OpenAi | Protocol::Anthropic)
38    {
39        return Err(LlmError::config(format!(
40            "Unsupported protocol '{}'. Use 'openai' or 'anthropic'.",
41            protocol.as_str()
42        )));
43    }
44
45    // 2. Query registry
46    let profile = MODEL_REGISTRY.lookup(&config.model, Some(&config.base_url), config.protocol);
47
48    tracing::info!(
49        model = %config.model,
50        base_url = %config.base_url,
51        protocol = ?profile.protocol,
52        provider = %profile.provider_name,
53        reasoning_mode = ?profile.reasoning_mode,
54        max_output_tokens = ?profile.capabilities.max_output_tokens,
55        "provider created from registry profile"
56    );
57
58    // 2. Build protocol implementation
59    let protocol_impl = build_protocol(&profile, config)?;
60
61    // 3. Assemble provider
62    let provider = ProfiledProvider::new(GenericProvider::new(protocol_impl), profile);
63    Ok(Arc::new(provider))
64}
65
66/// Build protocol implementation based on profile.
67fn build_protocol(
68    profile: &ModelProfile,
69    config: &LlmConfig,
70) -> Result<Box<dyn llm_trait::RawAdapter>, LlmError> {
71    match profile.protocol {
72        Protocol::OpenAi => Ok(Box::new(
73            OpenAiProtocol::from_config(config).with_model_profile(profile.clone()),
74        )),
75        Protocol::Anthropic => Ok(Box::new(
76            AnthropicProtocol::from_config(config).with_model_profile(profile.clone()),
77        )),
78        Protocol::OpenAiResponses => Err(LlmError::config(
79            "OpenAI Responses API is not yet supported. Use protocol 'openai' or 'anthropic'.",
80        )),
81    }
82}
83
84/// Convenience: create a provider with minimal parameters.
85pub fn create(
86    api_key: &str,
87    model: &str,
88    base_url: &str,
89) -> Result<Arc<dyn LlmProvider>, LlmError> {
90    let config = LlmConfig {
91        protocol: None,
92        api_key: api_key.to_string(),
93        model: model.to_string(),
94        base_url: base_url.to_string(),
95        options: std::collections::HashMap::new(),
96    };
97    create_provider(&config)
98}
99
100/// Create a provider from environment variables.
101pub fn from_env() -> Result<Arc<dyn LlmProvider>, LlmError> {
102    let config = LlmConfig::from_env()?;
103    create_provider(&config)
104}
105
106#[cfg(test)]
107mod tests {
108    use super::*;
109    use llm_trait::ReasoningMode;
110
111    #[test]
112    fn create_provider_openai() {
113        let config = LlmConfig {
114            protocol: None,
115            api_key: "sk-test".to_string(),
116            model: "gpt-4o".to_string(),
117            base_url: "https://api.openai.com/v1".to_string(),
118            options: Default::default(),
119        };
120        let provider = create_provider(&config).unwrap();
121        let info = provider.info();
122        assert_eq!(info.name, "openai");
123        assert_eq!(info.model, "gpt-4o");
124    }
125
126    #[test]
127    fn create_provider_anthropic() {
128        let config = LlmConfig {
129            protocol: None,
130            api_key: "sk-test".to_string(),
131            model: "claude-sonnet".to_string(),
132            base_url: "https://api.anthropic.com".to_string(),
133            options: Default::default(),
134        };
135        let provider = create_provider(&config).unwrap();
136        let info = provider.info();
137        assert_eq!(info.name, "anthropic");
138        assert_eq!(info.model, "claude-sonnet");
139    }
140
141    #[test]
142    fn create_provider_explicit_protocol() {
143        let config = LlmConfig {
144            protocol: Some(Protocol::Anthropic),
145            api_key: "sk-test".to_string(),
146            model: "test-model".to_string(),
147            base_url: "https://custom.api.com".to_string(),
148            options: Default::default(),
149        };
150        let provider = create_provider(&config).unwrap();
151        let info = provider.info();
152        assert_eq!(info.name, "anthropic");
153    }
154
155    #[test]
156    fn create_provider_deepseek() {
157        let config = LlmConfig {
158            protocol: None,
159            api_key: "sk-test".to_string(),
160            model: "deepseek-chat".to_string(),
161            base_url: "https://api.deepseek.com/v1".to_string(),
162            options: Default::default(),
163        };
164        let provider = create_provider(&config).unwrap();
165        let info = provider.info();
166        assert_eq!(info.name, "deepseek");
167    }
168
169    #[test]
170    fn create_provider_mimo_no_reasoning() {
171        let config = LlmConfig {
172            protocol: None,
173            api_key: "tp-test".to_string(),
174            model: "mimo-v2.5-pro".to_string(),
175            base_url: "https://api.example-mimo.com/v1".to_string(),
176            options: Default::default(),
177        };
178        let provider = create_provider(&config).unwrap();
179        let info = provider.info();
180        assert_eq!(info.name, "mimo");
181
182        // Verify MiMo's reasoning_mode is None
183        let profile = MODEL_REGISTRY.lookup(
184            "mimo-v2.5-pro",
185            Some("https://api.example-mimo.com/v1"),
186            None,
187        );
188        assert_eq!(profile.reasoning_mode, ReasoningMode::None);
189    }
190
191    #[test]
192    fn create_provider_qwen() {
193        let config = LlmConfig {
194            protocol: None,
195            api_key: "sk-test".to_string(),
196            model: "qwen-plus".to_string(),
197            base_url: "https://dashscope.aliyuncs.com/compatible-mode/v1".to_string(),
198            options: Default::default(),
199        };
200        let provider = create_provider(&config).unwrap();
201        let info = provider.info();
202        assert_eq!(info.name, "qwen");
203    }
204
205    #[test]
206    fn create_provider_openai_responses_returns_error() {
207        // Previously silently downgraded to OpenAI.
208        // Now returns a clear error for unsupported protocols.
209        let config = LlmConfig {
210            protocol: Some(Protocol::OpenAiResponses),
211            api_key: "sk-test".to_string(),
212            model: "gpt-4o".to_string(),
213            base_url: "https://api.openai.com/v1".to_string(),
214            options: Default::default(),
215        };
216        match create_provider(&config) {
217            Ok(_) => panic!("Expected error for OpenAiResponses protocol"),
218            Err(e) => assert!(
219                e.to_string().contains("Unsupported protocol"),
220                "Expected clear error about unsupported protocol, got: {}",
221                e
222            ),
223        }
224    }
225
226    #[test]
227    fn create_convenience() {
228        let provider = create("sk-test", "gpt-4o", "https://api.openai.com/v1").unwrap();
229        let info = provider.info();
230        assert_eq!(info.name, "openai");
231    }
232
233    #[test]
234    fn create_provider_options_max_tokens_ignored() {
235        // LlmConfig.options["max_tokens"] is never used by the factory.
236        // OpenAiProtocol::from_config reads it, but build_protocol calls ::new() instead.
237        use std::collections::HashMap;
238        let mut options = HashMap::new();
239        options.insert("max_tokens".to_string(), serde_json::json!(42));
240        let config = LlmConfig {
241            protocol: None,
242            api_key: "sk-test".to_string(),
243            model: "gpt-4o".to_string(),
244            base_url: "https://api.openai.com/v1".to_string(),
245            options,
246        };
247        let provider = create_provider(&config).unwrap();
248
249        // The provider was created, but options["max_tokens"] was ignored.
250        // We can't directly check max_tokens from the provider, but we can verify
251        // the provider was created successfully (no error from invalid max_tokens).
252        assert_eq!(provider.info().name, "openai");
253        // options["max_tokens"]=42 was silently ignored.
254    }
255}