use std::fmt;
use std::sync::{Arc, Mutex};
use super::{Memory, MemoryError};
use crate::message::{ContentBlock, Message};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct Budget {
pub max_tokens: Option<usize>,
pub max_rounds: Option<usize>,
}
impl Default for Budget {
fn default() -> Self {
Self {
max_tokens: None,
max_rounds: None,
}
}
}
impl Budget {
pub fn tokens(max_tokens: usize) -> Self {
Self {
max_tokens: Some(max_tokens),
max_rounds: None,
}
}
pub fn rounds(max_rounds: usize) -> Self {
Self {
max_tokens: None,
max_rounds: Some(max_rounds),
}
}
pub fn both(max_tokens: usize, max_rounds: usize) -> Self {
Self {
max_tokens: Some(max_tokens),
max_rounds: Some(max_rounds),
}
}
}
#[async_trait::async_trait]
pub trait TokenCounter: Send + Sync {
async fn count(&self, text: &str) -> Result<usize, MemoryError>;
}
#[derive(Debug, Default, Clone, Copy, PartialEq, Eq)]
pub struct CharTokenCounter;
#[async_trait::async_trait]
impl TokenCounter for CharTokenCounter {
async fn count(&self, text: &str) -> Result<usize, MemoryError> {
let (mut cjk, mut other) = (0usize, 0usize);
for c in text.chars() {
if is_cjk(c) {
cjk += 1;
} else {
other += 1;
}
}
Ok(cjk + other.div_ceil(4))
}
}
fn is_cjk(c: char) -> bool {
matches!(
c as u32,
0x3400..=0x4DBF
| 0x4E00..=0x9FFF
| 0xF900..=0xFAFF
| 0x20000..=0x2FA1F
)
}
pub(crate) async fn count_message(
counter: &dyn TokenCounter,
message: &Message,
) -> Result<usize, MemoryError> {
match message {
Message::System(s) => counter.count(s).await,
Message::User(blocks) => {
let mut total = 0usize;
for b in blocks {
match b {
ContentBlock::Text(t) => total += counter.count(t).await?,
ContentBlock::Image(_) | ContentBlock::Wire(_) => {}
}
}
Ok(total)
}
Message::Assistant {
content,
reasoning,
tool_calls,
} => {
let mut total = counter.count(content).await?;
if let Some(r) = reasoning {
total += counter.count(r).await?;
}
for tc in tool_calls {
total += counter.count(&tc.arguments).await?;
}
Ok(total)
}
Message::ToolResult { content, .. } => counter.count(content).await,
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct TrimResult {
pub messages: Vec<Message>,
pub replace: bool,
}
#[async_trait::async_trait]
pub trait TrimStrategy: Send + Sync {
async fn trim(
&self,
messages: &[Message],
budget: &Budget,
counter: &dyn TokenCounter,
) -> Result<TrimResult, MemoryError>;
async fn trim_with_counts(
&self,
messages: &[Message],
_counts: &[usize],
budget: &Budget,
counter: &dyn TokenCounter,
) -> Result<TrimResult, MemoryError> {
self.trim(messages, budget, counter).await
}
}
#[derive(Debug, Default, Clone, Copy, PartialEq, Eq)]
pub struct WindowDrop;
#[async_trait::async_trait]
impl TrimStrategy for WindowDrop {
async fn trim(
&self,
messages: &[Message],
budget: &Budget,
counter: &dyn TokenCounter,
) -> Result<TrimResult, MemoryError> {
let mut counts = Vec::with_capacity(messages.len());
for m in messages {
counts.push(count_message(counter, m).await?);
}
Ok(TrimResult {
messages: window_from_counts(messages, &counts, budget),
replace: false,
})
}
async fn trim_with_counts(
&self,
messages: &[Message],
counts: &[usize],
budget: &Budget,
_counter: &dyn TokenCounter,
) -> Result<TrimResult, MemoryError> {
Ok(TrimResult {
messages: window_from_counts(messages, counts, budget),
replace: false,
})
}
}
pub(crate) type Round = (usize, usize, usize);
pub(crate) fn split_rounds(messages: &[Message], counts: &[usize]) -> Vec<Round> {
debug_assert_eq!(messages.len(), counts.len());
if messages.is_empty() {
return Vec::new();
}
let mut rounds = Vec::new();
let mut start = 0usize;
let mut tokens = counts[0];
for (i, message) in messages.iter().enumerate().skip(1) {
if matches!(message, Message::User(_)) {
rounds.push((start, i, tokens));
start = i;
tokens = counts[i];
} else {
tokens += counts[i];
}
}
rounds.push((start, messages.len(), tokens));
rounds
}
pub(crate) fn keep_rounds(rounds: &[Round], budget: &Budget, reserved: usize) -> usize {
let mut keep = 0usize;
let mut sum = 0usize;
for &(_, _, t) in rounds.iter().rev() {
if keep > 0
&& let Some(limit) = budget.max_tokens
&& sum + t > limit.saturating_sub(reserved)
{
break;
}
sum += t;
keep += 1;
}
if let Some(limit) = budget.max_rounds {
keep = keep.min(limit.max(1));
}
keep
}
fn window_from_counts(messages: &[Message], counts: &[usize], budget: &Budget) -> Vec<Message> {
if messages.is_empty() {
return Vec::new();
}
let rounds = split_rounds(messages, counts);
let keep = keep_rounds(&rounds, budget, 0);
let start_index = rounds[rounds.len() - keep].0;
messages[start_index..].to_vec()
}
pub struct WindowMemory {
inner: Mutex<WindowInner>,
budget: Budget,
counter: Box<dyn TokenCounter>,
strategy: Arc<dyn TrimStrategy>,
}
struct WindowInner {
messages: Vec<Message>,
protected: Vec<bool>,
tokens: Vec<usize>,
total_tokens: Option<usize>,
user_count: usize,
}
impl fmt::Debug for WindowMemory {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
let message_count = self
.inner
.lock()
.expect("WindowMemory internal lock poisoned")
.messages
.len();
f.debug_struct("WindowMemory")
.field("budget", &self.budget)
.field("message_count", &message_count)
.field("counter", &"Box<dyn TokenCounter>")
.field("strategy", &"Arc<dyn TrimStrategy>")
.finish()
}
}
impl WindowMemory {
pub fn new(max_tokens: usize) -> Self {
Self {
inner: Mutex::new(WindowInner {
messages: Vec::new(),
protected: Vec::new(),
tokens: Vec::new(),
total_tokens: Some(0),
user_count: 0,
}),
budget: Budget::tokens(max_tokens),
counter: Box::new(CharTokenCounter),
strategy: Arc::new(WindowDrop),
}
}
pub fn with_max_rounds(mut self, max_rounds: usize) -> Self {
self.budget.max_rounds = Some(max_rounds);
self
}
pub fn with_token_counter(mut self, counter: Box<dyn TokenCounter>) -> Self {
self.counter = counter;
let mut inner = self
.inner
.lock()
.expect("WindowMemory internal lock poisoned");
inner.total_tokens = None;
drop(inner);
self
}
pub fn with_strategy(mut self, strategy: Arc<dyn TrimStrategy>) -> Self {
self.strategy = strategy;
self
}
pub fn set_max_tokens(&mut self, max_tokens: usize) {
self.budget.max_tokens = Some(max_tokens);
}
pub fn set_max_rounds(&mut self, max_rounds: Option<usize>) {
self.budget.max_rounds = max_rounds;
}
async fn ensure_counts(&self) -> Result<(), MemoryError> {
let missing: Option<Vec<Message>> = {
let inner = self
.inner
.lock()
.expect("WindowMemory internal lock poisoned");
if inner.total_tokens.is_none() {
Some(inner.messages.clone())
} else {
None
}
};
let Some(messages) = missing else {
return Ok(());
};
if messages.is_empty() {
let mut inner = self
.inner
.lock()
.expect("WindowMemory internal lock poisoned");
inner.total_tokens = Some(0);
return Ok(());
}
let mut total = 0usize;
let mut tokens = Vec::with_capacity(messages.len());
for m in &messages {
let t = count_message(&*self.counter, m).await?;
total += t;
tokens.push(t);
}
let user_count = messages
.iter()
.filter(|m| matches!(m, Message::User(_)))
.count();
let mut inner = self
.inner
.lock()
.expect("WindowMemory internal lock poisoned");
inner.total_tokens = Some(total);
inner.tokens = tokens;
inner.user_count = user_count;
Ok(())
}
}
#[async_trait::async_trait]
impl Memory for WindowMemory {
async fn record(&mut self, message: Message) -> Result<(), MemoryError> {
self.record_impl(message, false).await
}
async fn record_protected(&mut self, message: Message) -> Result<(), MemoryError> {
self.record_impl(message, true).await
}
async fn context(&self) -> Result<Vec<Message>, MemoryError> {
self.ensure_counts().await?;
let (over_budget, snapshot, protected, tokens) = {
let inner = self
.inner
.lock()
.expect("WindowMemory internal lock poisoned");
let total = inner
.total_tokens
.expect("counts guaranteed by ensure_counts");
let over = self.budget.max_tokens.is_some_and(|limit| total > limit)
|| self
.budget
.max_rounds
.is_some_and(|limit| inner.user_count > limit);
(
over,
inner.messages.clone(),
inner.protected.clone(),
inner.tokens.clone(),
)
};
if !over_budget {
return Ok(snapshot);
}
let protected_set: std::collections::HashSet<usize> =
protected_round_indices(&snapshot, &protected)
.into_iter()
.collect();
let mut kept: Vec<Message> = Vec::with_capacity(snapshot.len());
let mut candidates: Vec<Message> = Vec::new();
let mut candidate_tokens: Vec<usize> = Vec::new();
for (i, message) in snapshot.into_iter().enumerate() {
if protected_set.contains(&i) {
kept.push(message);
} else {
candidates.push(message);
candidate_tokens.push(tokens[i]);
}
}
let result = self
.strategy
.trim_with_counts(&candidates, &candidate_tokens, &self.budget, &*self.counter)
.await?;
if result.replace {
let protected_len = kept.len();
kept.extend(result.messages);
let mut inner = self
.inner
.lock()
.expect("WindowMemory internal lock poisoned");
inner.messages = kept.clone();
inner.protected = (0..kept.len()).map(|i| i < protected_len).collect();
inner.total_tokens = None;
inner.tokens.clear();
inner.user_count = 0;
} else {
kept.extend(result.messages);
}
Ok(kept)
}
}
impl WindowMemory {
async fn record_impl(&mut self, message: Message, protected: bool) -> Result<(), MemoryError> {
self.ensure_counts().await?;
let tokens = count_message(&*self.counter, &message).await?;
let mut inner = self
.inner
.lock()
.expect("WindowMemory internal lock poisoned");
let total = inner
.total_tokens
.as_mut()
.expect("counts guaranteed by ensure_counts");
*total += tokens;
if matches!(message, Message::User(_)) {
inner.user_count += 1;
}
inner.messages.push(message);
inner.protected.push(protected);
inner.tokens.push(tokens);
Ok(())
}
}
fn protected_round_indices(messages: &[Message], protected: &[bool]) -> Vec<usize> {
let mut result = Vec::new();
let mut round_start = 0usize;
for (i, message) in messages.iter().enumerate().skip(1) {
if matches!(message, Message::User(_)) {
if protected[round_start..i].iter().any(|&p| p) {
result.extend(round_start..i);
}
round_start = i;
}
}
if protected[round_start..].iter().any(|&p| p) {
result.extend(round_start..messages.len());
}
result
}
#[cfg(test)]
mod tests {
use super::*;
use crate::message::{ContentBlock, ToolCall};
fn tool_result(id: &str, content: &str) -> Message {
Message::ToolResult {
id: id.into(),
content: content.into(),
}
}
#[derive(Debug, Default)]
struct FakeSummarizer {
calls: std::sync::atomic::AtomicUsize,
last_input_len: std::sync::atomic::AtomicUsize,
last_input_has_summary: std::sync::atomic::AtomicBool,
}
#[async_trait::async_trait]
impl TrimStrategy for FakeSummarizer {
async fn trim(
&self,
messages: &[Message],
_budget: &Budget,
_counter: &dyn TokenCounter,
) -> Result<TrimResult, MemoryError> {
self.calls
.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
self.last_input_len
.store(messages.len(), std::sync::atomic::Ordering::Relaxed);
let has_summary = messages
.iter()
.any(|m| matches!(m, Message::System(s) if s.starts_with("prior summary")));
self.last_input_has_summary
.store(has_summary, std::sync::atomic::Ordering::Relaxed);
let result = match messages.iter().rposition(|m| matches!(m, Message::User(_))) {
Some(pos) => {
let mut result = Vec::with_capacity(messages.len() - pos + 1);
result.push(Message::system("prior summary"));
result.extend_from_slice(&messages[pos..]);
result
}
None => messages.to_vec(),
};
Ok(TrimResult {
messages: result,
replace: true,
})
}
}
#[tokio::test]
async fn within_budget_returns_all() {
let mut memory = WindowMemory::new(1000);
memory.record(Message::user("hello")).await.unwrap();
memory.record(Message::assistant("hello!")).await.unwrap();
let context = memory.context().await.unwrap();
assert_eq!(context.len(), 2);
assert_eq!(context[0], Message::user("hello"));
assert_eq!(context[1], Message::assistant("hello!"));
}
#[tokio::test]
async fn drops_earliest_rounds_when_over_budget() {
let mut memory = WindowMemory::new(3);
memory.record(Message::user("u1")).await.unwrap();
memory.record(Message::assistant("a1")).await.unwrap();
memory.record(Message::user("u2")).await.unwrap();
memory.record(Message::assistant("a2")).await.unwrap();
memory.record(Message::user("u3")).await.unwrap();
memory.record(Message::assistant("a3")).await.unwrap();
let context = memory.context().await.unwrap();
assert_eq!(context, vec![Message::user("u3"), Message::assistant("a3")]);
}
#[tokio::test]
async fn tool_messages_trimmed_with_assistant() {
let mut memory = WindowMemory::new(8);
memory.record(Message::user("calculate")).await.unwrap();
memory
.record(Message::Assistant {
content: "".into(),
reasoning: None,
tool_calls: vec![
ToolCall {
id: "t1".into(),
name: "calc".into(),
arguments: "1+1".into(),
},
ToolCall {
id: "t2".into(),
name: "calc".into(),
arguments: "2+2".into(),
},
],
})
.await
.unwrap();
memory.record(tool_result("t1", "2")).await.unwrap();
memory.record(tool_result("t2", "4")).await.unwrap();
memory.record(Message::user("continue")).await.unwrap();
memory.record(Message::assistant("okay")).await.unwrap();
let context = memory.context().await.unwrap();
assert_eq!(
context,
vec![Message::user("continue"), Message::assistant("okay")]
);
assert!(
!context.iter().any(
|m| matches!(m, Message::Assistant { tool_calls, .. } if !tool_calls.is_empty())
)
);
assert!(
!context
.iter()
.any(|m| matches!(m, Message::ToolResult { .. }))
);
}
#[tokio::test]
async fn recorded_system_can_be_trimmed() {
let mut memory = WindowMemory::new(3);
memory.record(Message::system("setup")).await.unwrap();
memory.record(Message::user("u1")).await.unwrap();
memory.record(Message::assistant("a1")).await.unwrap();
memory.record(Message::user("u2")).await.unwrap();
memory.record(Message::assistant("a2")).await.unwrap();
let context = memory.context().await.unwrap();
assert_eq!(context, vec![Message::user("u2"), Message::assistant("a2")]);
}
#[tokio::test]
async fn keeps_last_round_even_when_over_budget() {
let mut memory = WindowMemory::new(3);
memory
.record(Message::user(
"An extremely long user message, over budget in a single round",
))
.await
.unwrap();
memory.record(Message::assistant("Reply")).await.unwrap();
memory.record(Message::user("u2")).await.unwrap();
memory.record(Message::assistant("a2")).await.unwrap();
let context = memory.context().await.unwrap();
assert_eq!(context, vec![Message::user("u2"), Message::assistant("a2")]);
}
#[tokio::test]
async fn max_rounds_window() {
let mut memory = WindowMemory::new(1000).with_max_rounds(2);
for i in 1..=4 {
memory.record(Message::user(format!("u{i}"))).await.unwrap();
memory
.record(Message::assistant(format!("a{i}")))
.await
.unwrap();
}
let context = memory.context().await.unwrap();
assert_eq!(
context,
vec![
Message::user("u3"),
Message::assistant("a3"),
Message::user("u4"),
Message::assistant("a4")
]
);
}
#[tokio::test]
async fn tokens_and_rounds_take_smaller() {
let mut memory = WindowMemory::new(3).with_max_rounds(3);
for i in 1..=4 {
memory.record(Message::user(format!("u{i}"))).await.unwrap();
memory
.record(Message::assistant(format!("a{i}")))
.await
.unwrap();
}
let context = memory.context().await.unwrap();
assert_eq!(context, vec![Message::user("u4"), Message::assistant("a4")]);
}
#[tokio::test]
async fn raising_budget_restores_history() {
let mut memory = WindowMemory::new(3);
memory.record(Message::user("u1")).await.unwrap();
memory.record(Message::assistant("a1")).await.unwrap();
memory.record(Message::user("u2")).await.unwrap();
memory.record(Message::assistant("a2")).await.unwrap();
assert_eq!(memory.context().await.unwrap().len(), 2);
memory.set_max_tokens(1000);
assert_eq!(memory.context().await.unwrap().len(), 4); }
#[tokio::test]
async fn reasoning_counts_toward_budget() {
let mut memory = WindowMemory::new(11);
memory.record(Message::user("u1")).await.unwrap();
memory
.record(Message::assistant_with_reasoning(
"hi",
"a".repeat(40), ))
.await
.unwrap();
memory.record(Message::user("u2")).await.unwrap();
memory.record(Message::assistant("hi")).await.unwrap();
let context = memory.context().await.unwrap();
assert_eq!(context, vec![Message::user("u2"), Message::assistant("hi")]);
let mut memory2 = WindowMemory::new(11);
memory2.record(Message::user("u1")).await.unwrap();
memory2.record(Message::assistant("hi")).await.unwrap();
memory2.record(Message::user("u2")).await.unwrap();
memory2.record(Message::assistant("hi")).await.unwrap();
assert_eq!(memory2.context().await.unwrap().len(), 4);
}
#[tokio::test]
async fn custom_token_counter() {
#[derive(Debug, Default)]
struct OnePerMessage;
#[async_trait::async_trait]
impl TokenCounter for OnePerMessage {
async fn count(&self, _text: &str) -> Result<usize, MemoryError> {
Ok(1)
}
}
let mut memory = WindowMemory::new(2).with_token_counter(Box::new(OnePerMessage));
memory.record(Message::user("u1")).await.unwrap();
memory.record(Message::assistant("a1")).await.unwrap();
memory.record(Message::user("u2")).await.unwrap();
memory.record(Message::assistant("a2")).await.unwrap();
memory.record(Message::user("u3")).await.unwrap();
let context = memory.context().await.unwrap();
assert_eq!(context, vec![Message::user("u3")]);
}
#[tokio::test]
async fn materialize_replaces_storage_and_skips_strategy() {
let summarizer = Arc::new(FakeSummarizer::default());
let mut memory = WindowMemory::new(7).with_strategy(summarizer.clone());
for i in 1..=4 {
memory.record(Message::user(format!("u{i}"))).await.unwrap();
memory
.record(Message::assistant(format!("a{i}")))
.await
.unwrap();
}
let first = memory.context().await.unwrap();
assert_eq!(
summarizer.calls.load(std::sync::atomic::Ordering::Relaxed),
1
);
assert_eq!(
first,
vec![
Message::system("prior summary"),
Message::user("u4"),
Message::assistant("a4"),
]
);
let second = memory.context().await.unwrap();
assert_eq!(
summarizer.calls.load(std::sync::atomic::Ordering::Relaxed),
1
);
assert_eq!(second, first);
}
#[tokio::test]
async fn materialize_strategy_input_is_materialized_sequence() {
let summarizer = Arc::new(FakeSummarizer::default());
let mut memory = WindowMemory::new(7).with_strategy(summarizer.clone());
for i in 1..=4 {
memory.record(Message::user(format!("u{i}"))).await.unwrap();
memory
.record(Message::assistant(format!("a{i}")))
.await
.unwrap();
}
memory.context().await.unwrap();
memory.record(Message::user("u5")).await.unwrap();
memory.record(Message::assistant("a5")).await.unwrap();
memory.record(Message::user("u6")).await.unwrap();
memory.record(Message::assistant("a6")).await.unwrap();
memory.context().await.unwrap();
assert_eq!(
summarizer.calls.load(std::sync::atomic::Ordering::Relaxed),
2
);
assert!(
summarizer
.last_input_has_summary
.load(std::sync::atomic::Ordering::Relaxed)
);
}
#[tokio::test]
async fn projection_keeps_storage_unchanged() {
#[derive(Debug, Default)]
struct KeepLastRound;
#[async_trait::async_trait]
impl TrimStrategy for KeepLastRound {
async fn trim(
&self,
messages: &[Message],
_budget: &Budget,
_counter: &dyn TokenCounter,
) -> Result<TrimResult, MemoryError> {
let pos = messages
.iter()
.rposition(|m| matches!(m, Message::User(_)))
.unwrap_or(0);
Ok(TrimResult {
messages: messages[pos..].to_vec(),
replace: false,
})
}
}
let strategy = Arc::new(KeepLastRound);
let mut memory = WindowMemory::new(1).with_strategy(strategy);
memory.record(Message::user("u1")).await.unwrap();
memory.record(Message::assistant("a1")).await.unwrap();
memory.record(Message::user("u2")).await.unwrap();
memory.record(Message::assistant("a2")).await.unwrap();
let view = memory.context().await.unwrap();
assert_eq!(view, vec![Message::user("u2"), Message::assistant("a2")]);
memory.set_max_tokens(1000);
assert_eq!(memory.context().await.unwrap().len(), 4);
}
#[tokio::test]
async fn context_is_a_copy() {
let mut memory = WindowMemory::new(1000);
memory.record(Message::user("a")).await.unwrap();
let mut context = memory.context().await.unwrap();
context.push(Message::assistant("b"));
assert_eq!(memory.context().await.unwrap().len(), 1);
}
#[tokio::test]
async fn no_user_messages_returns_all() {
let mut memory = WindowMemory::new(1);
memory
.record(Message::assistant("assistant-only message"))
.await
.unwrap();
let context = memory.context().await.unwrap();
assert_eq!(context.len(), 1);
}
#[tokio::test]
async fn user_blocks_all_counted() {
let mut memory = WindowMemory::new(2);
memory
.record(Message::user_blocks(vec![
ContentBlock::Text("aaaa".into()),
ContentBlock::Text("bbbb".into()),
]))
.await
.unwrap();
memory.record(Message::assistant("hi")).await.unwrap();
let context = memory.context().await.unwrap();
assert_eq!(context.len(), 2); }
#[tokio::test]
async fn changing_counter_recounts() {
#[derive(Debug)]
struct CountingCounter {
calls: Arc<std::sync::atomic::AtomicUsize>,
}
impl CountingCounter {
fn shared(calls: Arc<std::sync::atomic::AtomicUsize>) -> Self {
Self { calls }
}
}
#[async_trait::async_trait]
impl TokenCounter for CountingCounter {
async fn count(&self, _text: &str) -> Result<usize, MemoryError> {
self.calls
.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
Ok(1)
}
}
let first_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
let mut memory = WindowMemory::new(100)
.with_token_counter(Box::new(CountingCounter::shared(first_calls.clone())));
memory.record(Message::user("u1")).await.unwrap();
memory.record(Message::assistant("a1")).await.unwrap();
memory.context().await.unwrap(); assert_eq!(first_calls.load(std::sync::atomic::Ordering::Relaxed), 2);
let second_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
let mut memory =
memory.with_token_counter(Box::new(CountingCounter::shared(second_calls.clone())));
memory.record(Message::user("u2")).await.unwrap();
assert_eq!(second_calls.load(std::sync::atomic::Ordering::Relaxed), 3);
}
#[tokio::test]
async fn token_count_failure_propagates() {
#[derive(Debug, Default)]
struct FailingCounter;
#[async_trait::async_trait]
impl TokenCounter for FailingCounter {
async fn count(&self, _text: &str) -> Result<usize, MemoryError> {
Err(MemoryError::TokenCount(
"remote counting API unavailable".into(),
))
}
}
let mut memory = WindowMemory::new(100).with_token_counter(Box::new(FailingCounter));
let err = memory.record(Message::user("hi")).await.unwrap_err();
assert!(matches!(err, MemoryError::TokenCount(_)));
}
#[tokio::test]
async fn protected_round_survives_trim() {
let mut memory = WindowMemory::new(3);
memory.record(Message::user("u1")).await.unwrap();
memory.record(Message::assistant("a1")).await.unwrap();
memory.record(Message::user("u2")).await.unwrap();
memory
.record(Message::Assistant {
content: "a2".into(),
reasoning: None,
tool_calls: vec![ToolCall {
id: "c1".into(),
name: "load_skill".into(),
arguments: r#"{"name":"x"}"#.into(),
}],
})
.await
.unwrap();
memory
.record_protected(tool_result("c1", "<skill_content>body</skill_content>"))
.await
.unwrap();
memory.record(Message::user("u3")).await.unwrap();
memory.record(Message::assistant("a3")).await.unwrap();
let context = memory.context().await.unwrap();
assert_eq!(
context,
vec![
Message::user("u2"),
Message::Assistant {
content: "a2".into(),
reasoning: None,
tool_calls: vec![ToolCall {
id: "c1".into(),
name: "load_skill".into(),
arguments: r#"{"name":"x"}"#.into(),
}],
},
tool_result("c1", "<skill_content>body</skill_content>"),
Message::user("u3"),
Message::assistant("a3"),
]
);
}
#[tokio::test]
async fn protected_message_never_pruned() {
let mut memory = WindowMemory::new(3);
memory.record(Message::user("u1")).await.unwrap();
memory.record(Message::assistant("a1")).await.unwrap();
memory.record(Message::user("u2")).await.unwrap();
memory.record(Message::assistant("a2")).await.unwrap();
memory
.record_protected(tool_result("c1", "skill body"))
.await
.unwrap();
let mut all_kept = true;
for _ in 0..5 {
let context = memory.context().await.unwrap();
if !context.iter().any(
|m| matches!(m, Message::ToolResult { content, .. } if content == "skill body"),
) {
all_kept = false;
break;
}
}
assert!(
all_kept,
"the protected message must survive repeated trims"
);
}
}