use std::io::{BufReader, BufWriter, Write};
use std::ops::{Deref, DerefMut};
use std::path::{Path, PathBuf};
use std::time::Duration;
use atomicwrites::OverwriteBehavior::AllowOverwrite;
use atomicwrites::{AtomicFile, Error as AtomicWriteError};
use fs_err::{File, tokio as tokio_fs};
use parking_lot::{Condvar, Mutex, RwLock, RwLockReadGuard, RwLockUpgradableReadGuard};
use serde::{Deserialize, Serialize};
use crate::common::tar_ext;
#[derive(Debug, Default)]
pub struct SaveOnDisk<T> {
change_notification: Condvar,
notification_lock: Mutex<()>,
data: RwLock<T>,
path: PathBuf,
}
#[derive(thiserror::Error, Debug)]
pub enum Error {
#[error("Failed to save structure on disk with error: {0}")]
AtomicWrite(#[from] AtomicWriteError<serde_json::Error>),
#[error("Failed to perform io operation: {0}")]
IoError(#[from] std::io::Error),
#[error("Failed to (de)serialize from/to json: {0}")]
JsonError(#[from] serde_json::Error),
#[error("Error in write closure: {0}")]
FromClosure(Box<dyn std::error::Error>),
}
impl<T: Serialize + for<'de> Deserialize<'de> + Clone> SaveOnDisk<T> {
pub fn load_or_init_default(path: impl Into<PathBuf>) -> Result<Self, Error>
where
T: Default,
{
Self::load_or_init(path, T::default)
}
pub fn load_or_init(path: impl Into<PathBuf>, init: impl FnOnce() -> T) -> Result<Self, Error> {
let path: PathBuf = path.into();
let data = if path.exists() {
let file = BufReader::new(File::open(&path)?);
serde_json::from_reader(file)?
} else {
init()
};
Ok(Self {
change_notification: Condvar::new(),
notification_lock: Default::default(),
data: RwLock::new(data),
path,
})
}
pub fn new(path: impl Into<PathBuf>, data: T) -> Result<Self, Error> {
let data = Self {
change_notification: Condvar::new(),
notification_lock: Default::default(),
data: RwLock::new(data),
path: path.into(),
};
data.save()?;
Ok(data)
}
#[must_use]
pub fn wait_for<F>(&self, check: F, timeout: Duration) -> bool
where
F: Fn(&T) -> bool,
{
let deadline = std::time::Instant::now() + timeout;
loop {
let remaining = deadline.saturating_duration_since(std::time::Instant::now());
if remaining.is_zero() {
return false;
}
let mut data_read_guard = self.data.read();
if check(&data_read_guard) {
return true;
}
let notification_guard = self.notification_lock.lock();
RwLockReadGuard::unlocked(&mut data_read_guard, || {
let mut guard = notification_guard;
self.change_notification.wait_for(&mut guard, remaining);
});
}
}
pub fn write_optional(&self, f: impl FnOnce(&T) -> Option<T>) -> Result<bool, Error> {
let read_data = self.data.upgradable_read();
let output_opt = f(&read_data);
if let Some(output) = output_opt {
Self::save_data_to(&self.path, &output)?;
let mut write_data = RwLockUpgradableReadGuard::upgrade(read_data);
*write_data = output;
drop(write_data);
self.notify_change();
Ok(true)
} else {
Ok(false)
}
}
pub fn write<O>(&self, f: impl FnOnce(&mut T) -> O) -> Result<O, Error> {
let read_data = self.data.upgradable_read();
let mut data_copy = (*read_data).clone();
let output = f(&mut data_copy);
Self::save_data_to(&self.path, &data_copy)?;
let mut write_data = RwLockUpgradableReadGuard::upgrade(read_data);
*write_data = data_copy;
drop(write_data);
self.notify_change();
Ok(output)
}
fn notify_change(&self) {
let _guard = self.notification_lock.lock();
self.change_notification.notify_all();
}
fn save_data_to(path: impl Into<PathBuf>, data: &T) -> Result<(), Error> {
let path: PathBuf = path.into();
AtomicFile::new(path, AllowOverwrite).write(|file| {
let mut writer = BufWriter::new(file);
serde_json::to_writer(&mut writer, data)?;
writer.flush().map_err(serde_json::Error::io)
})?;
Ok(())
}
pub fn save(&self) -> Result<(), Error> {
self.save_to(&self.path)
}
pub fn save_to(&self, path: impl Into<PathBuf>) -> Result<(), Error> {
Self::save_data_to(path, &self.data.read())
}
pub async fn save_to_tar(
&self,
tar: &tar_ext::BuilderExt,
path: impl AsRef<Path>,
) -> Result<(), Error> {
let data_bytes = serde_json::to_vec(self.data.read().deref())?;
tar.append_data(data_bytes, path.as_ref()).await?;
Ok(())
}
pub async fn delete(self) -> std::io::Result<()> {
tokio_fs::remove_file(self.path).await
}
}
impl<T> Deref for SaveOnDisk<T> {
type Target = RwLock<T>;
fn deref(&self) -> &Self::Target {
&self.data
}
}
impl<T> DerefMut for SaveOnDisk<T> {
fn deref_mut(&mut self) -> &mut Self::Target {
&mut self.data
}
}
#[cfg(test)]
mod tests {
use std::sync::Arc;
use std::thread;
use std::thread::sleep;
use std::time::Duration;
use fs_err as fs;
use tempfile::Builder;
use super::SaveOnDisk;
#[test]
fn saves_data() {
let dir = Builder::new().prefix("test").tempdir().unwrap();
let counter_file = dir.path().join("counter");
let counter: SaveOnDisk<u32> = SaveOnDisk::load_or_init_default(&counter_file).unwrap();
counter.write(|counter| *counter += 1).unwrap();
assert_eq!(*counter.read(), 1);
assert_eq!(
counter.read().to_string(),
fs::read_to_string(&counter_file).unwrap()
);
counter.write(|counter| *counter += 1).unwrap();
assert_eq!(*counter.read(), 2);
assert_eq!(
counter.read().to_string(),
fs::read_to_string(&counter_file).unwrap()
);
}
#[test]
fn loads_data() {
let dir = Builder::new().prefix("test").tempdir().unwrap();
let counter_file = dir.path().join("counter");
let counter: SaveOnDisk<u32> = SaveOnDisk::load_or_init_default(&counter_file).unwrap();
counter.write(|counter| *counter += 1).unwrap();
let counter: SaveOnDisk<u32> = SaveOnDisk::load_or_init_default(&counter_file).unwrap();
let value = *counter.read();
assert_eq!(value, 1)
}
#[test]
fn test_wait_for_condition_change() {
let dir = Builder::new().prefix("test").tempdir().unwrap();
let counter_file = dir.path().join("counter");
let counter: Arc<SaveOnDisk<u32>> =
Arc::new(SaveOnDisk::load_or_init_default(counter_file).unwrap());
let counter_copy = counter.clone();
let handle = thread::spawn(move || {
sleep(Duration::from_millis(200));
counter_copy.write(|counter| *counter += 3).unwrap();
sleep(Duration::from_millis(200));
counter_copy.write(|counter| *counter += 7).unwrap();
sleep(Duration::from_millis(200));
});
assert!(counter.wait_for(|counter| *counter > 5, Duration::from_secs(2)));
handle.join().unwrap();
}
#[test]
fn test_wait_for_condition_change_timeout() {
let dir = Builder::new().prefix("test").tempdir().unwrap();
let counter_file = dir.path().join("counter");
let counter: Arc<SaveOnDisk<u32>> =
Arc::new(SaveOnDisk::load_or_init_default(counter_file).unwrap());
let counter_copy = counter.clone();
let handle = thread::spawn(move || {
sleep(Duration::from_millis(200));
counter_copy.write(|counter| *counter += 3).unwrap();
sleep(Duration::from_millis(200));
counter_copy.write(|counter| *counter += 7).unwrap();
sleep(Duration::from_millis(200));
});
assert!(!counter.wait_for(|counter| *counter > 5, Duration::from_millis(300)));
handle.join().unwrap();
}
}