1use 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
14static MODEL_REGISTRY: std::sync::LazyLock<ModelRegistry> =
16 std::sync::LazyLock::new(ModelRegistry::builtin);
17
18pub 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 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 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 let protocol_impl = build_protocol(&profile, config)?;
60
61 let provider = ProfiledProvider::new(GenericProvider::new(protocol_impl), profile);
63 Ok(Arc::new(provider))
64}
65
66fn 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
84pub 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
100pub 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 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 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 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 assert_eq!(provider.info().name, "openai");
253 }
255}