use crate::verifying_client::state_store::{StateStore, WriteThroughCache};
use diem_crypto::hash::{CryptoHasher, HashValue};
use diem_types::{
transaction::Version,
trusted_state::{TrustedState, TrustedStateHasher},
};
use std::{
cmp::max_by_key,
fs::{self, File},
io::{self, Read, Seek, SeekFrom, Write},
path::Path,
sync::{Arc, Mutex},
};
#[derive(Debug, Clone)]
pub struct FileStateStore(Arc<WriteThroughCache<ACIDStateFiles>>);
#[derive(Debug)]
struct ACIDStateFiles(Mutex<[StateFile; 2]>);
#[derive(Debug)]
struct StateFile {
version: Option<Version>,
file: File,
}
impl FileStateStore {
pub fn new(dir: &Path) -> io::Result<Self> {
let store = ACIDStateFiles::new(dir)?;
let store_cache = WriteThroughCache::new(store)?;
Ok(Self(Arc::new(store_cache)))
}
}
impl StateStore for FileStateStore {
type Error = io::Error;
fn latest_state(&self) -> io::Result<Option<TrustedState>> {
self.0.latest_state()
}
fn latest_state_version(&self) -> io::Result<Option<u64>> {
self.0.latest_state_version()
}
fn store(&self, new_state: &TrustedState) -> io::Result<()> {
self.0.store(new_state)
}
}
impl ACIDStateFiles {
fn new(dir: &Path) -> io::Result<Self> {
let state_file0 = StateFile::open(&dir.join("trusted_state.0"))?;
let state_file1 = StateFile::open(&dir.join("trusted_state.1"))?;
let state_files = [state_file0, state_file1];
fsync_dir(dir)?;
Ok(Self(Mutex::new(state_files)))
}
}
impl StateStore for ACIDStateFiles {
type Error = io::Error;
fn latest_state(&self) -> io::Result<Option<TrustedState>> {
let mut state_files = self.0.lock().unwrap();
let state0 = state_files[0].read()?;
let state1 = state_files[1].read()?;
Ok(max_by_key(state0, state1, |opt| {
opt.as_ref().map(|s| s.version())
}))
}
fn store(&self, new_state: &TrustedState) -> io::Result<()> {
let mut state_files = self.0.lock().unwrap();
let newest_version = state_files
.iter_mut()
.max_by_key(|f| f.version)
.unwrap()
.version;
let oldest_state_file = state_files.iter_mut().min_by_key(|f| f.version).unwrap();
if Some(new_state.version()) > newest_version {
oldest_state_file.write(new_state)?;
}
Ok(())
}
}
fn decode_and_validate_checksum(mut buf: Vec<u8>) -> io::Result<(Vec<u8>, HashValue)> {
let offset = buf
.len()
.checked_sub(HashValue::LENGTH)
.ok_or_else(|| invalid_data("state file: empty or too small"))?;
let file_hash = HashValue::from_slice(&buf[offset..]).expect("cannot fail");
buf.truncate(offset);
let computed_hash = TrustedStateHasher::hash_all(&buf);
if file_hash != computed_hash {
Err(invalid_data(format!(
"state file: corrupt: file checksum ({:x}) != computed checksum ({:x})",
file_hash, computed_hash
)))
} else {
Ok((buf, file_hash))
}
}
fn decode_state(buf: Vec<u8>) -> io::Result<TrustedState> {
let (buf, _) = decode_and_validate_checksum(buf)?;
let state = bcs::from_bytes(&buf).map_err(invalid_data)?;
Ok(state)
}
fn encode_state(state: &TrustedState) -> io::Result<Vec<u8>> {
let mut buf = bcs::to_bytes(state).map_err(invalid_input)?;
let hash = TrustedStateHasher::hash_all(&buf);
buf.extend_from_slice(hash.as_ref());
Ok(buf)
}
fn read_file(file: &mut File) -> io::Result<Vec<u8>> {
let mut buf = Vec::new();
file.seek(SeekFrom::Start(0))?;
file.read_to_end(&mut buf)?;
Ok(buf)
}
fn write_file(file: &mut File, buf: &[u8]) -> io::Result<()> {
file.seek(SeekFrom::Start(0))?;
file.write_all(buf)?;
file.set_len(buf.len() as u64)?;
file.sync_all()?;
Ok(())
}
impl StateFile {
fn open(path: &Path) -> io::Result<Self> {
let file = fs::OpenOptions::new()
.read(true)
.write(true)
.create(true)
.open(path)?;
Ok(Self::new(file))
}
fn new(file: File) -> Self {
Self {
version: None,
file,
}
}
fn read(&mut self) -> io::Result<Option<TrustedState>> {
let buf = read_file(&mut self.file)?;
let maybe_state = decode_state(buf)
.map_err(|_err| () )
.ok();
self.version = maybe_state.as_ref().map(|s| s.version());
Ok(maybe_state)
}
fn write(&mut self, new_state: &TrustedState) -> io::Result<()> {
let buf = encode_state(new_state)?;
write_file(&mut self.file, &buf)?;
self.version = Some(new_state.version());
Ok(())
}
}
type BoxError = Box<dyn std::error::Error + Send + Sync + 'static>;
fn invalid_input(err: impl Into<BoxError>) -> io::Error {
io::Error::new(io::ErrorKind::InvalidInput, err)
}
fn invalid_data(err: impl Into<BoxError>) -> io::Error {
io::Error::new(io::ErrorKind::InvalidData, err)
}
fn fsync_dir<P: AsRef<Path>>(dir: P) -> io::Result<()> {
let mut open_opts = fs::OpenOptions::new();
open_opts.read(true);
#[cfg(windows)]
{
use std::os::windows::fs::OpenOptionsExt;
use winapi::winbase;
open_opts
.write(true)
.custom_flags(winbase::FILE_FLAG_BACKUP_SEMANTICS);
}
let fd = open_opts.open(dir)?;
fd.sync_all()?;
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
use proptest::{collection::vec, prelude::*, sample::Index};
use tempfile::{tempdir, tempfile};
fn max_state(idx: usize, states: &[TrustedState]) -> &TrustedState {
states[..=idx].iter().max_by_key(|s| s.version()).unwrap()
}
fn corrupt_file(corrupt_idx: &Index, file: &mut File) {
let len = file.metadata().unwrap().len();
if len > 0 {
let new_len = corrupt_idx.index(len as usize);
file.set_len(new_len as u64).unwrap();
}
}
proptest! {
#![proptest_config(ProptestConfig::with_cases(100))]
#[test]
fn state_file_read_corrupt(
state in any::<TrustedState>(),
corrupt_idx in any::<Index>(),
) {
let file = tempfile().unwrap();
let mut state_file = StateFile::new(file);
state_file.write(&state).unwrap();
assert_eq!(state_file.version, Some(state.version()));
let maybe_state = state_file.read().unwrap();
assert_eq!(maybe_state, Some(state));
corrupt_file(&corrupt_idx, &mut state_file.file);
let maybe_state = state_file.read().unwrap();
assert_eq!(maybe_state, None);
}
}
proptest! {
#![proptest_config(ProptestConfig::with_cases(25))]
#[test]
fn file_state_store(
states in vec(any::<TrustedState>(), 1..10),
corrupt_idx in any::<Index>(),
) {
let dir = tempdir().unwrap();
let state_store = FileStateStore::new(dir.path()).unwrap();
assert_eq!(None, state_store.latest_state().unwrap());
for (idx, state) in states.iter().enumerate() {
state_store.store(state).unwrap();
let store_max = state_store.latest_state().unwrap().unwrap();
let expected_max = max_state(idx, &states);
assert_eq!(expected_max, &store_max);
}
let store_max1 = state_store.latest_state().unwrap().unwrap();
drop(state_store);
let state_store = FileStateStore::new(dir.path()).unwrap();
let store_max2 = state_store.latest_state().unwrap().unwrap();
assert_eq!(store_max1, store_max2);
{
let mut state_files = state_store.0.as_inner().0.lock().unwrap();
if let Some(oldest_state_file) = state_files.iter_mut().min_by_key(|f| f.version).as_mut() {
corrupt_file(&corrupt_idx, &mut oldest_state_file.file);
}
}
drop(state_store);
let state_store = FileStateStore::new(dir.path()).unwrap();
let store_max3 = state_store.latest_state().unwrap().unwrap();
assert_eq!(store_max1, store_max3);
}
}
}