tokio_prompt_orchestrator/
session_manager.rs1use std::collections::{HashMap, VecDeque};
8use std::sync::{
9 atomic::{AtomicU64, Ordering},
10 RwLock,
11};
12use std::time::Instant;
13
14#[derive(Debug, Clone)]
18pub struct Message {
19 pub role: String,
21 pub content: String,
23 pub tokens: usize,
25 pub timestamp: Instant,
27}
28
29#[derive(Debug)]
33pub struct Session {
34 pub id: String,
36 pub model: String,
38 pub messages: VecDeque<Message>,
40 pub total_tokens: usize,
42 pub max_context_tokens: usize,
44 pub created_at: Instant,
46 pub last_active: Instant,
48}
49
50impl Session {
51 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 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 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 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 pub fn is_expired(&self, timeout_secs: u64) -> bool {
104 self.last_active.elapsed().as_secs() >= timeout_secs
105 }
106}
107
108pub 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 pub fn new() -> Self {
128 Self {
129 sessions: RwLock::new(HashMap::new()),
130 counter: AtomicU64::new(0),
131 }
132 }
133
134 fn gen_id(&self) -> String {
136 let seq = self.counter.fetch_add(1, Ordering::Relaxed);
137 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 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 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 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 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 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 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 let removed = mgr.expire_sessions(0);
245 assert_eq!(removed, 1);
246 assert_eq!(mgr.active_count(), 0);
247 }
248}