use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::Arc;
use yoagent::context::{total_tokens, CompactionStrategy, ContextConfig, DefaultCompaction};
use yoagent::llm_compaction::{LlmCompaction, SUMMARY_MARKER};
use yoagent::provider::{ModelConfig, ProviderError, StreamConfig, StreamEvent, StreamProvider};
use yoagent::types::*;
#[derive(Debug, Clone)]
pub struct SessionShape {
pub budget: usize,
pub tokens_per_turn: usize,
pub rounds: usize,
pub keep_first: usize,
pub keep_recent: usize,
}
impl SessionShape {
pub fn new(budget: usize, tokens_per_turn: usize, rounds: usize) -> Self {
Self {
budget,
tokens_per_turn,
rounds,
keep_first: 2,
keep_recent: 10,
}
}
fn context_config(&self) -> ContextConfig {
ContextConfig {
max_context_tokens: self.budget,
system_prompt_tokens: 0,
keep_first: self.keep_first,
keep_recent: self.keep_recent,
..Default::default()
}
}
}
#[derive(Debug, Default, Clone)]
pub struct Report {
pub cache_breaks: usize,
pub break_rounds: Vec<usize>,
pub peak_tokens: usize,
pub final_tokens: usize,
pub rounds_with_summary: usize,
}
impl Report {
pub fn mean_interval(&self) -> Option<f64> {
(self.cache_breaks > 1).then(|| {
let first = *self.break_rounds.first().unwrap() as f64;
let last = *self.break_rounds.last().unwrap() as f64;
(last - first) / (self.cache_breaks - 1) as f64
})
}
}
fn common_prefix(before: &[AgentMessage], after: &[AgentMessage]) -> usize {
before
.iter()
.zip(after.iter())
.take_while(|(a, b)| a == b)
.count()
}
fn text_turn(i: usize, tokens: usize) -> Vec<AgentMessage> {
let bulk = (tokens * 4) / 2;
vec![
AgentMessage::Llm(Message::User {
content: vec![Content::Text {
text: format!("u{i}: {}", "x".repeat(bulk)),
}],
timestamp: i as u64,
}),
AgentMessage::Llm(
Message::assistant(
vec![Content::Text {
text: format!("a{i}: {}", "y".repeat(bulk)),
}],
StopReason::Stop,
"harness",
"harness",
Usage::default(),
)
.with_timestamp(i as u64),
),
]
}
fn has_summary(messages: &[AgentMessage]) -> bool {
messages.iter().any(|m| {
matches!(m, AgentMessage::Llm(Message::User { content, .. })
if content.iter().any(|c| matches!(c, Content::Text { text }
if text.starts_with(SUMMARY_MARKER))))
})
}
pub async fn measure(strategy: &dyn CompactionStrategy, shape: &SessionShape) -> Report {
let config = shape.context_config();
let mut messages: Vec<AgentMessage> = Vec::new();
let mut report = Report::default();
for round in 0..shape.rounds {
messages.extend(text_turn(round, shape.tokens_per_turn));
let before = messages.clone();
messages = strategy.compact(std::mem::take(&mut messages), &config);
if common_prefix(&before, &messages) < before.len() {
report.cache_breaks += 1;
report.break_rounds.push(round);
}
if has_summary(&messages) {
report.rounds_with_summary += 1;
}
report.peak_tokens = report.peak_tokens.max(total_tokens(&messages));
for _ in 0..30 {
tokio::time::sleep(std::time::Duration::from_millis(1)).await;
}
}
report.final_tokens = total_tokens(&messages);
report
}
struct CountingProvider {
calls: Arc<AtomicUsize>,
}
#[async_trait::async_trait]
impl StreamProvider for CountingProvider {
async fn stream(
&self,
_config: StreamConfig,
_tx: tokio::sync::mpsc::UnboundedSender<StreamEvent>,
_cancel: tokio_util::sync::CancellationToken,
) -> Result<Message, ProviderError> {
self.calls.fetch_add(1, Ordering::SeqCst);
Ok(Message::assistant(
vec![Content::Text {
text: "## Goal\nHarness briefing.\n## Open items\nNone.".into(),
}],
StopReason::Stop,
"harness",
"harness",
Usage::default(),
))
}
}
fn llm_strategy() -> (LlmCompaction, Arc<AtomicUsize>) {
let calls = Arc::new(AtomicUsize::new(0));
let strategy = LlmCompaction::from_provider(
Arc::new(CountingProvider {
calls: Arc::clone(&calls),
}),
ModelConfig::mock(),
);
(strategy, calls)
}
#[tokio::test(flavor = "multi_thread")]
#[ignore = "benchmark: minutes to run; see the module docs for the command"]
async fn compaction_strategies_prefix_stability() {
println!(
"\n{:<10} {:<22} {:>7} {:>9} {:>8} {:>9} {:>9}",
"budget", "strategy", "breaks", "interval", "peak", "final", "requests"
);
for (budget, tokens_per_turn, rounds) in
[(20_000usize, 460usize, 120usize), (100_000, 460, 600)]
{
let shape = SessionShape::new(budget, tokens_per_turn, rounds);
let d = measure(&DefaultCompaction, &shape).await;
println!(
"{budget:<10} {:<22} {:>7} {:>9} {:>8} {:>9} {:>9}",
"DefaultCompaction",
d.cache_breaks,
d.mean_interval()
.map(|i| format!("{i:.1}"))
.unwrap_or_else(|| "-".into()),
d.peak_tokens,
d.final_tokens,
0,
);
let (llm, calls) = llm_strategy();
let l = measure(&llm, &shape).await;
println!(
"{budget:<10} {:<22} {:>7} {:>9} {:>8} {:>9} {:>9}",
"LlmCompaction",
l.cache_breaks,
l.mean_interval()
.map(|i| format!("{i:.1}"))
.unwrap_or_else(|| "-".into()),
l.peak_tokens,
l.final_tokens,
calls.load(Ordering::SeqCst),
);
}
println!();
}
#[tokio::test(flavor = "multi_thread")]
#[ignore = "benchmark: minutes to run; see the module docs for the command"]
async fn llm_compaction_across_turn_sizes() {
println!(
"\n{:<16} {:>7} {:>9} {:>9} {:>8}",
"tokens/turn", "breaks", "splices", "requests", "wasted"
);
for tokens_per_turn in [100usize, 400, 1_000, 3_000, 6_000] {
let (llm, calls) = llm_strategy();
let shape = SessionShape::new(20_000, tokens_per_turn, 60);
let r = measure(&llm, &shape).await;
let requests = calls.load(Ordering::SeqCst);
println!(
"{tokens_per_turn:<16} {:>7} {:>9} {:>9} {:>8}",
r.cache_breaks,
r.rounds_with_summary,
requests,
requests > 0 && r.rounds_with_summary == 0,
);
}
println!();
}
#[tokio::test(flavor = "multi_thread")]
async fn harness_detects_rewrites_and_only_rewrites() {
struct Untouched;
impl CompactionStrategy for Untouched {
fn compact(&self, m: Vec<AgentMessage>, _c: &ContextConfig) -> Vec<AgentMessage> {
m
}
}
struct RewritesEveryTurn;
impl CompactionStrategy for RewritesEveryTurn {
fn compact(&self, mut m: Vec<AgentMessage>, _c: &ContextConfig) -> Vec<AgentMessage> {
if !m.is_empty() {
m[0] = AgentMessage::Llm(Message::user("rewritten"));
}
m
}
}
let shape = SessionShape::new(1_000_000, 100, 6);
assert_eq!(
measure(&Untouched, &shape).await.cache_breaks,
0,
"append-only history must register no cache breaks"
);
let rewritten = measure(&RewritesEveryTurn, &shape).await;
assert!(
rewritten.cache_breaks >= 5,
"a strategy rewriting index 0 every turn must register breaks, got {}",
rewritten.cache_breaks
);
}