use std::collections::BTreeMap;
use std::collections::btree_map::Values;
use std::fmt;
use std::io;
use crate::store::NodeStore;
use crate::tree::{Hash, batch_mutate_owned};
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum Mutation {
Put { key: Vec<u8>, value: Vec<u8> },
Delete { key: Vec<u8> },
}
impl Mutation {
#[must_use]
pub fn key(&self) -> &[u8] {
match self {
Self::Put { key, .. } | Self::Delete { key } => key,
}
}
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum LookupResult {
BufferedValue(Vec<u8>),
BufferedDelete,
NotBuffered,
}
pub(crate) const RECOVERY_UNKNOWN_TAG_EXPECTED: u32 = 0xffff_fffe;
pub(crate) const RECOVERY_MALFORMED_PAYLOAD_EXPECTED: u32 = 0xffff_fffd;
#[derive(Debug)]
pub enum WalError {
Io(io::Error),
ChecksumMismatch { expected: u32, actual: u32 },
TreeError(String),
InvalidTag { found: u8 },
Truncated,
TrailingBytes { trailing: usize },
LengthOverflow,
InvalidFsyncPolicy { interval: usize },
MissingCommittedRoot { root: Hash },
}
impl fmt::Display for WalError {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::Io(error) => write!(formatter, "wal i/o error: {error}"),
Self::ChecksumMismatch {
expected: RECOVERY_UNKNOWN_TAG_EXPECTED,
actual,
} => write!(formatter, "wal corruption: unknown tag {actual:#04x}"),
Self::ChecksumMismatch {
expected: RECOVERY_MALFORMED_PAYLOAD_EXPECTED,
..
} => write!(formatter, "wal corruption: malformed payload fields"),
Self::ChecksumMismatch { expected, actual } => write!(
formatter,
"wal checksum mismatch: expected {expected:#010x}, got {actual:#010x}"
),
Self::TreeError(message) => write!(formatter, "wal tree error: {message}"),
Self::InvalidTag { found } => write!(formatter, "invalid wal tag: {found:#04x}"),
Self::Truncated => write!(formatter, "wal bytes ended before the frame was complete"),
Self::TrailingBytes { trailing } => {
write!(formatter, "wal bytes contain {trailing} trailing bytes")
}
Self::LengthOverflow => {
write!(formatter, "encoded wal length cannot fit on this platform")
}
Self::InvalidFsyncPolicy { interval } => write!(
formatter,
"invalid wal fsync policy: batched interval must be greater than zero, got {interval}"
),
Self::MissingCommittedRoot { root } => write!(
formatter,
"wal committed root is missing from the node store: {root}"
),
}
}
}
impl std::error::Error for WalError {
fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
match self {
Self::Io(error) => Some(error),
Self::ChecksumMismatch { .. }
| Self::TreeError(_)
| Self::InvalidTag { .. }
| Self::Truncated
| Self::TrailingBytes { .. }
| Self::LengthOverflow
| Self::InvalidFsyncPolicy { .. }
| Self::MissingCommittedRoot { .. } => None,
}
}
}
impl From<io::Error> for WalError {
fn from(error: io::Error) -> Self {
Self::Io(error)
}
}
#[derive(Clone, Debug, Default)]
pub struct WalBuffer {
mutations: BTreeMap<Vec<u8>, Mutation>,
}
impl WalBuffer {
#[must_use]
pub fn new() -> Self {
Self::default()
}
pub fn put<K: AsRef<[u8]>, V: AsRef<[u8]>>(&mut self, key: K, value: V) {
let key = key.as_ref().to_vec();
self.mutations.insert(
key.clone(),
Mutation::Put {
key,
value: value.as_ref().to_vec(),
},
);
}
pub fn delete<K: AsRef<[u8]>>(&mut self, key: K) {
let key = key.as_ref().to_vec();
self.mutations.insert(key.clone(), Mutation::Delete { key });
}
#[must_use]
pub fn snapshot_entry(&self, key: &[u8]) -> Option<Mutation> {
self.mutations.get(key).cloned()
}
pub fn restore_entry(&mut self, key: &[u8], prior: Option<Mutation>) {
match prior {
Some(mutation) => {
self.mutations.insert(key.to_vec(), mutation);
}
None => {
self.mutations.remove(key);
}
}
}
pub fn get<K: AsRef<[u8]>>(&self, key: K) -> LookupResult {
match self.mutations.get(key.as_ref()) {
Some(Mutation::Put { value, .. }) => LookupResult::BufferedValue(value.clone()),
Some(Mutation::Delete { .. }) => LookupResult::BufferedDelete,
None => LookupResult::NotBuffered,
}
}
#[must_use]
pub fn len(&self) -> usize {
self.mutations.len()
}
#[must_use]
pub fn is_empty(&self) -> bool {
self.mutations.is_empty()
}
pub fn iter(&self) -> Values<'_, Vec<u8>, Mutation> {
self.mutations.values()
}
pub fn commit<S>(&mut self, tree_root: Hash, store: &mut S) -> Result<Hash, WalError>
where
S: NodeStore + ?Sized,
{
let drained = std::mem::take(&mut self.mutations);
let batch: Vec<(Vec<u8>, Option<Vec<u8>>)> = drained
.values()
.map(|mutation| match mutation {
Mutation::Put { key, value } => (key.clone(), Some(value.clone())),
Mutation::Delete { key } => (key.clone(), None),
})
.collect();
match batch_mutate_owned(store, tree_root, batch) {
Ok(new_root) => Ok(new_root),
Err(error) => {
self.mutations = drained;
Err(WalError::TreeError(error.to_string()))
}
}
}
}
impl<'a> IntoIterator for &'a WalBuffer {
type Item = &'a Mutation;
type IntoIter = Values<'a, Vec<u8>, Mutation>;
fn into_iter(self) -> Self::IntoIter {
self.iter()
}
}
#[cfg(test)]
mod tests {
use super::{LookupResult, Mutation, WalBuffer, WalError};
use crate::store::NodeStore;
use crate::tree::{Hash, LeafNode, Node, NodeError, batch_mutate};
use std::cell::Cell;
use std::convert::Infallible;
#[derive(Debug, Default)]
struct CountingStore {
nodes: std::collections::HashMap<Hash, Vec<u8>>,
puts: Cell<usize>,
}
impl NodeStore for CountingStore {
type Error = Infallible;
fn get(&self, hash: &Hash) -> Result<Option<std::sync::Arc<Node>>, Self::Error> {
Ok(self
.nodes
.get(hash)
.and_then(|bytes| Node::deserialise(bytes).ok())
.map(std::sync::Arc::new))
}
fn put(&mut self, node: &Node) -> Result<Hash, Self::Error> {
self.puts.set(self.puts.get() + 1);
let hash = node.hash();
self.nodes.insert(hash, node.serialise());
Ok(hash)
}
}
impl CountingStore {
fn put_count(&self) -> usize {
self.puts.get()
}
}
fn store_node(store: &mut CountingStore, node: &Node) -> Hash {
match store.put(node) {
Ok(hash) => hash,
Err(infallible) => match infallible {},
}
}
fn empty_root(store: &mut CountingStore) -> Result<Hash, NodeError> {
let leaf = Node::Leaf(LeafNode::new(Vec::new())?);
Ok(store_node(store, &leaf))
}
#[test]
fn new_buffer_is_empty() {
let buffer = WalBuffer::new();
assert!(buffer.is_empty());
assert_eq!(buffer.len(), 0);
}
#[test]
fn put_overwrites_prior_mutation_for_same_key() {
let mut buffer = WalBuffer::new();
buffer.put(b"a", b"1");
buffer.put(b"a", b"2");
assert_eq!(buffer.len(), 1);
assert_eq!(buffer.get(b"a"), LookupResult::BufferedValue(b"2".to_vec()));
}
#[test]
fn delete_overwrites_prior_put_for_same_key() {
let mut buffer = WalBuffer::new();
buffer.put(b"a", b"1");
buffer.delete(b"a");
assert_eq!(buffer.len(), 1);
assert_eq!(buffer.get(b"a"), LookupResult::BufferedDelete);
}
#[test]
fn put_overwrites_prior_delete_for_same_key() {
let mut buffer = WalBuffer::new();
buffer.delete(b"a");
buffer.put(b"a", b"v");
assert_eq!(buffer.len(), 1);
assert_eq!(buffer.get(b"a"), LookupResult::BufferedValue(b"v".to_vec()));
}
#[test]
fn get_shadows_tree_per_variant() {
let mut buffer = WalBuffer::new();
assert_eq!(buffer.get(b"key"), LookupResult::NotBuffered);
buffer.put(b"key", b"val");
assert_eq!(
buffer.get(b"key"),
LookupResult::BufferedValue(b"val".to_vec())
);
buffer.put(b"key", b"v2");
assert_eq!(
buffer.get(b"key"),
LookupResult::BufferedValue(b"v2".to_vec())
);
buffer.delete(b"key");
assert_eq!(buffer.get(b"key"), LookupResult::BufferedDelete);
}
#[test]
fn iteration_is_ascending_key_order() {
let mut buffer = WalBuffer::new();
buffer.put(b"c", b"3");
buffer.put(b"a", b"1");
buffer.put(b"b", b"2");
let keys: Vec<&[u8]> = buffer.iter().map(Mutation::key).collect();
assert_eq!(
keys,
vec![b"a".as_slice(), b"b".as_slice(), b"c".as_slice()]
);
}
#[test]
fn mutation_clone_equals_original() {
let original = Mutation::Put {
key: b"k".to_vec(),
value: b"v".to_vec(),
};
assert_eq!(original.clone(), original);
}
#[test]
fn commit_clears_buffer_and_returns_new_root() -> Result<(), NodeError> {
let mut store = CountingStore::default();
let root = empty_root(&mut store)?;
let mut buffer = WalBuffer::new();
for index in 0..50u32 {
buffer.put(format!("key-{index:04}"), format!("value-{index}"));
}
let new_root = buffer
.commit(root, &mut store)
.map_err(|_| NodeError::Truncated)?;
assert!(buffer.is_empty());
assert_ne!(new_root, root);
Ok(())
}
#[test]
fn commit_triggers_exactly_one_batch_not_n_puts() -> Result<(), NodeError> {
let mut reference = CountingStore::default();
let ref_root = empty_root(&mut reference)?;
let baseline = reference.put_count();
let batch: Vec<(Vec<u8>, Option<Vec<u8>>)> = (0..50u32)
.map(|index| {
(
format!("key-{index:04}").into_bytes(),
Some(format!("value-{index}").into_bytes()),
)
})
.collect();
let expected_root = batch_mutate(&mut reference, ref_root, batch.as_slice())
.map_err(|_| NodeError::Truncated)?;
let batch_puts = reference.put_count() - baseline;
let mut store = CountingStore::default();
let root = empty_root(&mut store)?;
let commit_baseline = store.put_count();
let mut buffer = WalBuffer::new();
for index in 0..50u32 {
buffer.put(format!("key-{index:04}"), format!("value-{index}"));
}
let new_root = buffer
.commit(root, &mut store)
.map_err(|_| NodeError::Truncated)?;
let commit_puts = store.put_count() - commit_baseline;
assert_eq!(new_root, expected_root);
assert_eq!(commit_puts, batch_puts);
assert!(commit_puts < 50);
Ok(())
}
#[test]
fn commit_on_empty_buffer_returns_root_unchanged() -> Result<(), NodeError> {
let mut store = CountingStore::default();
let root = empty_root(&mut store)?;
let before = store.put_count();
let mut buffer = WalBuffer::new();
let result = buffer
.commit(root, &mut store)
.map_err(|_| NodeError::Truncated)?;
assert_eq!(result, root);
assert_eq!(store.put_count(), before);
Ok(())
}
#[test]
fn commit_failure_retains_buffer() {
#[derive(Debug)]
struct MissingRootStore;
#[derive(Debug)]
struct NeverHappens;
impl std::fmt::Display for NeverHappens {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(f, "never happens")
}
}
impl std::error::Error for NeverHappens {}
impl NodeStore for MissingRootStore {
type Error = NeverHappens;
fn get(&self, hash: &Hash) -> Result<Option<std::sync::Arc<Node>>, Self::Error> {
debug_assert_eq!(hash.as_bytes().len(), 32);
Ok(None)
}
fn put(&mut self, node: &Node) -> Result<Hash, Self::Error> {
Ok(node.hash())
}
}
let mut store = MissingRootStore;
let root = Hash::from_bytes([0; 32]);
let mut buffer = WalBuffer::new();
for index in 0..50u32 {
buffer.put(format!("key-{index:04}"), b"v");
}
let result = buffer.commit(root, &mut store);
assert!(matches!(result, Err(WalError::TreeError(_))));
assert_eq!(buffer.len(), 50);
}
#[test]
fn restore_entry_returns_prior_value_on_single_key_rollback() {
let mut buffer = WalBuffer::new();
buffer.put(b"k", b"old");
let prior = buffer.snapshot_entry(b"k");
buffer.put(b"k", b"new");
assert_eq!(
buffer.get(b"k"),
LookupResult::BufferedValue(b"new".to_vec())
);
buffer.restore_entry(b"k", prior);
assert_eq!(
buffer.get(b"k"),
LookupResult::BufferedValue(b"old".to_vec())
);
assert_eq!(buffer.len(), 1);
}
#[test]
fn restore_entry_removes_key_that_was_absent_before() {
let mut buffer = WalBuffer::new();
buffer.put(b"other", b"keep");
let prior = buffer.snapshot_entry(b"k");
assert!(prior.is_none());
buffer.put(b"k", b"new");
buffer.restore_entry(b"k", prior);
assert_eq!(buffer.get(b"k"), LookupResult::NotBuffered);
assert_eq!(
buffer.get(b"other"),
LookupResult::BufferedValue(b"keep".to_vec())
);
assert_eq!(buffer.len(), 1);
}
#[test]
fn wal_error_display_names_both_checksums() {
let error = WalError::ChecksumMismatch {
expected: 0xDEAD,
actual: 0xBEEF,
};
let rendered = error.to_string();
assert!(rendered.contains("dead"));
assert!(rendered.contains("beef"));
}
}