use std::io;
use crate::aae::tictac::{KeyEntry, Tree, TreeError};
pub const EXCHANGE_MAGIC: u32 = 0x7145_4145;
pub const MAX_PAYLOAD_LEN: u32 = 16 * 1024 * 1024;
pub const FRAME_HEADER_LEN: usize = 12;
pub const PHASE_ROOT_SYNC: u8 = 1;
pub const PHASE_TREE_SYNC: u8 = 2;
pub const PHASE_KEY_SYNC: u8 = 3;
#[derive(Debug, Clone, Eq, PartialEq)]
pub struct ExchangeFrame {
pub phase: u8,
pub reserved: u8,
pub flags: u16,
pub payload: Vec<u8>,
}
impl ExchangeFrame {
#[must_use]
pub fn new(phase: u8, payload: Vec<u8>) -> Self {
Self {
phase,
reserved: 0,
flags: 0,
payload,
}
}
pub fn to_wire(&self) -> Result<Vec<u8>, ExchangeError> {
let len = u32::try_from(self.payload.len())
.map_err(|_| ExchangeError::PayloadTooLarge(self.payload.len()))?;
if len > MAX_PAYLOAD_LEN {
return Err(ExchangeError::PayloadTooLarge(self.payload.len()));
}
let mut out = Vec::with_capacity(FRAME_HEADER_LEN + self.payload.len());
out.extend_from_slice(&EXCHANGE_MAGIC.to_be_bytes());
out.push(self.phase);
out.push(self.reserved);
out.extend_from_slice(&self.flags.to_be_bytes());
out.extend_from_slice(&len.to_be_bytes());
out.extend_from_slice(&self.payload);
Ok(out)
}
pub fn from_wire(bytes: &[u8]) -> Result<(Self, usize), ExchangeError> {
if bytes.len() < FRAME_HEADER_LEN {
return Err(ExchangeError::ShortFrame(bytes.len()));
}
let magic = u32::from_be_bytes(bytes[0..4].try_into().unwrap());
if magic != EXCHANGE_MAGIC {
return Err(ExchangeError::BadMagic(magic));
}
let phase = bytes[4];
let reserved = bytes[5];
let flags = u16::from_be_bytes(bytes[6..8].try_into().unwrap());
let payload_len = u32::from_be_bytes(bytes[8..12].try_into().unwrap());
if payload_len > MAX_PAYLOAD_LEN {
return Err(ExchangeError::PayloadTooLarge(payload_len as usize));
}
let total = FRAME_HEADER_LEN + payload_len as usize;
if bytes.len() < total {
return Err(ExchangeError::ShortFrame(bytes.len()));
}
let payload = bytes[FRAME_HEADER_LEN..total].to_vec();
Ok((
Self {
phase,
reserved,
flags,
payload,
},
total,
))
}
}
#[must_use]
pub fn encode_root_sync(roots: &[(u32, u64)]) -> Vec<u8> {
let mut out = Vec::with_capacity(4 + roots.len() * 12);
let n = u32::try_from(roots.len()).unwrap_or(u32::MAX);
out.extend_from_slice(&n.to_be_bytes());
for (idx, hash) in roots {
out.extend_from_slice(&idx.to_be_bytes());
out.extend_from_slice(&hash.to_be_bytes());
}
out
}
pub fn decode_root_sync(bytes: &[u8]) -> Result<Vec<(u32, u64)>, ExchangeError> {
if bytes.len() < 4 {
return Err(ExchangeError::BadPayload("root_sync truncated".to_string()));
}
let n = u32::from_be_bytes(bytes[0..4].try_into().unwrap()) as usize;
let needed = 4 + n * 12;
if bytes.len() < needed {
return Err(ExchangeError::BadPayload(format!(
"root_sync expected {needed} bytes, got {}",
bytes.len()
)));
}
let mut out = Vec::with_capacity(n);
let mut off = 4;
for _ in 0..n {
let idx = u32::from_be_bytes(bytes[off..off + 4].try_into().unwrap());
let hash = u64::from_be_bytes(bytes[off + 4..off + 12].try_into().unwrap());
out.push((idx, hash));
off += 12;
}
Ok(out)
}
#[must_use]
pub fn encode_tree_sync(time_bucket: u32, segments: &[(u32, u64)]) -> Vec<u8> {
let mut out = Vec::with_capacity(8 + segments.len() * 12);
out.extend_from_slice(&time_bucket.to_be_bytes());
let n = u32::try_from(segments.len()).unwrap_or(u32::MAX);
out.extend_from_slice(&n.to_be_bytes());
for (idx, hash) in segments {
out.extend_from_slice(&idx.to_be_bytes());
out.extend_from_slice(&hash.to_be_bytes());
}
out
}
pub fn decode_tree_sync(bytes: &[u8]) -> Result<(u32, Vec<(u32, u64)>), ExchangeError> {
if bytes.len() < 8 {
return Err(ExchangeError::BadPayload("tree_sync truncated".to_string()));
}
let tb = u32::from_be_bytes(bytes[0..4].try_into().unwrap());
let n = u32::from_be_bytes(bytes[4..8].try_into().unwrap()) as usize;
let needed = 8 + n * 12;
if bytes.len() < needed {
return Err(ExchangeError::BadPayload(format!(
"tree_sync expected {needed} bytes, got {}",
bytes.len()
)));
}
let mut out = Vec::with_capacity(n);
let mut off = 8;
for _ in 0..n {
let idx = u32::from_be_bytes(bytes[off..off + 4].try_into().unwrap());
let hash = u64::from_be_bytes(bytes[off + 4..off + 12].try_into().unwrap());
out.push((idx, hash));
off += 12;
}
Ok((tb, out))
}
#[must_use]
pub fn encode_key_sync(time_bucket: u32, segment: u32, entries: &[KeyEntry]) -> Vec<u8> {
let mut out = Vec::new();
out.extend_from_slice(&time_bucket.to_be_bytes());
out.extend_from_slice(&segment.to_be_bytes());
let n = u32::try_from(entries.len()).unwrap_or(u32::MAX);
out.extend_from_slice(&n.to_be_bytes());
for entry in entries {
write_lp(&mut out, &entry.bucket);
write_lp(&mut out, &entry.key);
write_lp(&mut out, &entry.vclock);
}
out
}
fn write_lp(out: &mut Vec<u8>, bytes: &[u8]) {
let n = u32::try_from(bytes.len()).unwrap_or(u32::MAX);
out.extend_from_slice(&n.to_be_bytes());
out.extend_from_slice(bytes);
}
fn read_lp(bytes: &[u8], off: &mut usize) -> Result<Vec<u8>, ExchangeError> {
if bytes.len() < *off + 4 {
return Err(ExchangeError::BadPayload(
"key_sync length-prefix truncated".to_string(),
));
}
let n = u32::from_be_bytes(bytes[*off..*off + 4].try_into().unwrap()) as usize;
*off += 4;
if bytes.len() < *off + n {
return Err(ExchangeError::BadPayload(format!(
"key_sync field expected {n} bytes, only {} remain",
bytes.len() - *off
)));
}
let v = bytes[*off..*off + n].to_vec();
*off += n;
Ok(v)
}
pub fn decode_key_sync(bytes: &[u8]) -> Result<(u32, u32, Vec<KeyEntry>), ExchangeError> {
if bytes.len() < 12 {
return Err(ExchangeError::BadPayload("key_sync truncated".to_string()));
}
let tb = u32::from_be_bytes(bytes[0..4].try_into().unwrap());
let seg = u32::from_be_bytes(bytes[4..8].try_into().unwrap());
let n = u32::from_be_bytes(bytes[8..12].try_into().unwrap()) as usize;
let mut off = 12;
let mut entries = Vec::with_capacity(n);
for _ in 0..n {
let bucket = read_lp(bytes, &mut off)?;
let key = read_lp(bytes, &mut off)?;
let vclock = read_lp(bytes, &mut off)?;
entries.push(KeyEntry {
bucket,
key,
vclock,
});
}
Ok((tb, seg, entries))
}
pub struct Exchange<'a, V: PeerView> {
local: &'a Tree,
remote: V,
}
pub trait PeerView {
fn roots(&self) -> Result<Vec<(u32, u64)>, ExchangeError>;
fn segments(&self, time_bucket: u32) -> Result<Vec<(u32, u64)>, ExchangeError>;
fn keys_in_segment(
&self,
time_bucket: u32,
segment: u32,
) -> Result<Vec<KeyEntry>, ExchangeError>;
}
#[derive(Debug, Clone, Copy)]
pub struct LocalPeerView<'a> {
tree: &'a Tree,
}
impl<'a> LocalPeerView<'a> {
#[must_use]
pub fn new(tree: &'a Tree) -> Self {
Self { tree }
}
}
impl PeerView for LocalPeerView<'_> {
fn roots(&self) -> Result<Vec<(u32, u64)>, ExchangeError> {
Ok(self.tree.roots())
}
fn segments(&self, time_bucket: u32) -> Result<Vec<(u32, u64)>, ExchangeError> {
self.tree.segments(time_bucket).map_err(ExchangeError::from)
}
fn keys_in_segment(
&self,
time_bucket: u32,
segment: u32,
) -> Result<Vec<KeyEntry>, ExchangeError> {
self.tree
.keys_in_segment(time_bucket, segment)
.map_err(ExchangeError::from)
}
}
#[derive(Debug, Clone, Eq, PartialEq)]
pub struct Divergence {
pub time_bucket: u32,
pub segment: u32,
pub local_only: Vec<KeyEntry>,
pub remote_only: Vec<KeyEntry>,
}
impl<'a, V: PeerView> Exchange<'a, V> {
#[must_use]
pub fn new(local: &'a Tree, remote: V) -> Self {
Self { local, remote }
}
pub fn run(&self) -> Result<Vec<Divergence>, ExchangeError> {
let local_roots = self.local.roots();
let remote_roots = self.remote.roots()?;
let dr = Tree::diverging_time_buckets(&local_roots, &remote_roots);
let mut out = Vec::new();
for tb in dr {
let local_segs = self.local.segments(tb)?;
let remote_segs = self.remote.segments(tb)?;
let ds = Tree::diverging_segments(&local_segs, &remote_segs);
for seg in ds {
let local_keys = self.local.keys_in_segment(tb, seg)?;
let remote_keys = self.remote.keys_in_segment(tb, seg)?;
let (local_only, remote_only) = symmetric_difference(&local_keys, &remote_keys);
if local_only.is_empty() && remote_only.is_empty() {
continue;
}
out.push(Divergence {
time_bucket: tb,
segment: seg,
local_only,
remote_only,
});
}
}
Ok(out)
}
}
fn symmetric_difference(a: &[KeyEntry], b: &[KeyEntry]) -> (Vec<KeyEntry>, Vec<KeyEntry>) {
let bset: std::collections::BTreeSet<&KeyEntry> = b.iter().collect();
let aset: std::collections::BTreeSet<&KeyEntry> = a.iter().collect();
let only_a = a.iter().filter(|e| !bset.contains(e)).cloned().collect();
let only_b = b.iter().filter(|e| !aset.contains(e)).cloned().collect();
(only_a, only_b)
}
#[derive(Debug, thiserror::Error)]
pub enum ExchangeError {
#[error("short exchange frame ({0} bytes)")]
ShortFrame(usize),
#[error("bad exchange magic 0x{0:08x}")]
BadMagic(u32),
#[error("exchange payload too large ({0} bytes)")]
PayloadTooLarge(usize),
#[error("exchange payload: {0}")]
BadPayload(String),
#[error("exchange tree: {0}")]
Tree(#[from] TreeError),
#[error("exchange io: {0}")]
Io(#[from] io::Error),
}
#[cfg(test)]
mod tests {
use super::*;
use crate::aae::tictac::TreeShape;
fn shape() -> TreeShape {
TreeShape {
n_time_buckets: 4,
n_segments: 32,
time_window_seconds: 60,
}
}
#[test]
fn frame_round_trips() {
let f = ExchangeFrame::new(PHASE_ROOT_SYNC, vec![1, 2, 3, 4]);
let wire = f.to_wire().unwrap();
let (parsed, n) = ExchangeFrame::from_wire(&wire).unwrap();
assert_eq!(n, wire.len());
assert_eq!(parsed, f);
}
#[test]
fn frame_rejects_bad_magic() {
let mut wire = ExchangeFrame::new(PHASE_ROOT_SYNC, vec![])
.to_wire()
.unwrap();
wire[0..4].copy_from_slice(&0xdead_beefu32.to_be_bytes());
assert!(matches!(
ExchangeFrame::from_wire(&wire),
Err(ExchangeError::BadMagic(_))
));
}
#[test]
fn root_sync_round_trips() {
let v = vec![(0u32, 0xaaaa_bbbb_ccccu64), (1, 0xdead), (2, 0)];
let wire = encode_root_sync(&v);
assert_eq!(decode_root_sync(&wire).unwrap(), v);
}
#[test]
fn tree_sync_round_trips() {
let v = vec![(0u32, 1u64), (5, 99)];
let wire = encode_tree_sync(7, &v);
let (tb, parsed) = decode_tree_sync(&wire).unwrap();
assert_eq!(tb, 7);
assert_eq!(parsed, v);
}
#[test]
fn key_sync_round_trips() {
let entries = vec![
KeyEntry {
bucket: b"users".to_vec(),
key: b"alice".to_vec(),
vclock: b"vc1".to_vec(),
},
KeyEntry {
bucket: b"posts".to_vec(),
key: b"42".to_vec(),
vclock: b"vc99".to_vec(),
},
];
let wire = encode_key_sync(3, 17, &entries);
let (tb, seg, parsed) = decode_key_sync(&wire).unwrap();
assert_eq!(tb, 3);
assert_eq!(seg, 17);
assert_eq!(parsed, entries);
}
#[test]
fn exchange_in_memory_surfaces_divergent_key() {
let mut a = Tree::new(shape());
let mut b = Tree::new(shape());
for i in 0..200u32 {
let k = format!("k{i}");
a.insert(b"users", k.as_bytes(), b"vc1", 0);
b.insert(b"users", k.as_bytes(), b"vc1", 0);
}
b.update(b"users", b"k7", b"vc1", b"vc2", 0, 0);
let view = LocalPeerView::new(&b);
let ex = Exchange::new(&a, view);
let divs = ex.run().unwrap();
assert!((1..=2).contains(&divs.len()), "unexpected divs: {divs:?}");
let local_keys: Vec<_> = divs.iter().flat_map(|d| d.local_only.iter()).collect();
let remote_keys: Vec<_> = divs.iter().flat_map(|d| d.remote_only.iter()).collect();
assert!(local_keys
.iter()
.any(|e| e.key == b"k7" && e.vclock == b"vc1"));
assert!(remote_keys
.iter()
.any(|e| e.key == b"k7" && e.vclock == b"vc2"));
}
}