use std::collections::BTreeMap;
use std::path::PathBuf;
use std::sync::{Arc, LazyLock, Mutex};
use serde::{Deserialize, Serialize};
use crate::channels::ReplyReference;
use crate::util::{UnwrapPoison, unix_millis};
const MAX_PERSISTED_ENTRIES: usize = 64;
const DRAFT_FILE_NAME: &str = "chat-draft.json";
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub(crate) struct DraftEntry {
#[serde(default)]
pub(crate) text: String,
#[serde(default)]
pub(crate) reply: Option<ReplyReference>,
#[serde(default)]
pub(crate) saved_at: u64,
}
impl DraftEntry {
#[must_use]
fn is_empty(&self) -> bool {
self.text.trim().is_empty() && self.reply.is_none()
}
}
type DraftFile = BTreeMap<String, BTreeMap<String, DraftEntry>>;
pub(crate) struct DraftStore {
path: Option<PathBuf>,
entries: Mutex<DraftFile>,
}
impl DraftStore {
#[must_use]
fn file_path() -> Option<PathBuf> {
let home = std::env::var("HOME").ok()?;
if home.is_empty() {
return None;
}
Some(PathBuf::from(home).join(".mahbot").join(DRAFT_FILE_NAME))
}
fn load(path: Option<PathBuf>) -> Self {
let entries = path
.as_deref()
.and_then(|p| std::fs::read_to_string(p).ok())
.and_then(|json| serde_json::from_str(&json).ok())
.unwrap_or_default();
Self {
path,
entries: Mutex::new(entries),
}
}
#[must_use]
pub(crate) fn global() -> &'static Arc<DraftStore> {
static GLOBAL: LazyLock<Arc<DraftStore>> =
LazyLock::new(|| Arc::new(DraftStore::load(DraftStore::file_path())));
&GLOBAL
}
#[cfg(test)]
#[must_use]
pub(crate) fn at(path: Option<PathBuf>) -> Arc<DraftStore> {
Arc::new(DraftStore::load(path))
}
pub(crate) fn set(
&self,
user: &str,
workspace: &str,
text: String,
reply: Option<ReplyReference>,
) {
if self.path.is_none() {
return;
}
let entry = DraftEntry {
text,
reply,
saved_at: unix_millis(),
};
if entry.is_empty() {
self.remove(user, workspace);
return;
}
self.entries
.lock()
.unwrap_poison()
.entry(user.to_string())
.or_default()
.insert(workspace.to_string(), entry);
}
pub(crate) fn remove(&self, user: &str, workspace: &str) {
if self.path.is_none() {
return;
}
let mut map = self.entries.lock().unwrap_poison();
if let Some(inner) = map.get_mut(user) {
inner.remove(workspace);
if inner.is_empty() {
map.remove(user);
}
}
}
#[must_use]
pub(crate) fn get(&self, user: &str, workspace: &str) -> Option<DraftEntry> {
self.path.as_ref()?;
let map = self.entries.lock().unwrap_poison();
map.get(user)?.get(workspace).cloned()
}
pub(crate) fn persist(&self) {
let Some(path) = &self.path else {
return;
};
let mut map = self.entries.lock().unwrap_poison();
prune(&mut map);
let json = serde_json::to_string(&*map).unwrap_or_default();
let tmp = path.with_extension("json.tmp");
if let Some(parent) = path.parent() {
let _ = std::fs::create_dir_all(parent);
}
let _ = std::fs::write(&tmp, json);
let _ = std::fs::rename(&tmp, path);
}
pub(crate) fn persist_async(self: Arc<Self>) {
if let Ok(handle) = tokio::runtime::Handle::try_current() {
handle.spawn_blocking(move || self.persist());
}
}
}
fn prune(map: &mut DraftFile) {
let total: usize = map.values().map(BTreeMap::len).sum();
if total <= MAX_PERSISTED_ENTRIES {
return;
}
let mut all: Vec<(String, String, u64)> = map
.iter()
.flat_map(|(user, ws)| {
ws.iter()
.map(move |(workspace, entry)| (user.clone(), workspace.clone(), entry.saved_at))
})
.collect();
all.sort_by_key(|(_, _, saved_at)| *saved_at);
let keep_from = all.len().saturating_sub(MAX_PERSISTED_ENTRIES);
let keep: std::collections::HashSet<(String, String)> = all[keep_from..]
.iter()
.map(|(user, workspace, _)| (user.clone(), workspace.clone()))
.collect();
map.retain(|user, ws| {
ws.retain(|workspace, _| keep.contains(&(user.clone(), workspace.clone())));
!ws.is_empty()
});
}
pub(crate) fn flush_global() {
DraftStore::global().persist();
}
#[cfg(test)]
mod tests {
use super::*;
fn temp_store() -> (tempfile::TempDir, Arc<DraftStore>) {
let dir = tempfile::tempdir().unwrap();
let store = DraftStore::at(Some(dir.path().join(DRAFT_FILE_NAME)));
(dir, store)
}
#[test]
fn round_trip_set_get_persist_reload() {
let (dir, store) = temp_store();
let path = dir.path().join(DRAFT_FILE_NAME);
store.set("alice", "ws1", "hello".to_string(), None);
assert_eq!(
store.get("alice", "ws1").map(|e| e.text).as_deref(),
Some("hello")
);
store.persist();
let reloaded = DraftStore::at(Some(path));
let entry = reloaded.get("alice", "ws1").expect("entry persisted");
assert_eq!(entry.text, "hello");
assert_eq!(entry.reply, None);
}
#[test]
fn corrupt_file_fails_open_empty() {
let dir = tempfile::tempdir().unwrap();
let path = dir.path().join(DRAFT_FILE_NAME);
std::fs::write(&path, "not valid json{").unwrap();
let store = DraftStore::at(Some(path));
assert_eq!(store.get("alice", "ws1"), None);
}
#[test]
fn empty_draft_removes_the_entry() {
let (_dir, store) = temp_store();
store.set("alice", "ws1", "hello".to_string(), None);
store.set("alice", "ws1", " ".to_string(), None);
assert_eq!(store.get("alice", "ws1"), None);
}
#[test]
fn pruning_keeps_the_newest_entries() {
let (dir, store) = temp_store();
let path = dir.path().join(DRAFT_FILE_NAME);
{
let mut map = store.entries.lock().unwrap_poison();
for i in 0..(MAX_PERSISTED_ENTRIES + 10) {
map.entry("user".to_string()).or_default().insert(
format!("ws{i}"),
DraftEntry {
text: format!("text{i}"),
reply: None,
saved_at: i as u64,
},
);
}
}
store.persist();
let reloaded = DraftStore::at(Some(path));
let total: usize = reloaded
.entries
.lock()
.unwrap_poison()
.values()
.map(BTreeMap::len)
.sum();
assert_eq!(total, MAX_PERSISTED_ENTRIES);
for i in 10..(MAX_PERSISTED_ENTRIES + 10) {
assert!(
reloaded.get("user", &format!("ws{i}")).is_some(),
"ws{i} must survive"
);
}
for i in 0..10 {
assert!(
reloaded.get("user", &format!("ws{i}")).is_none(),
"ws{i} must be pruned"
);
}
}
#[test]
fn persist_leaves_no_tmp_file_behind() {
let (dir, store) = temp_store();
let path = dir.path().join(DRAFT_FILE_NAME);
store.set("alice", "ws1", "hello".to_string(), None);
store.persist();
assert!(path.exists());
assert!(!path.with_extension("json.tmp").exists());
}
}