use std::sync::Arc;
use crate::core::language_models::BaseChatModel;
use crate::core::token_counter::{TiktokenCounter, TokenCounter};
use crate::memory::base::MemoryError;
use crate::schema::Message;
use super::trimmer::Strategy;
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::new().expect("tiktoken cl100k_base 加载失败")),
strategy: Strategy::Truncate,
}
}
pub fn with_strategy(max_tokens: usize, strategy: Strategy<M>) -> Self {
Self {
max_tokens,
counter: Arc::new(TiktokenCounter::new().expect("tiktoken cl100k_base 加载失败")),
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 base_tokens = self.counter.count_messages(&system_messages) as usize;
let msg_incremental_costs: Vec<usize> = other_messages
.iter()
.map(|m| {
let single_count = self.counter.count_messages(std::slice::from_ref(m)) as usize;
single_count.saturating_sub(2)
})
.collect();
let mut kept: Vec<Message> = Vec::new();
let mut running_tokens = base_tokens;
for (msg, cost) in other_messages
.into_iter()
.rev()
.zip(msg_incremental_costs.into_iter().rev())
{
if running_tokens + cost <= self.max_tokens {
running_tokens += cost;
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)
}
}