Skip to main content

tokio_prompt_orchestrator/
conversation.rs

1//! # Conversation History Manager
2//!
3//! Tracks multi-turn conversation history per [`SessionId`], automatically injects
4//! prior turns into outgoing prompts, and compresses old history when approaching
5//! a configurable token budget.
6//!
7//! ## Quick start
8//!
9//! ```rust,no_run
10//! use tokio_prompt_orchestrator::conversation::{ConversationManager, ConversationConfig};
11//! use tokio_prompt_orchestrator::SessionId;
12//!
13//! #[tokio::main]
14//! async fn main() {
15//!     let mgr = ConversationManager::new(ConversationConfig::default());
16//!     let sid = SessionId::new("user-42");
17//!
18//!     // Record a user turn
19//!     mgr.push_user(&sid, "What is backpressure?").await;
20//!
21//!     // Build a prompt that includes conversation history
22//!     let prompt = mgr.build_prompt(&sid, "Explain it with an example.").await;
23//!     println!("{prompt}");
24//!
25//!     // Record the assistant's reply
26//!     mgr.push_assistant(&sid, "Backpressure is…").await;
27//! }
28//! ```
29
30use std::{
31    collections::HashMap,
32    sync::Arc,
33    time::{Duration, SystemTime},
34};
35use tokio::sync::RwLock;
36
37use crate::SessionId;
38
39// ---------------------------------------------------------------------------
40// Types
41// ---------------------------------------------------------------------------
42
43/// Role of a conversation turn.
44#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
45#[serde(rename_all = "lowercase")]
46pub enum Role {
47    System,
48    User,
49    Assistant,
50    /// Tool / function result injected mid-conversation.
51    Tool,
52}
53
54impl std::fmt::Display for Role {
55    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
56        let s = match self {
57            Role::System => "system",
58            Role::User => "user",
59            Role::Assistant => "assistant",
60            Role::Tool => "tool",
61        };
62        f.write_str(s)
63    }
64}
65
66/// A single turn in a conversation.
67#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
68pub struct Turn {
69    pub id: String,
70    pub role: Role,
71    pub content: String,
72    pub timestamp: SystemTime,
73    /// Rough token estimate (4 chars ≈ 1 token).
74    pub token_estimate: usize,
75    /// Optional key–value tags for filtering or retrieval.
76    pub tags: Vec<String>,
77    /// True if this turn was synthesised during compression (not verbatim).
78    pub is_summary: bool,
79}
80
81impl Turn {
82    fn new(role: Role, content: impl Into<String>) -> Self {
83        let content = content.into();
84        let token_estimate = estimate_tokens(&content);
85        Turn {
86            id: new_id(),
87            role,
88            content,
89            timestamp: SystemTime::now(),
90            token_estimate,
91            tags: Vec::new(),
92            is_summary: false,
93        }
94    }
95}
96
97/// A single conversation thread.
98#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
99pub struct Conversation {
100    pub id: String,
101    pub session_id: String,
102    pub turns: Vec<Turn>,
103    pub total_tokens: usize,
104    pub created_at: SystemTime,
105    pub last_active: SystemTime,
106    /// Number of times history was compressed.
107    pub compressions: u32,
108    /// Metadata bag for callers to attach arbitrary data.
109    pub meta: HashMap<String, String>,
110}
111
112impl Conversation {
113    fn new(session_id: &str) -> Self {
114        let now = SystemTime::now();
115        Conversation {
116            id: new_id(),
117            session_id: session_id.to_owned(),
118            turns: Vec::new(),
119            total_tokens: 0,
120            created_at: now,
121            last_active: now,
122            compressions: 0,
123            meta: HashMap::new(),
124        }
125    }
126
127    fn push(&mut self, turn: Turn) {
128        self.total_tokens += turn.token_estimate;
129        self.last_active = SystemTime::now();
130        self.turns.push(turn);
131    }
132}
133
134// ---------------------------------------------------------------------------
135// Configuration
136// ---------------------------------------------------------------------------
137
138/// Configuration for [`ConversationManager`].
139#[derive(Debug, Clone)]
140pub struct ConversationConfig {
141    /// Maximum tokens allowed in a conversation before compression is triggered.
142    /// Default: 6 000.
143    pub max_tokens: usize,
144
145    /// Always keep the N most recent turns verbatim (never compressed).
146    /// Default: 4.
147    pub recency_keep: usize,
148
149    /// System prompt injected at the start of every built prompt.
150    /// If `None`, no system prompt is added.
151    pub system_prompt: Option<String>,
152
153    /// Conversations inactive for longer than this duration are evicted.
154    /// Default: 2 hours.
155    pub ttl: Duration,
156
157    /// Format used when assembling a multi-turn prompt.
158    /// Default: `ChatMl`.
159    pub format: PromptFormat,
160}
161
162impl Default for ConversationConfig {
163    fn default() -> Self {
164        ConversationConfig {
165            max_tokens: 6_000,
166            recency_keep: 4,
167            system_prompt: None,
168            ttl: Duration::from_secs(2 * 3600),
169            format: PromptFormat::ChatMl,
170        }
171    }
172}
173
174/// Wire format used when serialising conversation history into a prompt string.
175#[derive(Debug, Clone, PartialEq, Eq)]
176pub enum PromptFormat {
177    /// OpenAI / Anthropic style:
178    /// ```text
179    /// <|system|>…<|user|>…<|assistant|>…
180    /// ```
181    ChatMl,
182
183    /// Plain markdown with role headers:
184    /// ```text
185    /// ## User\n…\n## Assistant\n…
186    /// ```
187    Markdown,
188
189    /// Each turn on one line as `ROLE: content`.
190    Inline,
191}
192
193// ---------------------------------------------------------------------------
194// Manager
195// ---------------------------------------------------------------------------
196
197/// Shared, async conversation history manager.
198///
199/// Thread-safe: clone the `Arc` to share across tasks. All operations are
200/// O(1) amortised apart from [`build_prompt`](Self::build_prompt) which
201/// is O(turns).
202#[derive(Clone)]
203pub struct ConversationManager {
204    store: Arc<RwLock<HashMap<String, Conversation>>>,
205    config: Arc<ConversationConfig>,
206}
207
208impl ConversationManager {
209    /// Create a new manager with the given configuration.
210    pub fn new(config: ConversationConfig) -> Self {
211        ConversationManager {
212            store: Arc::new(RwLock::new(HashMap::new())),
213            config: Arc::new(config),
214        }
215    }
216
217    // -----------------------------------------------------------------------
218    // Mutation helpers
219    // -----------------------------------------------------------------------
220
221    /// Append a user turn to the conversation for `session`.
222    pub async fn push_user(&self, session: &SessionId, content: impl Into<String>) {
223        self.push(session, Role::User, content).await;
224    }
225
226    /// Append an assistant turn to the conversation for `session`.
227    pub async fn push_assistant(&self, session: &SessionId, content: impl Into<String>) {
228        self.push(session, Role::Assistant, content).await;
229    }
230
231    /// Append a tool-result turn.
232    pub async fn push_tool(&self, session: &SessionId, content: impl Into<String>) {
233        self.push(session, Role::Tool, content).await;
234    }
235
236    /// Append a system turn (overrides the global system prompt for this turn).
237    pub async fn push_system(&self, session: &SessionId, content: impl Into<String>) {
238        self.push(session, Role::System, content).await;
239    }
240
241    async fn push(&self, session: &SessionId, role: Role, content: impl Into<String>) {
242        let key = session.as_str().to_owned();
243        let turn = Turn::new(role, content);
244        let max_tokens = self.config.max_tokens;
245        let recency_keep = self.config.recency_keep;
246
247        let mut guard = self.store.write().await;
248        let conv = guard.entry(key).or_insert_with(|| Conversation::new(session.as_str()));
249        conv.push(turn);
250
251        if conv.total_tokens > max_tokens {
252            compress_conversation(conv, recency_keep);
253        }
254    }
255
256    // -----------------------------------------------------------------------
257    // Prompt assembly
258    // -----------------------------------------------------------------------
259
260    /// Build a full prompt string that includes prior conversation history
261    /// followed by `new_input` as the next user message.
262    ///
263    /// The returned string is ready to pass directly to a [`crate::ModelWorker`].
264    pub async fn build_prompt(&self, session: &SessionId, new_input: &str) -> String {
265        let key = session.as_str();
266        let guard = self.store.read().await;
267
268        let history: Vec<(Role, String)> = if let Some(conv) = guard.get(key) {
269            conv.turns
270                .iter()
271                .map(|t| (t.role.clone(), t.content.clone()))
272                .collect()
273        } else {
274            Vec::new()
275        };
276
277        drop(guard);
278        self.format_prompt(&history, new_input)
279    }
280
281    fn format_prompt(&self, history: &[(Role, String)], new_input: &str) -> String {
282        match self.config.format {
283            PromptFormat::ChatMl => {
284                let mut out = String::with_capacity(1024);
285                if let Some(sys) = &self.config.system_prompt {
286                    out.push_str("<|system|>\n");
287                    out.push_str(sys);
288                    out.push_str("\n<|end|>\n");
289                }
290                for (role, content) in history {
291                    match role {
292                        Role::System => {
293                            out.push_str("<|system|>\n");
294                            out.push_str(content);
295                            out.push_str("\n<|end|>\n");
296                        }
297                        Role::User => {
298                            out.push_str("<|user|>\n");
299                            out.push_str(content);
300                            out.push_str("\n<|end|>\n");
301                        }
302                        Role::Assistant => {
303                            out.push_str("<|assistant|>\n");
304                            out.push_str(content);
305                            out.push_str("\n<|end|>\n");
306                        }
307                        Role::Tool => {
308                            out.push_str("<|tool|>\n");
309                            out.push_str(content);
310                            out.push_str("\n<|end|>\n");
311                        }
312                    }
313                }
314                out.push_str("<|user|>\n");
315                out.push_str(new_input);
316                out.push_str("\n<|end|>\n<|assistant|>\n");
317                out
318            }
319
320            PromptFormat::Markdown => {
321                let mut out = String::with_capacity(1024);
322                if let Some(sys) = &self.config.system_prompt {
323                    out.push_str("## System\n\n");
324                    out.push_str(sys);
325                    out.push_str("\n\n---\n\n");
326                }
327                for (role, content) in history {
328                    let header = match role {
329                        Role::System => "## System",
330                        Role::User => "## User",
331                        Role::Assistant => "## Assistant",
332                        Role::Tool => "## Tool",
333                    };
334                    out.push_str(header);
335                    out.push_str("\n\n");
336                    out.push_str(content);
337                    out.push_str("\n\n");
338                }
339                out.push_str("## User\n\n");
340                out.push_str(new_input);
341                out.push_str("\n\n## Assistant\n\n");
342                out
343            }
344
345            PromptFormat::Inline => {
346                let mut out = String::with_capacity(1024);
347                if let Some(sys) = &self.config.system_prompt {
348                    out.push_str("system: ");
349                    out.push_str(sys);
350                    out.push('\n');
351                }
352                for (role, content) in history {
353                    out.push_str(&role.to_string());
354                    out.push_str(": ");
355                    out.push_str(content);
356                    out.push('\n');
357                }
358                out.push_str("user: ");
359                out.push_str(new_input);
360                out.push('\n');
361                out
362            }
363        }
364    }
365
366    // -----------------------------------------------------------------------
367    // Introspection
368    // -----------------------------------------------------------------------
369
370    /// Return a snapshot of the conversation for `session`, or `None` if no
371    /// history exists.
372    pub async fn get(&self, session: &SessionId) -> Option<Conversation> {
373        let guard = self.store.read().await;
374        guard.get(session.as_str()).cloned()
375    }
376
377    /// Return the number of turns currently stored for `session`.
378    pub async fn turn_count(&self, session: &SessionId) -> usize {
379        let guard = self.store.read().await;
380        guard
381            .get(session.as_str())
382            .map(|c| c.turns.len())
383            .unwrap_or(0)
384    }
385
386    /// Return the estimated token count for `session`.
387    pub async fn token_count(&self, session: &SessionId) -> usize {
388        let guard = self.store.read().await;
389        guard
390            .get(session.as_str())
391            .map(|c| c.total_tokens)
392            .unwrap_or(0)
393    }
394
395    /// Attach arbitrary metadata to a session's conversation.
396    pub async fn set_meta(&self, session: &SessionId, key: impl Into<String>, value: impl Into<String>) {
397        let mut guard = self.store.write().await;
398        let conv = guard
399            .entry(session.as_str().to_owned())
400            .or_insert_with(|| Conversation::new(session.as_str()));
401        conv.meta.insert(key.into(), value.into());
402    }
403
404    /// Clear the conversation history for `session`.
405    pub async fn clear(&self, session: &SessionId) {
406        let mut guard = self.store.write().await;
407        guard.remove(session.as_str());
408    }
409
410    /// Evict all conversations that have been inactive longer than the
411    /// configured TTL. Call this periodically (e.g. every hour) to reclaim
412    /// memory.
413    pub async fn evict_stale(&self) {
414        let ttl = self.config.ttl;
415        let mut guard = self.store.write().await;
416        guard.retain(|_, conv| {
417            conv.last_active
418                .elapsed()
419                .map(|d| d < ttl)
420                .unwrap_or(true)
421        });
422    }
423
424    /// Export a conversation as a JSON string for persistence or debugging.
425    pub async fn export_json(&self, session: &SessionId) -> Option<String> {
426        let guard = self.store.read().await;
427        guard
428            .get(session.as_str())
429            .and_then(|c| serde_json::to_string_pretty(c).ok())
430    }
431
432    /// Return all active session IDs.
433    pub async fn active_sessions(&self) -> Vec<String> {
434        let guard = self.store.read().await;
435        guard.keys().cloned().collect()
436    }
437
438    /// Total number of active conversations.
439    pub async fn len(&self) -> usize {
440        self.store.read().await.len()
441    }
442
443    /// True if there are no active conversations.
444    pub async fn is_empty(&self) -> bool {
445        self.store.read().await.is_empty()
446    }
447}
448
449// ---------------------------------------------------------------------------
450// Compression
451// ---------------------------------------------------------------------------
452
453/// In-place compression: replaces old turns with a single summary turn,
454/// preserving the most recent `recency_keep` turns verbatim.
455fn compress_conversation(conv: &mut Conversation, recency_keep: usize) {
456    let len = conv.turns.len();
457    if len <= recency_keep {
458        return;
459    }
460
461    let split = len.saturating_sub(recency_keep);
462    let old_turns: Vec<Turn> = conv.turns.drain(..split).collect();
463
464    // Build a compact summary of the compressed turns
465    let mut summary_parts = Vec::new();
466    let mut user_count = 0usize;
467    let mut assistant_count = 0usize;
468
469    for t in &old_turns {
470        match t.role {
471            Role::User => {
472                user_count += 1;
473                // Include significant user messages (> 20 tokens)
474                if t.token_estimate > 20 && summary_parts.len() < 5 {
475                    let excerpt = truncate_str(&t.content, 120);
476                    summary_parts.push(format!("User said: {excerpt}"));
477                }
478            }
479            Role::Assistant => {
480                assistant_count += 1;
481                if t.token_estimate > 30 && summary_parts.len() < 8 {
482                    let excerpt = truncate_str(&t.content, 150);
483                    summary_parts.push(format!("Assistant replied: {excerpt}"));
484                }
485            }
486            Role::Tool => {
487                if summary_parts.len() < 8 {
488                    let excerpt = truncate_str(&t.content, 80);
489                    summary_parts.push(format!("Tool returned: {excerpt}"));
490                }
491            }
492            Role::System => {}
493        }
494    }
495
496    let header = format!(
497        "[Conversation summary — {user_count} user turns, {assistant_count} assistant turns compressed]"
498    );
499
500    let summary_text = if summary_parts.is_empty() {
501        header
502    } else {
503        format!("{header}\n{}", summary_parts.join("\n"))
504    };
505
506    let summary_turn = Turn {
507        id: new_id(),
508        role: Role::System,
509        content: summary_text,
510        timestamp: SystemTime::now(),
511        token_estimate: estimate_tokens(&{
512            
513            "[Conversation summary]".to_owned()
514        }),
515        tags: vec!["summary".to_owned()],
516        is_summary: true,
517    };
518
519    // Recalculate total tokens
520    conv.total_tokens = summary_turn.token_estimate
521        + conv.turns.iter().map(|t| t.token_estimate).sum::<usize>();
522
523    conv.turns.insert(0, summary_turn);
524    conv.compressions += 1;
525}
526
527// ---------------------------------------------------------------------------
528// Utilities
529// ---------------------------------------------------------------------------
530
531/// Estimate token count for a string (4 chars ≈ 1 token is a well-known heuristic).
532pub fn estimate_tokens(s: &str) -> usize {
533    (s.len() / 4).max(1)
534}
535
536fn truncate_str(s: &str, max_chars: usize) -> String {
537    if s.len() <= max_chars {
538        s.to_owned()
539    } else {
540        let mut end = max_chars;
541        while !s.is_char_boundary(end) {
542            end -= 1;
543        }
544        format!("{}…", &s[..end])
545    }
546}
547
548fn new_id() -> String {
549    use std::time::{SystemTime, UNIX_EPOCH};
550    let t = SystemTime::now()
551        .duration_since(UNIX_EPOCH)
552        .unwrap_or_default()
553        .subsec_nanos();
554    format!("{:08x}", t ^ (t.wrapping_mul(0x9e37_79b9)))
555}
556
557// ---------------------------------------------------------------------------
558// Tests
559// ---------------------------------------------------------------------------
560
561#[cfg(test)]
562mod tests {
563    use super::*;
564
565    fn sid(s: &str) -> SessionId {
566        SessionId::new(s)
567    }
568
569    #[tokio::test]
570    async fn test_push_and_count() {
571        let mgr = ConversationManager::new(ConversationConfig::default());
572        let s = sid("s1");
573        mgr.push_user(&s, "Hello").await;
574        mgr.push_assistant(&s, "Hi there").await;
575        assert_eq!(mgr.turn_count(&s).await, 2);
576    }
577
578    #[tokio::test]
579    async fn test_build_prompt_chatml() {
580        let mgr = ConversationManager::new(ConversationConfig {
581            format: PromptFormat::ChatMl,
582            system_prompt: Some("You are helpful.".into()),
583            ..Default::default()
584        });
585        let s = sid("s2");
586        mgr.push_user(&s, "Q1").await;
587        mgr.push_assistant(&s, "A1").await;
588        let prompt = mgr.build_prompt(&s, "Q2").await;
589        assert!(prompt.contains("<|system|>"));
590        assert!(prompt.contains("Q1"));
591        assert!(prompt.contains("A1"));
592        assert!(prompt.contains("Q2"));
593        assert!(prompt.ends_with("<|assistant|>\n"));
594    }
595
596    #[tokio::test]
597    async fn test_build_prompt_markdown() {
598        let mgr = ConversationManager::new(ConversationConfig {
599            format: PromptFormat::Markdown,
600            ..Default::default()
601        });
602        let s = sid("s3");
603        mgr.push_user(&s, "What?").await;
604        let prompt = mgr.build_prompt(&s, "Why?").await;
605        assert!(prompt.contains("## User"));
606        assert!(prompt.contains("What?"));
607        assert!(prompt.contains("## Assistant"));
608    }
609
610    #[tokio::test]
611    async fn test_compression_triggered() {
612        let mgr = ConversationManager::new(ConversationConfig {
613            max_tokens: 50,
614            recency_keep: 2,
615            ..Default::default()
616        });
617        let s = sid("s4");
618        // Push enough content to exceed the 50-token budget
619        for i in 0..10 {
620            mgr.push_user(&s, format!("This is user message number {i} which has some content in it."))
621                .await;
622            mgr.push_assistant(&s, format!("This is the assistant reply to message {i}."))
623                .await;
624        }
625        let conv = mgr.get(&s).await.unwrap();
626        assert!(conv.compressions > 0, "should have triggered compression");
627    }
628
629    #[tokio::test]
630    async fn test_clear() {
631        let mgr = ConversationManager::new(ConversationConfig::default());
632        let s = sid("s5");
633        mgr.push_user(&s, "hello").await;
634        mgr.clear(&s).await;
635        assert_eq!(mgr.turn_count(&s).await, 0);
636    }
637
638    #[tokio::test]
639    async fn test_export_json() {
640        let mgr = ConversationManager::new(ConversationConfig::default());
641        let s = sid("s6");
642        mgr.push_user(&s, "test").await;
643        let json = mgr.export_json(&s).await.unwrap();
644        assert!(json.contains("\"content\""));
645        assert!(json.contains("test"));
646    }
647
648    #[tokio::test]
649    async fn test_evict_stale_leaves_active() {
650        let mgr = ConversationManager::new(ConversationConfig {
651            ttl: Duration::from_secs(3600),
652            ..Default::default()
653        });
654        let s = sid("s7");
655        mgr.push_user(&s, "hello").await;
656        mgr.evict_stale().await;
657        assert_eq!(mgr.len().await, 1);
658    }
659
660    #[test]
661    fn test_estimate_tokens() {
662        assert_eq!(estimate_tokens("hello"), 1);
663        assert_eq!(estimate_tokens("hello world"), 2);
664        assert!(estimate_tokens("a".repeat(100).as_str()) == 25);
665    }
666}