refactor claude message buffer

This commit is contained in:
Chris Beck
2025-12-14 12:20:29 -07:00
parent 4a1b1ecafb
commit c66228f398
2 changed files with 136 additions and 61 deletions
+15 -61
View File
@@ -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<MessageContent>,
messages: MessageBuffer,
tool_executor: Weak<dyn ToolExecutor>,
/// Timestamp of the last automatic compaction (not user-requested).
last_auto_compaction: Option<std::time::Instant>,
@@ -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::<usize>()
})
.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::<usize>();
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::<usize>() + arr.len().saturating_sub(1)
}
Value::Object(obj) => {
2 + obj
.iter()
.map(|(k, v)| k.len() + 3 + estimate_json_size(v))
.sum::<usize>()
+ obj.len().saturating_sub(1)
}
}
}
// ---- API Types ----
/// Request body for the Claude Messages API.
@@ -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<MessageContent>,
/// 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<MessageContent> {
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::<usize>() + arr.len().saturating_sub(1)
}
Value::Object(obj) => {
2 + obj
.iter()
.map(|(k, v)| k.len() + 3 + estimate_json_size(v))
.sum::<usize>()
+ obj.len().saturating_sub(1)
}
}
}