tokio_prompt_orchestrator/
context_mgr.rs1use std::{
34 collections::VecDeque,
35 time::{SystemTime, UNIX_EPOCH},
36};
37
38pub struct TokenCounter;
45
46impl TokenCounter {
47 pub fn count(text: &str) -> usize {
49 let words = text.split_whitespace().count();
50 words * 4 / 3
51 }
52}
53
54#[derive(Debug, Clone, PartialEq, Eq)]
58pub enum Role {
59 System,
61 User,
63 Assistant,
65 Tool,
67}
68
69impl std::fmt::Display for Role {
70 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
71 match self {
72 Role::System => write!(f, "system"),
73 Role::User => write!(f, "user"),
74 Role::Assistant => write!(f, "assistant"),
75 Role::Tool => write!(f, "tool"),
76 }
77 }
78}
79
80#[derive(Debug, Clone)]
84pub struct Message {
85 pub role: Role,
87 pub content: String,
89 pub token_count: usize,
91 pub timestamp: u64,
93}
94
95impl Message {
96 fn new(role: Role, content: &str) -> Self {
97 let token_count = TokenCounter::count(content);
98 let timestamp = SystemTime::now()
99 .duration_since(UNIX_EPOCH)
100 .unwrap_or_default()
101 .as_secs();
102 Self {
103 role,
104 content: content.to_string(),
105 token_count,
106 timestamp,
107 }
108 }
109}
110
111#[derive(Debug, Clone)]
115pub enum TruncationStrategy {
116 DropOldest,
118 SummarizeOldest,
120 KeepFirst {
123 n: usize,
125 },
126}
127
128#[derive(Debug, Clone)]
132pub struct ContextConfig {
133 pub max_tokens: usize,
135 pub system_prompt: Option<String>,
137 pub reserve_for_response: usize,
140 pub truncation_strategy: TruncationStrategy,
142}
143
144impl ContextConfig {
145 pub fn effective_budget(&self) -> usize {
147 self.max_tokens.saturating_sub(self.reserve_for_response)
148 }
149}
150
151#[derive(Debug, Clone)]
155pub struct ContextSummary {
156 pub message_count: usize,
158 pub total_tokens: usize,
160 pub user_turns: usize,
162 pub assistant_turns: usize,
164 pub oldest_message_age_secs: u64,
166}
167
168pub struct ConversationContext {
172 pub config: ContextConfig,
174 pub messages: VecDeque<Message>,
176 pub total_tokens: usize,
178}
179
180impl ConversationContext {
181 pub fn new(config: ContextConfig) -> Self {
184 let mut ctx = Self {
185 config,
186 messages: VecDeque::new(),
187 total_tokens: 0,
188 };
189 if let Some(sp) = ctx.config.system_prompt.clone() {
191 let msg = Message::new(Role::System, &sp);
192 ctx.total_tokens += msg.token_count;
193 ctx.messages.push_back(msg);
194 }
195 ctx
196 }
197
198 pub fn add_message(&mut self, role: Role, content: &str) {
200 let msg = Message::new(role, content);
201 self.total_tokens += msg.token_count;
202 self.messages.push_back(msg);
203 self.enforce_budget();
204 }
205
206 pub fn enforce_budget(&mut self) {
209 let budget = self.config.effective_budget();
210 if self.total_tokens <= budget {
211 return;
212 }
213 match self.config.truncation_strategy.clone() {
214 TruncationStrategy::DropOldest => {
215 self.drop_oldest(budget);
216 }
217 TruncationStrategy::SummarizeOldest => {
218 self.summarize_oldest(budget);
219 }
220 TruncationStrategy::KeepFirst { n } => {
221 self.keep_first(n, budget);
222 }
223 }
224 }
225
226 fn drop_oldest(&mut self, budget: usize) {
230 while self.total_tokens > budget {
231 let idx = self
233 .messages
234 .iter()
235 .position(|m| m.role != Role::System);
236 match idx {
237 Some(i) => {
238 if let Some(removed) = self.messages.remove(i) {
239 self.total_tokens = self.total_tokens.saturating_sub(removed.token_count);
240 }
241 }
242 None => break, }
244 }
245 }
246
247 fn summarize_oldest(&mut self, budget: usize) {
250 if self.total_tokens <= budget {
251 return;
252 }
253 let reserve = Message::new(Role::System, "[1 messages summarized]").token_count;
255 let mut summarized = 0usize;
256 while self.total_tokens + reserve > budget {
257 let idx = self
258 .messages
259 .iter()
260 .position(|m| m.role != Role::System);
261 match idx {
262 Some(i) => {
263 if let Some(removed) = self.messages.remove(i) {
264 self.total_tokens =
265 self.total_tokens.saturating_sub(removed.token_count);
266 summarized += 1;
267 }
268 }
269 None => break,
270 }
271 }
272 if summarized > 0 {
273 let insert_pos = self
275 .messages
276 .iter()
277 .rposition(|m| m.role == Role::System)
278 .map(|p| p + 1)
279 .unwrap_or(0);
280 let placeholder =
281 format!("[{} messages summarized]", summarized);
282 let msg = Message::new(Role::System, &placeholder);
283 self.total_tokens += msg.token_count;
284 self.messages.insert(insert_pos, msg);
285 }
286 }
287
288 fn keep_first(&mut self, n: usize, budget: usize) {
291 if self.total_tokens <= budget {
292 return;
293 }
294 let (system_msgs, non_system): (Vec<_>, Vec<_>) = self
296 .messages
297 .drain(..)
298 .partition(|m| m.role == Role::System);
299
300 let system_tokens: usize = system_msgs.iter().map(|m| m.token_count).sum();
302 let remaining_budget = budget.saturating_sub(system_tokens);
303
304 let first_n: Vec<_> = non_system.iter().take(n).cloned().collect();
306 let rest: Vec<_> = non_system.into_iter().skip(n).collect();
307
308 let first_n_tokens: usize = first_n.iter().map(|m| m.token_count).sum();
309 let tail_budget = remaining_budget.saturating_sub(first_n_tokens);
310
311 let mut tail: Vec<Message> = Vec::new();
313 let mut tail_tokens = 0usize;
314 for msg in rest.into_iter().rev() {
315 if tail_tokens + msg.token_count > tail_budget {
316 break;
317 }
318 tail_tokens += msg.token_count;
319 tail.push(msg);
320 }
321 tail.reverse();
322
323 self.messages.clear();
325 for m in system_msgs {
326 self.messages.push_back(m);
327 }
328 for m in first_n {
329 self.messages.push_back(m);
330 }
331 for m in tail {
332 self.messages.push_back(m);
333 }
334 self.total_tokens = self.messages.iter().map(|m| m.token_count).sum();
335 }
336
337 pub fn messages_for_api(&self) -> Vec<&Message> {
341 self.messages.iter().collect()
342 }
343
344 pub fn token_utilization(&self) -> f64 {
346 if self.config.max_tokens == 0 {
347 return 1.0;
348 }
349 self.total_tokens as f64 / self.config.max_tokens as f64
350 }
351
352 pub fn summary(&self) -> ContextSummary {
354 let now = SystemTime::now()
355 .duration_since(UNIX_EPOCH)
356 .unwrap_or_default()
357 .as_secs();
358
359 let user_turns = self.messages.iter().filter(|m| m.role == Role::User).count();
360 let assistant_turns = self
361 .messages
362 .iter()
363 .filter(|m| m.role == Role::Assistant)
364 .count();
365 let oldest_age = self
366 .messages
367 .front()
368 .map(|m| now.saturating_sub(m.timestamp))
369 .unwrap_or(0);
370
371 ContextSummary {
372 message_count: self.messages.len(),
373 total_tokens: self.total_tokens,
374 user_turns,
375 assistant_turns,
376 oldest_message_age_secs: oldest_age,
377 }
378 }
379}
380
381#[cfg(test)]
384mod tests {
385 use super::*;
386
387 fn basic_config(max_tokens: usize, strategy: TruncationStrategy) -> ContextConfig {
388 ContextConfig {
389 max_tokens,
390 system_prompt: None,
391 reserve_for_response: 0,
392 truncation_strategy: strategy,
393 }
394 }
395
396 const THREE_WORD_MSG: &str = "hello world foo";
398
399 #[test]
400 fn test_token_counter() {
401 assert_eq!(TokenCounter::count("hello world"), 2); assert_eq!(TokenCounter::count(""), 0);
403 assert_eq!(TokenCounter::count("one two three four five six"), 8);
405 }
406
407 #[test]
408 fn test_add_message_accumulates_tokens() {
409 let config = basic_config(10_000, TruncationStrategy::DropOldest);
410 let mut ctx = ConversationContext::new(config);
411 ctx.add_message(Role::User, THREE_WORD_MSG);
412 assert_eq!(ctx.total_tokens, 4);
414 ctx.add_message(Role::Assistant, THREE_WORD_MSG);
415 assert_eq!(ctx.total_tokens, 8);
416 }
417
418 #[test]
419 fn test_drop_oldest_enforces_budget() {
420 let config = basic_config(8, TruncationStrategy::DropOldest);
423 let mut ctx = ConversationContext::new(config);
424 ctx.add_message(Role::User, THREE_WORD_MSG); ctx.add_message(Role::Assistant, THREE_WORD_MSG); ctx.add_message(Role::User, THREE_WORD_MSG); assert!(ctx.total_tokens <= 8, "tokens={}", ctx.total_tokens);
429 assert_eq!(ctx.messages.len(), 2);
430 }
431
432 #[test]
433 fn test_summarize_oldest_inserts_placeholder() {
434 let config = basic_config(8, TruncationStrategy::SummarizeOldest);
435 let mut ctx = ConversationContext::new(config);
436 ctx.add_message(Role::User, THREE_WORD_MSG); ctx.add_message(Role::Assistant, THREE_WORD_MSG); ctx.add_message(Role::User, THREE_WORD_MSG); assert!(ctx.total_tokens <= 8, "tokens={}", ctx.total_tokens);
442 let has_placeholder = ctx
444 .messages
445 .iter()
446 .any(|m| m.content.contains("messages summarized"));
447 assert!(has_placeholder, "expected a summary placeholder");
448 }
449
450 #[test]
451 fn test_keep_first_strategy() {
452 let config = basic_config(12, TruncationStrategy::KeepFirst { n: 1 });
455 let mut ctx = ConversationContext::new(config);
456 ctx.add_message(Role::User, THREE_WORD_MSG); ctx.add_message(Role::Assistant, THREE_WORD_MSG); ctx.add_message(Role::User, THREE_WORD_MSG); ctx.add_message(Role::Assistant, THREE_WORD_MSG); assert!(ctx.total_tokens <= 12, "tokens={}", ctx.total_tokens);
461 }
462
463 #[test]
464 fn test_messages_for_api() {
465 let config = basic_config(10_000, TruncationStrategy::DropOldest);
466 let mut ctx = ConversationContext::new(config);
467 ctx.add_message(Role::User, "hello");
468 ctx.add_message(Role::Assistant, "world");
469 let msgs = ctx.messages_for_api();
470 assert_eq!(msgs.len(), 2);
471 }
472
473 #[test]
474 fn test_token_utilization() {
475 let config = ContextConfig {
476 max_tokens: 100,
477 system_prompt: None,
478 reserve_for_response: 0,
479 truncation_strategy: TruncationStrategy::DropOldest,
480 };
481 let mut ctx = ConversationContext::new(config);
482 ctx.add_message(Role::User, "hello world"); let util = ctx.token_utilization();
484 assert!(util > 0.0 && util <= 1.0, "util={}", util);
485 }
486
487 #[test]
488 fn test_system_prompt_preserved() {
489 let config = ContextConfig {
490 max_tokens: 16,
491 system_prompt: Some("You are helpful.".to_string()),
492 reserve_for_response: 0,
493 truncation_strategy: TruncationStrategy::DropOldest,
494 };
495 let mut ctx = ConversationContext::new(config);
496 for _ in 0..5 {
498 ctx.add_message(Role::User, "hello world foo");
499 ctx.add_message(Role::Assistant, "hello world foo");
500 }
501 let has_system = ctx.messages.iter().any(|m| m.role == Role::System);
503 assert!(has_system, "system prompt was dropped during truncation");
504 }
505
506 #[test]
507 fn test_summary() {
508 let config = basic_config(10_000, TruncationStrategy::DropOldest);
509 let mut ctx = ConversationContext::new(config);
510 ctx.add_message(Role::User, "question");
511 ctx.add_message(Role::Assistant, "answer");
512 let s = ctx.summary();
513 assert_eq!(s.user_turns, 1);
514 assert_eq!(s.assistant_turns, 1);
515 assert_eq!(s.message_count, 2);
516 }
517
518 #[test]
519 fn test_reserve_for_response_reduces_budget() {
520 let config = ContextConfig {
521 max_tokens: 20,
522 system_prompt: None,
523 reserve_for_response: 12,
524 truncation_strategy: TruncationStrategy::DropOldest,
525 };
526 assert_eq!(config.effective_budget(), 8);
528 let mut ctx = ConversationContext::new(config);
529 ctx.add_message(Role::User, THREE_WORD_MSG); ctx.add_message(Role::User, THREE_WORD_MSG); ctx.add_message(Role::User, THREE_WORD_MSG); assert!(ctx.total_tokens <= 8, "tokens={}", ctx.total_tokens);
534 }
535}