refactor claude message buffer
This commit is contained in:
@@ -1,6 +1,9 @@
|
|||||||
//! Claude API implementation of the Assistant trait.
|
//! Claude API implementation of the Assistant trait.
|
||||||
|
|
||||||
|
mod message_buffer;
|
||||||
|
|
||||||
use conf::Conf;
|
use conf::Conf;
|
||||||
|
use message_buffer::MessageBuffer;
|
||||||
use serde::{Deserialize, Serialize};
|
use serde::{Deserialize, Serialize};
|
||||||
use serde_json::Value;
|
use serde_json::Value;
|
||||||
use signal_gateway_assistant::{Assistant, AssistantResponse, ChatMessage, Tool, ToolExecutor};
|
use signal_gateway_assistant::{Assistant, AssistantResponse, ChatMessage, Tool, ToolExecutor};
|
||||||
@@ -96,7 +99,7 @@ pub struct ClaudeAssistant {
|
|||||||
compaction_prompt: String,
|
compaction_prompt: String,
|
||||||
/// Summary of previous conversation history, wrapped in XML tags.
|
/// Summary of previous conversation history, wrapped in XML tags.
|
||||||
summary: String,
|
summary: String,
|
||||||
messages: Vec<MessageContent>,
|
messages: MessageBuffer,
|
||||||
tool_executor: Weak<dyn ToolExecutor>,
|
tool_executor: Weak<dyn ToolExecutor>,
|
||||||
/// Timestamp of the last automatic compaction (not user-requested).
|
/// Timestamp of the last automatic compaction (not user-requested).
|
||||||
last_auto_compaction: Option<std::time::Instant>,
|
last_auto_compaction: Option<std::time::Instant>,
|
||||||
@@ -135,32 +138,15 @@ impl ClaudeAssistant {
|
|||||||
system_prompts,
|
system_prompts,
|
||||||
compaction_prompt,
|
compaction_prompt,
|
||||||
summary: String::new(),
|
summary: String::new(),
|
||||||
messages: Vec::new(),
|
messages: MessageBuffer::new(),
|
||||||
tool_executor,
|
tool_executor,
|
||||||
last_auto_compaction: None,
|
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.
|
/// Check if automatic compaction should be triggered and handle it.
|
||||||
async fn maybe_compact(&mut self) {
|
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 {
|
if buffer_chars <= self.config.compaction.trigger_chars as usize {
|
||||||
return;
|
return;
|
||||||
}
|
}
|
||||||
@@ -184,31 +170,19 @@ impl ClaudeAssistant {
|
|||||||
/// Drop oldest messages until buffer is under the trigger threshold.
|
/// Drop oldest messages until buffer is under the trigger threshold.
|
||||||
fn drop_oldest_messages(&mut self) {
|
fn drop_oldest_messages(&mut self) {
|
||||||
let target = self.config.compaction.trigger_chars as usize;
|
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 {
|
if before_chars <= target {
|
||||||
return;
|
return;
|
||||||
}
|
}
|
||||||
|
|
||||||
let mut current_chars = before_chars;
|
|
||||||
let mut dropped = 0;
|
let mut dropped = 0;
|
||||||
|
while self.messages.total_chars() > target && !self.messages.is_empty() {
|
||||||
while current_chars > target && !self.messages.is_empty() {
|
self.messages.pop_front();
|
||||||
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);
|
|
||||||
dropped += 1;
|
dropped += 1;
|
||||||
}
|
}
|
||||||
|
|
||||||
let after_chars = self.message_buffer_chars();
|
let after_chars = self.messages.total_chars();
|
||||||
warn!(
|
warn!(
|
||||||
"Compaction rate-limited: dropped {} oldest messages ({} -> {} chars)",
|
"Compaction rate-limited: dropped {} oldest messages ({} -> {} chars)",
|
||||||
dropped, before_chars, after_chars
|
dropped, before_chars, after_chars
|
||||||
@@ -228,7 +202,7 @@ impl ClaudeAssistant {
|
|||||||
}
|
}
|
||||||
|
|
||||||
let num_messages = self.messages.len();
|
let num_messages = self.messages.len();
|
||||||
let buffer_chars = self.message_buffer_chars();
|
let buffer_chars = self.messages.total_chars();
|
||||||
warn!(
|
warn!(
|
||||||
"Starting {} compaction: {} messages, {} chars",
|
"Starting {} compaction: {} messages, {} chars",
|
||||||
if is_automatic { "automatic" } else { "manual" },
|
if is_automatic { "automatic" } else { "manual" },
|
||||||
@@ -251,11 +225,12 @@ impl ClaudeAssistant {
|
|||||||
last.set_cached();
|
last.set_cached();
|
||||||
}
|
}
|
||||||
|
|
||||||
|
let messages = self.messages.make_contiguous();
|
||||||
let request_body = MessagesRequest {
|
let request_body = MessagesRequest {
|
||||||
model: &self.config.compaction.model,
|
model: &self.config.compaction.model,
|
||||||
max_tokens: self.config.compaction.max_tokens,
|
max_tokens: self.config.compaction.max_tokens,
|
||||||
system: &system,
|
system: &system,
|
||||||
messages: &self.messages,
|
messages,
|
||||||
tools: Vec::new(),
|
tools: Vec::new(),
|
||||||
};
|
};
|
||||||
|
|
||||||
@@ -363,11 +338,12 @@ impl ClaudeAssistant {
|
|||||||
return Ok(None);
|
return Ok(None);
|
||||||
}
|
}
|
||||||
|
|
||||||
|
let messages = self.messages.make_contiguous();
|
||||||
let request_body = MessagesRequest {
|
let request_body = MessagesRequest {
|
||||||
model: &self.config.claude_model,
|
model: &self.config.claude_model,
|
||||||
max_tokens: self.config.claude_max_tokens,
|
max_tokens: self.config.claude_max_tokens,
|
||||||
system: &system,
|
system: &system,
|
||||||
messages: &self.messages,
|
messages,
|
||||||
tools: tools.clone(),
|
tools: tools.clone(),
|
||||||
};
|
};
|
||||||
|
|
||||||
@@ -509,28 +485,6 @@ fn message_to_content(msg: ChatMessage) -> MessageContent {
|
|||||||
content: vec![ContentBlock::Text { text: text.into() }],
|
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 ----
|
// ---- API Types ----
|
||||||
|
|
||||||
/// Request body for the Claude Messages API.
|
/// 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)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
Reference in New Issue
Block a user