Skip to main content

forge_guard/ai/
providers.rs

1//! AI provider implementations — HTTP clients for OpenAI, Claude, and Ollama.
2//!
3//! Each provider implements the [`LlmProvider`] trait, which exposes a single
4//! `call` method that sends a system + user prompt pair to the upstream API
5//! and returns the raw text response.
6
7use crate::core::ForgeGuardError;
8use std::time::Duration;
9
10/// A generic LLM provider capable of completing a chat prompt.
11pub trait LlmProvider: Send + Sync {
12    /// Human-readable provider name (e.g. "openai", "claude", "ollama").
13    fn name(&self) -> &'static str;
14
15    /// The model identifier being used (e.g. "gpt-4", "claude-sonnet-4-20250514").
16    fn model_name(&self) -> &str;
17
18    /// Send a system + user prompt pair and return the raw response text.
19    fn call(&self, system_prompt: &str, user_prompt: &str) -> Result<String, ForgeGuardError>;
20}
21
22// ─────────────────────────────────────────────────────────────
23// OpenAI
24// ─────────────────────────────────────────────────────────────
25
26/// OpenAI GPT provider (compatible with GPT-4, GPT-4o, etc.).
27pub 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    /// Create a new OpenAI provider.
39    ///
40    /// `api_key` — OpenAI API key (sk-...). Reads `OPENAI_API_KEY` env var if empty.
41    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    /// Returns `true` when no API key has been configured.
55    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
124// ─────────────────────────────────────────────────────────────
125// Claude (Anthropic)
126// ─────────────────────────────────────────────────────────────
127
128/// Anthropic Claude provider.
129pub 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    /// Create a new Claude provider.
142    ///
143    /// Reads `ANTHROPIC_API_KEY` env var if `api_key` is empty.
144    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    /// Returns `true` when no API key has been configured.
158    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
226// ─────────────────────────────────────────────────────────────
227// Ollama (local)
228// ─────────────────────────────────────────────────────────────
229
230/// Local Ollama provider.
231pub struct OllamaProvider {
232    endpoint: String,
233    model: String,
234    temperature: f64,
235}
236
237impl OllamaProvider {
238    const TIMEOUT_SECS: u64 = 300; // local models can be slow
239
240    /// Create a new Ollama provider.
241    ///
242    /// `endpoint` defaults to `http://localhost:11434` when empty.
243    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
310/// Factory function: create a boxed provider from an [`AiAuditorConfig`].
311///
312/// API keys are read from environment variables by the provider constructors.
313pub 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// ── Tests ────────────────────────────────────────────────────
346#[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        // No API key set in test env — should be missing
356        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        // The error should be about connection, not about missing endpoint
378        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        // Connect to a port that won't have Ollama running
459        let p = OllamaProvider::new(Some("http://127.0.0.1:1".into()), "test-model", 0.1);
460        let result = p.call("system", "user");
461        // Should fail with connection error, not parse error
462        assert!(result.is_err());
463    }
464}