diff --git a/signal-gateway-assistant/claude/src/lib.rs b/signal-gateway-assistant/claude/src/lib.rs index ecd556f..777a852 100644 --- a/signal-gateway-assistant/claude/src/lib.rs +++ b/signal-gateway-assistant/claude/src/lib.rs @@ -1,6 +1,9 @@ //! Claude API implementation of the Assistant trait. +mod message_buffer; + use conf::Conf; +use message_buffer::MessageBuffer; use serde::{Deserialize, Serialize}; use serde_json::Value; use signal_gateway_assistant::{Assistant, AssistantResponse, ChatMessage, Tool, ToolExecutor}; @@ -96,7 +99,7 @@ pub struct ClaudeAssistant { compaction_prompt: String, /// Summary of previous conversation history, wrapped in XML tags. summary: String, - messages: Vec, + messages: MessageBuffer, tool_executor: Weak, /// Timestamp of the last automatic compaction (not user-requested). last_auto_compaction: Option, @@ -135,32 +138,15 @@ impl ClaudeAssistant { system_prompts, compaction_prompt, summary: String::new(), - messages: Vec::new(), + messages: MessageBuffer::new(), tool_executor, last_auto_compaction: None, }) } - /// 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() - } - /// Check if automatic compaction should be triggered and handle it. async fn maybe_compact(&mut self) { - let buffer_chars = self.message_buffer_chars(); + let buffer_chars = self.messages.total_chars(); if buffer_chars <= self.config.compaction.trigger_chars as usize { return; } @@ -184,31 +170,19 @@ impl ClaudeAssistant { /// Drop oldest messages until buffer is under the trigger threshold. fn drop_oldest_messages(&mut self) { let target = self.config.compaction.trigger_chars as usize; - let before_chars = self.message_buffer_chars(); + let before_chars = self.messages.total_chars(); if before_chars <= target { return; } - let mut current_chars = before_chars; let mut dropped = 0; - - while current_chars > target && !self.messages.is_empty() { - let msg_chars = self.messages[0] - .content - .iter() - .map(|block| match block { - ContentBlock::Text { text } => text.len(), - ContentBlock::ToolUse { input, .. } => estimate_json_size(input), - ContentBlock::ToolResult { content, .. } => content.len(), - }) - .sum::(); - current_chars = current_chars.saturating_sub(msg_chars); - self.messages.remove(0); + while self.messages.total_chars() > target && !self.messages.is_empty() { + self.messages.pop_front(); dropped += 1; } - let after_chars = self.message_buffer_chars(); + let after_chars = self.messages.total_chars(); warn!( "Compaction rate-limited: dropped {} oldest messages ({} -> {} chars)", dropped, before_chars, after_chars @@ -228,7 +202,7 @@ impl ClaudeAssistant { } let num_messages = self.messages.len(); - let buffer_chars = self.message_buffer_chars(); + let buffer_chars = self.messages.total_chars(); warn!( "Starting {} compaction: {} messages, {} chars", if is_automatic { "automatic" } else { "manual" }, @@ -251,11 +225,12 @@ impl ClaudeAssistant { last.set_cached(); } + let messages = self.messages.make_contiguous(); let request_body = MessagesRequest { model: &self.config.compaction.model, max_tokens: self.config.compaction.max_tokens, system: &system, - messages: &self.messages, + messages, tools: Vec::new(), }; @@ -363,11 +338,12 @@ impl ClaudeAssistant { return Ok(None); } + let messages = self.messages.make_contiguous(); let request_body = MessagesRequest { model: &self.config.claude_model, max_tokens: self.config.claude_max_tokens, system: &system, - messages: &self.messages, + messages, tools: tools.clone(), }; @@ -509,28 +485,6 @@ fn message_to_content(msg: ChatMessage) -> MessageContent { content: vec![ContentBlock::Text { text: text.into() }], } } - -/// Estimate the serialized size of a JSON Value without allocating. -fn estimate_json_size(value: &Value) -> usize { - match value { - Value::Null => 4, - Value::Bool(true) => 4, - Value::Bool(false) => 5, - Value::Number(n) => n.to_string().len(), - Value::String(s) => s.len() + 2, - 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)) - .sum::() - + obj.len().saturating_sub(1) - } - } -} - // ---- API Types ---- /// Request body for the Claude Messages API. diff --git a/signal-gateway-assistant/claude/src/message_buffer.rs b/signal-gateway-assistant/claude/src/message_buffer.rs new file mode 100644 index 0000000..208d1b1 --- /dev/null +++ b/signal-gateway-assistant/claude/src/message_buffer.rs @@ -0,0 +1,121 @@ +//! Message buffer with cached character count. + +use serde_json::Value; +use std::collections::VecDeque; +use std::fmt; + +use crate::{ContentBlock, MessageContent}; + +/// A buffer of messages with cached total character count. +/// +/// Uses `VecDeque` for efficient front removal during compaction. +pub struct MessageBuffer { + messages: VecDeque, + /// Cached total character count of all messages. + total_chars: usize, +} + +impl MessageBuffer { + /// Create a new empty message buffer. + pub fn new() -> Self { + Self { + messages: VecDeque::new(), + total_chars: 0, + } + } + + /// Push a message to the back of the buffer. + pub fn push(&mut self, msg: MessageContent) { + self.total_chars += message_chars(&msg); + self.messages.push_back(msg); + } + + /// Remove and return the first message, if any. + pub fn pop_front(&mut self) -> Option { + let msg = self.messages.pop_front()?; + self.total_chars = self.total_chars.saturating_sub(message_chars(&msg)); + Some(msg) + } + + /// Clear all messages from the buffer. + pub fn clear(&mut self) { + self.messages.clear(); + self.total_chars = 0; + } + + /// Returns true if the buffer is empty. + pub fn is_empty(&self) -> bool { + self.messages.is_empty() + } + + /// Returns the number of messages in the buffer. + pub fn len(&self) -> usize { + self.messages.len() + } + + /// Returns the cached total character count. + pub fn total_chars(&self) -> usize { + self.total_chars + } + + /// Returns a reference to the last message, if any. + pub fn last(&self) -> Option<&MessageContent> { + self.messages.back() + } + + /// Returns a slice of all messages for API requests. + /// + /// Note: This may require making the deque contiguous first. + pub fn make_contiguous(&mut self) -> &[MessageContent] { + self.messages.make_contiguous() + } +} + +impl Default for MessageBuffer { + fn default() -> Self { + Self::new() + } +} + +impl fmt::Debug for MessageBuffer { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.debug_struct("MessageBuffer") + .field("len", &self.messages.len()) + .field("total_chars", &self.total_chars) + .field("messages", &self.messages) + .finish() + } +} + +/// Calculate the character count for a single message. +fn message_chars(msg: &MessageContent) -> usize { + msg.content + .iter() + .map(|block| match block { + ContentBlock::Text { text } => text.len(), + ContentBlock::ToolUse { input, .. } => estimate_json_size(input), + ContentBlock::ToolResult { content, .. } => content.len(), + }) + .sum() +} + +/// Estimate the serialized JSON size of a value. +fn estimate_json_size(value: &Value) -> usize { + match value { + Value::Null => 4, + Value::Bool(true) => 4, + Value::Bool(false) => 5, + Value::Number(n) => n.to_string().len(), + Value::String(s) => s.len() + 2, + 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)) + .sum::() + + obj.len().saturating_sub(1) + } + } +}