use anyhow::{Result, anyhow};
use notify::{Event, EventKind, RecommendedWatcher, RecursiveMode, Watcher};
use serde::{Deserialize, Serialize};
use std::{
collections::{HashMap, hash_map::Entry},
fs::{self, OpenOptions},
path::PathBuf,
sync::{
Arc,
atomic::{AtomicBool, AtomicUsize, Ordering},
},
time::SystemTime,
};
#[derive(Debug, Clone, Serialize, Deserialize)]
struct SessionMetadata {
created_at: SystemTime,
last_used: SystemTime,
}
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
struct SessionEntry<T> {
data: T,
metadata: SessionMetadata,
}
impl<T> SessionEntry<T> {
fn update_last_used(&mut self) {
self.metadata.last_used = SystemTime::now();
}
}
impl Default for SessionMetadata {
fn default() -> Self {
let now = SystemTime::now();
Self {
created_at: now,
last_used: now,
}
}
}
#[derive(Debug)]
pub struct SessionStore<T> {
sessions: HashMap<String, SessionEntry<T>>,
storage_path: Option<PathBuf>,
needs_reload: Arc<AtomicBool>,
ignore_next_events: Arc<AtomicUsize>, _watcher: Option<RecommendedWatcher>, }
impl<T> SessionStore<T>
where
T: Serialize + for<'de> Deserialize<'de> + Clone + Default + PartialEq + Eq,
{
pub fn new(storage_path: Option<PathBuf>) -> Result<Self> {
let mut store = Self {
sessions: HashMap::new(),
storage_path: storage_path.clone(),
needs_reload: Arc::new(AtomicBool::new(false)),
ignore_next_events: Arc::new(AtomicUsize::new(0)),
_watcher: None,
};
if let Some(storage_path) = &storage_path {
if let Some(parent) = storage_path.parent() {
fs::create_dir_all(parent)?;
}
OpenOptions::new()
.append(true)
.create(true)
.open(storage_path)
.map_err(|_| anyhow!("could not open {}", storage_path.to_string_lossy()))?;
}
store.load()?;
if storage_path.is_some() {
store.setup_file_watching()?;
}
Ok(store)
}
fn setup_file_watching(&mut self) -> Result<()> {
let Some(storage_path) = &self.storage_path else {
return Ok(());
};
let needs_reload = Arc::clone(&self.needs_reload);
let ignore_next_events = Arc::clone(&self.ignore_next_events);
let watch_path = storage_path.clone();
let mut watcher = RecommendedWatcher::new(
move |res: Result<Event, notify::Error>| {
if let Ok(event) = res {
log::trace!("received {event:?}");
match event.kind {
EventKind::Modify(_) | EventKind::Create(_) => {
let current = ignore_next_events.load(Ordering::Relaxed);
if current > 0 {
let new_value = current.saturating_sub(1);
ignore_next_events.store(new_value, Ordering::Relaxed);
log::trace!(
"ignoring event from our own write (remaining: {new_value})"
);
return; }
log::trace!("marking needs_reload");
needs_reload.store(true, Ordering::Relaxed);
}
_ => {} }
}
},
notify::Config::default(),
)?;
watcher.watch(&watch_path, RecursiveMode::NonRecursive)?;
self._watcher = Some(watcher);
log::trace!("watching {}", watch_path.display());
Ok(())
}
fn check_and_reload(&mut self) -> Result<()> {
if self.needs_reload.load(Ordering::Relaxed) {
log::trace!("needs reload detected");
self.load()?;
self.needs_reload.store(false, Ordering::Relaxed);
}
Ok(())
}
pub fn get_or_create(&mut self, session_id: &str) -> Result<&T> {
self.check_and_reload()?;
let mut changed = false;
{
self.sessions
.entry(session_id.to_string())
.and_modify(|_e| {
changed = false;
})
.or_insert_with(|| {
changed = true; SessionEntry::default()
});
}
if changed {
self.save()?;
}
Ok(&self.sessions.get(session_id).unwrap().data)
}
pub fn get(&mut self, session_id: &str) -> Result<Option<&T>> {
self.check_and_reload()?;
Ok(self.sessions.get(session_id).map(|entry| &entry.data))
}
pub fn update(&mut self, session_id: &str, fun: impl FnOnce(&mut T)) -> Result<()> {
self.check_and_reload()?;
let mut changed = false;
{
match self.sessions.entry(session_id.to_string()) {
Entry::Occupied(mut entry) => {
let entry = entry.get_mut();
let before_data = entry.data.clone();
fun(&mut entry.data);
if before_data != entry.data {
entry.update_last_used();
changed = true;
}
}
Entry::Vacant(vacant) => {
let mut entry = SessionEntry::default();
fun(&mut entry.data);
entry.update_last_used();
changed = true; vacant.insert(entry);
}
}
}
if changed {
self.save()?;
}
Ok(())
}
pub fn set(&mut self, session_id: &str, data: T) -> Result<()> {
self.update(session_id, |existing| *existing = data)
}
fn load(&mut self) -> Result<()> {
if let Some(storage_path) = &self.storage_path
&& storage_path.exists()
{
log::trace!("reloading {}...", storage_path.display());
let contents = std::fs::read_to_string(storage_path)?;
if !contents.trim().is_empty()
&& let Ok(sessions) = serde_json::from_str(&contents)
{
log::debug!("reloaded {}", storage_path.display());
self.sessions = sessions;
}
}
Ok(())
}
fn save(&self) -> Result<()> {
if let Some(storage_path) = &self.storage_path {
self.ignore_next_events.store(2, Ordering::Relaxed);
log::trace!("saving");
let temp_path = storage_path.with_extension("tmp");
let contents = serde_json::to_string_pretty(&self.sessions)?;
std::fs::write(&temp_path, &contents)?;
std::fs::rename(temp_path, storage_path)?;
log::trace!("saved");
}
Ok(())
}
}