use std::collections::BTreeMap;
use std::sync::atomic::{AtomicU64, Ordering};
use serde::{Deserialize, Serialize};
use crate::datastore::noxu::{NoxuDatastore, NoxuDatastoreError};
use crate::ramp::{select, RampClock, RampItem, Timestamp};
const VER_TAG: &[u8] = b"V\0";
const PTR_TAG: &[u8] = b"L\0";
pub type RampReadResult = (BTreeMap<Vec<u8>, Vec<u8>>, u8);
#[derive(Debug, thiserror::Error)]
#[non_exhaustive]
pub enum RampError {
#[error("ramp store: {0}")]
Store(String),
#[error("ramp: dangling version pointer for key {key:?} at ts {ts}")]
DanglingVersion {
key: Vec<u8>,
ts: Timestamp,
},
#[error("ramp: corrupt version record for key {key:?} at ts {ts}")]
CorruptVersion {
key: Vec<u8>,
ts: Timestamp,
},
#[error("ramp: empty write set")]
EmptyWrite,
}
impl From<NoxuDatastoreError> for RampError {
fn from(e: NoxuDatastoreError) -> Self {
Self::Store(e.to_string())
}
}
pub trait RampStore {
fn put_version(&self, item: &RampItem) -> Result<(), RampError>;
fn commit_pointer(&self, key: &[u8], ts: Timestamp) -> Result<(), RampError>;
fn latest_visible(&self, key: &[u8]) -> Result<Option<Timestamp>, RampError>;
fn get_version(&self, key: &[u8], ts: Timestamp) -> Result<Option<RampItem>, RampError>;
}
fn version_key(key: &[u8], ts: Timestamp) -> Vec<u8> {
let mut out = Vec::with_capacity(VER_TAG.len() + key.len() + 1 + 8);
out.extend_from_slice(VER_TAG);
out.extend_from_slice(key);
out.push(0);
out.extend_from_slice(&ts.to_be_bytes());
out
}
fn pointer_key(key: &[u8]) -> Vec<u8> {
let mut out = Vec::with_capacity(PTR_TAG.len() + key.len());
out.extend_from_slice(PTR_TAG);
out.extend_from_slice(key);
out
}
fn encode_body(item: &RampItem) -> Vec<u8> {
let mut out = Vec::new();
let count = u32::try_from(item.siblings.len()).unwrap_or(u32::MAX);
out.extend_from_slice(&count.to_be_bytes());
for sib in &item.siblings {
let len = u32::try_from(sib.len()).unwrap_or(u32::MAX);
out.extend_from_slice(&len.to_be_bytes());
out.extend_from_slice(sib);
}
out.extend_from_slice(&item.value);
out
}
fn decode_body(bytes: &[u8]) -> Option<(Vec<Vec<u8>>, Vec<u8>)> {
let mut pos = 0usize;
let count = read_u32(bytes, &mut pos)? as usize;
let mut siblings = Vec::with_capacity(count);
for _ in 0..count {
let len = read_u32(bytes, &mut pos)? as usize;
let end = pos.checked_add(len)?;
if end > bytes.len() {
return None;
}
siblings.push(bytes[pos..end].to_vec());
pos = end;
}
Some((siblings, bytes[pos..].to_vec()))
}
fn read_u32(bytes: &[u8], pos: &mut usize) -> Option<u32> {
let end = pos.checked_add(4)?;
if end > bytes.len() {
return None;
}
let v = u32::from_be_bytes([
bytes[*pos],
bytes[*pos + 1],
bytes[*pos + 2],
bytes[*pos + 3],
]);
*pos = end;
Some(v)
}
impl RampStore for NoxuDatastore {
fn put_version(&self, item: &RampItem) -> Result<(), RampError> {
let vk = version_key(&item.key, item.ts);
self.put(&vk, &encode_body(item))?;
Ok(())
}
fn commit_pointer(&self, key: &[u8], ts: Timestamp) -> Result<(), RampError> {
let pk = pointer_key(key);
if let Some(cur) = self.get(&pk)? {
if cur.len() == 8 {
let mut buf = [0u8; 8];
buf.copy_from_slice(&cur);
if u64::from_be_bytes(buf) >= ts {
return Ok(());
}
}
}
self.put(&pk, &ts.to_be_bytes())?;
Ok(())
}
fn latest_visible(&self, key: &[u8]) -> Result<Option<Timestamp>, RampError> {
let pk = pointer_key(key);
match self.get(&pk)? {
Some(v) if v.len() == 8 => {
let mut buf = [0u8; 8];
buf.copy_from_slice(&v);
Ok(Some(u64::from_be_bytes(buf)))
}
_ => Ok(None),
}
}
fn get_version(&self, key: &[u8], ts: Timestamp) -> Result<Option<RampItem>, RampError> {
let vk = version_key(key, ts);
match self.get(&vk)? {
Some(body) => {
let (siblings, value) =
decode_body(&body).ok_or_else(|| RampError::CorruptVersion {
key: key.to_vec(),
ts,
})?;
Ok(Some(RampItem::new(key.to_vec(), ts, siblings, value)))
}
None => Ok(None),
}
}
}
#[derive(Clone, Debug, Eq, PartialEq)]
pub struct RampWrite {
pub key: Vec<u8>,
pub value: Vec<u8>,
}
pub struct RampCoordinator<S: RampStore> {
store: S,
clock: RampClock,
}
impl<S: RampStore> RampCoordinator<S> {
pub fn new(store: S, coordinator_id: u16) -> Self {
Self {
store,
clock: RampClock::new(coordinator_id),
}
}
pub fn store(&self) -> &S {
&self.store
}
pub fn write(&mut self, writes: &[RampWrite]) -> Result<Timestamp, RampError> {
if writes.is_empty() {
return Err(RampError::EmptyWrite);
}
let ts = self.clock.mint();
let all_keys: Vec<Vec<u8>> = writes.iter().map(|w| w.key.clone()).collect();
for w in writes {
let siblings: Vec<Vec<u8>> =
all_keys.iter().filter(|k| *k != &w.key).cloned().collect();
let item = RampItem::new(w.key.clone(), ts, siblings, w.value.clone());
self.store.put_version(&item)?;
}
for w in writes {
self.store.commit_pointer(&w.key, ts)?;
}
Ok(ts)
}
pub fn read(&self, keys: &[Vec<u8>]) -> Result<BTreeMap<Vec<u8>, Vec<u8>>, RampError> {
let (snapshot, _rounds) = self.read_with_rounds(keys)?;
Ok(snapshot)
}
pub fn read_with_rounds(&self, keys: &[Vec<u8>]) -> Result<RampReadResult, RampError> {
read_rounds(&self.store, keys)
}
}
static HTTP_COUNTER: AtomicU64 = AtomicU64::new(0);
fn next_http_ts(id: u16) -> Timestamp {
let n = HTTP_COUNTER.fetch_add(1, Ordering::Relaxed) + 1;
(u64::from(id) << 48) | (n & 0x0000_ffff_ffff_ffff)
}
pub fn ramp_write<S: RampStore>(
store: &S,
coordinator_id: u16,
writes: &[RampWrite],
) -> Result<Timestamp, RampError> {
if writes.is_empty() {
return Err(RampError::EmptyWrite);
}
let ts = next_http_ts(coordinator_id);
let all_keys: Vec<Vec<u8>> = writes.iter().map(|w| w.key.clone()).collect();
for w in writes {
let siblings: Vec<Vec<u8>> = all_keys.iter().filter(|k| *k != &w.key).cloned().collect();
let item = RampItem::new(w.key.clone(), ts, siblings, w.value.clone());
store.put_version(&item)?;
}
for w in writes {
store.commit_pointer(&w.key, ts)?;
}
Ok(ts)
}
pub fn ramp_read<S: RampStore>(store: &S, keys: &[Vec<u8>]) -> Result<RampReadResult, RampError> {
read_rounds(store, keys)
}
fn read_rounds<S: RampStore>(store: &S, keys: &[Vec<u8>]) -> Result<RampReadResult, RampError> {
let mut round1: Vec<RampItem> = Vec::with_capacity(keys.len());
for key in keys {
if let Some(ts) = store.latest_visible(key)? {
let item = store
.get_version(key, ts)?
.ok_or_else(|| RampError::DanglingVersion {
key: key.clone(),
ts,
})?;
round1.push(item);
}
}
let missing = select(&round1);
let two_rounds = if missing.is_empty() { 1 } else { 2 };
let mut chosen: BTreeMap<Vec<u8>, RampItem> = BTreeMap::new();
for item in round1 {
chosen.insert(item.key.clone(), item);
}
for (key, ts) in missing {
let item = store
.get_version(&key, ts)?
.ok_or_else(|| RampError::DanglingVersion {
key: key.clone(),
ts,
})?;
chosen.insert(key, item);
}
let snapshot = chosen
.into_iter()
.map(|(k, item)| (k, item.value))
.collect();
Ok((snapshot, two_rounds))
}
#[derive(Clone, Debug, Deserialize)]
pub struct HttpRampWriteRequest {
pub writes: Vec<HttpRampWrite>,
}
#[derive(Clone, Debug, Deserialize)]
pub struct HttpRampWrite {
pub key: String,
pub value: String,
}
#[derive(Clone, Debug, Serialize)]
pub struct HttpRampWriteResponse {
pub result: String,
pub ts: Timestamp,
pub keys: usize,
}
#[derive(Clone, Debug, Deserialize)]
pub struct HttpRampReadRequest {
pub keys: Vec<String>,
}
#[derive(Clone, Debug, Serialize)]
pub struct HttpRampReadResponse {
pub snapshot: BTreeMap<String, String>,
pub rounds: u8,
}
impl HttpRampWriteRequest {
#[must_use]
pub fn into_writes(self) -> Vec<RampWrite> {
self.writes
.into_iter()
.map(|w| RampWrite {
key: w.key.into_bytes(),
value: w.value.into_bytes(),
})
.collect()
}
}
impl HttpRampReadRequest {
#[must_use]
pub fn into_keys(self) -> Vec<Vec<u8>> {
self.keys.into_iter().map(String::into_bytes).collect()
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::cell::RefCell;
use std::collections::HashMap;
type MemVersion = (Vec<Vec<u8>>, Vec<u8>);
#[derive(Default)]
struct MemStore {
versions: RefCell<HashMap<(Vec<u8>, Timestamp), MemVersion>>,
pointers: RefCell<HashMap<Vec<u8>, Timestamp>>,
}
impl RampStore for MemStore {
fn put_version(&self, item: &RampItem) -> Result<(), RampError> {
self.versions.borrow_mut().insert(
(item.key.clone(), item.ts),
(item.siblings.clone(), item.value.clone()),
);
Ok(())
}
fn commit_pointer(&self, key: &[u8], ts: Timestamp) -> Result<(), RampError> {
let mut p = self.pointers.borrow_mut();
let e = p.entry(key.to_vec()).or_insert(ts);
if ts > *e {
*e = ts;
}
Ok(())
}
fn latest_visible(&self, key: &[u8]) -> Result<Option<Timestamp>, RampError> {
Ok(self.pointers.borrow().get(key).copied())
}
fn get_version(&self, key: &[u8], ts: Timestamp) -> Result<Option<RampItem>, RampError> {
Ok(self
.versions
.borrow()
.get(&(key.to_vec(), ts))
.map(|(sib, val)| RampItem::new(key.to_vec(), ts, sib.clone(), val.clone())))
}
}
fn w(key: &[u8], value: &[u8]) -> RampWrite {
RampWrite {
key: key.to_vec(),
value: value.to_vec(),
}
}
#[test]
fn body_round_trips() {
let item = RampItem::new(
b"a".to_vec(),
9,
vec![b"b".to_vec(), b"cc".to_vec()],
b"hello".to_vec(),
);
let bytes = encode_body(&item);
let (sibs, val) = decode_body(&bytes).expect("decode");
assert_eq!(sibs, item.siblings);
assert_eq!(val, item.value);
}
#[test]
fn decode_body_rejects_truncation() {
assert!(decode_body(&[0, 0, 0]).is_none());
assert!(decode_body(&[0, 0, 0, 1, 0, 0, 0, 5]).is_none());
}
#[test]
fn write_then_read_is_atomic() {
let mut c = RampCoordinator::new(MemStore::default(), 0);
c.write(&[w(b"a", b"1"), w(b"b", b"2")]).expect("write");
let snap = c.read(&[b"a".to_vec(), b"b".to_vec()]).expect("read");
assert_eq!(snap.get(b"a".as_slice()), Some(&b"1".to_vec()));
assert_eq!(snap.get(b"b".as_slice()), Some(&b"2".to_vec()));
}
#[test]
fn contention_free_read_is_one_round() {
let mut c = RampCoordinator::new(MemStore::default(), 0);
c.write(&[w(b"a", b"1"), w(b"b", b"2")]).expect("write");
let (_snap, rounds) = c
.read_with_rounds(&[b"a".to_vec(), b"b".to_vec()])
.expect("read");
assert_eq!(rounds, 1, "no concurrency -> single round");
}
#[test]
fn partial_write_triggers_second_round_repair() {
let store = MemStore::default();
store
.put_version(&RampItem::new(b"a".to_vec(), 1, vec![], b"a-old".to_vec()))
.unwrap();
store
.put_version(&RampItem::new(b"b".to_vec(), 1, vec![], b"b-old".to_vec()))
.unwrap();
store.commit_pointer(b"a", 1).unwrap();
store.commit_pointer(b"b", 1).unwrap();
store
.put_version(&RampItem::new(
b"a".to_vec(),
5,
vec![b"b".to_vec()],
b"a-new".to_vec(),
))
.unwrap();
store
.put_version(&RampItem::new(
b"b".to_vec(),
5,
vec![b"a".to_vec()],
b"b-new".to_vec(),
))
.unwrap();
store.commit_pointer(b"a", 5).unwrap();
let c = RampCoordinator::new(store, 0);
let (snap, rounds) = c
.read_with_rounds(&[b"a".to_vec(), b"b".to_vec()])
.expect("read");
assert_eq!(rounds, 2, "partial apply must force the second round");
assert_eq!(snap.get(b"a".as_slice()), Some(&b"a-new".to_vec()));
assert_eq!(snap.get(b"b".as_slice()), Some(&b"b-new".to_vec()));
}
#[test]
fn empty_write_is_rejected() {
let mut c = RampCoordinator::new(MemStore::default(), 0);
assert!(matches!(c.write(&[]), Err(RampError::EmptyWrite)));
}
}