use std::collections::{HashMap, HashSet};
use std::fs;
use std::io;
use std::net::IpAddr;
use std::path::{Path, PathBuf};
use std::sync::Mutex;
use serde::de::DeserializeOwned;
use serde::{Deserialize, Serialize};
use crate::clock::Timestamp;
pub type DatedEntries<K, V> = Vec<(K, (Timestamp, Option<V>))>;
#[derive(Clone, Debug, Deserialize, Serialize)]
#[serde(bound(
serialize = "K: Serialize, V: Serialize",
deserialize = "K: Deserialize<'de> + Eq + std::hash::Hash, V: Deserialize<'de>"
))]
pub struct PersistedState<K, V> {
pub entries: DatedEntries<K, V>,
pub members: HashSet<IpAddr>,
pub tombstone_acks: HashMap<K, HashMap<IpAddr, u64>>,
}
pub trait Persistence<K, V>: Send + Sync + 'static {
fn load(&self) -> io::Result<Option<PersistedState<K, V>>>;
fn save(&self, state: &PersistedState<K, V>) -> io::Result<()>;
}
pub struct InMemoryPersistence<K, V> {
state: Mutex<Option<PersistedState<K, V>>>,
}
impl<K, V> Default for InMemoryPersistence<K, V> {
fn default() -> Self {
InMemoryPersistence {
state: Mutex::new(None),
}
}
}
impl<K, V> InMemoryPersistence<K, V> {
pub fn new() -> Self {
Self::default()
}
}
impl<K, V> Persistence<K, V> for InMemoryPersistence<K, V>
where
K: Clone + Send + Sync + 'static,
V: Clone + Send + Sync + 'static,
{
fn load(&self) -> io::Result<Option<PersistedState<K, V>>> {
Ok(self.state.lock().unwrap().clone())
}
fn save(&self, state: &PersistedState<K, V>) -> io::Result<()> {
*self.state.lock().unwrap() = Some(state.clone());
Ok(())
}
}
#[derive(Clone, Debug)]
pub struct FileSnapshot {
path: PathBuf,
}
impl FileSnapshot {
pub fn new(path: impl AsRef<Path>) -> Self {
FileSnapshot {
path: path.as_ref().to_path_buf(),
}
}
fn tmp_path(&self) -> PathBuf {
let mut tmp = self.path.clone().into_os_string();
tmp.push(".tmp");
PathBuf::from(tmp)
}
}
impl<K, V> Persistence<K, V> for FileSnapshot
where
K: Serialize + DeserializeOwned + Eq + std::hash::Hash + Send + Sync + 'static,
V: Serialize + DeserializeOwned + Send + Sync + 'static,
{
fn load(&self) -> io::Result<Option<PersistedState<K, V>>> {
let bytes = match fs::read(&self.path) {
Ok(bytes) => bytes,
Err(err) if err.kind() == io::ErrorKind::NotFound => return Ok(None),
Err(err) => return Err(err),
};
let state = bincode::deserialize(&bytes)
.map_err(|err| io::Error::new(io::ErrorKind::InvalidData, err))?;
Ok(Some(state))
}
fn save(&self, state: &PersistedState<K, V>) -> io::Result<()> {
let bytes = bincode::serialize(state)
.map_err(|err| io::Error::new(io::ErrorKind::InvalidData, err))?;
let tmp = self.tmp_path();
{
use std::io::Write;
let mut file = fs::File::create(&tmp)?;
file.write_all(&bytes)?;
file.sync_all()?;
}
fs::rename(&tmp, &self.path)?;
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
fn sample_state() -> PersistedState<i32, String> {
let mut members = HashSet::new();
members.insert("127.0.0.1".parse().unwrap());
members.insert("127.0.0.2".parse().unwrap());
let mut acks = HashMap::new();
let mut key_acks = HashMap::new();
key_acks.insert("127.0.0.1".parse().unwrap(), 42u64);
acks.insert(7, key_acks);
PersistedState {
entries: vec![
(1, (Timestamp::new(1_000, 0, 7), Some("alive".to_string()))),
(2, (Timestamp::new(2_000, 1, 7), None)), ],
members,
tombstone_acks: acks,
}
}
fn assert_states_eq(a: &PersistedState<i32, String>, b: &PersistedState<i32, String>) {
assert_eq!(a.entries, b.entries);
assert_eq!(a.members, b.members);
assert_eq!(a.tombstone_acks, b.tombstone_acks);
}
#[test]
fn persisted_state_bincode_roundtrip() {
let state = sample_state();
let bytes = bincode::serialize(&state).unwrap();
let back: PersistedState<i32, String> = bincode::deserialize(&bytes).unwrap();
assert_states_eq(&back, &state);
}
#[test]
fn in_memory_roundtrips_within_process() {
let backend = InMemoryPersistence::<i32, String>::new();
assert!(backend.load().unwrap().is_none());
let state = sample_state();
backend.save(&state).unwrap();
assert_states_eq(&backend.load().unwrap().unwrap(), &state);
}
#[test]
fn in_memory_save_replaces_previous() {
let backend = InMemoryPersistence::<i32, String>::new();
let mut first = sample_state();
first.entries = vec![(1, (Timestamp::new(1, 0, 0), Some("first".to_string())))];
backend.save(&first).unwrap();
let mut second = sample_state();
second.entries = vec![(1, (Timestamp::new(2, 0, 0), Some("second".to_string())))];
backend.save(&second).unwrap();
let loaded = backend.load().unwrap().unwrap();
assert_eq!(loaded.entries[0].1 .1, Some("second".to_string()));
}
#[test]
fn file_snapshot_save_then_load() {
let dir = tempfile::tempdir().unwrap();
let backend = FileSnapshot::new(dir.path().join("snapshot.bin"));
assert!(Persistence::<i32, String>::load(&backend)
.unwrap()
.is_none());
let state = sample_state();
Persistence::<i32, String>::save(&backend, &state).unwrap();
let loaded = Persistence::<i32, String>::load(&backend)
.unwrap()
.expect("a snapshot was saved");
assert_states_eq(&loaded, &state);
}
#[test]
fn file_snapshot_save_is_atomic_replace() {
let dir = tempfile::tempdir().unwrap();
let backend = FileSnapshot::new(dir.path().join("snapshot.bin"));
let mut first = sample_state();
first.entries = vec![(1, (Timestamp::new(1, 0, 0), Some("first".to_string())))];
Persistence::<i32, String>::save(&backend, &first).unwrap();
let mut second = sample_state();
second.entries = vec![(1, (Timestamp::new(2, 0, 0), Some("second".to_string())))];
Persistence::<i32, String>::save(&backend, &second).unwrap();
assert!(!dir.path().join("snapshot.bin.tmp").exists());
let loaded = Persistence::<i32, String>::load(&backend).unwrap().unwrap();
assert_eq!(loaded.entries[0].1 .1, Some("second".to_string()));
}
}