Skip to main content

tokio_prompt_orchestrator/
session_manager.rs

1//! Session tracking and context-window management.
2//!
3//! [`SessionManager`] owns a collection of [`Session`] objects behind a
4//! `RwLock<HashMap>`.  Each session stores a rolling [`VecDeque`] of
5//! [`Message`] records and enforces a maximum context-token budget.
6
7use std::collections::{HashMap, VecDeque};
8use std::sync::{
9    atomic::{AtomicU64, Ordering},
10    RwLock,
11};
12use std::time::Instant;
13
14// ── Message ───────────────────────────────────────────────────────────────────
15
16/// A single turn in a conversation session.
17#[derive(Debug, Clone)]
18pub struct Message {
19    /// Conversation role, e.g. `"user"` or `"assistant"`.
20    pub role: String,
21    /// Raw text content of the message.
22    pub content: String,
23    /// Token count for this message.
24    pub tokens: usize,
25    /// Wall-clock time at which the message was recorded.
26    pub timestamp: Instant,
27}
28
29// ── Session ───────────────────────────────────────────────────────────────────
30
31/// A conversation session with bounded context-window management.
32#[derive(Debug)]
33pub struct Session {
34    /// Unique session identifier.
35    pub id: String,
36    /// Model associated with this session.
37    pub model: String,
38    /// Ordered message history (oldest first).
39    pub messages: VecDeque<Message>,
40    /// Running sum of tokens across all messages currently in `messages`.
41    pub total_tokens: usize,
42    /// Hard ceiling on context tokens; messages are trimmed to stay below this.
43    pub max_context_tokens: usize,
44    /// Wall-clock time at which the session was created.
45    pub created_at: Instant,
46    /// Wall-clock time of the last `add_message` call.
47    pub last_active: Instant,
48}
49
50impl Session {
51    /// Create a new empty session.
52    fn new(id: String, model: String, max_context_tokens: usize) -> Self {
53        let now = Instant::now();
54        Self {
55            id,
56            model,
57            messages: VecDeque::new(),
58            total_tokens: 0,
59            max_context_tokens,
60            created_at: now,
61            last_active: now,
62        }
63    }
64
65    /// Append a message and update the running token total.
66    ///
67    /// Trims the oldest messages first if adding `tokens` would overflow the
68    /// context window.
69    pub fn add_message(&mut self, role: String, content: String, tokens: usize) {
70        self.trim_to_fit(tokens);
71        self.total_tokens += tokens;
72        self.last_active = Instant::now();
73        self.messages.push_back(Message {
74            role,
75            content,
76            tokens,
77            timestamp: self.last_active,
78        });
79    }
80
81    /// Drop the oldest messages until `total_tokens + new_tokens` fits within
82    /// `max_context_tokens`.
83    pub fn trim_to_fit(&mut self, new_tokens: usize) {
84        while self.total_tokens + new_tokens > self.max_context_tokens {
85            if let Some(dropped) = self.messages.pop_front() {
86                self.total_tokens = self.total_tokens.saturating_sub(dropped.tokens);
87            } else {
88                break;
89            }
90        }
91    }
92
93    /// Fraction of the context window currently occupied (`0.0`–`1.0`).
94    pub fn token_utilization(&self) -> f64 {
95        if self.max_context_tokens == 0 {
96            return 0.0;
97        }
98        self.total_tokens as f64 / self.max_context_tokens as f64
99    }
100
101    /// Returns `true` if the session has not been active for longer than
102    /// `timeout_secs` seconds.
103    pub fn is_expired(&self, timeout_secs: u64) -> bool {
104        self.last_active.elapsed().as_secs() >= timeout_secs
105    }
106}
107
108// ── SessionManager ────────────────────────────────────────────────────────────
109
110/// Thread-safe store for active [`Session`] objects.
111///
112/// Session IDs are generated from a monotonically increasing counter combined
113/// with a timestamp so they are unique and roughly sortable by creation time.
114pub struct SessionManager {
115    sessions: RwLock<HashMap<String, Session>>,
116    counter: AtomicU64,
117}
118
119impl Default for SessionManager {
120    fn default() -> Self {
121        Self::new()
122    }
123}
124
125impl SessionManager {
126    /// Create an empty session manager.
127    pub fn new() -> Self {
128        Self {
129            sessions: RwLock::new(HashMap::new()),
130            counter: AtomicU64::new(0),
131        }
132    }
133
134    /// Generate a UUID-like session ID from the current timestamp + counter.
135    fn gen_id(&self) -> String {
136        let seq = self.counter.fetch_add(1, Ordering::Relaxed);
137        // Use elapsed nanos from a fixed reference + counter for uniqueness.
138        // This avoids a dependency on `uuid` while remaining collision-free
139        // within a single process.
140        let ts = std::time::SystemTime::now()
141            .duration_since(std::time::UNIX_EPOCH)
142            .unwrap_or_default()
143            .as_nanos();
144        format!("sess-{ts:x}-{seq:04x}")
145    }
146
147    /// Create a new session for `model` with the given context-token limit.
148    ///
149    /// Returns the new session ID.
150    pub fn create_session(&self, model: &str, max_context: usize) -> String {
151        let id = self.gen_id();
152        let session = Session::new(id.clone(), model.to_owned(), max_context);
153        self.sessions.write().unwrap_or_else(|e| e.into_inner()).insert(id.clone(), session);
154        id
155    }
156
157    /// Return a snapshot of the messages currently held in the session, or
158    /// `None` if the session does not exist.
159    pub fn get_context(&self, session_id: &str) -> Option<Vec<Message>> {
160        self.sessions
161            .read()
162            .unwrap_or_else(|e| e.into_inner())
163            .get(session_id)
164            .map(|s| s.messages.iter().cloned().collect())
165    }
166
167    /// Append a user turn and an assistant turn to the session.
168    ///
169    /// `user_msg_tokens` and `assistant_msg_tokens` carry placeholder content
170    /// strings; callers that need real content should use
171    /// [`Session::add_message`] directly after obtaining a write lock.
172    pub fn add_turn(
173        &self,
174        session_id: &str,
175        user_msg_tokens: usize,
176        assistant_msg_tokens: usize,
177    ) {
178        let mut guard = self.sessions.write().unwrap_or_else(|e| e.into_inner());
179        if let Some(session) = guard.get_mut(session_id) {
180            session.add_message("user".to_owned(), String::new(), user_msg_tokens);
181            session.add_message(
182                "assistant".to_owned(),
183                String::new(),
184                assistant_msg_tokens,
185            );
186        }
187    }
188
189    /// Remove sessions that have been inactive for longer than `timeout_secs`.
190    ///
191    /// Returns the number of sessions removed.
192    pub fn expire_sessions(&self, timeout_secs: u64) -> usize {
193        let mut guard = self.sessions.write().unwrap_or_else(|e| e.into_inner());
194        let before = guard.len();
195        guard.retain(|_, s| !s.is_expired(timeout_secs));
196        before - guard.len()
197    }
198
199    /// Return the number of active sessions.
200    pub fn active_count(&self) -> usize {
201        self.sessions.read().unwrap_or_else(|e| e.into_inner()).len()
202    }
203}
204
205#[cfg(test)]
206mod tests {
207    use super::*;
208
209    #[test]
210    fn create_and_retrieve_session() {
211        let mgr = SessionManager::new();
212        let id = mgr.create_session("gpt-4o", 4096);
213        assert_eq!(mgr.active_count(), 1);
214        let ctx = mgr.get_context(&id);
215        assert!(ctx.is_some());
216        assert!(ctx.unwrap().is_empty());
217    }
218
219    #[test]
220    fn add_turn_populates_context() {
221        let mgr = SessionManager::new();
222        let id = mgr.create_session("claude-3", 1000);
223        mgr.add_turn(&id, 50, 80);
224        let ctx = mgr.get_context(&id).unwrap();
225        assert_eq!(ctx.len(), 2);
226        assert_eq!(ctx[0].role, "user");
227        assert_eq!(ctx[1].role, "assistant");
228    }
229
230    #[test]
231    fn session_trim_to_fit() {
232        let mut session = Session::new("s1".into(), "m".into(), 100);
233        session.add_message("user".into(), "hello".into(), 60);
234        session.add_message("assistant".into(), "hi".into(), 60);
235        // Second add should have evicted the first message to stay within 100
236        assert!(session.total_tokens <= 100);
237    }
238
239    #[test]
240    fn expire_sessions() {
241        let mgr = SessionManager::new();
242        let _id = mgr.create_session("gpt-4o", 1000);
243        // A timeout of 0 seconds means every session is expired immediately.
244        let removed = mgr.expire_sessions(0);
245        assert_eq!(removed, 1);
246        assert_eq!(mgr.active_count(), 0);
247    }
248}