use std::collections::HashMap;
use std::sync::RwLock;
use std::time::{Duration, Instant};
use async_trait::async_trait;
use serde_json::Value;
use turbomcp_core::{Implementation, LogLevel, ProtocolVersion};
#[derive(Clone, Debug)]
#[non_exhaustive]
pub struct SessionState {
pub version: ProtocolVersion,
pub client_info: Implementation,
pub client_capabilities: Value,
pub log_level: Option<turbomcp_core::LogLevel>,
}
struct Entry {
state: SessionState,
last_seen: Instant,
}
pub struct SessionStore {
inner: RwLock<HashMap<String, Entry>>,
capacity: usize,
idle_timeout: Option<Duration>,
}
impl SessionStore {
pub const DEFAULT_CAPACITY: usize = 4096;
#[must_use]
pub fn with_capacity(capacity: usize) -> Self {
Self {
inner: RwLock::new(HashMap::new()),
capacity: capacity.max(1),
idle_timeout: None,
}
}
#[must_use]
pub fn with_idle_timeout(mut self, timeout: Option<Duration>) -> Self {
self.idle_timeout = timeout;
self
}
fn is_expired(&self, entry: &Entry, now: Instant) -> bool {
self.idle_timeout
.is_some_and(|t| now.duration_since(entry.last_seen) >= t)
}
pub fn insert(&self, id: impl Into<String>, state: SessionState) {
let id = id.into();
let mut map = self.inner.write().expect("session store lock poisoned");
let now = Instant::now();
if !map.contains_key(&id) && map.len() >= self.capacity {
if let Some(oldest) = map
.iter()
.min_by_key(|(_, e)| e.last_seen)
.map(|(k, _)| k.clone())
{
map.remove(&oldest);
}
}
map.insert(
id,
Entry {
state,
last_seen: now,
},
);
}
#[must_use]
pub fn get(&self, id: &str) -> Option<SessionState> {
let mut map = self.inner.write().expect("session store lock poisoned");
let now = Instant::now();
let entry = map.get_mut(id)?;
if self.is_expired(entry, now) {
map.remove(id);
return None;
}
entry.last_seen = now;
Some(entry.state.clone())
}
#[must_use]
pub fn sweep_expired(&self) -> Vec<String> {
if self.idle_timeout.is_none() {
return Vec::new();
}
let now = Instant::now();
let mut map = self.inner.write().expect("session store lock poisoned");
let expired: Vec<String> = map
.iter()
.filter(|(_, e)| self.is_expired(e, now))
.map(|(k, _)| k.clone())
.collect();
for id in &expired {
map.remove(id);
}
expired
}
#[must_use]
pub fn contains(&self, id: &str) -> bool {
self.inner
.read()
.expect("session store lock poisoned")
.contains_key(id)
}
pub fn set_log_level(&self, id: &str, level: turbomcp_core::LogLevel) -> bool {
let mut map = self.inner.write().expect("session store lock poisoned");
match map.get_mut(id) {
Some(entry) => {
entry.state.log_level = Some(level);
entry.last_seen = Instant::now();
true
}
None => false,
}
}
pub fn remove(&self, id: &str) -> bool {
self.inner
.write()
.expect("session store lock poisoned")
.remove(id)
.is_some()
}
#[must_use]
pub fn len(&self) -> usize {
self.inner
.read()
.expect("session store lock poisoned")
.len()
}
#[must_use]
pub fn is_empty(&self) -> bool {
self.len() == 0
}
}
impl Default for SessionStore {
fn default() -> Self {
Self::with_capacity(Self::DEFAULT_CAPACITY)
}
}
#[async_trait]
pub trait SessionBackend: Send + Sync {
async fn insert(&self, id: &str, state: SessionState);
async fn get(&self, id: &str) -> Option<SessionState>;
async fn set_log_level(&self, id: &str, level: LogLevel) -> bool;
async fn remove(&self, id: &str) -> bool;
async fn sweep_expired(&self) -> Vec<String>;
}
#[async_trait]
impl SessionBackend for SessionStore {
async fn insert(&self, id: &str, state: SessionState) {
SessionStore::insert(self, id, state);
}
async fn get(&self, id: &str) -> Option<SessionState> {
SessionStore::get(self, id)
}
async fn set_log_level(&self, id: &str, level: LogLevel) -> bool {
SessionStore::set_log_level(self, id, level)
}
async fn remove(&self, id: &str) -> bool {
SessionStore::remove(self, id)
}
async fn sweep_expired(&self) -> Vec<String> {
SessionStore::sweep_expired(self)
}
}
#[cfg(test)]
mod tests {
use super::*;
fn state() -> SessionState {
SessionState {
version: ProtocolVersion::V2025_11_25,
client_info: Implementation::new("test-client", "1.0"),
client_capabilities: serde_json::json!({}),
log_level: None,
}
}
#[test]
fn insert_get_remove_round_trip() {
let store = SessionStore::default();
store.insert("a", state());
assert!(store.contains("a"));
assert_eq!(store.get("a").unwrap().client_info.name, "test-client");
assert!(store.remove("a"));
assert!(store.get("a").is_none());
assert!(!store.remove("a"));
}
#[test]
fn capacity_evicts_least_recently_seen() {
let store = SessionStore::with_capacity(2);
store.insert("a", state());
store.insert("b", state());
let _ = store.get("a"); store.insert("c", state());
assert!(store.contains("a"));
assert!(!store.contains("b"));
assert!(store.contains("c"));
}
#[test]
fn idle_timeout_expires_on_get() {
let store =
SessionStore::with_capacity(8).with_idle_timeout(Some(Duration::from_millis(10)));
store.insert("a", state());
assert!(store.get("a").is_some()); std::thread::sleep(Duration::from_millis(25));
assert!(store.get("a").is_none(), "idle past the timeout → gone");
assert!(!store.contains("a"), "the expired entry was dropped");
}
#[test]
fn sweep_expired_reclaims_idle_sessions() {
let store =
SessionStore::with_capacity(8).with_idle_timeout(Some(Duration::from_millis(10)));
store.insert("a", state());
store.insert("b", state());
std::thread::sleep(Duration::from_millis(25));
store.insert("c", state()); let mut swept = store.sweep_expired();
swept.sort();
assert_eq!(swept, vec!["a".to_owned(), "b".to_owned()]);
assert!(store.contains("c"));
assert!(
store.sweep_expired().is_empty(),
"second sweep finds nothing"
);
}
#[test]
fn no_idle_timeout_never_sweeps() {
let store = SessionStore::with_capacity(8); store.insert("a", state());
std::thread::sleep(Duration::from_millis(5));
assert!(store.sweep_expired().is_empty());
assert!(store.contains("a"));
}
}