use std::sync::{Arc, Mutex};
use std::time::Duration;
use chrono::Utc;
use dashmap::DashMap;
use uuid::Uuid;
use super::types::Session;
type ExpireCallback = Arc<dyn Fn(Uuid) + Send + Sync>;
pub const DEFAULT_MAX_SESSIONS: usize = 10_000;
#[derive(Clone)]
pub struct SessionManager {
sessions: Arc<DashMap<Uuid, Session>>,
ttl: Duration,
max_turns: usize,
max_sessions: usize,
on_expire: Arc<Mutex<Vec<ExpireCallback>>>,
}
impl SessionManager {
pub fn new(ttl: Duration, max_turns: usize) -> Self {
Self::with_cap(ttl, max_turns, DEFAULT_MAX_SESSIONS)
}
pub fn with_cap(ttl: Duration, max_turns: usize, max_sessions: usize) -> Self {
let sessions = Arc::new(DashMap::new());
let on_expire: Arc<Mutex<Vec<ExpireCallback>>> = Arc::new(Mutex::new(Vec::new()));
let mgr = Self {
sessions,
ttl,
max_turns,
max_sessions,
on_expire,
};
mgr.spawn_sweeper();
mgr
}
fn enforce_cap(&self) {
if self.max_sessions == 0 {
return;
}
while self.sessions.len() >= self.max_sessions {
let oldest = self
.sessions
.iter()
.min_by_key(|e| e.value().last_access)
.map(|e| *e.key());
match oldest {
Some(id) => {
if self.sessions.remove(&id).is_some() {
tracing::warn!(
session_id = %id,
cap = self.max_sessions,
"session cap reached; evicting oldest-idle"
);
self.fire_expire(id);
}
}
None => break,
}
}
}
pub fn create(&self, agent_id: impl Into<String>) -> Session {
self.enforce_cap();
let session = Session::new(agent_id, self.max_turns);
self.sessions.insert(session.id, session.clone());
session
}
pub fn get(&self, id: Uuid) -> Option<Session> {
let mut entry = self.sessions.get_mut(&id)?;
entry.last_access = Utc::now();
Some(entry.clone())
}
pub fn get_or_create(&self, id: Uuid, agent_id: impl Into<String>) -> Session {
if !self.sessions.contains_key(&id) {
self.enforce_cap();
}
let agent_id = agent_id.into();
let max_turns = self.max_turns;
let mut entry = self
.sessions
.entry(id)
.or_insert_with(|| Session::with_id(id, agent_id, max_turns));
entry.last_access = Utc::now();
entry.clone()
}
pub fn update(&self, session: Session) -> bool {
let Some(mut entry) = self.sessions.get_mut(&session.id) else {
return false;
};
*entry = session;
true
}
pub fn push_message(&self, id: Uuid, interaction: super::types::Interaction) -> bool {
let Some(mut entry) = self.sessions.get_mut(&id) else {
return false;
};
entry.push(interaction);
true
}
pub fn delete(&self, id: Uuid) -> bool {
let existed = self.sessions.remove(&id).is_some();
if existed {
self.fire_expire(id);
}
existed
}
pub fn active_count(&self) -> usize {
self.sessions.len()
}
pub fn on_expire<F>(&self, f: F)
where
F: Fn(Uuid) + Send + Sync + 'static,
{
self.on_expire.lock().unwrap().push(Arc::new(f));
}
fn fire_expire(&self, id: Uuid) {
let callbacks: Vec<ExpireCallback> = self.on_expire.lock().unwrap().clone();
spawn_callbacks(&callbacks, id);
}
fn spawn_sweeper(&self) {
let sessions = Arc::clone(&self.sessions);
let on_expire = Arc::clone(&self.on_expire);
let ttl = self.ttl;
let interval = (ttl / 4).max(Duration::from_millis(10));
tokio::spawn(async move {
let mut ticker = tokio::time::interval(interval);
ticker.tick().await; loop {
ticker.tick().await;
let now = Utc::now();
let ttl_chrono =
chrono::Duration::from_std(ttl).unwrap_or(chrono::Duration::hours(24));
let expired: Vec<Uuid> = sessions
.iter()
.filter_map(|entry| {
if now.signed_duration_since(entry.value().last_access) >= ttl_chrono {
Some(*entry.key())
} else {
None
}
})
.collect();
if expired.is_empty() {
continue;
}
for id in &expired {
sessions.remove(id);
}
let callbacks: Vec<ExpireCallback> = on_expire.lock().unwrap().clone();
if callbacks.is_empty() {
continue;
}
for id in expired {
spawn_callbacks(&callbacks, id);
}
}
});
}
}
fn spawn_callbacks(callbacks: &[ExpireCallback], id: Uuid) {
for cb in callbacks {
let cb = cb.clone();
tokio::spawn(async move {
cb(id);
});
}
}