use std::collections::HashMap;
use std::sync::Arc;
use parking_lot::RwLock;
use super::offset_tracker::OffsetTracker;
use super::session::{Session, SessionId};
#[derive(Clone)]
pub struct SessionStore {
inner: Arc<RwLock<HashMap<String, Arc<Session>>>>,
offset_tracker: Arc<OffsetTracker>,
}
impl SessionStore {
pub fn new() -> Self {
Self {
inner: Arc::new(RwLock::new(HashMap::new())),
offset_tracker: Arc::new(OffsetTracker::new()),
}
}
pub fn get(&self, session_id: &str) -> Option<Arc<Session>> {
self.inner.read().get(session_id).cloned()
}
pub fn insert(&self, session: Arc<Session>) -> Option<Arc<Session>> {
let old = self.inner.write().insert(session.id.0.clone(), session);
if let Some(ref prev) = old {
self.offset_tracker.forget(&prev.id);
}
old
}
pub fn remove(&self, session_id: &str) -> bool {
let sid = SessionId(session_id.to_string());
self.offset_tracker.forget(&sid);
self.inner.write().remove(session_id).is_some()
}
pub fn offset_tracker(&self) -> &Arc<OffsetTracker> {
&self.offset_tracker
}
pub fn len(&self) -> usize {
self.inner.read().len()
}
pub fn is_empty(&self) -> bool {
self.inner.read().is_empty()
}
pub fn expire_sessions(&self, idle_timeout: std::time::Duration) -> usize {
let mut expired = Vec::new();
{
let guard = self.inner.read();
for (id, session) in guard.iter() {
let state = session.state();
match state {
crate::session::SessionState::Closed
| crate::session::SessionState::Draining => {
expired.push(id.clone());
}
_ => {
if session.idle() > idle_timeout {
expired.push(id.clone());
}
}
}
}
}
for id in &expired {
self.remove(id);
}
expired.len()
}
}
impl Default for SessionStore {
fn default() -> Self {
Self::new()
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::session::session::ClientId;
#[test]
fn insert_and_get() {
let store = SessionStore::new();
let s = Arc::new(Session::new(SessionId::new(), ClientId::new("c1")));
let id = s.id.0.clone();
store.insert(s.clone());
assert!(store.get(&id).is_some());
assert_eq!(store.len(), 1);
}
#[test]
fn remove_cleans_offsets() {
let store = SessionStore::new();
let s = Arc::new(Session::new(SessionId::new(), ClientId::new("c1")));
let id = s.id.0.clone();
let sid = s.id.clone();
store.insert(s);
store.offset_tracker().record(&sid, "t", 5);
assert_eq!(store.offset_tracker().get(&sid, "t"), Some(5));
assert!(store.remove(&id));
assert!(store.get(&id).is_none());
assert_eq!(store.offset_tracker().get(&sid, "t"), None);
}
}