use std::collections::BTreeMap;
use mongreldb_types::ids::{TabletId, TransactionId};
use serde::{Deserialize, Serialize};
pub type KeyBytes = Vec<u8>;
#[derive(Debug, Default, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct TabletKeySet {
pub keys: Vec<KeyBytes>,
}
impl TabletKeySet {
pub fn new() -> Self {
Self::default()
}
pub fn insert(&mut self, key: impl Into<KeyBytes>) {
let key = key.into();
if !self.keys.iter().any(|k| k == &key) {
self.keys.push(key);
}
}
pub fn intersects(&self, other: &Self) -> bool {
self.keys.iter().any(|k| other.keys.iter().any(|o| o == k))
}
}
#[derive(Debug, Default, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct DistAccessSet {
pub reads: BTreeMap<TabletId, TabletKeySet>,
pub writes: BTreeMap<TabletId, TabletKeySet>,
}
impl DistAccessSet {
pub fn read(&mut self, tablet: TabletId, key: impl Into<KeyBytes>) {
self.reads.entry(tablet).or_default().insert(key);
}
pub fn write(&mut self, tablet: TabletId, key: impl Into<KeyBytes>) {
self.writes.entry(tablet).or_default().insert(key);
}
pub fn writes_intersect_reads(&self, other: &Self) -> bool {
for (tablet, wset) in &self.writes {
if let Some(rset) = other.reads.get(tablet) {
if wset.intersects(rset) {
return true;
}
}
}
false
}
pub fn writes_intersect_writes(&self, other: &Self) -> bool {
for (tablet, wset) in &self.writes {
if let Some(ow) = other.writes.get(tablet) {
if wset.intersects(ow) {
return true;
}
}
}
false
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
pub enum DistCertOutcome {
Ok,
SerializationFailure,
}
#[derive(Debug, Default)]
pub struct DistSsiCertifier {
inflight: BTreeMap<TransactionId, DistAccessSet>,
committed: Vec<(TransactionId, DistAccessSet)>,
window: usize,
}
impl DistSsiCertifier {
pub fn new(window: usize) -> Self {
Self {
inflight: BTreeMap::new(),
committed: Vec::new(),
window: window.max(1),
}
}
pub fn track(&mut self, txn: TransactionId, access: DistAccessSet) {
self.inflight.insert(txn, access);
}
pub fn observe_read(&mut self, txn: TransactionId, tablet: TabletId, key: impl Into<KeyBytes>) {
self.inflight.entry(txn).or_default().read(tablet, key);
}
pub fn observe_write(
&mut self,
txn: TransactionId,
tablet: TabletId,
key: impl Into<KeyBytes>,
) {
self.inflight.entry(txn).or_default().write(tablet, key);
}
pub fn certify(&mut self, txn: TransactionId) -> DistCertOutcome {
let Some(mine) = self.inflight.remove(&txn) else {
return DistCertOutcome::Ok;
};
for (_other_id, other) in &self.committed {
if mine.writes_intersect_writes(other) {
return DistCertOutcome::SerializationFailure;
}
if other.writes_intersect_reads(&mine) {
if mine.writes_intersect_reads(other) {
return DistCertOutcome::SerializationFailure;
}
return DistCertOutcome::SerializationFailure;
}
if mine.writes_intersect_reads(other) && other.writes_intersect_reads(&mine) {
return DistCertOutcome::SerializationFailure;
}
}
self.committed.push((txn, mine));
if self.committed.len() > self.window {
let drop = self.committed.len() - self.window;
self.committed.drain(0..drop);
}
DistCertOutcome::Ok
}
pub fn abort(&mut self, txn: TransactionId) {
self.inflight.remove(&txn);
}
}
#[cfg(test)]
mod tests {
use super::*;
fn tid(n: u8) -> TabletId {
TabletId::from_bytes({
let mut b = [0u8; 16];
b[15] = n;
b
})
}
fn txn(n: u8) -> TransactionId {
TransactionId::from_bytes({
let mut b = [0u8; 16];
b[15] = n;
b
})
}
#[test]
fn multi_tablet_write_skew_aborts() {
let mut cert = DistSsiCertifier::new(32);
let t1 = txn(1);
let t2 = txn(2);
cert.observe_read(t1, tid(2), b"y");
cert.observe_write(t1, tid(1), b"x");
assert_eq!(cert.certify(t1), DistCertOutcome::Ok);
cert.observe_read(t2, tid(1), b"x");
cert.observe_write(t2, tid(2), b"y");
assert_eq!(cert.certify(t2), DistCertOutcome::SerializationFailure);
}
#[test]
fn multi_tablet_ww_aborts() {
let mut cert = DistSsiCertifier::new(8);
cert.observe_write(txn(1), tid(1), b"k");
assert_eq!(cert.certify(txn(1)), DistCertOutcome::Ok);
cert.observe_write(txn(2), tid(1), b"k");
assert_eq!(cert.certify(txn(2)), DistCertOutcome::SerializationFailure);
}
#[test]
fn disjoint_writes_commit() {
let mut cert = DistSsiCertifier::new(8);
cert.observe_write(txn(1), tid(1), b"a");
assert_eq!(cert.certify(txn(1)), DistCertOutcome::Ok);
cert.observe_write(txn(2), tid(2), b"b");
assert_eq!(cert.certify(txn(2)), DistCertOutcome::Ok);
}
}