use crate::core::ForgeGuardError;
use std::time::Duration;
pub trait LlmProvider: Send + Sync {
fn name(&self) -> &'static str;
fn model_name(&self) -> &str;
fn call(&self, system_prompt: &str, user_prompt: &str) -> Result<String, ForgeGuardError>;
}
pub struct OpenAiProvider {
api_key: String,
model: String,
temperature: f64,
max_tokens: u32,
}
impl OpenAiProvider {
const BASE_URL: &'static str = "https://api.openai.com/v1/chat/completions";
const TIMEOUT_SECS: u64 = 120;
pub fn new(model: &str, api_key: Option<String>, temperature: f64, max_tokens: u32) -> Self {
let key = api_key
.filter(|k| !k.is_empty())
.or_else(|| std::env::var("OPENAI_API_KEY").ok())
.unwrap_or_default();
Self {
api_key: key,
model: model.to_owned(),
temperature,
max_tokens,
}
}
pub fn missing_key(&self) -> bool {
self.api_key.is_empty()
}
}
impl LlmProvider for OpenAiProvider {
fn name(&self) -> &'static str {
"openai"
}
fn model_name(&self) -> &str {
&self.model
}
fn call(&self, system_prompt: &str, user_prompt: &str) -> Result<String, ForgeGuardError> {
if self.api_key.is_empty() {
return Err(ForgeGuardError::Config(
"OPENAI_API_KEY not set. Set the environment variable or pass --ai-api-key.".into(),
));
}
let client = reqwest::blocking::Client::builder()
.timeout(Duration::from_secs(Self::TIMEOUT_SECS))
.build()
.map_err(|e| ForgeGuardError::Internal(format!("HTTP client build: {e}")))?;
let body = serde_json::json!({
"model": self.model,
"messages": [
{"role": "system", "content": system_prompt},
{"role": "user", "content": user_prompt}
],
"temperature": self.temperature,
"max_tokens": self.max_tokens,
"response_format": {"type": "json_object"}
});
let resp = client
.post(Self::BASE_URL)
.header("Authorization", format!("Bearer {}", self.api_key))
.header("Content-Type", "application/json")
.json(&body)
.send()
.map_err(|e| ForgeGuardError::Rpc(format!("OpenAI request failed: {e}")))?;
let status = resp.status();
let text = resp
.text()
.map_err(|e| ForgeGuardError::Rpc(format!("OpenAI read response: {e}")))?;
if !status.is_success() {
return Err(ForgeGuardError::Rpc(format!(
"OpenAI API error (HTTP {status}): {text}"
)));
}
let parsed: serde_json::Value = serde_json::from_str(&text)
.map_err(|e| ForgeGuardError::Parse(format!("OpenAI JSON parse: {e}")))?;
let content = parsed["choices"][0]["message"]["content"]
.as_str()
.ok_or_else(|| {
ForgeGuardError::Parse("OpenAI response missing choices[0].message.content".into())
})?;
Ok(content.to_owned())
}
}
pub struct ClaudeProvider {
api_key: String,
model: String,
temperature: f64,
max_tokens: u32,
}
impl ClaudeProvider {
const BASE_URL: &'static str = "https://api.anthropic.com/v1/messages";
const ANTHROPIC_VERSION: &'static str = "2023-06-01";
const TIMEOUT_SECS: u64 = 120;
pub fn new(model: &str, api_key: Option<String>, temperature: f64, max_tokens: u32) -> Self {
let key = api_key
.filter(|k| !k.is_empty())
.or_else(|| std::env::var("ANTHROPIC_API_KEY").ok())
.unwrap_or_default();
Self {
api_key: key,
model: model.to_owned(),
temperature,
max_tokens,
}
}
pub fn missing_key(&self) -> bool {
self.api_key.is_empty()
}
}
impl LlmProvider for ClaudeProvider {
fn name(&self) -> &'static str {
"claude"
}
fn model_name(&self) -> &str {
&self.model
}
fn call(&self, system_prompt: &str, user_prompt: &str) -> Result<String, ForgeGuardError> {
if self.api_key.is_empty() {
return Err(ForgeGuardError::Config(
"ANTHROPIC_API_KEY not set. Set the environment variable or pass --ai-api-key."
.into(),
));
}
let client = reqwest::blocking::Client::builder()
.timeout(Duration::from_secs(Self::TIMEOUT_SECS))
.build()
.map_err(|e| ForgeGuardError::Internal(format!("HTTP client build: {e}")))?;
let body = serde_json::json!({
"model": self.model,
"max_tokens": self.max_tokens,
"system": system_prompt,
"messages": [
{"role": "user", "content": user_prompt}
],
"temperature": self.temperature
});
let resp = client
.post(Self::BASE_URL)
.header("x-api-key", &self.api_key)
.header("anthropic-version", Self::ANTHROPIC_VERSION)
.header("Content-Type", "application/json")
.json(&body)
.send()
.map_err(|e| ForgeGuardError::Rpc(format!("Claude request failed: {e}")))?;
let status = resp.status();
let text = resp
.text()
.map_err(|e| ForgeGuardError::Rpc(format!("Claude read response: {e}")))?;
if !status.is_success() {
return Err(ForgeGuardError::Rpc(format!(
"Claude API error (HTTP {status}): {text}"
)));
}
let parsed: serde_json::Value = serde_json::from_str(&text)
.map_err(|e| ForgeGuardError::Parse(format!("Claude JSON parse: {e}")))?;
let content = parsed["content"][0]["text"].as_str().ok_or_else(|| {
ForgeGuardError::Parse("Claude response missing content[0].text".into())
})?;
Ok(content.to_owned())
}
}
pub struct OllamaProvider {
endpoint: String,
model: String,
temperature: f64,
}
impl OllamaProvider {
const TIMEOUT_SECS: u64 = 300;
pub fn new(endpoint: Option<String>, model: &str, temperature: f64) -> Self {
Self {
endpoint: endpoint.unwrap_or_else(|| "http://localhost:11434".into()),
model: model.to_owned(),
temperature,
}
}
}
impl LlmProvider for OllamaProvider {
fn name(&self) -> &'static str {
"ollama"
}
fn model_name(&self) -> &str {
&self.model
}
fn call(&self, system_prompt: &str, user_prompt: &str) -> Result<String, ForgeGuardError> {
let url = format!("{}/api/chat", self.endpoint.trim_end_matches('/'));
let client = reqwest::blocking::Client::builder()
.timeout(Duration::from_secs(Self::TIMEOUT_SECS))
.build()
.map_err(|e| ForgeGuardError::Internal(format!("HTTP client build: {e}")))?;
let body = serde_json::json!({
"model": self.model,
"messages": [
{"role": "system", "content": system_prompt},
{"role": "user", "content": user_prompt}
],
"stream": false,
"options": {
"temperature": self.temperature
}
});
let resp = client
.post(&url)
.header("Content-Type", "application/json")
.json(&body)
.send()
.map_err(|e| ForgeGuardError::Rpc(format!("Ollama request failed: {e}")))?;
let status = resp.status();
let text = resp
.text()
.map_err(|e| ForgeGuardError::Rpc(format!("Ollama read response: {e}")))?;
if !status.is_success() {
return Err(ForgeGuardError::Rpc(format!(
"Ollama API error (HTTP {status}): {text}"
)));
}
let parsed: serde_json::Value = serde_json::from_str(&text)
.map_err(|e| ForgeGuardError::Parse(format!("Ollama JSON parse: {e}")))?;
let content = parsed["message"]["content"].as_str().ok_or_else(|| {
ForgeGuardError::Parse("Ollama response missing message.content".into())
})?;
Ok(content.to_owned())
}
}
pub fn create_provider(
provider_type: &str,
model: &str,
temperature: f64,
max_tokens: u32,
api_key: Option<String>,
ollama_endpoint: Option<String>,
) -> Result<Box<dyn LlmProvider>, ForgeGuardError> {
match provider_type.to_lowercase().as_str() {
"openai" | "gpt" => Ok(Box::new(OpenAiProvider::new(
model,
api_key,
temperature,
max_tokens,
))),
"claude" | "anthropic" => Ok(Box::new(ClaudeProvider::new(
model,
api_key,
temperature,
max_tokens,
))),
"ollama" => Ok(Box::new(OllamaProvider::new(
ollama_endpoint,
model,
temperature,
))),
other => Err(ForgeGuardError::Config(format!(
"Unknown AI provider: '{other}'. Supported: openai, claude, ollama"
))),
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_openai_provider_creation() {
let p = OpenAiProvider::new("gpt-5", None, 0.1, 4000);
assert_eq!(p.name(), "openai");
assert_eq!(p.model_name(), "gpt-5");
assert!(p.missing_key());
}
#[test]
fn test_claude_provider_creation() {
let p = ClaudeProvider::new("claude-5-sonnet-20260701", None, 0.1, 4000);
assert_eq!(p.name(), "claude");
assert_eq!(p.model_name(), "claude-5-sonnet-20260701");
assert!(p.missing_key());
}
#[test]
fn test_ollama_provider_creation() {
let p = OllamaProvider::new(Some("http://localhost:11434".into()), "llama3", 0.1);
assert_eq!(p.name(), "ollama");
assert_eq!(p.model_name(), "llama3");
}
#[test]
fn test_ollama_default_endpoint() {
let p = OllamaProvider::new(None, "codellama", 0.0);
match p.call("test", "test") {
Err(e) => {
let msg = e.to_string();
assert!(!msg.contains("endpoint"), "should use default endpoint");
}
Ok(_) => panic!("expected error with no Ollama running"),
}
}
#[test]
fn test_create_provider_openai() {
match create_provider("openai", "gpt-5", 0.2, 3000, None, None) {
Ok(p) => {
assert_eq!(p.name(), "openai");
assert_eq!(p.model_name(), "gpt-5");
}
Err(e) => panic!("Expected Ok, got: {e}"),
}
}
#[test]
fn test_create_provider_claude() {
match create_provider("claude", "claude-5-opus-20260701", 0.3, 5000, None, None) {
Ok(p) => {
assert_eq!(p.name(), "claude");
assert_eq!(p.model_name(), "claude-5-opus-20260701");
}
Err(e) => panic!("Expected Ok, got: {e}"),
}
}
#[test]
fn test_create_provider_ollama() {
match create_provider(
"ollama",
"llama3.1",
0.1,
4000,
None,
Some("http://ollama:11434".into()),
) {
Ok(p) => {
assert_eq!(p.name(), "ollama");
assert_eq!(p.model_name(), "llama3.1");
}
Err(e) => panic!("Expected Ok, got: {e}"),
}
}
#[test]
fn test_create_provider_unknown() {
match create_provider("nonexistent", "x", 0.1, 1000, None, None) {
Err(e) => {
let msg = e.to_string();
assert!(msg.contains("Unknown AI provider"));
}
Ok(_) => panic!("Expected Err"),
}
}
#[test]
fn test_openai_call_without_key() {
let p = OpenAiProvider::new("gpt-5", Some("".into()), 0.1, 1000);
match p.call("system", "user") {
Err(e) => assert!(e.to_string().contains("OPENAI_API_KEY")),
Ok(_) => panic!("Expected Err"),
}
}
#[test]
fn test_claude_call_without_key() {
let p = ClaudeProvider::new("claude-5-sonnet-20260701", Some("".into()), 0.1, 1000);
match p.call("system", "user") {
Err(e) => assert!(e.to_string().contains("ANTHROPIC_API_KEY")),
Ok(_) => panic!("Expected Err"),
}
}
#[test]
fn test_ollama_call_no_server() {
let p = OllamaProvider::new(Some("http://127.0.0.1:1".into()), "test-model", 0.1);
let result = p.call("system", "user");
assert!(result.is_err());
}
}