#[cfg(feature = "vertex")]
use anyhow::Context;
use anyhow::{Result, bail};
use async_trait::async_trait;
use serde::{Deserialize, Serialize};
use std::sync::Arc;
use std::time::Duration;
use crate::settings::{AiSettings, 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>,
},
}
impl AiResponseFormat {
pub fn format_json_schema_instruction(&self) -> Option<String> {
match self {
AiResponseFormat::Json {
schema: Some(schema),
} => {
if let Ok(schema_str) = serde_json::to_string(schema) {
Some(format!(
"RESPONSE FORMAT: You MUST respond with ONLY a valid JSON object matching this schema: {schema_str}. \
Do not include any explanation, markdown, or code fences — output raw JSON."
))
} else {
Some(
"RESPONSE FORMAT: You MUST respond with ONLY a valid JSON object. \
Do not include any explanation, markdown, or code fences — output raw JSON."
.to_string(),
)
}
}
AiResponseFormat::Json { schema: None } => Some(
"RESPONSE FORMAT: You MUST respond with ONLY a valid JSON object. \
Do not include any explanation, markdown, or code fences — output raw JSON."
.to_string(),
),
AiResponseFormat::Text => None,
}
}
}
#[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>,
#[serde(default)]
pub truncated: bool,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[serde(tag = "class", rename_all = "snake_case")]
pub enum AiErrorClass {
Fatal,
RateLimit {
#[serde(rename = "retry_after_secs", with = "duration_secs")]
retry_after: Duration,
},
Transient {
#[serde(rename = "retry_after_secs", with = "duration_secs")]
retry_after: Duration,
},
}
mod duration_secs {
use serde::{Deserialize, Deserializer, Serializer};
use std::time::Duration;
pub fn serialize<S: Serializer>(duration: &Duration, serializer: S) -> Result<S::Ok, S::Error> {
serializer.serialize_u64(duration.as_secs())
}
pub fn deserialize<'de, D: Deserializer<'de>>(deserializer: D) -> Result<Duration, D::Error> {
let secs = u64::deserialize(deserializer)?;
Ok(Duration::from_secs(secs))
}
}
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
pub(crate) struct RemoteAiErrorPayload {
pub message: String,
#[serde(flatten)]
pub class: AiErrorClass,
}
impl RemoteAiErrorPayload {
pub fn new(message: String, class: AiErrorClass) -> Self {
Self { message, class }
}
pub fn into_error(self) -> RemoteAiError {
RemoteAiError {
message: self.message,
class: self.class,
}
}
}
#[derive(Debug, Clone, thiserror::Error)]
#[error("Remote AI Error: {message}")]
pub struct RemoteAiError {
pub message: String,
pub class: AiErrorClass,
}
pub(crate) const DEFAULT_RETRY_AFTER: Duration = Duration::from_secs(60);
pub trait ClassifyAiError {
fn ai_error_class(&self) -> AiErrorClass;
}
impl ClassifyAiError for RemoteAiError {
fn ai_error_class(&self) -> AiErrorClass {
self.class
}
}
pub(crate) fn classify_status_code(status: reqwest::StatusCode) -> Option<AiErrorClass> {
match status {
reqwest::StatusCode::TOO_MANY_REQUESTS => Some(AiErrorClass::RateLimit {
retry_after: DEFAULT_RETRY_AFTER,
}),
reqwest::StatusCode::INTERNAL_SERVER_ERROR
| reqwest::StatusCode::BAD_GATEWAY
| reqwest::StatusCode::SERVICE_UNAVAILABLE
| reqwest::StatusCode::GATEWAY_TIMEOUT => Some(AiErrorClass::Transient {
retry_after: DEFAULT_RETRY_AFTER,
}),
status if status.as_u16() == 529 => Some(AiErrorClass::Transient {
retry_after: DEFAULT_RETRY_AFTER,
}),
_ => None,
}
}
pub fn classify_ai_error(error: &anyhow::Error) -> AiErrorClass {
if let Some(e) = error.downcast_ref::<RemoteAiError>() {
return e.ai_error_class();
}
if let Some(e) = error.downcast_ref::<openai::OpenAiCompatError>() {
return e.ai_error_class();
}
if let Some(e) = error.downcast_ref::<claude::ClaudeError>() {
return e.ai_error_class();
}
if let Some(e) = error.downcast_ref::<claude_cli::ClaudeCliError>() {
return e.ai_error_class();
}
if let Some(e) = error.downcast_ref::<gemini::GeminiError>() {
return e.ai_error_class();
}
if let Some(e) = error.downcast_ref::<crate::worker::prompts::ReviewError>() {
return e.ai_error_class();
}
AiErrorClass::Fatal
}
#[allow(dead_code)]
pub(crate) fn decode_stdio_ai_response(line: &str) -> Result<AiResponse> {
let resp_msg: serde_json::Value = serde_json::from_str(line)?;
match resp_msg["type"].as_str() {
Some("ai_response") => Ok(serde_json::from_value(resp_msg["payload"].clone())?),
Some("error") => {
let payload: RemoteAiErrorPayload =
serde_json::from_value(resp_msg["payload"].clone())?;
Err(payload.into_error().into())
}
_ => bail!("Unexpected response type: {:?}", resp_msg["type"]),
}
}
#[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>> {
create_provider_from_ai(&settings.ai)
}
pub fn create_provider_from_ai(ai: &AiSettings) -> Result<Arc<dyn AiProvider>> {
match ai.provider.to_lowercase().as_str() {
"gemini" => {
let model = ai.model.clone();
Ok(Arc::new(gemini::GeminiClient::new(model)))
}
"stdio-gemini" => Ok(Arc::new(gemini::StdioGeminiClient::new())),
"claude" => {
let model = ai.model.clone();
let enable_caching = ai.claude.as_ref().map(|c| c.prompt_caching).unwrap_or(true); let claude = 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::new())),
#[cfg(feature = "bedrock")]
"bedrock" => {
let model = ai.model.clone();
let bedrock = 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 ai.provider.to_lowercase().as_str() {
"openai" => openai::OpenAiProviderType::OpenAi,
_ => openai::OpenAiProviderType::OpenAiCompatible,
};
let base_url = ai
.openai_compat
.as_ref()
.and_then(|c| c.base_url.clone())
.unwrap_or_else(|| {
openai::OpenAiCompatClient::default_base_url_for_model(&ai.model)
});
let context_window = ai
.openai_compat
.as_ref()
.and_then(|c| c.context_window_size)
.unwrap_or_else(|| {
openai::OpenAiCompatClient::default_context_window_for_model(&ai.model)
});
let max_tokens = ai
.openai_compat
.as_ref()
.and_then(|c| c.max_tokens)
.unwrap_or(4096);
let provider = openai::OpenAiCompatClient::new(
base_url,
provider_type,
ai.model.clone(),
context_window,
max_tokens,
ai.api_timeout_secs,
)?;
Ok(Arc::new(provider))
}
"claude-cli" => {
let cfg = ai.claude_cli.as_ref();
Ok(Arc::new(claude_cli::ClaudeCliProvider {
model: ai.model.clone(),
effort: cfg.and_then(|c| c.effort.clone()),
}))
}
"devin-cli" => {
let cfg = ai.devin_cli.as_ref();
let model = if ai.model.is_empty() {
None
} else {
Some(ai.model.clone())
};
Ok(Arc::new(devin_cli::DevinCliProvider {
model,
agent_config: cfg.and_then(|c| c.agent_config.clone()),
config: cfg.and_then(|c| c.config.clone()),
}))
}
"codex-cli" => Ok(Arc::new(codex_cli::CodexCliProvider {
model: ai.model.clone(),
})),
"copilot-cli" => Ok(Arc::new(copilot_cli::CopilotCliProvider {
model: ai.model.clone(),
})),
"kiro-cli" => {
let cfg = ai.kiro_cli.as_ref();
Ok(Arc::new(kiro_cli::KiroCliProvider {
model: ai.model.clone(),
binary: cfg
.map(|c| c.binary.clone())
.unwrap_or_else(|| "kiro-cli".to_string()),
agent: cfg.and_then(|c| c.agent.clone()),
context_window_size: cfg.map(|c| c.context_window_size).unwrap_or(200_000),
timeout_secs: ai.api_timeout_secs,
}))
}
#[cfg(feature = "vertex")]
"vertex" => {
let model = ai.model.clone();
let vertex = 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 copilot_cli;
pub mod devin_cli;
pub mod gemini;
pub mod kiro_cli;
pub mod openai;
pub mod proxy;
pub mod quota;
pub mod session;
pub mod token_budget;
pub mod truncator;
#[cfg(feature = "vertex")]
pub mod vertex;
pub use session::{ErrorAction, LlmSession, SessionRunner, ValidationError};
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);
}
}
_ => {}
}
}
#[derive(Debug, serde::Serialize, serde::Deserialize)]
pub(crate) struct IpcEnvelope {
#[serde(rename = "type")]
pub msg_type: String,
pub tx_id: u64,
pub payload: serde_json::Value,
}
pub(crate) struct IpcRegistry {
next_tx_id: std::sync::atomic::AtomicU64,
pending: tokio::sync::Mutex<
std::collections::HashMap<
u64,
tokio::sync::oneshot::Sender<Result<AiResponse, RemoteAiError>>,
>,
>,
}
impl IpcRegistry {
pub fn new() -> Self {
Self {
next_tx_id: std::sync::atomic::AtomicU64::new(1),
pending: tokio::sync::Mutex::new(std::collections::HashMap::new()),
}
}
pub fn next_id(&self) -> u64 {
self.next_tx_id
.fetch_add(1, std::sync::atomic::Ordering::SeqCst)
}
pub async fn register(
&self,
tx_id: u64,
tx: tokio::sync::oneshot::Sender<Result<AiResponse, RemoteAiError>>,
) {
let mut map = self.pending.lock().await;
if map.insert(tx_id, tx).is_some() {
eprintln!(
"CRITICAL PROTOCOL ERROR: Duplicate transaction ID {} registered!",
tx_id
);
std::process::exit(1);
}
}
pub async fn dispatch(&self, tx_id: u64, result: Result<AiResponse, RemoteAiError>) {
let mut map = self.pending.lock().await;
if let Some(sender) = map.remove(&tx_id) {
let _ = sender.send(result);
} else {
eprintln!(
"CRITICAL PROTOCOL ERROR: Unsolicited response received for tx_id: {}!",
tx_id
);
std::process::exit(1);
}
}
pub async fn abort_all(&self, err: RemoteAiError) {
let mut map = self.pending.lock().await;
for (_tx_id, sender) in map.drain() {
let _ = sender.send(Err(err.clone()));
}
}
}
impl Default for IpcRegistry {
fn default() -> Self {
Self::new()
}
}
pub(crate) struct AtomicWriter {
writer: tokio::sync::Mutex<tokio::io::Stdout>,
}
impl AtomicWriter {
pub fn new() -> Self {
Self {
writer: tokio::sync::Mutex::new(tokio::io::stdout()),
}
}
pub async fn write_line(&self, line: &str) -> Result<()> {
use tokio::io::AsyncWriteExt;
let mut lock = self.writer.lock().await;
lock.write_all(line.as_bytes()).await?;
lock.write_all(b"\n").await?;
lock.flush().await?;
Ok(())
}
}
impl Default for AtomicWriter {
fn default() -> Self {
Self::new()
}
}
pub(crate) fn start_stdin_reader(registry: std::sync::Weak<IpcRegistry>) {
tokio::spawn(async move {
use tokio::io::{AsyncBufReadExt, BufReader};
let stdin = tokio::io::stdin();
let reader = BufReader::new(stdin);
let mut lines = reader.lines();
while let Ok(Some(line)) = lines.next_line().await {
let active_registry = match registry.upgrade() {
Some(r) => r,
None => {
tracing::info!("IPC Registry dropped, shutting down stdin reader task.");
break;
}
};
if let Ok(envelope) = serde_json::from_str::<IpcEnvelope>(&line) {
match envelope.msg_type.as_str() {
"ai_response" => {
if let Ok(payload) = serde_json::from_value::<AiResponse>(envelope.payload)
{
active_registry.dispatch(envelope.tx_id, Ok(payload)).await;
} else {
eprintln!(
"CRITICAL PROTOCOL ERROR: Failed to parse payload as AiResponse for tx_id {}",
envelope.tx_id
);
std::process::exit(1);
}
}
"error" => {
if let Ok(payload) =
serde_json::from_value::<RemoteAiErrorPayload>(envelope.payload)
{
active_registry
.dispatch(envelope.tx_id, Err(payload.into_error()))
.await;
} else {
eprintln!(
"CRITICAL PROTOCOL ERROR: Failed to parse payload as RemoteAiErrorPayload for tx_id {}",
envelope.tx_id
);
std::process::exit(1);
}
}
unknown => {
eprintln!("CRITICAL PROTOCOL ERROR: Unknown message type: {}", unknown);
std::process::exit(1);
}
}
} else {
eprintln!(
"CRITICAL PROTOCOL ERROR: Received malformed JSON on stdin: {}",
line
);
std::process::exit(1);
}
}
if let Some(active_registry) = registry.upgrade() {
active_registry
.abort_all(RemoteAiError {
message: "IPC channel disconnected (stdin closed)".to_string(),
class: AiErrorClass::Fatal,
})
.await;
}
});
}
#[cfg(test)]
mod tests {
use super::*;
use crate::worker::prompts::ReviewError;
use anyhow::anyhow;
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(())
}
fn assert_remote_ai_error_payload_wire_format(
class: AiErrorClass,
expected: serde_json::Value,
) -> Result<()> {
let payload = RemoteAiErrorPayload::new("x".to_string(), class);
let serialized = serde_json::to_value(&payload)?;
assert_eq!(serialized, expected);
let round_trip: RemoteAiErrorPayload = serde_json::from_value(serialized)?;
assert_eq!(round_trip, payload);
Ok(())
}
#[test]
fn test_remote_ai_error_payload_rate_limit_wire_format() -> Result<()> {
assert_remote_ai_error_payload_wire_format(
AiErrorClass::RateLimit {
retry_after: Duration::from_secs(42),
},
json!({
"message": "x",
"class": "rate_limit",
"retry_after_secs": 42
}),
)
}
#[test]
fn test_remote_ai_error_payload_transient_wire_format() -> Result<()> {
assert_remote_ai_error_payload_wire_format(
AiErrorClass::Transient {
retry_after: Duration::from_secs(15),
},
json!({
"message": "x",
"class": "transient",
"retry_after_secs": 15
}),
)
}
#[test]
fn test_remote_ai_error_payload_fatal_wire_format() -> Result<()> {
assert_remote_ai_error_payload_wire_format(
AiErrorClass::Fatal,
json!({
"message": "x",
"class": "fatal"
}),
)
}
#[test]
fn test_remote_ai_error_payload_requires_retry_after() {
let err = serde_json::from_value::<RemoteAiErrorPayload>(json!({
"message": "x",
"class": "rate_limit"
}))
.unwrap_err();
assert!(err.to_string().contains("retry_after_secs"));
}
fn assert_typed_stdio_error_downcasts(class: AiErrorClass) -> Result<()> {
let raw_json = json!({
"type": "error",
"payload": RemoteAiErrorPayload::new("typed failure".to_string(), class)
});
let serialized = serde_json::to_string(&raw_json)?;
let err = decode_stdio_ai_response(&serialized).unwrap_err();
let remote = err
.downcast_ref::<RemoteAiError>()
.expect("typed payload should downcast to RemoteAiError");
assert_eq!(remote.message, "typed failure");
assert_eq!(remote.class, class);
assert_eq!(err.to_string(), "Remote AI Error: typed failure");
Ok(())
}
#[test]
fn test_decode_stdio_ai_response_typed_rate_limit_error_payload() -> Result<()> {
assert_typed_stdio_error_downcasts(AiErrorClass::RateLimit {
retry_after: Duration::from_secs(60),
})
}
#[test]
fn test_decode_stdio_ai_response_typed_transient_error_payload() -> Result<()> {
assert_typed_stdio_error_downcasts(AiErrorClass::Transient {
retry_after: Duration::from_secs(5),
})
}
#[test]
fn test_decode_stdio_ai_response_typed_fatal_error_payload() -> Result<()> {
assert_typed_stdio_error_downcasts(AiErrorClass::Fatal)
}
#[test]
fn test_decode_stdio_ai_response_rejects_non_object_error_payload() -> Result<()> {
let raw_json = json!({
"type": "error",
"payload": "Rate limit exceeded, retry after 60s"
});
let serialized = serde_json::to_string(&raw_json)?;
let err = decode_stdio_ai_response(&serialized).unwrap_err();
assert!(err.to_string().contains("invalid type"));
assert!(err.downcast_ref::<RemoteAiError>().is_none());
Ok(())
}
#[test]
fn test_decode_stdio_ai_response_malformed_typed_error_payload() -> Result<()> {
let raw_json = json!({
"type": "error",
"payload": {
"message": "try again later",
"class": "transient"
}
});
let serialized = serde_json::to_string(&raw_json)?;
let err = decode_stdio_ai_response(&serialized).unwrap_err();
assert!(err.to_string().contains("retry_after_secs"));
Ok(())
}
#[test]
fn test_remote_ai_error_classifies_from_payload_class() {
let retry_after = Duration::from_secs(42);
let err = RemoteAiError {
message: "remote rate limit".to_string(),
class: AiErrorClass::RateLimit { retry_after },
};
assert_eq!(
err.ai_error_class(),
AiErrorClass::RateLimit { retry_after }
);
}
#[test]
fn test_classify_status_code_rate_limit() {
assert_eq!(
classify_status_code(reqwest::StatusCode::TOO_MANY_REQUESTS),
Some(AiErrorClass::RateLimit {
retry_after: DEFAULT_RETRY_AFTER,
})
);
}
#[test]
fn test_classify_status_code_server_error() {
assert_eq!(
classify_status_code(reqwest::StatusCode::SERVICE_UNAVAILABLE),
Some(AiErrorClass::Transient {
retry_after: DEFAULT_RETRY_AFTER,
})
);
}
#[test]
fn test_classify_status_code_other() {
assert_eq!(classify_status_code(reqwest::StatusCode::BAD_REQUEST), None);
}
fn assert_ai_error_class(error: impl Into<anyhow::Error>, expected: AiErrorClass) {
let error = error.into();
assert_eq!(classify_ai_error(&error), expected);
}
#[test]
fn test_classify_ai_error_remote_rate_limit() {
let retry_after = Duration::from_secs(42);
assert_ai_error_class(
RemoteAiError {
message: "remote rate limit".to_string(),
class: AiErrorClass::RateLimit { retry_after },
},
AiErrorClass::RateLimit { retry_after },
);
}
#[test]
fn test_classify_ai_error_remote_transient() {
let retry_after = Duration::from_secs(15);
assert_ai_error_class(
RemoteAiError {
message: "remote transient".to_string(),
class: AiErrorClass::Transient { retry_after },
},
AiErrorClass::Transient { retry_after },
);
}
#[test]
fn test_classify_ai_error_remote_fatal() {
assert_ai_error_class(
RemoteAiError {
message: "remote fatal".to_string(),
class: AiErrorClass::Fatal,
},
AiErrorClass::Fatal,
);
}
#[test]
fn test_classify_ai_error_provider_cascade() {
assert_ai_error_class(
openai::OpenAiCompatError::RateLimitExceeded(Duration::from_secs(7)),
AiErrorClass::RateLimit {
retry_after: Duration::from_secs(7),
},
);
assert_ai_error_class(
claude::ClaudeError::OverloadedError(Duration::from_secs(9)),
AiErrorClass::Transient {
retry_after: Duration::from_secs(9),
},
);
assert_ai_error_class(
gemini::GeminiError::TransientError(Duration::from_secs(11), "busy".to_string()),
AiErrorClass::Transient {
retry_after: Duration::from_secs(11),
},
);
assert_ai_error_class(
ReviewError::FormatRejection("bad response".to_string()),
AiErrorClass::Fatal,
);
}
#[test]
fn test_classify_ai_error_unrelated_error_is_fatal() {
assert_ai_error_class(anyhow!("totally unrelated error"), AiErrorClass::Fatal);
}
#[test]
fn test_classify_ai_error_string_shaped_remote_error_is_fatal() {
assert_ai_error_class(
anyhow!("Remote AI Error: rate limit exceeded"),
AiErrorClass::Fatal,
);
}
#[test]
fn test_classify_status_code_not_implemented_is_fatal() {
assert_eq!(
classify_status_code(reqwest::StatusCode::NOT_IMPLEMENTED),
None
);
}
#[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(())
}
}