1use crate::core::ForgeGuardError;
8use std::time::Duration;
9
10pub trait LlmProvider: Send + Sync {
12 fn name(&self) -> &'static str;
14
15 fn model_name(&self) -> &str;
17
18 fn call(&self, system_prompt: &str, user_prompt: &str) -> Result<String, ForgeGuardError>;
20}
21
22pub struct OpenAiProvider {
28 api_key: String,
29 model: String,
30 temperature: f64,
31 max_tokens: u32,
32}
33
34impl OpenAiProvider {
35 const BASE_URL: &'static str = "https://api.openai.com/v1/chat/completions";
36 const TIMEOUT_SECS: u64 = 120;
37
38 pub fn new(model: &str, api_key: Option<String>, temperature: f64, max_tokens: u32) -> Self {
42 let key = api_key
43 .filter(|k| !k.is_empty())
44 .or_else(|| std::env::var("OPENAI_API_KEY").ok())
45 .unwrap_or_default();
46 Self {
47 api_key: key,
48 model: model.to_owned(),
49 temperature,
50 max_tokens,
51 }
52 }
53
54 pub fn missing_key(&self) -> bool {
56 self.api_key.is_empty()
57 }
58}
59
60impl LlmProvider for OpenAiProvider {
61 fn name(&self) -> &'static str {
62 "openai"
63 }
64
65 fn model_name(&self) -> &str {
66 &self.model
67 }
68
69 fn call(&self, system_prompt: &str, user_prompt: &str) -> Result<String, ForgeGuardError> {
70 if self.api_key.is_empty() {
71 return Err(ForgeGuardError::Config(
72 "OPENAI_API_KEY not set. Set the environment variable or pass --ai-api-key.".into(),
73 ));
74 }
75
76 let client = reqwest::blocking::Client::builder()
77 .timeout(Duration::from_secs(Self::TIMEOUT_SECS))
78 .build()
79 .map_err(|e| ForgeGuardError::Internal(format!("HTTP client build: {e}")))?;
80
81 let body = serde_json::json!({
82 "model": self.model,
83 "messages": [
84 {"role": "system", "content": system_prompt},
85 {"role": "user", "content": user_prompt}
86 ],
87 "temperature": self.temperature,
88 "max_tokens": self.max_tokens,
89 "response_format": {"type": "json_object"}
90 });
91
92 let resp = client
93 .post(Self::BASE_URL)
94 .header("Authorization", format!("Bearer {}", self.api_key))
95 .header("Content-Type", "application/json")
96 .json(&body)
97 .send()
98 .map_err(|e| ForgeGuardError::Rpc(format!("OpenAI request failed: {e}")))?;
99
100 let status = resp.status();
101 let text = resp
102 .text()
103 .map_err(|e| ForgeGuardError::Rpc(format!("OpenAI read response: {e}")))?;
104
105 if !status.is_success() {
106 return Err(ForgeGuardError::Rpc(format!(
107 "OpenAI API error (HTTP {status}): {text}"
108 )));
109 }
110
111 let parsed: serde_json::Value = serde_json::from_str(&text)
112 .map_err(|e| ForgeGuardError::Parse(format!("OpenAI JSON parse: {e}")))?;
113
114 let content = parsed["choices"][0]["message"]["content"]
115 .as_str()
116 .ok_or_else(|| {
117 ForgeGuardError::Parse("OpenAI response missing choices[0].message.content".into())
118 })?;
119
120 Ok(content.to_owned())
121 }
122}
123
124pub struct ClaudeProvider {
130 api_key: String,
131 model: String,
132 temperature: f64,
133 max_tokens: u32,
134}
135
136impl ClaudeProvider {
137 const BASE_URL: &'static str = "https://api.anthropic.com/v1/messages";
138 const ANTHROPIC_VERSION: &'static str = "2023-06-01";
139 const TIMEOUT_SECS: u64 = 120;
140
141 pub fn new(model: &str, api_key: Option<String>, temperature: f64, max_tokens: u32) -> Self {
145 let key = api_key
146 .filter(|k| !k.is_empty())
147 .or_else(|| std::env::var("ANTHROPIC_API_KEY").ok())
148 .unwrap_or_default();
149 Self {
150 api_key: key,
151 model: model.to_owned(),
152 temperature,
153 max_tokens,
154 }
155 }
156
157 pub fn missing_key(&self) -> bool {
159 self.api_key.is_empty()
160 }
161}
162
163impl LlmProvider for ClaudeProvider {
164 fn name(&self) -> &'static str {
165 "claude"
166 }
167
168 fn model_name(&self) -> &str {
169 &self.model
170 }
171
172 fn call(&self, system_prompt: &str, user_prompt: &str) -> Result<String, ForgeGuardError> {
173 if self.api_key.is_empty() {
174 return Err(ForgeGuardError::Config(
175 "ANTHROPIC_API_KEY not set. Set the environment variable or pass --ai-api-key."
176 .into(),
177 ));
178 }
179
180 let client = reqwest::blocking::Client::builder()
181 .timeout(Duration::from_secs(Self::TIMEOUT_SECS))
182 .build()
183 .map_err(|e| ForgeGuardError::Internal(format!("HTTP client build: {e}")))?;
184
185 let body = serde_json::json!({
186 "model": self.model,
187 "max_tokens": self.max_tokens,
188 "system": system_prompt,
189 "messages": [
190 {"role": "user", "content": user_prompt}
191 ],
192 "temperature": self.temperature
193 });
194
195 let resp = client
196 .post(Self::BASE_URL)
197 .header("x-api-key", &self.api_key)
198 .header("anthropic-version", Self::ANTHROPIC_VERSION)
199 .header("Content-Type", "application/json")
200 .json(&body)
201 .send()
202 .map_err(|e| ForgeGuardError::Rpc(format!("Claude request failed: {e}")))?;
203
204 let status = resp.status();
205 let text = resp
206 .text()
207 .map_err(|e| ForgeGuardError::Rpc(format!("Claude read response: {e}")))?;
208
209 if !status.is_success() {
210 return Err(ForgeGuardError::Rpc(format!(
211 "Claude API error (HTTP {status}): {text}"
212 )));
213 }
214
215 let parsed: serde_json::Value = serde_json::from_str(&text)
216 .map_err(|e| ForgeGuardError::Parse(format!("Claude JSON parse: {e}")))?;
217
218 let content = parsed["content"][0]["text"].as_str().ok_or_else(|| {
219 ForgeGuardError::Parse("Claude response missing content[0].text".into())
220 })?;
221
222 Ok(content.to_owned())
223 }
224}
225
226pub struct OllamaProvider {
232 endpoint: String,
233 model: String,
234 temperature: f64,
235}
236
237impl OllamaProvider {
238 const TIMEOUT_SECS: u64 = 300; pub fn new(endpoint: Option<String>, model: &str, temperature: f64) -> Self {
244 Self {
245 endpoint: endpoint.unwrap_or_else(|| "http://localhost:11434".into()),
246 model: model.to_owned(),
247 temperature,
248 }
249 }
250}
251
252impl LlmProvider for OllamaProvider {
253 fn name(&self) -> &'static str {
254 "ollama"
255 }
256
257 fn model_name(&self) -> &str {
258 &self.model
259 }
260
261 fn call(&self, system_prompt: &str, user_prompt: &str) -> Result<String, ForgeGuardError> {
262 let url = format!("{}/api/chat", self.endpoint.trim_end_matches('/'));
263
264 let client = reqwest::blocking::Client::builder()
265 .timeout(Duration::from_secs(Self::TIMEOUT_SECS))
266 .build()
267 .map_err(|e| ForgeGuardError::Internal(format!("HTTP client build: {e}")))?;
268
269 let body = serde_json::json!({
270 "model": self.model,
271 "messages": [
272 {"role": "system", "content": system_prompt},
273 {"role": "user", "content": user_prompt}
274 ],
275 "stream": false,
276 "options": {
277 "temperature": self.temperature
278 }
279 });
280
281 let resp = client
282 .post(&url)
283 .header("Content-Type", "application/json")
284 .json(&body)
285 .send()
286 .map_err(|e| ForgeGuardError::Rpc(format!("Ollama request failed: {e}")))?;
287
288 let status = resp.status();
289 let text = resp
290 .text()
291 .map_err(|e| ForgeGuardError::Rpc(format!("Ollama read response: {e}")))?;
292
293 if !status.is_success() {
294 return Err(ForgeGuardError::Rpc(format!(
295 "Ollama API error (HTTP {status}): {text}"
296 )));
297 }
298
299 let parsed: serde_json::Value = serde_json::from_str(&text)
300 .map_err(|e| ForgeGuardError::Parse(format!("Ollama JSON parse: {e}")))?;
301
302 let content = parsed["message"]["content"].as_str().ok_or_else(|| {
303 ForgeGuardError::Parse("Ollama response missing message.content".into())
304 })?;
305
306 Ok(content.to_owned())
307 }
308}
309
310pub fn create_provider(
314 provider_type: &str,
315 model: &str,
316 temperature: f64,
317 max_tokens: u32,
318 api_key: Option<String>,
319 ollama_endpoint: Option<String>,
320) -> Result<Box<dyn LlmProvider>, ForgeGuardError> {
321 match provider_type.to_lowercase().as_str() {
322 "openai" | "gpt" => Ok(Box::new(OpenAiProvider::new(
323 model,
324 api_key,
325 temperature,
326 max_tokens,
327 ))),
328 "claude" | "anthropic" => Ok(Box::new(ClaudeProvider::new(
329 model,
330 api_key,
331 temperature,
332 max_tokens,
333 ))),
334 "ollama" => Ok(Box::new(OllamaProvider::new(
335 ollama_endpoint,
336 model,
337 temperature,
338 ))),
339 other => Err(ForgeGuardError::Config(format!(
340 "Unknown AI provider: '{other}'. Supported: openai, claude, ollama"
341 ))),
342 }
343}
344
345#[cfg(test)]
347mod tests {
348 use super::*;
349
350 #[test]
351 fn test_openai_provider_creation() {
352 let p = OpenAiProvider::new("gpt-5", None, 0.1, 4000);
353 assert_eq!(p.name(), "openai");
354 assert_eq!(p.model_name(), "gpt-5");
355 assert!(p.missing_key());
357 }
358
359 #[test]
360 fn test_claude_provider_creation() {
361 let p = ClaudeProvider::new("claude-5-sonnet-20260701", None, 0.1, 4000);
362 assert_eq!(p.name(), "claude");
363 assert_eq!(p.model_name(), "claude-5-sonnet-20260701");
364 assert!(p.missing_key());
365 }
366
367 #[test]
368 fn test_ollama_provider_creation() {
369 let p = OllamaProvider::new(Some("http://localhost:11434".into()), "llama3", 0.1);
370 assert_eq!(p.name(), "ollama");
371 assert_eq!(p.model_name(), "llama3");
372 }
373
374 #[test]
375 fn test_ollama_default_endpoint() {
376 let p = OllamaProvider::new(None, "codellama", 0.0);
377 match p.call("test", "test") {
379 Err(e) => {
380 let msg = e.to_string();
381 assert!(!msg.contains("endpoint"), "should use default endpoint");
382 }
383 Ok(_) => panic!("expected error with no Ollama running"),
384 }
385 }
386
387 #[test]
388 fn test_create_provider_openai() {
389 match create_provider("openai", "gpt-5", 0.2, 3000, None, None) {
390 Ok(p) => {
391 assert_eq!(p.name(), "openai");
392 assert_eq!(p.model_name(), "gpt-5");
393 }
394 Err(e) => panic!("Expected Ok, got: {e}"),
395 }
396 }
397
398 #[test]
399 fn test_create_provider_claude() {
400 match create_provider("claude", "claude-5-opus-20260701", 0.3, 5000, None, None) {
401 Ok(p) => {
402 assert_eq!(p.name(), "claude");
403 assert_eq!(p.model_name(), "claude-5-opus-20260701");
404 }
405 Err(e) => panic!("Expected Ok, got: {e}"),
406 }
407 }
408
409 #[test]
410 fn test_create_provider_ollama() {
411 match create_provider(
412 "ollama",
413 "llama3.1",
414 0.1,
415 4000,
416 None,
417 Some("http://ollama:11434".into()),
418 ) {
419 Ok(p) => {
420 assert_eq!(p.name(), "ollama");
421 assert_eq!(p.model_name(), "llama3.1");
422 }
423 Err(e) => panic!("Expected Ok, got: {e}"),
424 }
425 }
426
427 #[test]
428 fn test_create_provider_unknown() {
429 match create_provider("nonexistent", "x", 0.1, 1000, None, None) {
430 Err(e) => {
431 let msg = e.to_string();
432 assert!(msg.contains("Unknown AI provider"));
433 }
434 Ok(_) => panic!("Expected Err"),
435 }
436 }
437
438 #[test]
439 fn test_openai_call_without_key() {
440 let p = OpenAiProvider::new("gpt-5", Some("".into()), 0.1, 1000);
441 match p.call("system", "user") {
442 Err(e) => assert!(e.to_string().contains("OPENAI_API_KEY")),
443 Ok(_) => panic!("Expected Err"),
444 }
445 }
446
447 #[test]
448 fn test_claude_call_without_key() {
449 let p = ClaudeProvider::new("claude-5-sonnet-20260701", Some("".into()), 0.1, 1000);
450 match p.call("system", "user") {
451 Err(e) => assert!(e.to_string().contains("ANTHROPIC_API_KEY")),
452 Ok(_) => panic!("Expected Err"),
453 }
454 }
455
456 #[test]
457 fn test_ollama_call_no_server() {
458 let p = OllamaProvider::new(Some("http://127.0.0.1:1".into()), "test-model", 0.1);
460 let result = p.call("system", "user");
461 assert!(result.is_err());
463 }
464}