Skip to main content

agent_base/engine/
context.rs

1use crate::types::{ChatMessage, SessionId};
2
3/// Find the first non-ephemeral System message (the system prompt).
4pub fn first_system_prompt(messages: &[ChatMessage]) -> Option<ChatMessage> {
5    messages.iter().find_map(|msg| match msg {
6        ChatMessage::System {
7            content,
8            ephemeral: false,
9        } => Some(ChatMessage::system(content.clone())),
10        _ => None,
11    })
12}
13
14/// Estimate total tokens across a message list using `ContextWindowManager::message_tokens`.
15pub fn estimate_messages_tokens(messages: &[ChatMessage]) -> usize {
16    messages
17        .iter()
18        .map(ContextWindowManager::message_tokens)
19        .sum()
20}
21
22// ── Context Window Manager ──────────────────────────────────────────────────
23
24#[derive(Clone, Debug)]
25pub struct ContextWindowManager {
26    pub max_tokens: usize,
27    /// Always keep first N messages (typically system prompt)
28    pub keep_first_n: usize,
29    /// Always keep last N messages
30    pub keep_last_n: usize,
31}
32
33impl Default for ContextWindowManager {
34    fn default() -> Self {
35        Self {
36            max_tokens: 128_000,
37            keep_first_n: 1,
38            keep_last_n: 20,
39        }
40    }
41}
42
43impl ContextWindowManager {
44    /// OpenAI Vision API fixed token overhead per image
45    const IMAGE_OVERHEAD_TOKENS: usize = 85;
46
47    pub fn new(max_tokens: usize) -> Self {
48        Self {
49            max_tokens,
50            ..Default::default()
51        }
52    }
53
54    pub fn with_keep_first_n(mut self, n: usize) -> Self {
55        self.keep_first_n = n;
56        self
57    }
58
59    pub fn with_keep_last_n(mut self, n: usize) -> Self {
60        self.keep_last_n = n;
61        self
62    }
63
64    /// Simple token estimation: ~4 chars/token for Latin, ~1.5 for CJK
65    /// Mixed text uses a compromise of 3 chars/token
66    pub fn estimate_tokens(text: &str) -> usize {
67        if text.is_empty() {
68            return 0;
69        }
70        let chars = text.chars().count();
71        let cjk_count = text.chars().filter(|c| is_cjk(*c)).count();
72        let latin_count = chars - cjk_count;
73        // CJK: ~1.5 chars/token, Latin: ~4 chars/token
74        (cjk_count as f64 / 1.5 + latin_count as f64 / 4.0).ceil() as usize
75    }
76
77    pub(crate) fn message_tokens(msg: &ChatMessage) -> usize {
78        match msg {
79            ChatMessage::System { content, .. } => Self::estimate_tokens(content),
80            ChatMessage::User {
81                content, images, ..
82            } => {
83                let mut tokens = Self::estimate_tokens(content);
84                for img in images {
85                    match img {
86                        crate::types::ImageAttachment::Url { url, detail: _ } => {
87                            tokens += Self::estimate_tokens(url);
88                        }
89                        crate::types::ImageAttachment::Base64 {
90                            data,
91                            media_type,
92                            detail: _,
93                        } => {
94                            tokens += data.len() / 4;
95                            if let Some(mt) = media_type {
96                                tokens += Self::estimate_tokens(mt);
97                            }
98                        }
99                    }
100                    tokens += Self::IMAGE_OVERHEAD_TOKENS;
101                }
102                tokens
103            }
104            ChatMessage::Assistant {
105                content,
106                reasoning_content,
107                tool_calls,
108                thinking_signature: _,
109            } => {
110                let mut tokens = content.as_deref().map(Self::estimate_tokens).unwrap_or(0);
111                if let Some(rc) = reasoning_content {
112                    tokens += Self::estimate_tokens(rc);
113                }
114                if let Some(tc) = tool_calls {
115                    for t in tc {
116                        tokens += Self::estimate_tokens(&t.name);
117                        tokens += Self::estimate_tokens(&t.arguments);
118                        tokens += Self::estimate_tokens(&t.id);
119                    }
120                }
121                tokens
122            }
123            ChatMessage::Tool {
124                tool_call_id,
125                content,
126                ..
127            } => Self::estimate_tokens(tool_call_id) + Self::estimate_tokens(content),
128            ChatMessage::Custom { role, data } => {
129                Self::estimate_tokens(role) + Self::estimate_tokens(&data.to_string())
130            }
131        }
132    }
133
134    /// Trim message list to keep total tokens under `max_tokens`。
135    ///
136    /// Trimming strategy:
137    /// - Always keep the first `keep_first_n` messages (typically system prompt)
138    /// - Always keep the last `keep_last_n` messages (recent conversation)
139    /// - Remove oldest messages from the middle until within budget
140    pub fn trim(&self, messages: &mut Vec<ChatMessage>) {
141        if messages.is_empty() || self.max_tokens == 0 {
142            return;
143        }
144
145        let total_tokens: usize = messages.iter().map(Self::message_tokens).sum();
146        if total_tokens <= self.max_tokens {
147            return;
148        }
149
150        let keep_first = self.keep_first_n.min(messages.len());
151        let keep_last = self
152            .keep_last_n
153            .min(messages.len().saturating_sub(keep_first));
154
155        // Trimmable range: [keep_first, messages.len() - keep_last)
156        let trim_start = keep_first;
157        let trim_end = messages.len().saturating_sub(keep_last);
158        if trim_start >= trim_end {
159            return;
160        }
161
162        let mut current_tokens: usize = total_tokens;
163        let remove_idx = trim_start;
164        let mut trim_end = trim_end;
165
166        while current_tokens > self.max_tokens && remove_idx < trim_end {
167            let removed = Self::message_tokens(&messages[remove_idx]);
168            messages.remove(remove_idx);
169            current_tokens = current_tokens.saturating_sub(removed);
170            trim_end = messages.len().saturating_sub(keep_last);
171        }
172    }
173}
174
175fn is_cjk(c: char) -> bool {
176    matches!(
177        c,
178        '\u{4E00}'..='\u{9FFF}'   // CJK Unified Ideographs
179        | '\u{3400}'..='\u{4DBF}' // CJK Unified Ideographs Extension A
180        | '\u{3000}'..='\u{303F}' // CJK Symbols and Punctuation
181        | '\u{FF00}'..='\u{FFEF}' // Halfwidth and Fullwidth Forms
182        | '\u{3040}'..='\u{309F}' // Hiragana
183        | '\u{30A0}'..='\u{30FF}' // Katakana
184        | '\u{AC00}'..='\u{D7AF}' // Hangul Syllables
185    )
186}
187
188// ── Inline Context Compaction ──────────────────────────────────────────────
189
190/// What a successful [`ContextCompaction::compact`] actually did to the
191/// history. agent-base stays strategy-free: it only relays the kind into
192/// the log line, so a nudge append never reads as a compaction and a real
193/// replacement always does (issue #33).
194#[derive(Debug, Clone, Copy, PartialEq, Eq)]
195pub enum CompactionKind {
196    /// One budget nudge appended at the end; prior messages untouched.
197    Reminder,
198    /// Last-turn nudge appended; prior messages untouched.
199    Fallback,
200    /// Messages replaced wholesale (window rotation or compression).
201    Reset,
202}
203
204impl CompactionKind {
205    /// The grep-able log line emitted when the compactor returns this kind.
206    /// The three strings are deliberately distinct: a Reminder/Fallback append
207    /// must NOT contain "compaction" (nothing was compacted — one message was
208    /// appended), while a Reset must (the history was replaced).
209    pub fn trigger_log(self) -> &'static str {
210        match self {
211            Self::Reminder => "context reminder appended",
212            Self::Fallback => "context fallback appended",
213            Self::Reset => "inline compaction triggered",
214        }
215    }
216
217    /// The completion log line, emitted after the compactor's messages are
218    /// installed. `None` for appends — they are atomic, there is nothing to
219    /// complete, and a second line would only invite misreading.
220    pub fn completion_log(self) -> Option<&'static str> {
221        match self {
222            Self::Reminder | Self::Fallback => None,
223            Self::Reset => Some("inline compaction completed (window reset)"),
224        }
225    }
226}
227
228/// Successful compaction outcome: what happened plus the messages to install.
229pub struct CompactionOutcome {
230    pub kind: CompactionKind,
231    pub messages: Vec<ChatMessage>,
232}
233
234/// Trait for inline context compaction within the react loop.
235///
236/// Implemented by agent-works's `ContextCompactor`. agent-base defines
237/// the trait to avoid circular dependencies (agent-works depends on agent-base).
238///
239/// The react loop calls [`compact`](Self::compact) after tool execution when
240/// the estimated token count exceeds a configurable threshold. This prevents
241/// context window overflow without requiring the LLM call to fail first.
242#[async_trait::async_trait]
243pub trait ContextCompaction: Send + Sync {
244    /// Compact a message history.
245    ///
246    /// Takes the current messages and returns `Some(outcome)` if an action
247    /// was performed — the outcome carries the action kind (for logging) and
248    /// the messages to install. Returns `None` if compaction was skipped
249    /// (below threshold, too few messages, or disabled).
250    ///
251    /// The react loop handles reading/writing the session — the compactor
252    /// only transforms the message list.
253    async fn compact(
254        &self,
255        session_id: &SessionId,
256        messages: &[ChatMessage],
257    ) -> Option<CompactionOutcome>;
258
259    /// Estimate current token count for the session.
260    ///
261    /// Returns `None` if the implementation cannot estimate (falls back to
262    /// `ContextWindowManager::estimate_tokens` on the react loop side).
263    fn token_count_hint(&self, session_id: &SessionId) -> Option<usize>;
264}
265
266#[cfg(test)]
267mod tests {
268    use super::*;
269    use crate::types::{ImageAttachment, ToolCallMessage};
270
271    #[test]
272    fn test_estimate_tokens_empty() {
273        assert_eq!(ContextWindowManager::estimate_tokens(""), 0);
274    }
275
276    #[test]
277    fn test_estimate_tokens_english() {
278        let text = "Hello world this is a test";
279        let tokens = ContextWindowManager::estimate_tokens(text);
280        // ~28 chars / 4 ≈ 7
281        assert!(tokens > 0 && tokens <= 15);
282    }
283
284    #[test]
285    fn test_estimate_tokens_cjk() {
286        // 4 CJK chars / 1.5 -> ceil(2.667) = 3
287        assert_eq!(ContextWindowManager::estimate_tokens("你好世界"), 3);
288    }
289
290    #[test]
291    fn test_estimate_tokens_mixed() {
292        // 2 CJK + 5 latin: 2/1.5 + 5/4 = 1.333 + 1.25 = 2.583 -> 3
293        assert_eq!(ContextWindowManager::estimate_tokens("你好hello"), 3);
294    }
295
296    #[test]
297    fn test_message_tokens_user_with_url_image() {
298        let msg = ChatMessage::user_with_images(
299            "pic",
300            vec![ImageAttachment::Url {
301                url: "http://x/a.png".into(),
302                detail: None,
303            }],
304        );
305        let base = ContextWindowManager::message_tokens(&ChatMessage::user("pic"));
306        let t = ContextWindowManager::message_tokens(&msg);
307        assert!(t > base);
308    }
309
310    #[test]
311    fn test_message_tokens_user_with_base64_image() {
312        // with media_type
313        let msg = ChatMessage::user_with_images(
314            "pic",
315            vec![ImageAttachment::Base64 {
316                data: "abcd".into(),
317                media_type: Some("image/png".into()),
318                detail: None,
319            }],
320        );
321        let base = ContextWindowManager::message_tokens(&ChatMessage::user("pic"));
322        assert!(ContextWindowManager::message_tokens(&msg) > base);
323
324        // without media_type
325        let msg = ChatMessage::user_with_images(
326            "pic",
327            vec![ImageAttachment::Base64 {
328                data: "abcd".into(),
329                media_type: None,
330                detail: None,
331            }],
332        );
333        assert!(ContextWindowManager::message_tokens(&msg) > base);
334    }
335
336    #[test]
337    fn test_message_tokens_assistant_reasoning_and_tool_calls() {
338        let msg = ChatMessage::Assistant {
339            content: Some("ans".into()),
340            reasoning_content: Some("thinking".into()),
341            tool_calls: Some(vec![ToolCallMessage {
342                id: "tc1".into(),
343                name: "echo".into(),
344                arguments: "{}".into(),
345            }]),
346            thinking_signature: None,
347        };
348        let t = ContextWindowManager::message_tokens(&msg);
349        assert!(t > 0);
350    }
351
352    #[test]
353    fn test_message_tokens_tool_and_custom() {
354        let tool = ChatMessage::tool("tc1", "done");
355        assert!(ContextWindowManager::message_tokens(&tool) > 0);
356
357        let custom = ChatMessage::Custom {
358            role: "artifact".into(),
359            data: serde_json::json!({"x": 1}),
360        };
361        assert!(ContextWindowManager::message_tokens(&custom) > 0);
362    }
363
364    #[test]
365    fn test_trim_no_trim_needed() {
366        let mgr = ContextWindowManager::new(1000);
367        let mut msgs = vec![
368            ChatMessage::system("You are a helpful assistant."),
369            ChatMessage::user("Hello"),
370            ChatMessage::assistant("Hi there!"),
371        ];
372        let original_len = msgs.len();
373        mgr.trim(&mut msgs);
374        assert_eq!(msgs.len(), original_len);
375    }
376
377    #[test]
378    fn test_trim_keeps_first_and_last() {
379        let mgr = ContextWindowManager::new(8)
380            .with_keep_first_n(1)
381            .with_keep_last_n(2);
382        let mut msgs = vec![
383            ChatMessage::system("system"),
384            ChatMessage::user("message number one"),
385            ChatMessage::assistant("message number two"),
386            ChatMessage::user("message number three"),
387            ChatMessage::assistant("message number four"),
388            ChatMessage::user("message number five"),
389            ChatMessage::assistant("message number six"),
390        ];
391        mgr.trim(&mut msgs);
392        assert_eq!(msgs.len(), 3);
393        assert!(matches!(msgs[0], ChatMessage::System { .. }));
394    }
395}
396
397#[cfg(test)]
398mod proptest_tests {
399    use super::*;
400    use proptest::prelude::*;
401
402    proptest! {
403        #[test]
404        fn estimate_tokens_never_panics(text in ".*") {
405            let tokens = ContextWindowManager::estimate_tokens(&text);
406            // tokens should be non-negative (usize) and reasonable
407            assert!(tokens <= text.len() + 1); // at most 1 token per byte + ceil
408        }
409
410        #[test]
411        fn estimate_tokens_empty_is_zero(text in "[a-z\u{4e00}-\u{9fff}]{0,100}") {
412            if text.is_empty() {
413                assert_eq!(ContextWindowManager::estimate_tokens(&text), 0);
414            } else {
415                assert!(ContextWindowManager::estimate_tokens(&text) > 0);
416            }
417        }
418
419        #[test]
420        fn estimate_tokens_cjk_higher_than_latin_same_len(
421            cjk_text in "[\u{4e00}-\u{9fff}]{1,50}",
422            latin_text in "[a-z]{1,50}",
423        ) {
424            // Pad to same char length
425            let max_len = cjk_text.chars().count().max(latin_text.chars().count());
426            let cjk_padded: String = cjk_text.chars().cycle().take(max_len).collect();
427            let latin_padded: String = latin_text.chars().cycle().take(max_len).collect();
428            let cjk_tokens = ContextWindowManager::estimate_tokens(&cjk_padded);
429            let latin_tokens = ContextWindowManager::estimate_tokens(&latin_padded);
430            // CJK ~1.5 chars/token, Latin ~4 chars/token → CJK uses more tokens
431            assert!(cjk_tokens >= latin_tokens,
432                "CJK ({}) should use >= tokens than Latin ({}) for {} chars",
433                cjk_tokens, latin_tokens, max_len);
434        }
435
436        #[test]
437        fn trim_preserves_system_prefix(
438            num_messages in 2usize..15,
439            max_tokens in 5usize..50,
440        ) {
441            let mgr = ContextWindowManager {
442                max_tokens,
443                keep_first_n: 1,
444                keep_last_n: 0,
445            };
446            let mut msgs = vec![ChatMessage::system("system prompt")];
447            for i in 0..num_messages {
448                msgs.push(ChatMessage::user(format!("message {}", i)));
449            }
450            mgr.trim(&mut msgs);
451            // System message should always be preserved
452            assert!(!msgs.is_empty());
453            assert!(matches!(msgs[0], ChatMessage::System { .. }));
454        }
455
456        #[test]
457        fn trim_result_within_budget(
458            num_messages in 3usize..15,
459            max_tokens in 10usize..100,
460        ) {
461            let mgr = ContextWindowManager {
462                max_tokens,
463                keep_first_n: 1,
464                keep_last_n: 1,
465            };
466            let mut msgs = vec![ChatMessage::system("sys")];
467            for i in 0..num_messages {
468                msgs.push(ChatMessage::user(format!("msg {}", i)));
469            }
470            let total_before: usize = msgs.iter().map(ContextWindowManager::message_tokens).sum();
471            // Only test when we actually exceed budget
472            if total_before > max_tokens {
473                mgr.trim(&mut msgs);
474                let total_after: usize = msgs.iter().map(ContextWindowManager::message_tokens).sum();
475                // After trimming, should be within budget (or couldn't trim more)
476                assert!(total_after <= total_before);
477            }
478        }
479    }
480
481    // ── CompactionKind log dispatch (issue #33) ──
482
483    #[test]
484    fn compaction_kind_log_lines_are_grep_able() {
485        use super::{CompactionKind, CompactionOutcome};
486
487        // A nudge append must never read as a compaction in the log...
488        for kind in [CompactionKind::Reminder, CompactionKind::Fallback] {
489            let trigger = kind.trigger_log();
490            assert!(
491                !trigger.contains("compaction"),
492                "{kind:?} trigger log must not claim compaction: {trigger}"
493            );
494            assert!(
495                kind.completion_log().is_none(),
496                "{kind:?} append is atomic — no completion line"
497            );
498        }
499        // ...while a real replacement must say compaction.
500        let reset = CompactionKind::Reset;
501        assert!(reset.trigger_log().contains("compaction"));
502        let completion = reset.completion_log().expect("reset has a completion line");
503        assert!(completion.contains("compaction"));
504
505        // The three kinds are distinguishable by grep in session.log.
506        let all = [
507            CompactionKind::Reminder.trigger_log(),
508            CompactionKind::Fallback.trigger_log(),
509            reset.trigger_log(),
510        ];
511        for (i, a) in all.iter().enumerate() {
512            for b in all.iter().skip(i + 1) {
513                assert_ne!(a, b, "trigger log lines must be distinct");
514            }
515        }
516
517        // Outcome struct: kind + messages travel together.
518        let outcome = CompactionOutcome {
519            kind: CompactionKind::Reminder,
520            messages: vec![ChatMessage::user("nudge")],
521        };
522        assert_eq!(outcome.kind, CompactionKind::Reminder);
523        assert_eq!(outcome.messages.len(), 1);
524    }
525}