use crate::llm::{Message, Role};
use serde::{Deserialize, Serialize};
use tracing::debug;
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
#[non_exhaustive]
pub enum MemoryStrategy {
#[default]
Full,
SlidingWindow,
HeadTail,
}
#[non_exhaustive]
pub struct ConversationBuffer {
messages: Vec<Message>,
max_messages: usize,
strategy: MemoryStrategy,
agent_id: String,
}
impl ConversationBuffer {
pub fn new(agent_id: impl Into<String>) -> Self {
Self {
messages: Vec::new(),
max_messages: 0,
strategy: MemoryStrategy::Full,
agent_id: agent_id.into(),
}
}
pub fn with_sliding_window(agent_id: impl Into<String>, max_messages: usize) -> Self {
Self {
messages: Vec::new(),
max_messages,
strategy: MemoryStrategy::SlidingWindow,
agent_id: agent_id.into(),
}
}
pub fn with_head_tail(agent_id: impl Into<String>, max_messages: usize) -> Self {
Self {
messages: Vec::new(),
max_messages,
strategy: MemoryStrategy::HeadTail,
agent_id: agent_id.into(),
}
}
pub fn add_user(&mut self, content: impl Into<String>) {
self.push(Message::new(Role::User, content));
}
pub fn add_assistant(&mut self, content: impl Into<String>) {
self.push(Message::new(Role::Assistant, content));
}
pub fn push(&mut self, message: Message) {
self.messages.push(message);
self.trim();
}
#[must_use]
pub fn messages(&self) -> &[Message] {
&self.messages
}
#[must_use]
pub fn to_vec(&self) -> Vec<Message> {
self.messages.clone()
}
#[must_use]
pub fn len(&self) -> usize {
self.messages.len()
}
#[must_use]
pub fn is_empty(&self) -> bool {
self.messages.is_empty()
}
pub fn clear(&mut self) {
self.messages.clear();
}
fn trim(&mut self) {
if self.max_messages == 0 || self.messages.len() <= self.max_messages {
return;
}
match self.strategy {
MemoryStrategy::Full => {} MemoryStrategy::SlidingWindow => {
let excess = self.messages.len() - self.max_messages;
debug!(
agent_id = %self.agent_id,
evicted = excess,
"sliding window: evicting oldest messages"
);
self.messages.drain(..excess);
}
MemoryStrategy::HeadTail => {
if self.max_messages == 1 {
let last = self.messages.len() - 1;
self.messages.drain(..last);
} else {
let keep_tail = self.max_messages - 1;
let remove_start = 1;
let remove_end = self.messages.len() - keep_tail;
if remove_end > remove_start {
debug!(
agent_id = %self.agent_id,
evicted = remove_end - remove_start,
"head-tail: evicting middle messages"
);
self.messages.drain(remove_start..remove_end);
}
}
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn full_strategy_keeps_all() {
let mut buf = ConversationBuffer::new("agent-1");
for i in 0..100 {
buf.add_user(format!("msg {i}"));
}
assert_eq!(buf.len(), 100);
}
#[test]
fn sliding_window_evicts_oldest() {
let mut buf = ConversationBuffer::with_sliding_window("agent-1", 4);
buf.add_user("a");
buf.add_assistant("b");
buf.add_user("c");
buf.add_assistant("d");
assert_eq!(buf.len(), 4);
buf.add_user("e");
assert_eq!(buf.len(), 4);
assert_eq!(buf.messages()[0].content, "b");
assert_eq!(buf.messages()[3].content, "e");
}
#[test]
fn head_tail_keeps_first_and_last() {
let mut buf = ConversationBuffer::with_head_tail("agent-1", 4);
for i in 0..10 {
buf.add_user(format!("msg-{i}"));
}
assert_eq!(buf.len(), 4);
assert_eq!(buf.messages()[0].content, "msg-0");
assert_eq!(buf.messages()[1].content, "msg-7");
assert_eq!(buf.messages()[2].content, "msg-8");
assert_eq!(buf.messages()[3].content, "msg-9");
}
#[test]
fn head_tail_max_one_keeps_last() {
let mut buf = ConversationBuffer::with_head_tail("agent-1", 1);
buf.add_user("first");
buf.add_user("second");
buf.add_user("third");
assert_eq!(buf.len(), 1);
assert_eq!(buf.messages()[0].content, "third");
}
#[test]
fn empty_buffer() {
let buf = ConversationBuffer::new("agent-1");
assert!(buf.is_empty());
assert_eq!(buf.len(), 0);
assert!(buf.messages().is_empty());
}
#[test]
fn clear_removes_all() {
let mut buf = ConversationBuffer::new("agent-1");
buf.add_user("hello");
buf.add_assistant("hi");
assert_eq!(buf.len(), 2);
buf.clear();
assert!(buf.is_empty());
}
#[test]
fn to_vec_clones_messages() {
let mut buf = ConversationBuffer::new("agent-1");
buf.add_user("hello");
let vec = buf.to_vec();
assert_eq!(vec.len(), 1);
assert_eq!(vec[0].content, "hello");
}
#[test]
fn sliding_window_no_trim_under_limit() {
let mut buf = ConversationBuffer::with_sliding_window("agent-1", 10);
buf.add_user("a");
buf.add_assistant("b");
assert_eq!(buf.len(), 2);
}
#[test]
fn strategy_serde_roundtrip() {
let strategies = [
MemoryStrategy::Full,
MemoryStrategy::SlidingWindow,
MemoryStrategy::HeadTail,
];
for s in &strategies {
let json = serde_json::to_string(s).unwrap();
let restored: MemoryStrategy = serde_json::from_str(&json).unwrap();
assert_eq!(*s, restored);
}
}
}