diff --git a/signal-gateway/src/claude/mod.rs b/signal-gateway/src/claude/mod.rs index 0399645..fef5c7f 100644 --- a/signal-gateway/src/claude/mod.rs +++ b/signal-gateway/src/claude/mod.rs @@ -45,6 +45,27 @@ pub struct ClaudeConfig { /// Maximum tool use iterations before giving up. #[conf(long, env, default_value = "10")] pub claude_max_iterations: u32, + /// Compaction configuration. + #[conf(flatten, prefix)] + pub compaction: CompactionConfig, +} + +/// Configuration for message buffer compaction. +#[derive(Clone, Conf, Debug)] +#[conf(serde)] +pub struct CompactionConfig { + /// Path to file containing the compaction prompt. + #[conf(long, env)] + pub compaction_prompt_file: PathBuf, + /// Model to use for compaction (typically a faster/cheaper model). + #[conf(long, env, default_value = "claude-sonnet-4-5-20250929")] + pub compaction_model: String, + /// Maximum tokens for the compaction response. + #[conf(long, env, default_value = "2048")] + pub compaction_max_tokens: u32, + /// Trigger compaction when message buffer exceeds this many characters. + #[conf(long, env, default_value = "50000")] + pub compaction_trigger_chars: u32, } /// Error type for Claude API operations. diff --git a/signal-gateway/src/claude/worker.rs b/signal-gateway/src/claude/worker.rs index 1bd9cf5..e57571f 100644 --- a/signal-gateway/src/claude/worker.rs +++ b/signal-gateway/src/claude/worker.rs @@ -8,7 +8,7 @@ use serde_json::Value; use std::path::PathBuf; use std::sync::Weak; use tokio::sync::{mpsc, oneshot}; -use tracing::info; +use tracing::{error, info, warn}; /// Sent with inputs to claude that claude is expected to respond to. The sender /// gives the worker a way to return the results to the caller asynchronously. @@ -87,8 +87,8 @@ pub struct ClaudeWorker { client: reqwest::Client, api_key: String, system_prompts: Vec, - // FIXME: use this and append to system prompt within tags - #[allow(dead_code)] + compaction_prompt: String, + /// Summary of previous conversation history, wrapped in XML tags. summary: String, messages: Vec, tool_executor: Weak, @@ -99,7 +99,7 @@ pub struct ClaudeWorker { impl ClaudeWorker { /// Create a new Claude worker. /// - /// Reads the API key and system prompts from the configured files. + /// Reads the API key, system prompts, and compaction prompt from the configured files. pub fn new( config: ClaudeConfig, tool_executor: Weak, @@ -120,11 +120,18 @@ impl ClaudeWorker { }) .collect::>()?; + let compaction_prompt = { + let path = &config.compaction.compaction_prompt_file; + std::fs::read_to_string(path) + .map_err(|e| ClaudeError::SystemPromptRead(path.clone(), e))? + }; + Ok(Self { config, client: reqwest::Client::new(), api_key, system_prompts, + compaction_prompt, summary: String::new(), messages: Default::default(), tool_executor, @@ -146,6 +153,13 @@ impl ClaudeWorker { Input::Chat(msg) => { let (mc, maybe_sender) = msg.into_content_and_sender(); self.messages.push(mc); + + // Check if we need to trigger compaction + let buffer_chars = self.message_buffer_chars(); + if buffer_chars > self.config.compaction.compaction_trigger_chars as usize { + self.handle_compact().await; + } + if let Some(sender) = maybe_sender { let result = self.handle_request().await; // If handle_request was interrupted by stop request, go on to drain the queues @@ -171,6 +185,23 @@ impl ClaudeWorker { } } + /// Calculate total characters in the message buffer. + fn message_buffer_chars(&self) -> usize { + self.messages + .iter() + .map(|m| { + m.content + .iter() + .map(|block| match block { + ContentBlock::Text { text } => text.len(), + ContentBlock::ToolUse { input, .. } => estimate_json_size(input), + ContentBlock::ToolResult { content, .. } => content.len(), + }) + .sum::() + }) + .sum() + } + /// Handle a stop request by draining queues and sending errors to pending requests. fn handle_stop(&mut self) { // Drain the stop_rx queue @@ -200,10 +231,94 @@ impl ClaudeWorker { } } - /// Perform compaction + /// Perform compaction by summarizing messages and storing the result. async fn handle_compact(&mut self) { - // FIXME: we should actually try to summarize messages using an api request, and then store it, before tossing messages - self.messages.clear(); + if self.messages.is_empty() { + return; + } + + let num_messages = self.messages.len(); + let buffer_chars = self.message_buffer_chars(); + warn!( + "Starting compaction: {} messages, {} chars", + num_messages, buffer_chars + ); + + // Build system content: compaction prompt first, then other system prompts, then existing summary + let mut system: Vec = Vec::new(); + system.push(SystemContent::text(&self.compaction_prompt)); + for prompt in &self.system_prompts { + system.push(SystemContent::text(prompt)); + } + if !self.summary.is_empty() { + system.push(SystemContent::text(&self.summary)); + } + // Mark the last one as cached + if let Some(last) = system.last_mut() { + *last = std::mem::take(last).cached(); + } + + let request_body = MessagesRequest { + model: &self.config.compaction.compaction_model, + max_tokens: self.config.compaction.compaction_max_tokens, + system: &system, + messages: &self.messages, + tools: Vec::new(), + }; + + let result = self + .client + .post(&self.config.claude_api_url) + .header("x-api-key", &self.api_key) + .header("anthropic-version", ANTHROPIC_API_VERSION) + .header("content-type", "application/json") + .json(&request_body) + .send() + .await; + + match result { + Ok(response) if response.status().is_success() => { + match response.json::().await { + Ok(parsed) => { + // Extract text from response + let summary_text: String = parsed + .content + .into_iter() + .filter_map(|block| { + if let ContentBlock::Text { text } = block { + Some(text) + } else { + None + } + }) + .collect::>() + .join("\n"); + + // Wrap in XML tags and store + self.summary = + format!("\n{}\n", summary_text); + self.messages.clear(); + + info!( + "Compaction complete: summarized {} messages into {} chars", + num_messages, + self.summary.len() + ); + } + Err(e) => { + error!("Compaction failed to parse response: {}", e); + } + } + } + Ok(response) => { + let status = response.status(); + let body = response.text().await.unwrap_or_default(); + error!("Compaction API error ({}): {}", status, body); + } + Err(e) => { + error!("Compaction request failed: {}", e); + } + } } /// Log the message buffer for debugging. @@ -228,20 +343,19 @@ impl ClaudeWorker { info!("Claude request: {}", text); } - // Build system content blocks, caching only the last one - let system: Vec = self + // Build system content blocks: system prompts + summary (if any), caching only the last one + let mut system: Vec = self .system_prompts .iter() - .enumerate() - .map(|(i, text)| { - let content = SystemContent::text(text); - if i == self.system_prompts.len() - 1 { - content.cached() - } else { - content - } - }) + .map(|text| SystemContent::text(text)) .collect(); + if !self.summary.is_empty() { + system.push(SystemContent::text(&self.summary)); + } + // Mark the last one as cached + if let Some(last) = system.last_mut() { + *last = std::mem::take(last).cached(); + } for iteration in 0..max_iterations { // Check for stop before making API call @@ -341,6 +455,27 @@ impl ClaudeWorker { } } +/// Estimate the serialized size of a JSON Value without allocating. +fn estimate_json_size(value: &Value) -> usize { + match value { + Value::Null => 4, // "null" + Value::Bool(true) => 4, // "true" + Value::Bool(false) => 5, // "false" + Value::Number(n) => n.to_string().len(), // Numbers are small, ok to alloc + Value::String(s) => s.len() + 2, // quotes + Value::Array(arr) => { + 2 + arr.iter().map(estimate_json_size).sum::() + arr.len().saturating_sub(1) + } + Value::Object(obj) => { + 2 + obj + .iter() + .map(|(k, v)| k.len() + 3 + estimate_json_size(v)) // "key":value + .sum::() + + obj.len().saturating_sub(1) // commas + } + } +} + /// Request body for the Claude Messages API. #[derive(Serialize)] struct MessagesRequest<'a> { @@ -353,7 +488,7 @@ struct MessagesRequest<'a> { } /// A content block in the system prompt array. -#[derive(Clone, Serialize)] +#[derive(Clone, Default, Serialize)] struct SystemContent { #[serde(rename = "type")] content_type: &'static str,