use std::collections::{BTreeMap, BTreeSet};
use crate::datatypes::set::{OrSet, Tag};
use crate::datatypes::{ActorId, PnCounter};
const FORMAT_V1: u8 = 1;
pub const TAG_COUNTER: u8 = 1;
pub const TAG_SET: u8 = 2;
#[derive(Clone, Debug, Eq, PartialEq, thiserror::Error)]
#[non_exhaustive]
pub enum CrdtSerialError {
#[error("crdt serial: truncated payload")]
Truncated,
#[error("crdt serial: unsupported format version {0}")]
BadVersion(u8),
#[error("crdt serial: type tag {found} does not match expected {expected}")]
TypeMismatch {
found: u8,
expected: u8,
},
#[error("crdt serial: unknown type tag {0}")]
UnknownTag(u8),
#[error("crdt serial: {0} trailing bytes")]
Trailing(usize),
}
fn put_u64(out: &mut Vec<u8>, v: u64) {
out.extend_from_slice(&v.to_be_bytes());
}
fn put_bytes(out: &mut Vec<u8>, b: &[u8]) {
put_u64(out, b.len() as u64);
out.extend_from_slice(b);
}
fn put_actor(out: &mut Vec<u8>, a: &ActorId) {
put_bytes(out, a.dc.as_bytes());
put_bytes(out, a.peer.as_bytes());
}
struct Reader<'a> {
buf: &'a [u8],
pos: usize,
}
impl<'a> Reader<'a> {
fn new(buf: &'a [u8]) -> Self {
Self { buf, pos: 0 }
}
fn u8(&mut self) -> Result<u8, CrdtSerialError> {
let b = *self.buf.get(self.pos).ok_or(CrdtSerialError::Truncated)?;
self.pos += 1;
Ok(b)
}
fn u64(&mut self) -> Result<u64, CrdtSerialError> {
let end = self.pos + 8;
let slice = self
.buf
.get(self.pos..end)
.ok_or(CrdtSerialError::Truncated)?;
let mut a = [0u8; 8];
a.copy_from_slice(slice);
self.pos = end;
Ok(u64::from_be_bytes(a))
}
fn bytes(&mut self) -> Result<Vec<u8>, CrdtSerialError> {
let n = usize::try_from(self.u64()?).map_err(|_| CrdtSerialError::Truncated)?;
let end = self.pos + n;
let slice = self
.buf
.get(self.pos..end)
.ok_or(CrdtSerialError::Truncated)?;
self.pos = end;
Ok(slice.to_vec())
}
fn string(&mut self) -> Result<String, CrdtSerialError> {
String::from_utf8(self.bytes()?).map_err(|_| CrdtSerialError::Truncated)
}
fn actor(&mut self) -> Result<ActorId, CrdtSerialError> {
let dc = self.string()?;
let peer = self.string()?;
Ok(ActorId::new(dc, peer))
}
fn done(&self) -> Result<(), CrdtSerialError> {
let rem = self.buf.len() - self.pos;
if rem == 0 {
Ok(())
} else {
Err(CrdtSerialError::Trailing(rem))
}
}
}
#[must_use]
pub fn counter_to_bytes(c: &PnCounter) -> Vec<u8> {
let mut out = Vec::with_capacity(32);
out.push(FORMAT_V1);
out.push(TAG_COUNTER);
let (pos, neg) = c.columns();
put_u64(&mut out, pos.len() as u64);
for (actor, n) in pos {
put_actor(&mut out, actor);
put_u64(&mut out, *n);
}
put_u64(&mut out, neg.len() as u64);
for (actor, n) in neg {
put_actor(&mut out, actor);
put_u64(&mut out, *n);
}
out
}
pub fn counter_from_bytes(buf: &[u8]) -> Result<PnCounter, CrdtSerialError> {
let mut r = Reader::new(buf);
check_header(&mut r, TAG_COUNTER)?;
let mut pos = BTreeMap::new();
let np = r.u64()?;
for _ in 0..np {
let a = r.actor()?;
let n = r.u64()?;
pos.insert(a, n);
}
let mut neg = BTreeMap::new();
let nn = r.u64()?;
for _ in 0..nn {
let a = r.actor()?;
let n = r.u64()?;
neg.insert(a, n);
}
r.done()?;
Ok(PnCounter::from_columns(pos, neg))
}
#[must_use]
pub fn set_to_bytes(s: &OrSet) -> Vec<u8> {
let mut out = Vec::with_capacity(64);
out.push(FORMAT_V1);
out.push(TAG_SET);
let elements = s.raw_elements();
put_u64(&mut out, elements.len() as u64);
for (element, (adds, removes)) in elements {
put_bytes(&mut out, &element);
put_tags(&mut out, &adds);
put_tags(&mut out, &removes);
}
let counters = s.raw_actor_counters();
put_u64(&mut out, counters.len() as u64);
for (actor, n) in counters {
put_actor(&mut out, &actor);
put_u64(&mut out, n);
}
out
}
fn put_tags(out: &mut Vec<u8>, tags: &BTreeSet<Tag>) {
put_u64(out, tags.len() as u64);
for t in tags {
put_actor(out, &t.actor);
put_u64(out, t.counter);
}
}
fn read_tags(r: &mut Reader<'_>) -> Result<BTreeSet<Tag>, CrdtSerialError> {
let n = r.u64()?;
let mut set = BTreeSet::new();
for _ in 0..n {
let actor = r.actor()?;
let counter = r.u64()?;
set.insert(Tag { actor, counter });
}
Ok(set)
}
pub fn set_from_bytes(buf: &[u8]) -> Result<OrSet, CrdtSerialError> {
let mut r = Reader::new(buf);
check_header(&mut r, TAG_SET)?;
let ne = r.u64()?;
let mut elements: BTreeMap<Vec<u8>, (BTreeSet<Tag>, BTreeSet<Tag>)> = BTreeMap::new();
for _ in 0..ne {
let element = r.bytes()?;
let adds = read_tags(&mut r)?;
let removes = read_tags(&mut r)?;
elements.insert(element, (adds, removes));
}
let nc = r.u64()?;
let mut counters = BTreeMap::new();
for _ in 0..nc {
let actor = r.actor()?;
let n = r.u64()?;
counters.insert(actor, n);
}
r.done()?;
Ok(OrSet::from_raw(elements, counters))
}
pub fn peek_tag(buf: &[u8]) -> Result<u8, CrdtSerialError> {
let mut r = Reader::new(buf);
let version = r.u8()?;
if version != FORMAT_V1 {
return Err(CrdtSerialError::BadVersion(version));
}
r.u8()
}
fn check_header(r: &mut Reader<'_>, expected_tag: u8) -> Result<(), CrdtSerialError> {
let version = r.u8()?;
if version != FORMAT_V1 {
return Err(CrdtSerialError::BadVersion(version));
}
let tag = r.u8()?;
if tag == expected_tag {
Ok(())
} else if tag == TAG_COUNTER || tag == TAG_SET {
Err(CrdtSerialError::TypeMismatch {
found: tag,
expected: expected_tag,
})
} else {
Err(CrdtSerialError::UnknownTag(tag))
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::datatypes::Crdt;
fn aid(peer: &str) -> ActorId {
ActorId::new("dc1", peer)
}
#[test]
fn counter_round_trips() {
let mut c = PnCounter::new();
c.increment(&aid("a"), 5);
c.increment(&aid("b"), 3);
c.decrement(&aid("a"), 2);
let bytes = counter_to_bytes(&c);
assert_eq!(peek_tag(&bytes).unwrap(), TAG_COUNTER);
let back = counter_from_bytes(&bytes).unwrap();
assert_eq!(back, c);
assert_eq!(back.value(), c.value());
}
#[test]
fn counter_merge_after_round_trip_sums() {
let mut a = PnCounter::new();
a.increment(&aid("a"), 1);
let mut b = PnCounter::new();
b.increment(&aid("b"), 1);
let mut a2 = counter_from_bytes(&counter_to_bytes(&a)).unwrap();
let b2 = counter_from_bytes(&counter_to_bytes(&b)).unwrap();
a2.merge(&b2);
assert_eq!(a2.value(), 2);
a2.merge(&b2);
assert_eq!(a2.value(), 2);
}
#[test]
fn set_round_trips_and_merges_to_union() {
let mut s = OrSet::new();
s.add(&aid("a"), b"x".to_vec());
let mut t = OrSet::new();
t.add(&aid("b"), b"y".to_vec());
let mut s2 = set_from_bytes(&set_to_bytes(&s)).unwrap();
let t2 = set_from_bytes(&set_to_bytes(&t)).unwrap();
s2.merge(&t2);
let v = s2.value();
assert!(v.contains(b"x".as_slice()));
assert!(v.contains(b"y".as_slice()));
}
#[test]
fn type_mismatch_is_rejected() {
let mut c = PnCounter::new();
c.increment(&aid("a"), 1);
let bytes = counter_to_bytes(&c);
let err = set_from_bytes(&bytes).unwrap_err();
assert!(matches!(
err,
CrdtSerialError::TypeMismatch {
found: TAG_COUNTER,
expected: TAG_SET
}
));
}
#[test]
fn truncated_is_rejected() {
let mut c = PnCounter::new();
c.increment(&aid("a"), 1);
let bytes = counter_to_bytes(&c);
assert!(counter_from_bytes(&bytes[..bytes.len() - 3]).is_err());
}
}