1use std::sync::Arc;
32
33use async_trait::async_trait;
34
35use crate::error::Result;
36use crate::memory::{ContextProvider, SessionContext};
37use crate::types::{Content, Message, Role};
38
39pub trait Tokenizer: Send + Sync {
42 fn count_tokens(&self, text: &str) -> usize;
44}
45
46#[derive(Debug, Clone, Copy, Default)]
49pub struct ApproxTokenizer;
50
51impl Tokenizer for ApproxTokenizer {
52 fn count_tokens(&self, text: &str) -> usize {
53 text.chars().count().div_ceil(4)
54 }
55}
56
57pub fn count_message_tokens(tokenizer: &dyn Tokenizer, message: &Message) -> usize {
60 message
61 .contents
62 .iter()
63 .filter_map(Content::as_text)
64 .map(|text| tokenizer.count_tokens(text))
65 .sum()
66}
67
68pub trait CompactionStrategy: Send + Sync {
74 fn compact(&self, messages: &[Message], tokenizer: &dyn Tokenizer) -> Vec<Message>;
76}
77
78fn leading_system_count(messages: &[Message]) -> usize {
80 messages
81 .iter()
82 .take_while(|m| m.role == Role::system())
83 .count()
84}
85
86#[derive(Debug, Clone, Copy)]
89pub struct Truncation {
90 pub max_messages: usize,
91}
92
93impl Truncation {
94 pub fn new(max_messages: usize) -> Self {
95 Self { max_messages }
96 }
97}
98
99impl CompactionStrategy for Truncation {
100 fn compact(&self, messages: &[Message], _tokenizer: &dyn Tokenizer) -> Vec<Message> {
101 if messages.len() <= self.max_messages {
102 return messages.to_vec();
103 }
104 let sys_count = leading_system_count(messages);
105 let mut out: Vec<Message> = messages[..sys_count].to_vec();
106
107 if sys_count >= self.max_messages {
108 out.truncate(self.max_messages);
111 return out;
112 }
113
114 let remaining_budget = self.max_messages - sys_count;
115 let rest = &messages[sys_count..];
116 let start = rest.len().saturating_sub(remaining_budget);
117 out.extend_from_slice(&rest[start..]);
118 out
119 }
120}
121
122#[derive(Debug, Clone, Copy)]
125pub struct SlidingWindow {
126 pub window: usize,
127}
128
129impl SlidingWindow {
130 pub fn new(window: usize) -> Self {
131 Self { window }
132 }
133}
134
135impl CompactionStrategy for SlidingWindow {
136 fn compact(&self, messages: &[Message], _tokenizer: &dyn Tokenizer) -> Vec<Message> {
137 let sys_count = leading_system_count(messages);
138 let mut out: Vec<Message> = messages[..sys_count].to_vec();
139 let rest = &messages[sys_count..];
140 let start = rest.len().saturating_sub(self.window);
141 out.extend_from_slice(&rest[start..]);
142 out
143 }
144}
145
146#[derive(Debug, Clone, Copy)]
151pub struct TokenBudget {
152 pub max_tokens: usize,
153}
154
155impl TokenBudget {
156 pub fn new(max_tokens: usize) -> Self {
157 Self { max_tokens }
158 }
159}
160
161impl CompactionStrategy for TokenBudget {
162 fn compact(&self, messages: &[Message], tokenizer: &dyn Tokenizer) -> Vec<Message> {
163 let sys_count = leading_system_count(messages);
164 let system_prefix = &messages[..sys_count];
165 let rest = &messages[sys_count..];
166
167 let mut used: usize = system_prefix
168 .iter()
169 .map(|m| count_message_tokens(tokenizer, m))
170 .sum();
171
172 let mut kept_rest: Vec<&Message> = Vec::new();
178 for message in rest.iter().rev() {
179 let cost = count_message_tokens(tokenizer, message);
180 if !kept_rest.is_empty() && used + cost > self.max_tokens {
181 break;
182 }
183 used += cost;
184 kept_rest.push(message);
185 }
186 kept_rest.reverse();
187
188 let mut out: Vec<Message> = system_prefix.to_vec();
189 out.extend(kept_rest.into_iter().cloned());
190 out
191 }
192}
193
194fn has_tool_result(message: &Message) -> bool {
197 message
198 .contents
199 .iter()
200 .any(|c| matches!(c, Content::FunctionResult(_)))
201}
202
203#[derive(Debug, Clone, Copy)]
209pub struct SelectiveToolResult {
210 pub keep_last: usize,
211}
212
213impl SelectiveToolResult {
214 pub fn new(keep_last: usize) -> Self {
215 Self { keep_last }
216 }
217}
218
219impl CompactionStrategy for SelectiveToolResult {
220 fn compact(&self, messages: &[Message], _tokenizer: &dyn Tokenizer) -> Vec<Message> {
221 let tool_result_count = messages.iter().filter(|m| has_tool_result(m)).count();
222 let mut strip_budget = tool_result_count.saturating_sub(self.keep_last);
223
224 let mut out = Vec::with_capacity(messages.len());
225 for message in messages {
226 if has_tool_result(message) && strip_budget > 0 {
227 strip_budget -= 1;
228 let contents: Vec<Content> = message
229 .contents
230 .iter()
231 .filter(|c| !matches!(c, Content::FunctionResult(_)))
232 .cloned()
233 .collect();
234 if contents.is_empty() {
235 continue;
236 }
237 let mut stripped = message.clone();
238 stripped.contents = contents;
239 out.push(stripped);
240 } else {
241 out.push(message.clone());
242 }
243 }
244 out
245 }
246}
247
248pub fn compact(
251 messages: &[Message],
252 strategy: &dyn CompactionStrategy,
253 tokenizer: &dyn Tokenizer,
254) -> Vec<Message> {
255 strategy.compact(messages, tokenizer)
256}
257
258pub struct CompactionProvider {
272 strategy: Arc<dyn CompactionStrategy>,
273 tokenizer: Box<dyn Tokenizer>,
274}
275
276impl CompactionProvider {
277 pub fn new(strategy: impl CompactionStrategy + 'static) -> Self {
280 Self::with_tokenizer(strategy, ApproxTokenizer)
281 }
282
283 pub fn with_tokenizer(
285 strategy: impl CompactionStrategy + 'static,
286 tokenizer: impl Tokenizer + 'static,
287 ) -> Self {
288 Self {
289 strategy: Arc::new(strategy),
290 tokenizer: Box::new(tokenizer),
291 }
292 }
293}
294
295#[async_trait]
296impl ContextProvider for CompactionProvider {
297 async fn before_run(&self, ctx: &mut SessionContext) -> Result<()> {
300 ctx.messages = self.strategy.compact(&ctx.messages, &*self.tokenizer);
301 Ok(())
302 }
303
304 }
308
309#[cfg(test)]
310mod tests {
311 use super::*;
312 use crate::types::FunctionResultContent;
313 use serde_json::json;
314
315 fn text(role: Role, s: &str) -> Message {
316 Message::new(role, s)
317 }
318
319 fn tool_result_message(call_id: &str, result: &str) -> Message {
320 Message::with_contents(
321 Role::tool(),
322 vec![Content::FunctionResult(FunctionResultContent::new(
323 call_id,
324 Some(json!(result)),
325 ))],
326 )
327 }
328
329 #[test]
332 fn approx_tokenizer_uses_four_chars_per_token_ceiling() {
333 let t = ApproxTokenizer;
334 assert_eq!(t.count_tokens(""), 0);
335 assert_eq!(t.count_tokens("abcd"), 1);
336 assert_eq!(t.count_tokens("abcde"), 2); assert_eq!(t.count_tokens("abcdefgh"), 2);
338 assert_eq!(t.count_tokens("abcdefghi"), 3); }
340
341 #[test]
342 fn count_message_tokens_sums_text_content() {
343 let t = ApproxTokenizer;
344 let msg = Message::with_contents(
345 Role::user(),
346 vec![Content::text("abcd"), Content::text("abcdefgh")],
347 );
348 assert_eq!(count_message_tokens(&t, &msg), 3);
350 }
351
352 #[test]
355 fn truncation_keeps_most_recent_messages() {
356 let messages = vec![
357 text(Role::user(), "1"),
358 text(Role::assistant(), "2"),
359 text(Role::user(), "3"),
360 text(Role::assistant(), "4"),
361 ];
362 let strategy = Truncation::new(2);
363 let out = compact(&messages, &strategy, &ApproxTokenizer);
364 assert_eq!(out.len(), 2);
365 assert_eq!(out[0].text(), "3");
366 assert_eq!(out[1].text(), "4");
367 }
368
369 #[test]
370 fn truncation_preserves_leading_system_messages() {
371 let messages = vec![
372 text(Role::system(), "sys"),
373 text(Role::user(), "1"),
374 text(Role::assistant(), "2"),
375 text(Role::user(), "3"),
376 text(Role::assistant(), "4"),
377 ];
378 let strategy = Truncation::new(2);
379 let out = compact(&messages, &strategy, &ApproxTokenizer);
380 assert_eq!(out.len(), 2);
382 assert_eq!(out[0].role, Role::system());
383 assert_eq!(out[0].text(), "sys");
384 assert_eq!(out[1].text(), "4");
385 }
386
387 #[test]
388 fn truncation_preserves_multiple_leading_system_messages() {
389 let messages = vec![
390 text(Role::system(), "sys1"),
391 text(Role::system(), "sys2"),
392 text(Role::user(), "1"),
393 text(Role::assistant(), "2"),
394 ];
395 let strategy = Truncation::new(3);
396 let out = compact(&messages, &strategy, &ApproxTokenizer);
397 assert_eq!(out.len(), 3);
398 assert_eq!(out[0].text(), "sys1");
399 assert_eq!(out[1].text(), "sys2");
400 assert_eq!(out[2].text(), "2");
401 }
402
403 #[test]
404 fn truncation_noop_when_under_budget() {
405 let messages = vec![text(Role::user(), "1"), text(Role::assistant(), "2")];
406 let strategy = Truncation::new(10);
407 let out = compact(&messages, &strategy, &ApproxTokenizer);
408 assert_eq!(out, messages);
409 }
410
411 #[test]
414 fn sliding_window_keeps_system_plus_last_n_non_system() {
415 let messages = vec![
416 text(Role::system(), "sys"),
417 text(Role::user(), "1"),
418 text(Role::assistant(), "2"),
419 text(Role::user(), "3"),
420 ];
421 let strategy = SlidingWindow::new(2);
422 let out = compact(&messages, &strategy, &ApproxTokenizer);
423 assert_eq!(out.len(), 3);
424 assert_eq!(out[0].text(), "sys");
425 assert_eq!(out[1].text(), "2");
426 assert_eq!(out[2].text(), "3");
427 }
428
429 #[test]
430 fn sliding_window_with_no_system_message() {
431 let messages = vec![
432 text(Role::user(), "1"),
433 text(Role::assistant(), "2"),
434 text(Role::user(), "3"),
435 ];
436 let strategy = SlidingWindow::new(1);
437 let out = compact(&messages, &strategy, &ApproxTokenizer);
438 assert_eq!(out.len(), 1);
439 assert_eq!(out[0].text(), "3");
440 }
441
442 struct FixedTokenizer(usize);
447 impl Tokenizer for FixedTokenizer {
448 fn count_tokens(&self, _text: &str) -> usize {
449 self.0
450 }
451 }
452
453 #[test]
454 fn token_budget_keeps_only_what_fits_from_the_newest_backward() {
455 let messages = vec![
456 text(Role::user(), "1"),
457 text(Role::assistant(), "2"),
458 text(Role::user(), "3"),
459 text(Role::assistant(), "4"),
460 ];
461 let tokenizer = FixedTokenizer(10);
463 let strategy = TokenBudget::new(25);
464 let out = compact(&messages, &strategy, &tokenizer);
465 assert_eq!(out.len(), 2);
466 assert_eq!(out[0].text(), "3");
467 assert_eq!(out[1].text(), "4");
468 }
469
470 #[test]
471 fn token_budget_preserves_leading_system_message_and_counts_it() {
472 let messages = vec![
473 text(Role::system(), "sys"),
474 text(Role::user(), "1"),
475 text(Role::assistant(), "2"),
476 text(Role::user(), "3"),
477 ];
478 let tokenizer = FixedTokenizer(10);
479 let strategy = TokenBudget::new(20);
481 let out = compact(&messages, &strategy, &tokenizer);
482 assert_eq!(out.len(), 2);
483 assert_eq!(out[0].role, Role::system());
484 assert_eq!(out[1].text(), "3");
485 }
486
487 #[test]
488 fn token_budget_keeps_at_least_the_newest_message_even_if_it_alone_exceeds_budget() {
489 let messages = vec![text(Role::user(), "1"), text(Role::assistant(), "2")];
490 let tokenizer = FixedTokenizer(100);
491 let strategy = TokenBudget::new(1);
492 let out = compact(&messages, &strategy, &tokenizer);
493 assert_eq!(out.len(), 1);
494 assert_eq!(out[0].text(), "2");
495 }
496
497 #[test]
498 fn token_budget_keeps_everything_when_it_all_fits() {
499 let messages = vec![text(Role::user(), "1"), text(Role::assistant(), "2")];
500 let tokenizer = FixedTokenizer(1);
501 let strategy = TokenBudget::new(1000);
502 let out = compact(&messages, &strategy, &tokenizer);
503 assert_eq!(out, messages);
504 }
505
506 #[test]
509 fn selective_tool_result_strips_stale_results_and_keeps_recent_ones() {
510 let messages = vec![
511 text(Role::user(), "ask 1"),
512 tool_result_message("c1", "result 1"),
513 text(Role::user(), "ask 2"),
514 tool_result_message("c2", "result 2"),
515 text(Role::user(), "ask 3"),
516 tool_result_message("c3", "result 3"),
517 ];
518 let strategy = SelectiveToolResult::new(1);
519 let out = compact(&messages, &strategy, &ApproxTokenizer);
520
521 assert_eq!(out.len(), 4);
524 assert_eq!(out[0].text(), "ask 1");
525 assert_eq!(out[1].text(), "ask 2");
526 assert_eq!(out[2].text(), "ask 3");
527 assert!(has_tool_result(&out[3]));
528 assert_eq!(out[3].function_results()[0].call_id, "c3");
529 }
530
531 #[test]
532 fn selective_tool_result_keeps_text_alongside_a_stripped_tool_result() {
533 let mixed = Message::with_contents(
534 Role::tool(),
535 vec![
536 Content::text("some accompanying text"),
537 Content::FunctionResult(FunctionResultContent::new("c1", Some(json!("r1")))),
538 ],
539 );
540 let messages = vec![
541 mixed,
542 tool_result_message("c2", "result 2"),
543 tool_result_message("c3", "result 3"),
544 ];
545 let strategy = SelectiveToolResult::new(2);
546 let out = compact(&messages, &strategy, &ApproxTokenizer);
547
548 assert_eq!(out.len(), 3);
551 assert_eq!(out[0].text(), "some accompanying text");
552 assert!(!has_tool_result(&out[0]));
553 assert!(has_tool_result(&out[1]));
554 assert!(has_tool_result(&out[2]));
555 }
556
557 #[test]
558 fn selective_tool_result_noop_when_keep_last_covers_all() {
559 let messages = vec![
560 tool_result_message("c1", "result 1"),
561 tool_result_message("c2", "result 2"),
562 ];
563 let strategy = SelectiveToolResult::new(5);
564 let out = compact(&messages, &strategy, &ApproxTokenizer);
565 assert_eq!(out, messages);
566 }
567
568 #[test]
569 fn selective_tool_result_ignores_messages_without_tool_results() {
570 let messages = vec![
571 text(Role::system(), "sys"),
572 text(Role::user(), "hi"),
573 text(Role::assistant(), "hello"),
574 ];
575 let strategy = SelectiveToolResult::new(0);
576 let out = compact(&messages, &strategy, &ApproxTokenizer);
577 assert_eq!(out, messages);
578 }
579
580 #[tokio::test]
583 async fn compaction_provider_before_run_replaces_ctx_messages_with_compacted_subset() {
584 let provider = CompactionProvider::new(Truncation::new(2));
585 let mut ctx = SessionContext::new(vec![]);
586 ctx.messages = vec![
587 text(Role::user(), "1"),
588 text(Role::assistant(), "2"),
589 text(Role::user(), "3"),
590 text(Role::assistant(), "4"),
591 ];
592 provider.before_run(&mut ctx).await.unwrap();
593 assert_eq!(ctx.messages.len(), 2);
594 assert_eq!(ctx.messages[0].text(), "3");
595 assert_eq!(ctx.messages[1].text(), "4");
596 }
597
598 #[tokio::test]
599 async fn compaction_provider_with_tokenizer_uses_the_supplied_tokenizer() {
600 struct FixedTokenizer(usize);
601 impl Tokenizer for FixedTokenizer {
602 fn count_tokens(&self, _text: &str) -> usize {
603 self.0
604 }
605 }
606 let provider = CompactionProvider::with_tokenizer(TokenBudget::new(25), FixedTokenizer(10));
607 let mut ctx = SessionContext::new(vec![]);
608 ctx.messages = vec![
609 text(Role::user(), "1"),
610 text(Role::assistant(), "2"),
611 text(Role::user(), "3"),
612 text(Role::assistant(), "4"),
613 ];
614 provider.before_run(&mut ctx).await.unwrap();
615 assert_eq!(ctx.messages.len(), 2);
617 assert_eq!(ctx.messages[0].text(), "3");
618 assert_eq!(ctx.messages[1].text(), "4");
619 }
620
621 #[tokio::test]
622 async fn compaction_provider_after_run_is_a_noop() {
623 let provider = CompactionProvider::new(Truncation::new(1));
624 provider
625 .after_run(&[Message::new(Role::user(), "hi")], &[], None)
626 .await
627 .unwrap();
628 }
629}