#[cfg(feature = "vertex")]
use anyhow::Context;
use anyhow::{Result, bail};
use async_trait::async_trait;
use serde::{Deserialize, Serialize};
use std::sync::Arc;
use crate::settings::Settings;
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
#[serde(rename_all = "lowercase")]
pub enum AiRole {
System,
User,
Assistant,
Tool,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct AiMessage {
pub role: AiRole,
pub content: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub thought: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub thought_signature: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub tool_calls: Option<Vec<ToolCall>>,
#[serde(skip_serializing_if = "Option::is_none")]
pub tool_call_id: Option<String>,
}
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
pub struct ToolCall {
pub id: String,
pub function_name: String,
pub arguments: serde_json::Value,
#[serde(skip_serializing_if = "Option::is_none")]
pub thought_signature: Option<String>,
}
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
#[serde(rename_all = "lowercase")]
pub enum AiResponseFormat {
Text,
Json {
#[serde(skip_serializing_if = "Option::is_none")]
schema: Option<serde_json::Value>,
},
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct AiTool {
pub name: String,
pub description: String,
pub parameters: serde_json::Value,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct AiRequest {
#[serde(skip_serializing_if = "Option::is_none")]
pub system: Option<String>,
pub messages: Vec<AiMessage>,
#[serde(skip_serializing_if = "Option::is_none")]
pub tools: Option<Vec<AiTool>>,
#[serde(skip_serializing_if = "Option::is_none")]
pub temperature: Option<f32>,
#[serde(skip_serializing_if = "Option::is_none")]
pub response_format: Option<AiResponseFormat>,
#[serde(skip_serializing_if = "Option::is_none")]
pub context_tag: Option<String>,
}
tokio::task_local! {
pub static LOG_CONTEXT: String;
}
pub fn get_log_prefix() -> String {
LOG_CONTEXT.try_with(|c| c.clone()).unwrap_or_default()
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct AiResponse {
pub content: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub thought: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub thought_signature: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub tool_calls: Option<Vec<ToolCall>>,
pub usage: Option<AiUsage>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct AiUsage {
pub prompt_tokens: usize,
pub completion_tokens: usize,
pub total_tokens: usize,
#[serde(skip_serializing_if = "Option::is_none")]
pub cached_tokens: Option<usize>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ProviderCapabilities {
pub model_name: String,
pub context_window_size: usize,
}
#[derive(Debug, Clone, Default)]
pub struct CacheStats {
pub hits_this_session: u64,
pub hits_prev_session: u64,
pub tokens_saved_this_session: u64,
pub tokens_saved_prev_session: u64,
}
#[async_trait]
pub trait AiProvider: Send + Sync {
async fn generate_content(&self, request: AiRequest) -> Result<AiResponse>;
fn estimate_tokens(&self, request: &AiRequest) -> usize;
fn get_capabilities(&self) -> ProviderCapabilities;
fn cache_stats(&self) -> Option<CacheStats> {
None
}
}
pub async fn create_provider_cached(
settings: &Settings,
enable_cache: bool,
cache_ttl_days: u64,
) -> Result<Arc<dyn AiProvider>> {
let provider = create_provider(settings)?;
if enable_cache {
let cache_path = std::path::Path::new(&settings.database.url)
.parent()
.unwrap_or(std::path::Path::new("."))
.join("response_cache.db");
let cached =
cache::CachingAiProvider::new(provider, &cache_path.to_string_lossy(), cache_ttl_days)
.await?;
Ok(Arc::new(cached))
} else {
Ok(provider)
}
}
pub fn create_provider(settings: &Settings) -> Result<Arc<dyn AiProvider>> {
match settings.ai.provider.to_lowercase().as_str() {
"gemini" => {
let model = settings.ai.model.clone();
Ok(Arc::new(gemini::GeminiClient::new(model)))
}
"stdio-gemini" => Ok(Arc::new(gemini::StdioGeminiClient)),
"claude" => {
let model = settings.ai.model.clone();
let enable_caching = settings
.ai
.claude
.as_ref()
.map(|c| c.prompt_caching)
.unwrap_or(true); let claude = settings.ai.claude.as_ref();
let max_tokens = claude.map(|c| c.max_tokens).unwrap_or(4096);
let base_url = claude
.and_then(|c| c.base_url.clone())
.unwrap_or_else(claude::ClaudeClient::default_base_url);
let thinking = claude.and_then(|c| c.thinking.clone());
let effort = claude.and_then(|c| c.effort.clone());
Ok(Arc::new(claude::ClaudeClient::new(
model,
enable_caching,
max_tokens,
base_url,
thinking,
effort,
)))
}
"stdio-claude" => Ok(Arc::new(claude::StdioClaudeClient)),
#[cfg(feature = "bedrock")]
"bedrock" => {
let model = settings.ai.model.clone();
let bedrock = settings.ai.bedrock.as_ref();
let region = bedrock.and_then(|b| b.region.clone());
let enable_caching = bedrock.map(|b| b.prompt_caching).unwrap_or(true);
let max_tokens = bedrock.map(|b| b.max_tokens).unwrap_or(8192);
let thinking = bedrock.and_then(|b| b.thinking.clone());
let effort = bedrock.and_then(|b| b.effort.clone());
Ok(Arc::new(bedrock::BedrockClient::new(
model,
region,
enable_caching,
max_tokens,
thinking,
effort,
)))
}
#[cfg(not(feature = "bedrock"))]
"bedrock" => bail!("bedrock provider requires the 'bedrock' feature"),
"openai" | "openai-compatible" => {
let provider_type = match settings.ai.provider.to_lowercase().as_str() {
"openai" => openai::OpenAiProviderType::OpenAi,
_ => openai::OpenAiProviderType::OpenAiCompatible,
};
let base_url = settings
.ai
.openai_compat
.as_ref()
.and_then(|c| c.base_url.clone())
.unwrap_or_else(|| {
openai::OpenAiCompatClient::default_base_url_for_model(&settings.ai.model)
});
let context_window = settings
.ai
.openai_compat
.as_ref()
.and_then(|c| c.context_window_size)
.unwrap_or_else(|| {
openai::OpenAiCompatClient::default_context_window_for_model(&settings.ai.model)
});
let max_tokens = settings
.ai
.openai_compat
.as_ref()
.and_then(|c| c.max_tokens)
.unwrap_or(4096);
Ok(Arc::new(openai::OpenAiCompatClient::new(
base_url,
provider_type,
settings.ai.model.clone(),
context_window,
max_tokens,
settings.ai.api_timeout_secs,
)))
}
"claude-cli" => Ok(Arc::new(claude_cli::ClaudeCliProvider {
model: settings.ai.model.clone(),
})),
"codex-cli" => Ok(Arc::new(codex_cli::CodexCliProvider {
model: settings.ai.model.clone(),
})),
#[cfg(feature = "vertex")]
"vertex" => {
let model = settings.ai.model.clone();
let vertex = settings.ai.vertex.as_ref();
let project_id = vertex
.and_then(|v| v.project_id.clone())
.or_else(|| std::env::var("ANTHROPIC_VERTEX_PROJECT_ID").ok())
.context(
"Vertex AI requires project_id in [ai.vertex] \
or ANTHROPIC_VERTEX_PROJECT_ID env var",
)?;
let region = vertex
.and_then(|v| v.region.clone())
.or_else(|| std::env::var("CLOUD_ML_REGION").ok())
.unwrap_or_else(|| "us-east5".to_string());
let enable_caching = vertex.map(|v| v.prompt_caching).unwrap_or(true);
let max_tokens = vertex.map(|v| v.max_tokens).unwrap_or(8192);
let thinking = vertex.and_then(|v| v.thinking.clone());
let effort = vertex.and_then(|v| v.effort.clone());
Ok(Arc::new(vertex::VertexClient::new(
model,
project_id,
region,
enable_caching,
max_tokens,
thinking,
effort,
)?))
}
#[cfg(not(feature = "vertex"))]
"vertex" => bail!("vertex provider requires the 'vertex' feature"),
p => bail!("Unsupported AI provider: {}", p),
}
}
#[cfg(feature = "bedrock")]
pub mod bedrock;
pub mod cache;
pub mod claude;
pub mod claude_cli;
pub mod codex_cli;
pub mod gemini;
pub mod openai;
pub mod proxy;
pub mod quota;
pub mod token_budget;
pub mod truncator;
#[cfg(feature = "vertex")]
pub mod vertex;
pub fn scrub_thought_signatures(val: &mut serde_json::Value) {
match val {
serde_json::Value::Object(map) => {
map.remove("thought_signature");
map.remove("thoughtSignature");
for (_, v) in map.iter_mut() {
scrub_thought_signatures(v);
}
}
serde_json::Value::Array(arr) => {
for v in arr.iter_mut() {
scrub_thought_signatures(v);
}
}
_ => {}
}
}
#[cfg(test)]
mod tests {
use super::*;
use serde_json::json;
#[test]
fn test_ai_request_contract() -> Result<()> {
let request = AiRequest {
system: None,
messages: vec![AiMessage {
role: AiRole::User,
content: Some("Hello".to_string()),
thought: None,
thought_signature: None,
tool_calls: None,
tool_call_id: None,
}],
tools: None,
temperature: Some(0.5),
response_format: Some(AiResponseFormat::Text),
context_tag: None,
};
let msg = json!({
"type": "ai_request",
"payload": request
});
let serialized = serde_json::to_string(&msg)?;
let deserialized: serde_json::Value = serde_json::from_str(&serialized)?;
assert_eq!(deserialized["type"], "ai_request");
assert_eq!(deserialized["payload"]["temperature"], 0.5);
assert_eq!(deserialized["payload"]["messages"][0]["role"], "user");
assert_eq!(deserialized["payload"]["messages"][0]["content"], "Hello");
Ok(())
}
#[test]
fn test_ai_response_contract() -> Result<()> {
let raw_json = json!({
"type": "ai_response",
"payload": {
"content": "AI response text",
"tool_calls": [
{
"id": "call_1",
"function_name": "my_tool",
"arguments": {"a": 1},
"thought_signature": "sig_123"
}
],
"usage": {
"prompt_tokens": 100,
"completion_tokens": 50,
"total_tokens": 150
}
}
});
let serialized = serde_json::to_string(&raw_json)?;
let deserialized: serde_json::Value = serde_json::from_str(&serialized)?;
assert_eq!(deserialized["type"], "ai_response");
let payload: AiResponse = serde_json::from_value(deserialized["payload"].clone())?;
assert_eq!(payload.content.as_deref(), Some("AI response text"));
let tool_calls = payload.tool_calls.unwrap();
assert_eq!(tool_calls.len(), 1);
assert_eq!(tool_calls[0].id, "call_1");
assert_eq!(tool_calls[0].function_name, "my_tool");
assert_eq!(tool_calls[0].arguments["a"], 1);
assert_eq!(tool_calls[0].thought_signature.as_deref(), Some("sig_123"));
let usage = payload.usage.unwrap();
assert_eq!(usage.prompt_tokens, 100);
assert_eq!(usage.completion_tokens, 50);
assert_eq!(usage.total_tokens, 150);
Ok(())
}
#[test]
fn test_create_provider() -> Result<()> {
let mut settings = Settings::new().expect("Failed to load settings");
settings.ai.provider = "gemini".to_string();
settings.ai.model = "gemini-1.5-flash".to_string();
let provider = create_provider(&settings)?;
assert_eq!(provider.get_capabilities().model_name, "gemini-1.5-flash");
settings.ai.provider = "stdio-gemini".to_string();
let provider = create_provider(&settings)?;
assert_eq!(provider.get_capabilities().model_name, "stdio-gemini");
settings.ai.provider = "openai".to_string();
settings.ai.model = "gpt-4o".to_string();
let provider = create_provider(&settings)?;
assert_eq!(provider.get_capabilities().model_name, "gpt-4o");
settings.ai.provider = "unknown".to_string();
let result = create_provider(&settings);
assert!(result.is_err());
Ok(())
}
}