tokio_prompt_orchestrator/
session_mgr.rs1use 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#[derive(Debug, Error, Clone, PartialEq, Eq)]
41pub enum SessionError {
42 #[error("session {0} not found")]
44 NotFound(u64),
45}
46
47#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
53pub enum Role {
54 System,
56 User,
58 Assistant,
60}
61
62#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
64pub struct Message {
65 pub role: Role,
67 pub content: String,
69 pub timestamp: DateTime<Utc>,
71 pub tokens: Option<u32>,
73}
74
75impl Message {
76 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#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
86pub struct Session {
87 pub id: u64,
89 pub created_at: DateTime<Utc>,
91 pub messages: Vec<Message>,
93 pub metadata: HashMap<String, String>,
95}
96
97impl Session {
98 pub fn total_tokens(&self) -> usize {
100 self.messages.iter().map(|m| m.effective_tokens()).sum()
101 }
102}
103
104#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
110pub struct SessionStats {
111 pub total_sessions: usize,
113 pub active_sessions: usize,
115 pub avg_messages_per_session: f64,
117 pub total_messages: usize,
119}
120
121#[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 pub fn new() -> Self {
146 Self::default()
147 }
148
149 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 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 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 pub fn delete(&self, session_id: u64) -> bool {
207 self.sessions.remove(&session_id).is_some()
208 }
209
210 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 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 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 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 let (system_msgs, non_system): (Vec<Message>, Vec<Message>) = entry
287 .messages
288 .drain(..)
289 .partition(|m| m.role == Role::System);
290
291 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 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#[cfg(test)]
334mod tests {
335 use super::*;
336
337 fn mgr() -> SessionManager {
338 SessionManager::new()
339 }
340
341 #[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 #[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 #[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 #[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 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 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 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 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 let ctx = m.get_context(id, 0).unwrap();
481 assert!(ctx.iter().all(|msg| msg.role == Role::System));
482 }
483
484 #[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 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 #[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 #[test]
554 fn test_clone_shares_state() {
555 let m1 = mgr();
556 let m2 = m1.clone();
557 let id = m1.create(None);
558 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}