Skip to main content

tokio_prompt_orchestrator/
session_mgr.rs

1//! # Conversation Session Manager
2//!
3//! Manages multi-turn conversation sessions with full CRUD, context windowing,
4//! and automatic summarisation when token budgets are exceeded.
5//!
6//! ## Overview
7//!
8//! [`SessionManager`] stores sessions in a lock-free [`DashMap`], assigning
9//! each a unique `u64` ID.  Each [`Session`] holds an ordered list of
10//! [`Message`]s with per-message token counts.  Context retrieval trims the
11//! oldest non-system messages to stay within a caller-supplied token budget.
12//!
13//! ## Example
14//!
15//! ```rust
16//! use tokio_prompt_orchestrator::session_mgr::{SessionManager, Role};
17//!
18//! # #[tokio::main]
19//! # async fn main() {
20//! let mgr = SessionManager::default();
21//! let id = mgr.create(Some("You are a helpful assistant.".into()));
22//! mgr.append(id, Role::User, "Hello!".into()).unwrap();
23//! let ctx = mgr.get_context(id, 4096).unwrap();
24//! assert_eq!(ctx.len(), 2); // system + user
25//! # }
26//! ```
27
28use chrono::{DateTime, Utc};
29use dashmap::DashMap;
30use std::collections::HashMap;
31use std::sync::atomic::{AtomicU64, Ordering};
32use std::sync::Arc;
33use thiserror::Error;
34
35// ============================================================================
36// Error type
37// ============================================================================
38
39/// Errors produced by [`SessionManager`] operations.
40#[derive(Debug, Error, Clone, PartialEq, Eq)]
41pub enum SessionError {
42    /// The requested session ID does not exist.
43    #[error("session {0} not found")]
44    NotFound(u64),
45}
46
47// ============================================================================
48// Domain types
49// ============================================================================
50
51/// The speaker role of a conversation message.
52#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
53pub enum Role {
54    /// A system-level instruction prepended to the conversation.
55    System,
56    /// A message sent by the human user.
57    User,
58    /// A reply generated by the AI assistant.
59    Assistant,
60}
61
62/// A single message within a [`Session`].
63#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
64pub struct Message {
65    /// Who produced this message.
66    pub role: Role,
67    /// Raw text content.
68    pub content: String,
69    /// Wall-clock time when the message was appended.
70    pub timestamp: DateTime<Utc>,
71    /// Estimated token count, if known.
72    pub tokens: Option<u32>,
73}
74
75impl Message {
76    /// Best-effort token count: use provided value or fall back to word count.
77    fn effective_tokens(&self) -> usize {
78        self.tokens
79            .map(|t| t as usize)
80            .unwrap_or_else(|| self.content.split_whitespace().count())
81    }
82}
83
84/// A conversation session holding an ordered sequence of [`Message`]s.
85#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
86pub struct Session {
87    /// Unique numeric identifier.
88    pub id: u64,
89    /// When this session was first created.
90    pub created_at: DateTime<Utc>,
91    /// Ordered message history (system message, if present, is always first).
92    pub messages: Vec<Message>,
93    /// Arbitrary key-value metadata attached by the caller.
94    pub metadata: HashMap<String, String>,
95}
96
97impl Session {
98    /// Total estimated token count across all messages.
99    pub fn total_tokens(&self) -> usize {
100        self.messages.iter().map(|m| m.effective_tokens()).sum()
101    }
102}
103
104// ============================================================================
105// Stats
106// ============================================================================
107
108/// Aggregate statistics across all sessions in a [`SessionManager`].
109#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
110pub struct SessionStats {
111    /// Number of sessions currently stored.
112    pub total_sessions: usize,
113    /// Alias for `total_sessions` (all stored sessions are considered active).
114    pub active_sessions: usize,
115    /// Mean message count across all sessions.
116    pub avg_messages_per_session: f64,
117    /// Sum of message counts across all sessions.
118    pub total_messages: usize,
119}
120
121// ============================================================================
122// SessionManager
123// ============================================================================
124
125/// Concurrent, lock-free store of [`Session`]s keyed by `u64` ID.
126///
127/// All methods are `&self` and safe to call from multiple Tokio tasks.
128#[derive(Debug, Clone)]
129pub struct SessionManager {
130    sessions: Arc<DashMap<u64, Session>>,
131    next_id: Arc<AtomicU64>,
132}
133
134impl Default for SessionManager {
135    fn default() -> Self {
136        Self {
137            sessions: Arc::new(DashMap::new()),
138            next_id: Arc::new(AtomicU64::new(1)),
139        }
140    }
141}
142
143impl SessionManager {
144    /// Create a new manager with an empty store.
145    pub fn new() -> Self {
146        Self::default()
147    }
148
149    /// Create a new session, optionally pre-loading a system message.
150    ///
151    /// Returns the new session's ID.
152    pub fn create(&self, system_prompt: Option<String>) -> u64 {
153        let id = self.next_id.fetch_add(1, Ordering::Relaxed);
154        let mut messages = Vec::new();
155        if let Some(prompt) = system_prompt {
156            let tokens = prompt.split_whitespace().count() as u32;
157            messages.push(Message {
158                role: Role::System,
159                content: prompt,
160                timestamp: Utc::now(),
161                tokens: Some(tokens),
162            });
163        }
164        let session = Session {
165            id,
166            created_at: Utc::now(),
167            messages,
168            metadata: HashMap::new(),
169        };
170        self.sessions.insert(id, session);
171        id
172    }
173
174    /// Append a message to an existing session.
175    ///
176    /// Token count is estimated as word count when not supplied.
177    pub fn append(
178        &self,
179        session_id: u64,
180        role: Role,
181        content: String,
182    ) -> Result<(), SessionError> {
183        let mut entry = self
184            .sessions
185            .get_mut(&session_id)
186            .ok_or(SessionError::NotFound(session_id))?;
187        let tokens = content.split_whitespace().count() as u32;
188        entry.messages.push(Message {
189            role,
190            content,
191            timestamp: Utc::now(),
192            tokens: Some(tokens),
193        });
194        Ok(())
195    }
196
197    /// Return a copy of the session, or `Err(SessionError::NotFound)`.
198    pub fn get(&self, session_id: u64) -> Result<Session, SessionError> {
199        self.sessions
200            .get(&session_id)
201            .map(|r| r.clone())
202            .ok_or(SessionError::NotFound(session_id))
203    }
204
205    /// Delete a session, returning `true` if it existed.
206    pub fn delete(&self, session_id: u64) -> bool {
207        self.sessions.remove(&session_id).is_some()
208    }
209
210    /// Return messages for a session trimmed to fit within `max_tokens`.
211    ///
212    /// Strategy:
213    /// 1. The system message (if any) is always retained.
214    /// 2. Non-system messages are trimmed from the **oldest** end until the
215    ///    total fits within `max_tokens`.
216    pub fn get_context(
217        &self,
218        session_id: u64,
219        max_tokens: usize,
220    ) -> Result<Vec<Message>, SessionError> {
221        let session = self
222            .sessions
223            .get(&session_id)
224            .ok_or(SessionError::NotFound(session_id))?;
225
226        // Partition: system vs non-system
227        let mut system: Vec<Message> = session
228            .messages
229            .iter()
230            .filter(|m| m.role == Role::System)
231            .cloned()
232            .collect();
233        let non_system: Vec<Message> = session
234            .messages
235            .iter()
236            .filter(|m| m.role != Role::System)
237            .cloned()
238            .collect();
239
240        let system_tokens: usize = system.iter().map(|m| m.effective_tokens()).sum();
241        let budget = max_tokens.saturating_sub(system_tokens);
242
243        // Walk non-system from newest to oldest, accumulating until full.
244        let mut kept: Vec<Message> = Vec::new();
245        let mut used = 0usize;
246        for msg in non_system.iter().rev() {
247            let t = msg.effective_tokens();
248            if used + t > budget {
249                break;
250            }
251            used += t;
252            kept.push(msg.clone());
253        }
254        kept.reverse();
255        system.extend(kept);
256        Ok(system)
257    }
258
259    /// If the session's total tokens exceed `threshold_tokens`, call
260    /// `summarizer` on the oldest non-system messages and replace them with a
261    /// single assistant summary message.
262    pub fn summarize_if_needed(
263        &self,
264        session_id: u64,
265        threshold_tokens: usize,
266        summarizer: &dyn Fn(&str) -> String,
267    ) -> Result<(), SessionError> {
268        let total = {
269            let s = self
270                .sessions
271                .get(&session_id)
272                .ok_or(SessionError::NotFound(session_id))?;
273            s.total_tokens()
274        };
275
276        if total <= threshold_tokens {
277            return Ok(());
278        }
279
280        let mut entry = self
281            .sessions
282            .get_mut(&session_id)
283            .ok_or(SessionError::NotFound(session_id))?;
284
285        // Separate system messages from the rest.
286        let (system_msgs, non_system): (Vec<Message>, Vec<Message>) = entry
287            .messages
288            .drain(..)
289            .partition(|m| m.role == Role::System);
290
291        // Summarise the non-system portion.
292        let combined: String = non_system
293            .iter()
294            .map(|m| format!("{:?}: {}", m.role, m.content))
295            .collect::<Vec<_>>()
296            .join("\n");
297        let summary = summarizer(&combined);
298        let summary_tokens = summary.split_whitespace().count() as u32;
299        let summary_msg = Message {
300            role: Role::Assistant,
301            content: summary,
302            timestamp: Utc::now(),
303            tokens: Some(summary_tokens),
304        };
305
306        entry.messages = system_msgs;
307        entry.messages.push(summary_msg);
308        Ok(())
309    }
310
311    /// Collect aggregate statistics across all stored sessions.
312    pub fn stats(&self) -> SessionStats {
313        let total_sessions = self.sessions.len();
314        let total_messages: usize = self.sessions.iter().map(|r| r.messages.len()).sum();
315        let avg_messages_per_session = if total_sessions == 0 {
316            0.0
317        } else {
318            total_messages as f64 / total_sessions as f64
319        };
320        SessionStats {
321            total_sessions,
322            active_sessions: total_sessions,
323            avg_messages_per_session,
324            total_messages,
325        }
326    }
327}
328
329// ============================================================================
330// Unit tests
331// ============================================================================
332
333#[cfg(test)]
334mod tests {
335    use super::*;
336
337    fn mgr() -> SessionManager {
338        SessionManager::new()
339    }
340
341    // --- create ---
342
343    #[test]
344    fn test_create_no_system() {
345        let m = mgr();
346        let id = m.create(None);
347        let s = m.get(id).unwrap();
348        assert!(s.messages.is_empty());
349    }
350
351    #[test]
352    fn test_create_with_system() {
353        let m = mgr();
354        let id = m.create(Some("sys".into()));
355        let s = m.get(id).unwrap();
356        assert_eq!(s.messages.len(), 1);
357        assert_eq!(s.messages[0].role, Role::System);
358    }
359
360    #[test]
361    fn test_create_returns_unique_ids() {
362        let m = mgr();
363        let a = m.create(None);
364        let b = m.create(None);
365        assert_ne!(a, b);
366    }
367
368    #[test]
369    fn test_create_increments_ids() {
370        let m = mgr();
371        let a = m.create(None);
372        let b = m.create(None);
373        assert!(b > a);
374    }
375
376    // --- append ---
377
378    #[test]
379    fn test_append_ok() {
380        let m = mgr();
381        let id = m.create(None);
382        assert!(m.append(id, Role::User, "hello".into()).is_ok());
383    }
384
385    #[test]
386    fn test_append_not_found() {
387        let m = mgr();
388        let err = m.append(99, Role::User, "x".into()).unwrap_err();
389        assert_eq!(err, SessionError::NotFound(99));
390    }
391
392    #[test]
393    fn test_append_multiple_roles() {
394        let m = mgr();
395        let id = m.create(Some("sys".into()));
396        m.append(id, Role::User, "q".into()).unwrap();
397        m.append(id, Role::Assistant, "a".into()).unwrap();
398        let s = m.get(id).unwrap();
399        assert_eq!(s.messages.len(), 3);
400    }
401
402    // --- delete ---
403
404    #[test]
405    fn test_delete_existing() {
406        let m = mgr();
407        let id = m.create(None);
408        assert!(m.delete(id));
409    }
410
411    #[test]
412    fn test_delete_missing() {
413        let m = mgr();
414        assert!(!m.delete(999));
415    }
416
417    #[test]
418    fn test_delete_removes_session() {
419        let m = mgr();
420        let id = m.create(None);
421        m.delete(id);
422        assert!(m.get(id).is_err());
423    }
424
425    // --- get_context ---
426
427    #[test]
428    fn test_get_context_not_found() {
429        let m = mgr();
430        assert!(m.get_context(0, 100).is_err());
431    }
432
433    #[test]
434    fn test_get_context_all_fit() {
435        let m = mgr();
436        let id = m.create(Some("sys".into()));
437        m.append(id, Role::User, "hello world".into()).unwrap();
438        // 1 (sys) + 2 (user) = 3 tokens — well within budget
439        let ctx = m.get_context(id, 100).unwrap();
440        assert_eq!(ctx.len(), 2);
441    }
442
443    #[test]
444    fn test_get_context_trims_oldest() {
445        let m = mgr();
446        let id = m.create(None);
447        // Add 10 messages of ~10 words each → 100 tokens total
448        for i in 0..10u32 {
449            m.append(
450                id,
451                Role::User,
452                format!("word1 word2 word3 word4 word5 word6 word7 word8 word9 word{i}"),
453            )
454            .unwrap();
455        }
456        // Budget only allows the last few messages
457        let ctx = m.get_context(id, 25).unwrap();
458        assert!(ctx.len() < 10, "should have trimmed some old messages");
459    }
460
461    #[test]
462    fn test_get_context_always_keeps_system() {
463        let m = mgr();
464        let id = m.create(Some("system prompt here".into()));
465        for _ in 0..20 {
466            m.append(id, Role::User, "a b c d e f g h i j".into())
467                .unwrap();
468        }
469        let ctx = m.get_context(id, 5).unwrap();
470        // System message should be present even if nothing else fits
471        assert!(ctx.iter().any(|msg| msg.role == Role::System));
472    }
473
474    #[test]
475    fn test_get_context_zero_budget() {
476        let m = mgr();
477        let id = m.create(Some("sys".into()));
478        m.append(id, Role::User, "hi".into()).unwrap();
479        // Zero budget — only system can survive
480        let ctx = m.get_context(id, 0).unwrap();
481        assert!(ctx.iter().all(|msg| msg.role == Role::System));
482    }
483
484    // --- summarize_if_needed ---
485
486    #[test]
487    fn test_summarize_no_op_below_threshold() {
488        let m = mgr();
489        let id = m.create(Some("sys".into()));
490        m.append(id, Role::User, "hi".into()).unwrap();
491        m.summarize_if_needed(id, 10_000, &|text| format!("SUMMARY: {}", &text[..10.min(text.len())]))
492            .unwrap();
493        // Unchanged
494        assert_eq!(m.get(id).unwrap().messages.len(), 2);
495    }
496
497    #[test]
498    fn test_summarize_triggers_above_threshold() {
499        let m = mgr();
500        let id = m.create(Some("system".into()));
501        for _ in 0..50 {
502            m.append(id, Role::User, "a b c d e f g h".into()).unwrap();
503        }
504        let before = m.get(id).unwrap().messages.len();
505        m.summarize_if_needed(id, 10, &|_| "Summary text".into())
506            .unwrap();
507        let after = m.get(id).unwrap().messages.len();
508        assert!(after < before, "should have condensed messages");
509    }
510
511    #[test]
512    fn test_summarize_not_found() {
513        let m = mgr();
514        let result = m.summarize_if_needed(999, 100, &|_| "s".into());
515        assert!(result.is_err());
516    }
517
518    // --- stats ---
519
520    #[test]
521    fn test_stats_empty() {
522        let m = mgr();
523        let s = m.stats();
524        assert_eq!(s.total_sessions, 0);
525        assert_eq!(s.total_messages, 0);
526        assert_eq!(s.avg_messages_per_session, 0.0);
527    }
528
529    #[test]
530    fn test_stats_counts() {
531        let m = mgr();
532        let id1 = m.create(None);
533        let id2 = m.create(None);
534        m.append(id1, Role::User, "a".into()).unwrap();
535        m.append(id1, Role::User, "b".into()).unwrap();
536        m.append(id2, Role::User, "c".into()).unwrap();
537        let s = m.stats();
538        assert_eq!(s.total_sessions, 2);
539        assert_eq!(s.total_messages, 3);
540        assert!((s.avg_messages_per_session - 1.5).abs() < f64::EPSILON);
541    }
542
543    #[test]
544    fn test_stats_active_equals_total() {
545        let m = mgr();
546        m.create(None);
547        let s = m.stats();
548        assert_eq!(s.active_sessions, s.total_sessions);
549    }
550
551    // --- clone / Arc sharing ---
552
553    #[test]
554    fn test_clone_shares_state() {
555        let m1 = mgr();
556        let m2 = m1.clone();
557        let id = m1.create(None);
558        // m2 sees the session created by m1
559        assert!(m2.get(id).is_ok());
560    }
561
562    #[test]
563    fn test_concurrent_creates_unique() {
564        let m = mgr();
565        let ids: Vec<u64> = (0..100).map(|_| m.create(None)).collect();
566        let mut sorted = ids.clone();
567        sorted.sort_unstable();
568        sorted.dedup();
569        assert_eq!(sorted.len(), 100);
570    }
571}