use super::{
Candidate, FileMode, HistoryIndex, NodeFile, QueueRecord, State, TxMeta, WalFile, WalIntent,
WalLog, ensure_eof, open_regular_file,
};
use crate::{
Error, HistoryEntry, MergePair, Node, NodeData, NodeId, ObjectId, Owner, Provenance, Result,
TransactionId, WriterId,
};
use chrono::{DateTime, Utc};
use sha2::{Digest, Sha256};
use std::io::Read;
pub(super) trait BinaryRecord: Sized {
fn encode(&self, encoder: &mut Encoder) -> Result<()>;
fn decode(decoder: &mut Decoder<'_>) -> Result<Self>;
}
pub(super) fn encode_record<T: BinaryRecord>(magic: &[u8; 8], value: &T) -> Result<Vec<u8>> {
let mut encoder = Encoder::new();
encoder.bytes(magic);
encoder.u64(0);
value.encode(&mut encoder)?;
let mut bytes = encoder.finish();
let payload_length = bytes
.len()
.checked_sub(16)
.ok_or_else(|| Error::corrupt("encoded record length underflow"))?;
let payload_length = u64::try_from(payload_length)
.map_err(|_| Error::corrupt("encoded record length does not fit u64"))?;
bytes[8..16].copy_from_slice(&payload_length.to_be_bytes());
bytes.extend_from_slice(&Sha256::digest(&bytes));
Ok(bytes)
}
pub(super) fn read_record<T: BinaryRecord>(path: &std::path::Path, magic: &[u8; 8]) -> Result<T> {
let mut file = open_regular_file(path)?;
let file_length = file.metadata()?.len();
if file_length < 48 {
return Err(Error::corrupt(format!(
"unknown record format at {}",
path.display()
)));
}
let mut header = [0_u8; 16];
file.read_exact(&mut header)?;
if &header[..8] != magic {
return Err(Error::corrupt(format!(
"unknown record format at {}",
path.display()
)));
}
let length = u64::from_be_bytes(
header[8..]
.try_into()
.map_err(|_| Error::corrupt("record length field is truncated"))?,
);
let expected_file_length = length
.checked_add(16 + 32)
.ok_or_else(|| Error::corrupt("record length overflow"))?;
if file_length != expected_file_length {
return Err(Error::corrupt(format!(
"record length mismatch at {}",
path.display()
)));
}
let length =
usize::try_from(length).map_err(|_| Error::corrupt("record length does not fit"))?;
let mut payload = Vec::new();
payload
.try_reserve_exact(length)
.map_err(|_| Error::Io(std::io::Error::other("record allocation failed")))?;
payload.resize(length, 0);
file.read_exact(&mut payload)?;
let mut stored_checksum = [0_u8; 32];
file.read_exact(&mut stored_checksum)?;
ensure_eof(&mut file, "record")?;
let mut checksum = Sha256::new();
checksum.update(header);
checksum.update(&payload);
if checksum.finalize()[..] != stored_checksum {
return Err(Error::corrupt(format!(
"record checksum mismatch at {}",
path.display()
)));
}
let mut decoder = Decoder::new(&payload);
let value = T::decode(&mut decoder)?;
decoder.finish()?;
Ok(value)
}
pub(super) struct Encoder {
bytes: Vec<u8>,
}
impl Encoder {
fn new() -> Self {
Self { bytes: Vec::new() }
}
fn finish(self) -> Vec<u8> {
self.bytes
}
fn bytes(&mut self, value: &[u8]) {
self.bytes.extend_from_slice(value);
}
fn u8(&mut self, value: u8) {
self.bytes.push(value);
}
fn bool(&mut self, value: bool) {
self.u8(u8::from(value));
}
fn u32(&mut self, value: u32) {
self.bytes(&value.to_be_bytes());
}
fn u64(&mut self, value: u64) {
self.bytes(&value.to_be_bytes());
}
fn i64(&mut self, value: i64) {
self.bytes(&value.to_be_bytes());
}
fn usize(&mut self, value: usize) -> Result<()> {
self.u64(
u64::try_from(value)
.map_err(|_| Error::corrupt("record vector length does not fit u64"))?,
);
Ok(())
}
fn string(&mut self, value: &str) -> Result<()> {
self.u32(
u32::try_from(value.len())
.map_err(|_| Error::corrupt("record string length does not fit u32"))?,
);
self.bytes(value.as_bytes());
Ok(())
}
fn datetime(&mut self, value: DateTime<Utc>) {
self.i64(value.timestamp());
self.u32(value.timestamp_subsec_nanos());
}
fn node_id(&mut self, value: NodeId) {
self.bytes(&value.0);
}
fn object_id(&mut self, value: ObjectId) {
self.bytes(&value.0);
}
fn transaction_id(&mut self, value: TransactionId) {
self.bytes(&value.0);
}
fn writer_id(&mut self, value: WriterId) {
self.bytes(&value.0);
}
fn provenance(&mut self, value: &Provenance) -> Result<()> {
self.string(&value.author)?;
self.string(&value.source)?;
self.datetime(value.source_created_at);
self.string(&value.data)
}
fn owner(&mut self, value: Owner) {
match value {
Owner::Unowned => self.u8(0),
Owner::SelfNode => self.u8(1),
Owner::Node(id) => {
self.u8(2);
self.node_id(id);
}
}
}
fn node_ids(&mut self, values: &[NodeId]) -> Result<()> {
self.usize(values.len())?;
for value in values {
self.node_id(*value);
}
Ok(())
}
fn object_ids(&mut self, values: &[ObjectId]) -> Result<()> {
self.usize(values.len())?;
for value in values {
self.object_id(*value);
}
Ok(())
}
fn transaction_ids(&mut self, values: &[TransactionId]) -> Result<()> {
self.usize(values.len())?;
for value in values {
self.transaction_id(*value);
}
Ok(())
}
fn writer_ids(&mut self, values: &[WriterId]) -> Result<()> {
self.usize(values.len())?;
for value in values {
self.writer_id(*value);
}
Ok(())
}
fn node_data(&mut self, value: &NodeData) -> Result<()> {
self.string(&value.short_name)?;
self.string(&value.short_description)?;
self.string(&value.long_description)?;
self.owner(value.owner);
self.node_ids(&value.fixed_connections)?;
self.node_ids(&value.recent_connections)?;
self.object_ids(&value.objects)
}
fn merge_pairs(&mut self, values: &[MergePair]) -> Result<()> {
self.usize(values.len())?;
for value in values {
self.transaction_id(value.first);
self.transaction_id(value.second);
}
Ok(())
}
}
pub(super) struct Decoder<'a> {
input: &'a [u8],
offset: usize,
}
impl<'a> Decoder<'a> {
fn new(input: &'a [u8]) -> Self {
Self { input, offset: 0 }
}
fn finish(&self) -> Result<()> {
if self.offset == self.input.len() {
Ok(())
} else {
Err(Error::corrupt("trailing bytes in canonical disk record"))
}
}
fn remaining(&self) -> usize {
self.input.len() - self.offset
}
fn take(&mut self, length: usize) -> Result<&'a [u8]> {
let end = self
.offset
.checked_add(length)
.ok_or_else(|| Error::corrupt("record field length overflow"))?;
if end > self.input.len() {
return Err(Error::corrupt("truncated canonical disk record"));
}
let bytes = &self.input[self.offset..end];
self.offset = end;
Ok(bytes)
}
fn array<const N: usize>(&mut self) -> Result<[u8; N]> {
self.take(N)?
.try_into()
.map_err(|_| Error::corrupt("invalid fixed-width disk field"))
}
fn u8(&mut self) -> Result<u8> {
Ok(self.array::<1>()?[0])
}
fn bool(&mut self) -> Result<bool> {
match self.u8()? {
0 => Ok(false),
1 => Ok(true),
_ => Err(Error::corrupt("invalid canonical boolean")),
}
}
fn u32(&mut self) -> Result<u32> {
Ok(u32::from_be_bytes(self.array()?))
}
fn u64(&mut self) -> Result<u64> {
Ok(u64::from_be_bytes(self.array()?))
}
fn i64(&mut self) -> Result<i64> {
Ok(i64::from_be_bytes(self.array()?))
}
fn count(&mut self, minimum_item_bytes: usize) -> Result<usize> {
let count = usize::try_from(self.u64()?)
.map_err(|_| Error::corrupt("record vector length does not fit usize"))?;
let minimum = count
.checked_mul(minimum_item_bytes)
.ok_or_else(|| Error::corrupt("record vector length overflow"))?;
if minimum > self.remaining() {
return Err(Error::corrupt(
"record vector cannot fit inside its enclosing record",
));
}
Ok(count)
}
fn string(&mut self) -> Result<String> {
let length = self.u32()? as usize;
let value = std::str::from_utf8(self.take(length)?)
.map_err(|_| Error::corrupt("record string is not UTF-8"))?;
Ok(value.to_owned())
}
fn datetime(&mut self) -> Result<DateTime<Utc>> {
DateTime::from_timestamp(self.i64()?, self.u32()?)
.ok_or_else(|| Error::corrupt("record timestamp is outside the supported range"))
}
fn node_id(&mut self) -> Result<NodeId> {
NodeId::from_bytes(self.array()?)
.map_err(|_| Error::corrupt("record contains an invalid node ID"))
}
fn object_id(&mut self) -> Result<ObjectId> {
ObjectId::from_bytes(self.array()?)
.map_err(|_| Error::corrupt("record contains an invalid object ID"))
}
fn transaction_id(&mut self) -> Result<TransactionId> {
Ok(TransactionId(self.array()?))
}
fn writer_id(&mut self) -> Result<WriterId> {
WriterId::from_verifying_key(self.array()?)
.map_err(|_| Error::corrupt("record contains an invalid writer ID"))
}
fn provenance(&mut self) -> Result<Provenance> {
Ok(Provenance {
author: self.string()?,
source: self.string()?,
source_created_at: self.datetime()?,
data: self.string()?,
})
}
fn owner(&mut self) -> Result<Owner> {
match self.u8()? {
0 => Ok(Owner::Unowned),
1 => Ok(Owner::SelfNode),
2 => Ok(Owner::Node(self.node_id()?)),
_ => Err(Error::corrupt("unknown canonical owner tag")),
}
}
fn node_ids(&mut self) -> Result<Vec<NodeId>> {
let count = self.count(6)?;
let mut values = Vec::with_capacity(count);
for _ in 0..count {
values.push(self.node_id()?);
}
Ok(values)
}
fn object_ids(&mut self) -> Result<Vec<ObjectId>> {
let count = self.count(6)?;
let mut values = Vec::with_capacity(count);
for _ in 0..count {
values.push(self.object_id()?);
}
Ok(values)
}
fn transaction_ids(&mut self) -> Result<Vec<TransactionId>> {
let count = self.count(32)?;
let mut values = Vec::with_capacity(count);
for _ in 0..count {
values.push(self.transaction_id()?);
}
Ok(values)
}
fn writer_ids(&mut self) -> Result<Vec<WriterId>> {
let count = self.count(32)?;
let mut values = Vec::with_capacity(count);
for _ in 0..count {
values.push(self.writer_id()?);
}
Ok(values)
}
fn node_data(&mut self) -> Result<NodeData> {
Ok(NodeData {
short_name: self.string()?,
short_description: self.string()?,
long_description: self.string()?,
owner: self.owner()?,
fixed_connections: self.node_ids()?,
recent_connections: self.node_ids()?,
objects: self.object_ids()?,
})
}
fn merge_pairs(&mut self) -> Result<Vec<MergePair>> {
let count = self.count(64)?;
let mut values = Vec::with_capacity(count);
for _ in 0..count {
values.push(MergePair {
first: self.transaction_id()?,
second: self.transaction_id()?,
});
}
Ok(values)
}
}
impl BinaryRecord for State {
fn encode(&self, encoder: &mut Encoder) -> Result<()> {
encoder.u32(self.format);
encoder.u64(self.generation);
encoder.u64(self.log_offset);
encoder.transaction_ids(&self.heads)?;
encoder.writer_ids(&self.writers_by_priority)
}
fn decode(decoder: &mut Decoder<'_>) -> Result<Self> {
Ok(Self {
format: decoder.u32()?,
generation: decoder.u64()?,
log_offset: decoder.u64()?,
heads: decoder.transaction_ids()?,
writers_by_priority: decoder.writer_ids()?,
})
}
}
impl Candidate {
fn encode_binary(&self, encoder: &mut Encoder) -> Result<()> {
encoder.transaction_id(self.transaction);
encoder.writer_id(self.writer);
encoder.datetime(self.committed_at);
encoder.provenance(&self.provenance)?;
encoder.node_data(&self.data)
}
fn decode_binary(decoder: &mut Decoder<'_>) -> Result<Self> {
Ok(Self {
transaction: decoder.transaction_id()?,
writer: decoder.writer_id()?,
committed_at: decoder.datetime()?,
provenance: decoder.provenance()?,
data: decoder.node_data()?,
})
}
}
impl BinaryRecord for NodeFile {
fn encode(&self, encoder: &mut Encoder) -> Result<()> {
encoder.u32(self.format);
encoder.node_id(self.id);
encoder.transaction_id(self.visible_transaction);
encoder.usize(self.frontier.len())?;
for candidate in &self.frontier {
candidate.encode_binary(encoder)?;
}
Ok(())
}
fn decode(decoder: &mut Decoder<'_>) -> Result<Self> {
let format = decoder.u32()?;
let id = decoder.node_id()?;
let visible_transaction = decoder.transaction_id()?;
let count = decoder.count(137)?;
let mut frontier = Vec::with_capacity(count);
for _ in 0..count {
frontier.push(Candidate::decode_binary(decoder)?);
}
let visible = frontier
.iter()
.find(|candidate| candidate.transaction == visible_transaction)
.ok_or_else(|| Error::corrupt("node record has no visible frontier candidate"))?;
let node = Node {
id,
data: visible.data.clone(),
last_author: visible.provenance.author.clone(),
committed_at: visible.committed_at,
};
Ok(Self {
format,
id,
node,
visible_transaction,
frontier,
})
}
}
impl BinaryRecord for HistoryIndex {
fn encode(&self, encoder: &mut Encoder) -> Result<()> {
encoder.u32(self.format);
encoder.node_id(self.node_id);
encoder.transaction_ids(&self.frontier)?;
match self.visible {
Some(value) => {
encoder.u8(1);
encoder.transaction_id(value);
}
None => encoder.u8(0),
}
encoder.transaction_ids(&self.creations)
}
fn decode(decoder: &mut Decoder<'_>) -> Result<Self> {
let format = decoder.u32()?;
let node_id = decoder.node_id()?;
let frontier = decoder.transaction_ids()?;
let visible = match decoder.u8()? {
0 => None,
1 => Some(decoder.transaction_id()?),
_ => return Err(Error::corrupt("invalid canonical optional transaction ID")),
};
let creations = decoder.transaction_ids()?;
Ok(Self {
format,
node_id,
frontier,
visible,
creations,
})
}
}
impl BinaryRecord for HistoryEntry {
fn encode(&self, encoder: &mut Encoder) -> Result<()> {
if !self.active || self.updated == self.created {
return Err(Error::corrupt(
"history entry status cannot be canonically encoded",
));
}
let data = self
.data
.as_ref()
.ok_or_else(|| Error::corrupt("history entry has no node data"))?;
encoder.transaction_id(self.transaction_id);
encoder.writer_id(self.writer);
encoder.datetime(self.committed_at);
encoder.provenance(&self.provenance)?;
encoder.bool(self.created);
encoder.node_data(data)?;
encoder.merge_pairs(&self.merge_pairs)
}
fn decode(decoder: &mut Decoder<'_>) -> Result<Self> {
let transaction_id = decoder.transaction_id()?;
let writer = decoder.writer_id()?;
let committed_at = decoder.datetime()?;
let provenance = decoder.provenance()?;
let created = decoder.bool()?;
let data = decoder.node_data()?;
Ok(Self {
transaction_id,
writer,
committed_at,
provenance,
active: true,
created,
updated: !created,
data: Some(data),
merge_pairs: decoder.merge_pairs()?,
})
}
}
impl BinaryRecord for TxMeta {
fn encode(&self, encoder: &mut Encoder) -> Result<()> {
encoder.u32(self.format);
encoder.transaction_id(self.id);
encoder.transaction_ids(&self.parents)?;
encoder.u64(self.dag_generation);
Ok(())
}
fn decode(decoder: &mut Decoder<'_>) -> Result<Self> {
Ok(Self {
format: decoder.u32()?,
id: decoder.transaction_id()?,
parents: decoder.transaction_ids()?,
dag_generation: decoder.u64()?,
})
}
}
impl BinaryRecord for QueueRecord {
fn encode(&self, encoder: &mut Encoder) -> Result<()> {
encoder.u32(self.format);
encoder.transaction_id(self.transaction);
Ok(())
}
fn decode(decoder: &mut Decoder<'_>) -> Result<Self> {
Ok(Self {
format: decoder.u32()?,
transaction: decoder.transaction_id()?,
})
}
}
impl FileMode {
fn encode_binary(self, encoder: &mut Encoder) {
encoder.u8(match self {
Self::CreateNew => 0,
Self::Replace => 1,
});
}
fn decode_binary(decoder: &mut Decoder<'_>) -> Result<Self> {
match decoder.u8()? {
0 => Ok(Self::CreateNew),
1 => Ok(Self::Replace),
_ => Err(Error::corrupt("unknown canonical WAL file mode")),
}
}
}
impl WalFile {
fn encode_binary(&self, encoder: &mut Encoder) -> Result<()> {
encoder.string(&self.staged)?;
encoder.string(&self.destination)?;
encoder.bytes(&self.sha256);
self.mode.encode_binary(encoder);
Ok(())
}
fn decode_binary(decoder: &mut Decoder<'_>) -> Result<Self> {
Ok(Self {
staged: decoder.string()?,
destination: decoder.string()?,
sha256: decoder.array()?,
mode: FileMode::decode_binary(decoder)?,
})
}
}
impl WalLog {
fn encode_binary(&self, encoder: &mut Encoder) -> Result<()> {
encoder.string(&self.staged)?;
encoder.transaction_id(self.transaction);
encoder.u64(self.length);
encoder.bytes(&self.sha256);
Ok(())
}
fn decode_binary(decoder: &mut Decoder<'_>) -> Result<Self> {
Ok(Self {
staged: decoder.string()?,
transaction: decoder.transaction_id()?,
length: decoder.u64()?,
sha256: decoder.array()?,
})
}
}
impl BinaryRecord for WalIntent {
fn encode(&self, encoder: &mut Encoder) -> Result<()> {
encoder.u32(self.format);
encoder.u64(self.expected_generation);
encoder.u64(self.expected_log_offset);
match &self.log {
Some(log) => {
encoder.u8(1);
log.encode_binary(encoder)?;
}
None => encoder.u8(0),
}
encoder.usize(self.files.len())?;
for file in &self.files {
file.encode_binary(encoder)?;
}
Ok(())
}
fn decode(decoder: &mut Decoder<'_>) -> Result<Self> {
let format = decoder.u32()?;
let expected_generation = decoder.u64()?;
let expected_log_offset = decoder.u64()?;
let log = match decoder.u8()? {
0 => None,
1 => Some(WalLog::decode_binary(decoder)?),
_ => return Err(Error::corrupt("invalid canonical optional WAL log")),
};
let file_count = decoder.count(4 + 4 + 32 + 1)?;
let mut files = Vec::new();
for _ in 0..file_count {
files.push(WalFile::decode_binary(decoder)?);
}
Ok(Self {
format,
expected_generation,
expected_log_offset,
log,
files,
})
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn state_encoding_is_canonical_and_round_trips() {
let key = [9; 32];
let value = State {
format: 3,
generation: 4,
log_offset: 5,
heads: vec![TransactionId::from_bytes([7; 32])],
writers_by_priority: vec![WriterId::from_signing_key(&key)],
};
let first = encode_record(b"KWSTATE3", &value).unwrap();
let second = encode_record(b"KWSTATE3", &value).unwrap();
assert_eq!(first, second);
let root = tempfile::tempdir().unwrap();
let path = root.path().join("state");
std::fs::write(&path, first).unwrap();
let decoded: State = read_record(&path, b"KWSTATE3").unwrap();
assert_eq!(decoded.generation, value.generation);
assert_eq!(decoded.log_offset, value.log_offset);
assert_eq!(decoded.heads, value.heads);
assert_eq!(decoded.writers_by_priority, value.writers_by_priority);
let mut corrupt = second;
corrupt[20] ^= 1;
std::fs::write(&path, corrupt).unwrap();
assert!(read_record::<State>(&path, b"KWSTATE3").is_err());
}
}