Skip to main content

agent_framework_core/
compaction.rs

1//! Conversation-history compaction.
2//!
3//! Rust equivalent of (a self-contained subset of) upstream `_compaction.py`
4//! (see `UPSTREAM_DRIFT.md` §9). Upstream's `_compaction.py` is a large
5//! (1500+ line), annotation-driven system: it groups messages into logical
6//! spans (system / user / assistant-text / tool-call), stamps grouping and
7//! token-count metadata onto `Message.additional_properties`, and never
8//! deletes history — it flags messages `_excluded` and lets the client filter
9//! them out when building the payload sent to the model. It ships seven
10//! strategies (`Truncation`, `SlidingWindow`, `SelectiveToolCall`,
11//! `ToolResult`, LLM-backed `Summarization`, `TokenBudgetComposed`,
12//! `ContextWindow`) plus a `CompactionProvider(ContextProvider)` that wires a
13//! strategy into the client's `get_response` loop.
14//!
15//! This module intentionally delivers a smaller, dependency-free surface:
16//! the [`Tokenizer`] and [`CompactionStrategy`] abstractions upstream
17//! defines, plus four concrete, non-LLM strategies that mirror upstream's
18//! `Truncation`, `SlidingWindow`, `ContextWindow`/`TokenBudget`, and
19//! `ToolResult` (renamed [`SelectiveToolResult`] here to avoid confusion with
20//! `Content::FunctionResult`... "tool result" is the plain-English name).
21//! Compaction here works by *returning a reduced list* rather than annotating
22//! messages in place — simpler, and sufficient for the strategies included.
23//! Wiring a strategy into the client's `get_response` loop (upstream's
24//! `CompactionProvider`) is intentionally out of scope for this change; see
25//! `UPSTREAM_DRIFT.md` §9.
26//!
27//! Compaction never errors on content: given any message list it returns a
28//! (possibly unchanged) retained subset that satisfies the strategy's
29//! constraint.
30
31use 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
39/// Counts tokens for a piece of text. Rust equivalent of upstream
40/// `TokenizerProtocol`.
41pub trait Tokenizer: Send + Sync {
42    /// Count the tokens represented by `text`.
43    fn count_tokens(&self, text: &str) -> usize;
44}
45
46/// A dependency-free default tokenizer using a ~4-characters-per-token
47/// heuristic. Mirrors upstream's `CharacterEstimatorTokenizer`.
48#[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
57/// Sum the token counts of a message's text content (text and reasoning
58/// content items) using `tokenizer`.
59pub 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
68/// A strategy that reduces a message list to fit some constraint.
69///
70/// Compaction never errors on content — it always returns *some* retained
71/// subset of `messages`, in original order. Rust equivalent of upstream
72/// `CompactionStrategy`.
73pub trait CompactionStrategy: Send + Sync {
74    /// Return the retained messages (in original order) after compaction.
75    fn compact(&self, messages: &[Message], tokenizer: &dyn Tokenizer) -> Vec<Message>;
76}
77
78/// Returns the number of leading messages with `Role::system()`.
79fn leading_system_count(messages: &[Message]) -> usize {
80    messages
81        .iter()
82        .take_while(|m| m.role == Role::system())
83        .count()
84}
85
86/// Keep the most recent `max_messages`, always preserving any leading system
87/// message(s) at the front. Mirrors upstream's `Truncation` strategy.
88#[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            // The system prefix alone already fills (or exceeds) the budget;
109            // keep just the system prefix, truncated to the budget.
110            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/// Keep leading system message(s) + the last `window` non-system messages.
123/// Mirrors upstream's `SlidingWindow` strategy.
124#[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/// Keep leading system message(s), then walk from the newest message
147/// backward accumulating token counts, keeping messages until adding the
148/// next would exceed `max_tokens`. Returns the kept messages in original
149/// order. Mirrors upstream's `ContextWindow`/token-budget strategy.
150#[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        // Walk from newest to oldest over the non-system tail, keeping
173        // messages until adding the next would exceed the budget. The
174        // newest non-system message is always kept, even if it alone (plus
175        // the system prefix) exceeds the budget — compaction never reduces
176        // a non-empty tail to nothing.
177        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
194/// Whether a message carries any `Content::FunctionResult` (tool-result)
195/// content.
196fn has_tool_result(message: &Message) -> bool {
197    message
198        .contents
199        .iter()
200        .any(|c| matches!(c, Content::FunctionResult(_)))
201}
202
203/// Drop `Content::FunctionResult` (tool-result) content from all but the last
204/// `keep_last` messages that carry tool results — they are the bulkiest and
205/// least useful once stale. Text and other content is left intact. Messages
206/// that become empty after stripping are dropped entirely. Mirrors upstream's
207/// `ToolResult` strategy.
208#[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
248/// Convenience free function: compact `messages` with `strategy` and
249/// `tokenizer`.
250pub fn compact(
251    messages: &[Message],
252    strategy: &dyn CompactionStrategy,
253    tokenizer: &dyn Tokenizer,
254) -> Vec<Message> {
255    strategy.compact(messages, tokenizer)
256}
257
258/// A [`ContextProvider`] that compacts the accumulated message list —
259/// typically the run's history, once a [`HistoryProvider`](crate::history::HistoryProvider)
260/// has prepended it in `before_run` — down to fit a [`CompactionStrategy`]'s
261/// constraint before it reaches the model. Rust equivalent of (a subset of)
262/// upstream's `CompactionProvider` (see module docs and `UPSTREAM_DRIFT.md`
263/// §9).
264///
265/// Register it via [`AgentBuilder::with_compaction`](crate::agent::AgentBuilder::with_compaction),
266/// which attaches it as one of the agent's own context providers — those run
267/// *after* the session's (which is where a history provider, auto-attached
268/// or explicit, lives — see [`Agent::combined_providers`](crate::agent::Agent)),
269/// so compaction always sees the full, history-prepended message list for the
270/// run.
271pub struct CompactionProvider {
272    strategy: Arc<dyn CompactionStrategy>,
273    tokenizer: Box<dyn Tokenizer>,
274}
275
276impl CompactionProvider {
277    /// A compaction provider using `strategy` with the default
278    /// [`ApproxTokenizer`].
279    pub fn new(strategy: impl CompactionStrategy + 'static) -> Self {
280        Self::with_tokenizer(strategy, ApproxTokenizer)
281    }
282
283    /// A compaction provider using `strategy` and an explicit `tokenizer`.
284    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    /// Replace `ctx.messages` (the accumulated history + any earlier
298    /// provider-injected messages) with the strategy's compacted subset.
299    async fn before_run(&self, ctx: &mut SessionContext) -> Result<()> {
300        ctx.messages = self.strategy.compact(&ctx.messages, &*self.tokenizer);
301        Ok(())
302    }
303
304    // `after_run` is intentionally a no-op (the default from `ContextProvider`):
305    // compaction only shapes the outgoing request, it never observes or
306    // records the run's outcome.
307}
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    // ---- ApproxTokenizer -------------------------------------------------
330
331    #[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); // ceil(5/4) = 2
337        assert_eq!(t.count_tokens("abcdefgh"), 2);
338        assert_eq!(t.count_tokens("abcdefghi"), 3); // ceil(9/4) = 3
339    }
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        // 1 + 2 = 3
349        assert_eq!(count_message_tokens(&t, &msg), 3);
350    }
351
352    // ---- Truncation --------------------------------------------------------
353
354    #[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        // system preserved + 1 most recent (budget of 2 total)
381        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    // ---- SlidingWindow -------------------------------------------------
412
413    #[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    // ---- TokenBudget --------------------------------------------------
443
444    /// A tokenizer with a fixed per-message-call cost, for deterministic
445    /// tests independent of exact text length.
446    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        // Each message costs a fixed 10 tokens; budget for 2 messages.
462        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        // System (10) + budget for one more message (<=20 total).
480        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    // ---- SelectiveToolResult --------------------------------------------
507
508    #[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        // The two oldest tool-result messages become empty and are dropped;
522        // the newest tool-result message is kept intact.
523        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        // First message's tool result is stripped (only the two most recent
549        // tool-result-bearing messages are kept intact), but its text survives.
550        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    // ---- CompactionProvider ---------------------------------------------
581
582    #[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        // Budget of 25 with a fixed 10-token cost per message keeps 2 messages.
616        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}