use std::sync::Arc;
use crate::core::language_models::BaseChatModel;
use crate::core::token_counter::{TiktokenCounter, TokenCounter};
use crate::schema::Message;
use super::base::MemoryError;
const DEFAULT_SUMMARY_PROMPT: &str = "\
Summarize the following conversation concisely, preserving key facts, \
decisions, and context. Write the summary in the same language as the conversation.
Conversation:
{conversation}
Summary:";
#[derive(Debug)]
pub enum Strategy<M: BaseChatModel = crate::language_models::OpenAIChat> {
Truncate,
Summarize {
llm: Arc<M>,
summary_prompt: String,
},
}
impl<M: BaseChatModel> Strategy<M> {
pub fn summarize(llm: M) -> Self {
Strategy::Summarize {
llm: Arc::new(llm),
summary_prompt: DEFAULT_SUMMARY_PROMPT.to_string(),
}
}
pub fn summarize_with_prompt(llm: M, prompt: impl Into<String>) -> Self {
Strategy::Summarize {
llm: Arc::new(llm),
summary_prompt: prompt.into(),
}
}
}
pub struct ContextWindow<M: BaseChatModel = crate::language_models::OpenAIChat> {
max_tokens: usize,
counter: Arc<dyn TokenCounter>,
strategy: Strategy<M>,
}
impl<M: BaseChatModel> ContextWindow<M> {
pub fn new(max_tokens: usize) -> Self {
Self {
max_tokens,
counter: Arc::new(TiktokenCounter::default()),
strategy: Strategy::Truncate,
}
}
pub fn with_strategy(max_tokens: usize, strategy: Strategy<M>) -> Self {
Self {
max_tokens,
counter: Arc::new(TiktokenCounter::default()),
strategy,
}
}
pub fn with_max_tokens(max_tokens: usize) -> Self {
Self::new(max_tokens)
}
pub fn with_counter(mut self, counter: Arc<dyn TokenCounter>) -> Self {
self.counter = counter;
self
}
pub fn max_tokens(&self) -> usize {
self.max_tokens
}
pub async fn fit(&self, messages: Vec<Message>) -> Result<Vec<Message>, MemoryError> {
let total_tokens = self.counter.count_messages(&messages) as usize;
if total_tokens <= self.max_tokens {
return Ok(messages);
}
match &self.strategy {
Strategy::Truncate => self.truncate(messages),
Strategy::Summarize {
llm,
summary_prompt,
} => self.summarize(messages, llm, summary_prompt).await,
}
}
fn truncate(&self, messages: Vec<Message>) -> Result<Vec<Message>, MemoryError> {
let mut system_messages: Vec<Message> = Vec::new();
let mut other_messages: Vec<Message> = Vec::new();
for msg in messages {
if matches!(msg.message_type, crate::schema::MessageType::System) {
system_messages.push(msg);
} else {
other_messages.push(msg);
}
}
let mut kept: Vec<Message> = Vec::new();
for msg in other_messages.into_iter().rev() {
let mut candidate = system_messages.clone();
candidate.push(msg.clone());
candidate.extend(kept.iter().cloned());
let tokens = self.counter.count_messages(&candidate) as usize;
if tokens <= self.max_tokens {
kept.push(msg);
} else {
break;
}
}
kept.reverse();
let mut result = system_messages;
result.extend(kept);
Ok(result)
}
async fn summarize(
&self,
messages: Vec<Message>,
llm: &Arc<M>,
summary_prompt: &str,
) -> Result<Vec<Message>, MemoryError> {
let mut system_messages: Vec<Message> = Vec::new();
let mut other_messages: Vec<Message> = Vec::new();
for msg in messages {
if matches!(msg.message_type, crate::schema::MessageType::System) {
system_messages.push(msg);
} else {
other_messages.push(msg);
}
}
if other_messages.is_empty() {
return Ok(system_messages);
}
let mut keep_from_idx = other_messages.len();
for i in 0..other_messages.len() {
let recent = &other_messages[i..];
let mut candidate = system_messages.clone();
candidate.push(Message::system("summary placeholder"));
candidate.extend(recent.iter().cloned());
let tokens = self.counter.count_messages(&candidate) as usize;
if tokens <= self.max_tokens {
keep_from_idx = i;
break;
}
}
if keep_from_idx >= other_messages.len() {
return self.truncate(system_messages);
}
let to_summarize = &other_messages[..keep_from_idx];
let to_keep = &other_messages[keep_from_idx..];
if to_summarize.is_empty() {
let mut result = system_messages;
result.extend(to_keep.to_vec());
return Ok(result);
}
let conversation_text = to_summarize
.iter()
.map(|msg| {
let role = match msg.message_type {
crate::schema::MessageType::Human => "Human",
crate::schema::MessageType::AI => "AI",
crate::schema::MessageType::System => "System",
crate::schema::MessageType::Tool { .. } => "Tool",
};
format!("{}: {}", role, msg.content)
})
.collect::<Vec<_>>()
.join("\n");
let prompt = summary_prompt.replace("{conversation}", &conversation_text);
let summary_messages = vec![Message::human(&prompt)];
let result = llm
.invoke(summary_messages, None)
.await
.map_err(|e| MemoryError::SaveError(format!("LLM summarization failed: {}", e)))?;
let summary_message = Message::system(format!("[Conversation Summary] {}", result.content));
let mut final_messages = system_messages;
final_messages.push(summary_message);
final_messages.extend(to_keep.to_vec());
let final_tokens = self.counter.count_messages(&final_messages) as usize;
if final_tokens > self.max_tokens {
return self.truncate(final_messages);
}
Ok(final_messages)
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::core::language_models::{BaseLanguageModel, LLMResult};
use crate::core::runnables::{Runnable, RunnableConfig};
use crate::language_models::openai::{OpenAIChat, OpenAIConfig};
use crate::schema::MessageType;
use async_trait::async_trait;
use futures_util::Stream;
use std::pin::Pin;
use tokio::sync::Mutex;
#[derive(Debug)]
struct CharTokenCounter;
impl TokenCounter for CharTokenCounter {
fn count_tokens(&self, text: &str) -> u32 {
text.len() as u32
}
fn count_messages(&self, messages: &[Message]) -> u32 {
let mut total = 0u32;
for msg in messages {
total += 4; total += self.count_tokens(&msg.content);
}
total += 2; total
}
}
fn char_counter() -> Arc<dyn TokenCounter> {
Arc::new(CharTokenCounter)
}
#[derive(Debug)]
struct MockLLM {
responses: Arc<Mutex<Vec<String>>>,
}
impl MockLLM {
fn new(responses: Vec<String>) -> Self {
Self {
responses: Arc::new(Mutex::new(responses)),
}
}
}
impl BaseLanguageModel<Vec<Message>, LLMResult> for MockLLM {
fn model_name(&self) -> &str {
"mock-llm"
}
fn get_num_tokens(&self, text: &str) -> usize {
text.len()
}
fn with_temperature(self, _temp: f32) -> Self
where
Self: Sized,
{
self
}
fn with_max_tokens(self, _max: usize) -> Self
where
Self: Sized,
{
self
}
}
#[async_trait]
impl Runnable<Vec<Message>, LLMResult> for MockLLM {
type Error = std::convert::Infallible;
async fn invoke(
&self,
_input: Vec<Message>,
_config: Option<RunnableConfig>,
) -> Result<LLMResult, Self::Error> {
let mut responses = self.responses.lock().await;
let content = responses.pop().unwrap_or_else(|| "Summary".to_string());
Ok(LLMResult {
content,
model: "mock-llm".to_string(),
token_usage: None,
tool_calls: None,
})
}
}
#[async_trait]
impl BaseChatModel for MockLLM {
async fn chat(
&self,
messages: Vec<Message>,
config: Option<RunnableConfig>,
) -> Result<LLMResult, Self::Error> {
self.invoke(messages, config).await
}
async fn stream_chat(
&self,
_messages: Vec<Message>,
_config: Option<RunnableConfig>,
) -> Result<Pin<Box<dyn Stream<Item = Result<String, Self::Error>> + Send>>, Self::Error>
{
unimplemented!("stream_chat not needed for tests")
}
}
fn make_messages(contents: &[(&str, &str)]) -> Vec<Message> {
contents
.iter()
.map(|(role, content)| match *role {
"system" => Message::system(*content),
"human" => Message::human(*content),
"ai" => Message::ai(*content),
_ => Message::human(*content),
})
.collect()
}
#[test]
fn test_new_creates_truncate_strategy() {
let cw: ContextWindow<OpenAIChat> = ContextWindow::new(4096);
assert_eq!(cw.max_tokens(), 4096);
}
#[test]
fn test_with_max_tokens() {
let cw: ContextWindow<OpenAIChat> = ContextWindow::with_max_tokens(8192);
assert_eq!(cw.max_tokens(), 8192);
}
#[tokio::test]
async fn test_fit_under_limit_returns_as_is() {
let cw: ContextWindow<OpenAIChat> = ContextWindow::new(1000)
.with_counter(char_counter());
let messages = make_messages(&[
("human", "Hello"),
("ai", "Hi there"),
]);
let result = cw.fit(messages).await.unwrap();
assert_eq!(result.len(), 2);
}
#[tokio::test]
async fn test_fit_empty_messages() {
let cw: ContextWindow<OpenAIChat> = ContextWindow::new(100)
.with_counter(char_counter());
let result = cw.fit(vec![]).await.unwrap();
assert!(result.is_empty());
}
#[tokio::test]
async fn test_truncate_preserves_system_messages() {
let cw: ContextWindow<OpenAIChat> = ContextWindow::new(30)
.with_counter(char_counter());
let messages = make_messages(&[
("system", "You are"),
("human", "Q1?"),
("ai", "A1!"),
("human", "Q2?"),
("ai", "A2!"),
]);
let result = cw.fit(messages).await.unwrap();
assert!(result.iter().any(|m| matches!(m.message_type, MessageType::System)));
assert!(result.iter().any(|m| m.content == "A2!"));
}
#[tokio::test]
async fn test_truncate_drops_oldest_first() {
let cw: ContextWindow<OpenAIChat> = ContextWindow::new(25)
.with_counter(char_counter());
let messages = make_messages(&[
("system", "Sys"),
("human", "Old question here"),
("ai", "Old answer here"),
("human", "New"),
("ai", "Ans"),
]);
let result = cw.fit(messages).await.unwrap();
assert!(result.iter().any(|m| m.content == "Sys"));
assert!(result.iter().any(|m| m.content == "Ans"));
assert!(!result.iter().any(|m| m.content == "Old question here"));
}
#[tokio::test]
async fn test_truncate_only_system_messages() {
let cw: ContextWindow<OpenAIChat> = ContextWindow::new(20)
.with_counter(char_counter());
let messages = make_messages(&[
("system", "Hello"),
]);
let result = cw.fit(messages).await.unwrap();
assert_eq!(result.len(), 1);
assert_eq!(result[0].content, "Hello");
}
#[tokio::test]
async fn test_truncate_system_only_over_budget() {
let cw: ContextWindow<OpenAIChat> = ContextWindow::new(5)
.with_counter(char_counter());
let messages = make_messages(&[
("system", "Very long system prompt that exceeds budget"),
]);
let result = cw.fit(messages).await.unwrap();
assert_eq!(result.len(), 1);
}
#[tokio::test]
async fn test_summarize_replaces_old_messages() {
let mock_llm = MockLLM::new(vec!["S.".to_string()]);
let cw = ContextWindow::with_strategy(50, Strategy::summarize(mock_llm))
.with_counter(char_counter());
let messages = make_messages(&[
("system", "S"),
("human", "Q1"),
("ai", "A1"),
("human", "Q2"),
("ai", "A2"),
("human", "Q3"),
("ai", "A3"),
("human", "Q4"),
("ai", "A4"),
]);
let result = cw.fit(messages).await.unwrap();
assert!(result.iter().any(|m| m.content == "S"));
let summary_msgs: Vec<&Message> = result
.iter()
.filter(|m| m.content.starts_with("[Conversation Summary]"))
.collect();
assert_eq!(summary_msgs.len(), 1);
assert!(summary_msgs[0].content.contains("S"));
}
#[tokio::test]
async fn test_summarize_preserves_recent_messages() {
let mock_llm = MockLLM::new(vec!["S.".to_string()]);
let cw = ContextWindow::with_strategy(50, Strategy::summarize(mock_llm))
.with_counter(char_counter());
let messages = make_messages(&[
("system", "S"),
("human", "Q1"),
("ai", "A1"),
("human", "Q2"),
("ai", "A2"),
("human", "Q3"),
("ai", "A3"),
("human", "Q4"),
("ai", "A4"),
]);
let result = cw.fit(messages).await.unwrap();
assert!(result.iter().any(|m| m.content == "Q4"));
assert!(result.iter().any(|m| m.content == "A4"));
}
#[tokio::test]
async fn test_summarize_with_custom_prompt() {
let mock_llm = MockLLM::new(vec!["O.".to_string()]);
let cw = ContextWindow::with_strategy(
50,
Strategy::summarize_with_prompt(
mock_llm,
"Please compress: {conversation}\nCompressed:",
),
)
.with_counter(char_counter());
let messages = make_messages(&[
("system", "S"),
("human", "Q1"),
("ai", "A1"),
("human", "Q2"),
("ai", "A2"),
("human", "Q3"),
("ai", "A3"),
("human", "Q4"),
("ai", "A4"),
]);
let result = cw.fit(messages).await.unwrap();
let summary_msgs: Vec<&Message> = result
.iter()
.filter(|m| m.content.starts_with("[Conversation Summary]"))
.collect();
assert_eq!(summary_msgs.len(), 1);
assert!(summary_msgs[0].content.contains("O"));
}
#[tokio::test]
async fn test_summarize_no_non_system_messages() {
let mock_llm = MockLLM::new(vec!["Should not be called".to_string()]);
let cw: ContextWindow<MockLLM> = ContextWindow::with_strategy(50, Strategy::summarize(mock_llm))
.with_counter(char_counter());
let messages = make_messages(&[
("system", "S"),
]);
let result = cw.fit(messages).await.unwrap();
assert_eq!(result.len(), 1);
assert_eq!(result[0].content, "S");
}
#[tokio::test]
async fn test_strategy_truncate_enum() {
let cw = ContextWindow::with_strategy(100, Strategy::<OpenAIChat>::Truncate)
.with_counter(char_counter());
let messages = make_messages(&[
("human", "Hello"),
("ai", "World"),
]);
let result = cw.fit(messages).await.unwrap();
assert_eq!(result.len(), 2);
}
#[test]
fn test_strategy_summarize_new() {
let config = OpenAIConfig::default();
let llm = OpenAIChat::new(config);
let strategy: Strategy<OpenAIChat> = Strategy::summarize(llm);
if let Strategy::Summarize { summary_prompt, .. } = &strategy {
assert!(summary_prompt.contains("{conversation}"));
} else {
panic!("Expected Summarize variant");
}
}
#[test]
fn test_strategy_summarize_with_custom_prompt() {
let config = OpenAIConfig::default();
let llm = OpenAIChat::new(config);
let custom = "Custom: {conversation} ->";
let strategy: Strategy<OpenAIChat> = Strategy::summarize_with_prompt(llm, custom);
if let Strategy::Summarize { summary_prompt, .. } = &strategy {
assert_eq!(summary_prompt, custom);
} else {
panic!("Expected Summarize variant");
}
}
#[tokio::test]
async fn test_fit_with_real_tiktoken_counter() {
let cw: ContextWindow<OpenAIChat> = ContextWindow::new(4096);
let messages = make_messages(&[
("system", "You are a helpful assistant."),
("human", "Hello!"),
("ai", "Hi there! How can I help you?"),
]);
let result = cw.fit(messages).await.unwrap();
assert_eq!(result.len(), 3);
}
#[tokio::test]
async fn test_truncate_preserves_order() {
let cw: ContextWindow<OpenAIChat> = ContextWindow::new(40)
.with_counter(char_counter());
let messages = make_messages(&[
("system", "Sys"),
("human", "Old"),
("ai", "OldA"),
("human", "New"),
("ai", "NewA"),
]);
let result = cw.fit(messages).await.unwrap();
let types: Vec<&str> = result.iter().map(|m| m.type_str()).collect();
assert_eq!(types[0], "system");
for i in 1..types.len() {
if i + 1 < types.len() {
}
}
}
#[tokio::test]
async fn test_summarize_fallback_to_truncate() {
let mock_llm = MockLLM::new(vec!["A very long summary that will not fit in the small budget.".to_string()]);
let cw: ContextWindow<MockLLM> = ContextWindow::with_strategy(20, Strategy::summarize(mock_llm))
.with_counter(char_counter());
let messages = make_messages(&[
("system", "S"),
("human", "Q1"),
("ai", "A1"),
("human", "Q2"),
("ai", "A2"),
]);
let result = cw.fit(messages).await.unwrap();
assert!(!result.is_empty());
}
}