1use crate::types::{ChatMessage, SessionId};
2
3pub 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
14pub fn estimate_messages_tokens(messages: &[ChatMessage]) -> usize {
16 messages
17 .iter()
18 .map(ContextWindowManager::message_tokens)
19 .sum()
20}
21
22#[derive(Clone, Debug)]
25pub struct ContextWindowManager {
26 pub max_tokens: usize,
27 pub keep_first_n: usize,
29 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 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 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_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 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 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}' | '\u{3400}'..='\u{4DBF}' | '\u{3000}'..='\u{303F}' | '\u{FF00}'..='\u{FFEF}' | '\u{3040}'..='\u{309F}' | '\u{30A0}'..='\u{30FF}' | '\u{AC00}'..='\u{D7AF}' )
186}
187
188#[derive(Debug, Clone, Copy, PartialEq, Eq)]
195pub enum CompactionKind {
196 Reminder,
198 Fallback,
200 Reset,
202}
203
204impl CompactionKind {
205 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 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
228pub struct CompactionOutcome {
230 pub kind: CompactionKind,
231 pub messages: Vec<ChatMessage>,
232}
233
234#[async_trait::async_trait]
243pub trait ContextCompaction: Send + Sync {
244 async fn compact(
254 &self,
255 session_id: &SessionId,
256 messages: &[ChatMessage],
257 ) -> Option<CompactionOutcome>;
258
259 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 assert!(tokens > 0 && tokens <= 15);
282 }
283
284 #[test]
285 fn test_estimate_tokens_cjk() {
286 assert_eq!(ContextWindowManager::estimate_tokens("你好世界"), 3);
288 }
289
290 #[test]
291 fn test_estimate_tokens_mixed() {
292 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 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 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 assert!(tokens <= text.len() + 1); }
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 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 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 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 if total_before > max_tokens {
473 mgr.trim(&mut msgs);
474 let total_after: usize = msgs.iter().map(ContextWindowManager::message_tokens).sum();
475 assert!(total_after <= total_before);
477 }
478 }
479 }
480
481 #[test]
484 fn compaction_kind_log_lines_are_grep_able() {
485 use super::{CompactionKind, CompactionOutcome};
486
487 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 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 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 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}