Skip to main content

tokio_prompt_orchestrator/
context_mgr.rs

1//! # Context Manager
2//!
3//! Multi-turn LLM conversation context manager with token budget enforcement
4//! and configurable truncation strategies.
5//!
6//! ## Overview
7//!
8//! [`ConversationContext`] maintains an ordered deque of [`Message`] values and
9//! automatically enforces a token budget when new messages are added.  Three
10//! truncation strategies are supported: dropping the oldest messages, keeping
11//! only the first N turns plus recent history, and replacing oldest messages
12//! with a summary placeholder.
13//!
14//! ## Example
15//!
16//! ```rust
17//! use tokio_prompt_orchestrator::context_mgr::{
18//!     ConversationContext, ContextConfig, Role, TruncationStrategy,
19//! };
20//!
21//! let config = ContextConfig {
22//!     max_tokens: 1000,
23//!     system_prompt: Some("You are a helpful assistant.".into()),
24//!     reserve_for_response: 100,
25//!     truncation_strategy: TruncationStrategy::DropOldest,
26//! };
27//! let mut ctx = ConversationContext::new(config);
28//! ctx.add_message(Role::User, "Hello!");
29//! ctx.add_message(Role::Assistant, "Hi there!");
30//! assert!(!ctx.messages_for_api().is_empty());
31//! ```
32
33use std::{
34    collections::VecDeque,
35    time::{SystemTime, UNIX_EPOCH},
36};
37
38// ── TokenCounter ──────────────────────────────────────────────────────────────
39
40/// Simple word-count token estimator.
41///
42/// Approximates token count as `words * 4 / 3` (~1.33 tokens per word), which
43/// is a reasonable heuristic for most LLM tokenisers.
44pub struct TokenCounter;
45
46impl TokenCounter {
47    /// Estimate the number of tokens in `text`.
48    pub fn count(text: &str) -> usize {
49        let words = text.split_whitespace().count();
50        words * 4 / 3
51    }
52}
53
54// ── Role ──────────────────────────────────────────────────────────────────────
55
56/// The role of a message participant in a conversation.
57#[derive(Debug, Clone, PartialEq, Eq)]
58pub enum Role {
59    /// A system-level instruction (not counted toward turn limits).
60    System,
61    /// A human turn.
62    User,
63    /// A model turn.
64    Assistant,
65    /// A tool/function response.
66    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// ── Message ───────────────────────────────────────────────────────────────────
81
82/// A single message in a conversation.
83#[derive(Debug, Clone)]
84pub struct Message {
85    /// Who sent the message.
86    pub role: Role,
87    /// The message text.
88    pub content: String,
89    /// Estimated token count (computed at insertion time).
90    pub token_count: usize,
91    /// Unix timestamp (seconds) of when this message was added.
92    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// ── TruncationStrategy ────────────────────────────────────────────────────────
112
113/// How to reduce the context when the token budget is exceeded.
114#[derive(Debug, Clone)]
115pub enum TruncationStrategy {
116    /// Remove the oldest non-system messages one by one until within budget.
117    DropOldest,
118    /// Replace the oldest non-system messages with a single summary placeholder.
119    SummarizeOldest,
120    /// Keep the system prompt, the first `n` user/assistant turns, and fill
121    /// the remainder with the most recent messages.
122    KeepFirst {
123        /// Number of early turns to preserve.
124        n: usize,
125    },
126}
127
128// ── ContextConfig ─────────────────────────────────────────────────────────────
129
130/// Configuration for a [`ConversationContext`].
131#[derive(Debug, Clone)]
132pub struct ContextConfig {
133    /// Hard token budget for the entire context window.
134    pub max_tokens: usize,
135    /// Optional system prompt prepended to every context window.
136    pub system_prompt: Option<String>,
137    /// Token budget reserved for the model's response.  Effective budget is
138    /// `max_tokens - reserve_for_response`.
139    pub reserve_for_response: usize,
140    /// Strategy to apply when the token budget is exceeded.
141    pub truncation_strategy: TruncationStrategy,
142}
143
144impl ContextConfig {
145    /// Effective token budget after reserving space for the model response.
146    pub fn effective_budget(&self) -> usize {
147        self.max_tokens.saturating_sub(self.reserve_for_response)
148    }
149}
150
151// ── ContextSummary ────────────────────────────────────────────────────────────
152
153/// A lightweight summary of the current context state.
154#[derive(Debug, Clone)]
155pub struct ContextSummary {
156    /// Number of messages currently in the context (including system).
157    pub message_count: usize,
158    /// Total estimated token usage.
159    pub total_tokens: usize,
160    /// Number of user turns.
161    pub user_turns: usize,
162    /// Number of assistant turns.
163    pub assistant_turns: usize,
164    /// Age in seconds of the oldest message (0 if no messages).
165    pub oldest_message_age_secs: u64,
166}
167
168// ── ConversationContext ────────────────────────────────────────────────────────
169
170/// Multi-turn conversation context with automatic token budget enforcement.
171pub struct ConversationContext {
172    /// Configuration (strategy, budgets, system prompt).
173    pub config: ContextConfig,
174    /// Ordered message history (front = oldest).
175    pub messages: VecDeque<Message>,
176    /// Running total of estimated tokens across all messages.
177    pub total_tokens: usize,
178}
179
180impl ConversationContext {
181    /// Create a new context.  If the config has a `system_prompt`, it is
182    /// added as the first message immediately.
183    pub fn new(config: ContextConfig) -> Self {
184        let mut ctx = Self {
185            config,
186            messages: VecDeque::new(),
187            total_tokens: 0,
188        };
189        // Pre-load the system prompt.
190        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    /// Append a new message and enforce the token budget.
199    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    /// Apply the configured [`TruncationStrategy`] if the effective budget is
207    /// exceeded.
208    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    // -- internal truncation helpers ------------------------------------------
227
228    /// Drop the oldest non-system messages until within `budget`.
229    fn drop_oldest(&mut self, budget: usize) {
230        while self.total_tokens > budget {
231            // Find the first non-system message index.
232            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, // only system message left, nothing more to drop
243            }
244        }
245    }
246
247    /// Replace the oldest non-system messages with a single summary placeholder
248    /// until within `budget`.
249    fn summarize_oldest(&mut self, budget: usize) {
250        if self.total_tokens <= budget {
251            return;
252        }
253        // Reserve room for the placeholder itself so the result fits the budget.
254        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            // Insert the placeholder right after the last system message.
274            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    /// Keep the system prompt + first `n` non-system messages + most recent
289    /// messages, dropping middle messages, until within `budget`.
290    fn keep_first(&mut self, n: usize, budget: usize) {
291        if self.total_tokens <= budget {
292            return;
293        }
294        // Collect system messages and non-system messages separately.
295        let (system_msgs, non_system): (Vec<_>, Vec<_>) = self
296            .messages
297            .drain(..)
298            .partition(|m| m.role == Role::System);
299
300        // Determine tokens used by system messages.
301        let system_tokens: usize = system_msgs.iter().map(|m| m.token_count).sum();
302        let remaining_budget = budget.saturating_sub(system_tokens);
303
304        // Split non-system into first-n and the rest.
305        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        // Fill tail from the most recent messages.
312        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        // Reassemble.
324        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    // -- public query API --------------------------------------------------------
338
339    /// Return the messages that should be sent to the LLM API.
340    pub fn messages_for_api(&self) -> Vec<&Message> {
341        self.messages.iter().collect()
342    }
343
344    /// Token utilisation ratio in `[0.0, 1.0]`.
345    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    /// Return a lightweight summary snapshot.
353    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// ── Tests ─────────────────────────────────────────────────────────────────────
382
383#[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    // Helper: 3 words → TokenCounter::count = 3 * 4 / 3 = 4 tokens.
397    const THREE_WORD_MSG: &str = "hello world foo";
398
399    #[test]
400    fn test_token_counter() {
401        assert_eq!(TokenCounter::count("hello world"), 2); // 2 * 4/3 = 2
402        assert_eq!(TokenCounter::count(""), 0);
403        // 6 words → 6 * 4 / 3 = 8
404        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        // 3 words → 4 tokens
413        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        // Budget: 8 tokens.  Each 3-word message = 4 tokens.
421        // After 3 messages the total would be 12 > 8 so oldest must be dropped.
422        let config = basic_config(8, TruncationStrategy::DropOldest);
423        let mut ctx = ConversationContext::new(config);
424        ctx.add_message(Role::User, THREE_WORD_MSG); // 4 tokens
425        ctx.add_message(Role::Assistant, THREE_WORD_MSG); // 4 tokens → total 8
426        ctx.add_message(Role::User, THREE_WORD_MSG); // 4 tokens → would be 12, trigger drop
427        // One message should have been dropped to get back to ≤ 8.
428        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);     // 4 tokens
437        ctx.add_message(Role::Assistant, THREE_WORD_MSG); // 4 tokens → 8
438        ctx.add_message(Role::User, THREE_WORD_MSG);     // 4 → 12, trigger summarize
439
440        // Total should be within budget.
441        assert!(ctx.total_tokens <= 8, "tokens={}", ctx.total_tokens);
442        // There should be a placeholder message.
443        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        // max_tokens=12, n=1 (keep 1 first turn + recent).
453        // Each message = 4 tokens.
454        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);     // msg A
457        ctx.add_message(Role::Assistant, THREE_WORD_MSG); // msg B
458        ctx.add_message(Role::User, THREE_WORD_MSG);     // msg C
459        ctx.add_message(Role::Assistant, THREE_WORD_MSG); // msg D → 16 tokens > 12
460        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"); // 2 words → 2 tokens
483        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        // Fill with user/assistant pairs until truncation kicks in.
497        for _ in 0..5 {
498            ctx.add_message(Role::User, "hello world foo");
499            ctx.add_message(Role::Assistant, "hello world foo");
500        }
501        // System message must never be dropped.
502        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        // effective budget = 20 - 12 = 8 tokens.
527        assert_eq!(config.effective_budget(), 8);
528        let mut ctx = ConversationContext::new(config);
529        // Each 3-word message ≈ 4 tokens.
530        ctx.add_message(Role::User, THREE_WORD_MSG); // 4
531        ctx.add_message(Role::User, THREE_WORD_MSG); // 4 → 8 (at budget)
532        ctx.add_message(Role::User, THREE_WORD_MSG); // 4 → 12, must drop
533        assert!(ctx.total_tokens <= 8, "tokens={}", ctx.total_tokens);
534    }
535}