#![forbid(unsafe_code)]
#![deny(missing_docs)]
#![deny(unreachable_pub)]
use async_trait::async_trait;
use behest_core::message::Message;
use behest_core::token::estimate_messages_tokens;
use serde::{Deserialize, Serialize};
#[async_trait]
pub trait ConversationMemory: Send + Sync {
async fn load(&self, session_id: &str) -> Result<Vec<Message>, String>;
async fn append(&self, session_id: &str, messages: Vec<Message>) -> Result<(), String>;
async fn active_window(&self, session_id: &str) -> Result<Vec<Message>, String>;
async fn clear(&self, session_id: &str) -> Result<(), String>;
}
#[async_trait]
pub trait DemotionHook: Send + Sync {
async fn demote(&self, session_id: &str, messages: Vec<Message>) -> Result<usize, String>;
}
#[async_trait]
pub trait Compactor: Send + Sync {
async fn compact(&self, messages: Vec<Message>) -> Result<String, String>;
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct MemoryPolicy {
pub max_window_tokens: usize,
pub max_window_messages: usize,
pub keep_tail_turns: usize,
pub demotion_enabled: bool,
pub compaction_enabled: bool,
pub compaction_trigger_ratio: f64,
}
impl Default for MemoryPolicy {
fn default() -> Self {
Self {
max_window_tokens: 100_000,
max_window_messages: 200,
keep_tail_turns: 3,
demotion_enabled: true,
compaction_enabled: true,
compaction_trigger_ratio: 0.8,
}
}
}
#[derive(Debug, Clone)]
pub enum MemoryEvent {
Demoted {
count: usize,
session_id: String,
},
Compacted {
original_count: usize,
summary_length: usize,
},
Restored {
count: usize,
},
}
pub struct ActiveWindow {
messages: Vec<Message>,
policy: MemoryPolicy,
demotion_hook: Option<Box<dyn DemotionHook>>,
compactor: Option<Box<dyn Compactor>>,
}
impl ActiveWindow {
#[must_use]
pub fn new(policy: MemoryPolicy) -> Self {
Self {
messages: Vec::new(),
policy,
demotion_hook: None,
compactor: None,
}
}
pub fn with_demotion_hook(mut self, hook: Box<dyn DemotionHook>) -> Self {
self.demotion_hook = Some(hook);
self
}
pub fn with_compactor(mut self, compactor: Box<dyn Compactor>) -> Self {
self.compactor = Some(compactor);
self
}
pub fn push(&mut self, msg: Message) -> Vec<MemoryEvent> {
self.messages.push(msg);
self.trim_if_needed()
}
#[must_use]
pub fn messages(&self) -> &[Message] {
&self.messages
}
pub fn trim_if_needed(&mut self) -> Vec<MemoryEvent> {
let token_count = estimate_messages_tokens(&self.messages);
let over_token_limit = token_count > self.policy.max_window_tokens;
let over_message_limit = self.messages.len() > self.policy.max_window_messages;
if !over_token_limit && !over_message_limit {
return Vec::new();
}
let keep_tail = self.policy.keep_tail_turns.min(self.messages.len());
let tail = self.messages.split_off(self.messages.len() - keep_tail);
let evicted = std::mem::replace(&mut self.messages, tail);
let mut events = Vec::new();
if self.policy.demotion_enabled {
events.push(MemoryEvent::Demoted {
count: evicted.len(),
session_id: String::new(),
});
}
if self.policy.compaction_enabled && !evicted.is_empty() {
events.push(MemoryEvent::Compacted {
original_count: evicted.len(),
summary_length: 0,
});
}
events
}
}
pub struct InMemoryConversationMemory {
messages: std::sync::RwLock<std::collections::HashMap<String, Vec<Message>>>,
}
impl InMemoryConversationMemory {
#[must_use]
pub fn new() -> Self {
Self {
messages: std::sync::RwLock::new(std::collections::HashMap::new()),
}
}
}
impl Default for InMemoryConversationMemory {
fn default() -> Self {
Self::new()
}
}
#[async_trait]
impl ConversationMemory for InMemoryConversationMemory {
async fn load(&self, session_id: &str) -> Result<Vec<Message>, String> {
let map = self
.messages
.read()
.map_err(|e| format!("lock error: {e}"))?;
Ok(map.get(session_id).cloned().unwrap_or_default())
}
async fn append(&self, session_id: &str, messages: Vec<Message>) -> Result<(), String> {
let mut map = self
.messages
.write()
.map_err(|e| format!("lock error: {e}"))?;
map.entry(session_id.to_string())
.or_default()
.extend(messages);
Ok(())
}
async fn active_window(&self, session_id: &str) -> Result<Vec<Message>, String> {
self.load(session_id).await
}
async fn clear(&self, session_id: &str) -> Result<(), String> {
let mut map = self
.messages
.write()
.map_err(|e| format!("lock error: {e}"))?;
map.remove(session_id);
Ok(())
}
}
#[cfg(test)]
#[allow(clippy::unwrap_used)]
mod tests {
use super::*;
#[test]
fn active_window_trims_on_token_limit() {
let policy = MemoryPolicy {
max_window_tokens: 10,
max_window_messages: 100,
..Default::default()
};
let mut window = ActiveWindow::new(policy);
let long_msg = Message::user_text(
"This is a fairly long message that should consume many tokens in the estimation",
);
let events = window.push(long_msg);
assert!(!events.is_empty(), "should trigger demotion/compaction");
}
#[test]
fn active_window_trims_on_message_limit() {
let policy = MemoryPolicy {
max_window_tokens: 100_000,
max_window_messages: 2,
keep_tail_turns: 1,
..Default::default()
};
let mut window = ActiveWindow::new(policy);
window.push(Message::user_text("msg1"));
window.push(Message::user_text("msg2"));
let events = window.push(Message::user_text("msg3"));
assert!(!events.is_empty());
}
#[test]
fn active_window_respects_tail_turns() {
let policy = MemoryPolicy {
max_window_tokens: 10,
max_window_messages: 100,
keep_tail_turns: 2,
..Default::default()
};
let mut window = ActiveWindow::new(policy);
window.push(Message::user_text("a"));
window.push(Message::user_text("b"));
window.push(Message::user_text("c"));
window.push(Message::user_text("d"));
let msgs = window.messages();
assert_eq!(msgs.len(), 2);
}
#[test]
fn empty_window_no_trim() {
let policy = MemoryPolicy::default();
let window = ActiveWindow::new(policy);
assert!(window.messages().is_empty());
}
#[tokio::test]
async fn in_memory_conversation_memory() {
let mem = InMemoryConversationMemory::new();
let sid = "test-session";
let msgs = mem.load(sid).await.unwrap();
assert!(msgs.is_empty());
mem.append(sid, vec![Message::user_text("hello")])
.await
.unwrap();
let msgs = mem.load(sid).await.unwrap();
assert_eq!(msgs.len(), 1);
mem.clear(sid).await.unwrap();
let msgs = mem.load(sid).await.unwrap();
assert!(msgs.is_empty());
}
}