use std::collections::VecDeque;
use async_trait::async_trait;
use crate::{chat::ChatMessage, error::LLMError};
use super::{MemoryProvider, MemoryType};
#[derive(Debug, Clone)]
pub enum TrimStrategy {
Drop,
Summarize,
}
#[derive(Debug, Clone)]
pub struct SlidingWindowMemory {
messages: VecDeque<ChatMessage>,
window_size: usize,
trim_strategy: TrimStrategy,
needs_summary: bool,
}
impl SlidingWindowMemory {
pub fn new(window_size: usize) -> Self {
Self::with_strategy(window_size, TrimStrategy::Drop)
}
pub fn with_strategy(window_size: usize, strategy: TrimStrategy) -> Self {
if window_size == 0 {
panic!("Window size must be greater than 0");
}
Self {
messages: VecDeque::with_capacity(window_size),
window_size,
trim_strategy: strategy,
needs_summary: false,
}
}
pub fn window_size(&self) -> usize {
self.window_size
}
pub fn messages(&self) -> Vec<ChatMessage> {
Vec::from(self.messages.clone())
}
pub fn recent_messages(&self, limit: usize) -> Vec<ChatMessage> {
let len = self.messages.len();
let start = len.saturating_sub(limit);
self.messages.range(start..).cloned().collect()
}
pub fn needs_summary(&self) -> bool {
self.needs_summary
}
pub fn mark_for_summary(&mut self) {
self.needs_summary = true;
}
pub fn replace_with_summary(&mut self, summary: String) {
self.messages.clear();
self.messages.push_back(
crate::chat::ChatMessage::assistant()
.content(summary)
.build(),
);
self.needs_summary = false;
}
}
#[async_trait]
impl MemoryProvider for SlidingWindowMemory {
async fn remember(&mut self, message: &ChatMessage) -> Result<(), LLMError> {
if self.messages.len() >= self.window_size {
match self.trim_strategy {
TrimStrategy::Drop => {
self.messages.pop_front();
}
TrimStrategy::Summarize => {
self.mark_for_summary();
}
}
}
self.messages.push_back(message.clone());
Ok(())
}
async fn recall(
&self,
_query: &str,
limit: Option<usize>,
) -> Result<Vec<ChatMessage>, LLMError> {
let limit = limit.unwrap_or(self.messages.len());
Ok(self.recent_messages(limit))
}
async fn clear(&mut self) -> Result<(), LLMError> {
self.messages.clear();
Ok(())
}
fn memory_type(&self) -> MemoryType {
MemoryType::SlidingWindow
}
fn size(&self) -> usize {
self.messages.len()
}
fn needs_summary(&self) -> bool {
self.needs_summary
}
fn mark_for_summary(&mut self) {
self.needs_summary = true;
}
fn replace_with_summary(&mut self, summary: String) {
self.messages.clear();
self.messages.push_back(
crate::chat::ChatMessage::assistant()
.content(summary)
.build(),
);
self.needs_summary = false;
}
}