use std::cmp::Ordering;
use std::sync::Arc;
use bytes::Bytes;
use tokio::sync::mpsc;
use crate::aae::exchange::Divergence;
use crate::aae::metrics::AaeMetrics;
use crate::aae::tictac::KeyEntry;
use crate::datatypes::Itc;
#[derive(Debug, Clone, Copy, Eq, PartialEq)]
pub enum RepairDirection {
PushToRemote,
PullFromRemote,
}
#[derive(Debug, Clone, Eq, PartialEq)]
pub enum RepairOutcome {
Winner {
value: Bytes,
clock: Itc,
},
Siblings(Vec<(Bytes, Itc)>),
}
impl RepairOutcome {
#[must_use]
pub fn resolve_with_warn(self, key: &[u8]) -> (Bytes, Itc) {
match self {
Self::Winner { value, clock } => (value, clock),
Self::Siblings(siblings) => {
tracing::warn!(
target: "dyniak::aae::repair",
key = %String::from_utf8_lossy(key),
siblings = siblings.len(),
"sibling-aware merge: concurrent clocks; falling back to lex-largest value"
);
siblings
.into_iter()
.max_by(|a, b| a.0.cmp(&b.0).then_with(|| a.1.encode().cmp(&b.1.encode())))
.expect("invariant: Siblings carries at least two entries")
}
}
}
}
#[derive(Debug, Clone, Eq, PartialEq)]
pub struct RepairTask {
pub peer_idx: u32,
pub bucket: Vec<u8>,
pub key: Vec<u8>,
pub vclock: Vec<u8>,
pub direction: RepairDirection,
}
impl RepairTask {
#[must_use]
pub fn evaluate(replicas: &[(Bytes, Itc)]) -> RepairOutcome {
let n = replicas.len();
let mut keep = vec![true; n];
for i in 0..n {
if !keep[i] {
continue;
}
for j in 0..n {
if i == j || !keep[j] {
continue;
}
match replicas[i].1.partial_cmp_event(&replicas[j].1) {
Some(Ordering::Less) => {
keep[i] = false;
break;
}
Some(Ordering::Equal) if i > j => {
keep[i] = false;
break;
}
_ => {}
}
}
}
let surviving: Vec<(Bytes, Itc)> = replicas
.iter()
.zip(keep.iter())
.filter_map(|((v, c), &k)| {
if k {
Some((v.clone(), c.clone()))
} else {
None
}
})
.collect();
if surviving.len() == 1 {
let (value, clock) = surviving.into_iter().next().expect("len == 1");
RepairOutcome::Winner { value, clock }
} else {
RepairOutcome::Siblings(surviving)
}
}
}
#[derive(Debug, Clone)]
pub enum Outcome {
Repaired(RepairTask),
AmbiguousClock {
bucket: Vec<u8>,
key: Vec<u8>,
local: Vec<u8>,
remote: Vec<u8>,
},
PeerUnavailable {
peer_idx: u32,
task: RepairTask,
},
}
pub trait ClockOrder: Send + Sync {
fn compare(&self, a: &[u8], b: &[u8]) -> Option<Ordering>;
}
#[derive(Debug, Clone, Copy, Default)]
pub struct LexicographicOrder;
impl ClockOrder for LexicographicOrder {
fn compare(&self, a: &[u8], b: &[u8]) -> Option<Ordering> {
if a == b {
return Some(Ordering::Equal);
}
match a.len().cmp(&b.len()) {
Ordering::Equal => Some(a.cmp(b)),
ord => Some(ord),
}
}
}
#[derive(Debug, Clone, Copy, Default)]
pub struct ItcOrder;
impl ClockOrder for ItcOrder {
fn compare(&self, a: &[u8], b: &[u8]) -> Option<Ordering> {
let a_stamp = Itc::decode(a)?;
let b_stamp = Itc::decode(b)?;
a_stamp.partial_cmp_event(&b_stamp)
}
}
pub trait RepairSink: Send + Sync {
fn submit(&self, task: RepairTask) -> Result<(), RepairTask>;
}
pub struct MpscRepairSink {
tx: mpsc::Sender<RepairTask>,
}
impl MpscRepairSink {
#[must_use]
pub fn new(tx: mpsc::Sender<RepairTask>) -> Self {
Self { tx }
}
}
impl RepairSink for MpscRepairSink {
fn submit(&self, task: RepairTask) -> Result<(), RepairTask> {
match self.tx.try_send(task) {
Ok(()) => Ok(()),
Err(mpsc::error::TrySendError::Full(t) | mpsc::error::TrySendError::Closed(t)) => {
Err(t)
}
}
}
}
pub struct RepairScheduler {
sink: Arc<dyn RepairSink>,
order: Arc<dyn ClockOrder>,
peer_idx: u32,
metrics: Option<Arc<AaeMetrics>>,
metrics_dc: String,
metrics_rack: String,
}
impl RepairScheduler {
#[must_use]
pub fn new(peer_idx: u32, sink: Arc<dyn RepairSink>, order: Arc<dyn ClockOrder>) -> Self {
Self {
sink,
order,
peer_idx,
metrics: None,
metrics_dc: String::new(),
metrics_rack: String::new(),
}
}
#[must_use]
pub fn with_metrics(mut self, metrics: Arc<AaeMetrics>, dc: &str, rack: &str) -> Self {
self.metrics = Some(metrics);
dc.clone_into(&mut self.metrics_dc);
rack.clone_into(&mut self.metrics_rack);
self
}
pub fn resolve_all(&self, divergences: &[Divergence]) -> Vec<Outcome> {
let mut out = Vec::new();
for d in divergences {
out.extend(self.resolve(d));
}
if let Some(m) = self.metrics.as_ref() {
let dispatched = u64::try_from(
out.iter()
.filter(|o| matches!(o, Outcome::Repaired(_)))
.count(),
)
.unwrap_or(u64::MAX);
m.record_repair_dispatched(
self.peer_idx,
&self.metrics_dc,
&self.metrics_rack,
dispatched,
);
}
out
}
pub fn resolve(&self, divergence: &Divergence) -> Vec<Outcome> {
let mut out = Vec::new();
let mut remote_by_key: std::collections::BTreeMap<(Vec<u8>, Vec<u8>), &KeyEntry> =
std::collections::BTreeMap::new();
for r in &divergence.remote_only {
remote_by_key.insert((r.bucket.clone(), r.key.clone()), r);
}
let mut local_seen: std::collections::BTreeSet<(Vec<u8>, Vec<u8>)> =
std::collections::BTreeSet::new();
for l in &divergence.local_only {
let id = (l.bucket.clone(), l.key.clone());
local_seen.insert(id.clone());
if let Some(r) = remote_by_key.remove(&id) {
match self.order.compare(&l.vclock, &r.vclock) {
Some(Ordering::Greater) => {
out.push(self.enact(l, RepairDirection::PushToRemote));
}
Some(Ordering::Less) => {
out.push(self.enact(r, RepairDirection::PullFromRemote));
}
Some(Ordering::Equal) => {
}
None => {
out.push(Outcome::AmbiguousClock {
bucket: l.bucket.clone(),
key: l.key.clone(),
local: l.vclock.clone(),
remote: r.vclock.clone(),
});
}
}
} else {
out.push(self.enact(l, RepairDirection::PushToRemote));
}
}
for (_, r) in remote_by_key {
out.push(self.enact(r, RepairDirection::PullFromRemote));
}
out
}
fn enact(&self, entry: &KeyEntry, direction: RepairDirection) -> Outcome {
let task = RepairTask {
peer_idx: self.peer_idx,
bucket: entry.bucket.clone(),
key: entry.key.clone(),
vclock: entry.vclock.clone(),
direction,
};
match self.sink.submit(task.clone()) {
Ok(()) => Outcome::Repaired(task),
Err(t) => Outcome::PeerUnavailable {
peer_idx: self.peer_idx,
task: t,
},
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::aae::exchange::Divergence;
fn stamp_with_events(ticks: u64) -> Itc {
let mut s = Itc::seed();
for _ in 0..ticks {
s.event();
}
s
}
fn forked_concurrent(ticks_a: u64, ticks_b: u64) -> (Itc, Itc) {
let (mut a, mut b) = Itc::seed().fork();
for _ in 0..ticks_a {
a.event();
}
for _ in 0..ticks_b {
b.event();
}
(a, b)
}
#[test]
fn lexicographic_order_picks_longer() {
let lo = LexicographicOrder;
assert_eq!(lo.compare(b"vc2", b"vc1"), Some(Ordering::Greater));
assert_eq!(
lo.compare(b"vc11", b"vc2"),
Some(Ordering::Greater),
"longer wins over shorter"
);
assert_eq!(lo.compare(b"x", b"x"), Some(Ordering::Equal));
}
#[test]
fn itc_order_decodes_and_compares_stamps() {
let a = stamp_with_events(1);
let b = stamp_with_events(2);
let order = ItcOrder;
let a_bytes = a.encode();
let b_bytes = b.encode();
assert_eq!(order.compare(&a_bytes, &b_bytes), Some(Ordering::Less));
assert_eq!(order.compare(&b_bytes, &a_bytes), Some(Ordering::Greater));
assert_eq!(order.compare(&a_bytes, &a_bytes), Some(Ordering::Equal));
}
#[test]
fn itc_order_concurrent_stamps_compare_none() {
let (a, b) = forked_concurrent(1, 1);
let order = ItcOrder;
assert_eq!(order.compare(&a.encode(), &b.encode()), None);
}
#[test]
fn itc_order_unparseable_input_returns_none() {
let order = ItcOrder;
assert_eq!(order.compare(b"not-an-itc", b"\x00\x00\x00\x00"), None);
}
#[test]
fn repair_for_divergent_key_reaches_channel() {
let (tx, mut rx) = mpsc::channel::<RepairTask>(8);
let sink: Arc<dyn RepairSink> = Arc::new(MpscRepairSink::new(tx));
let order: Arc<dyn ClockOrder> = Arc::new(LexicographicOrder);
let sched = RepairScheduler::new(7, sink, order);
let div = Divergence {
time_bucket: 0,
segment: 11,
local_only: vec![KeyEntry {
bucket: b"users".to_vec(),
key: b"alice".to_vec(),
vclock: b"vc2".to_vec(),
}],
remote_only: vec![KeyEntry {
bucket: b"users".to_vec(),
key: b"alice".to_vec(),
vclock: b"vc1".to_vec(),
}],
};
let outcomes = sched.resolve(&div);
assert_eq!(outcomes.len(), 1);
match &outcomes[0] {
Outcome::Repaired(task) => {
assert_eq!(task.peer_idx, 7);
assert_eq!(task.key, b"alice");
assert_eq!(task.vclock, b"vc2");
assert_eq!(task.direction, RepairDirection::PushToRemote);
}
other => panic!("expected Repaired, got {other:?}"),
}
let received = rx.try_recv().unwrap();
assert_eq!(received.key, b"alice");
assert_eq!(received.vclock, b"vc2");
}
#[test]
fn repair_local_only_pushes_to_remote() {
let (tx, _rx) = mpsc::channel::<RepairTask>(8);
let sink: Arc<dyn RepairSink> = Arc::new(MpscRepairSink::new(tx));
let order: Arc<dyn ClockOrder> = Arc::new(LexicographicOrder);
let sched = RepairScheduler::new(3, sink, order);
let div = Divergence {
time_bucket: 0,
segment: 1,
local_only: vec![KeyEntry {
bucket: b"b".to_vec(),
key: b"k".to_vec(),
vclock: b"vc".to_vec(),
}],
remote_only: vec![],
};
let outcomes = sched.resolve(&div);
assert_eq!(outcomes.len(), 1);
assert!(matches!(
&outcomes[0],
Outcome::Repaired(t) if t.direction == RepairDirection::PushToRemote
));
}
#[test]
fn closed_channel_surfaces_peer_unavailable() {
let (tx, rx) = mpsc::channel::<RepairTask>(1);
drop(rx);
let sink: Arc<dyn RepairSink> = Arc::new(MpscRepairSink::new(tx));
let order: Arc<dyn ClockOrder> = Arc::new(LexicographicOrder);
let sched = RepairScheduler::new(99, sink, order);
let div = Divergence {
time_bucket: 0,
segment: 1,
local_only: vec![KeyEntry {
bucket: b"b".to_vec(),
key: b"k".to_vec(),
vclock: b"vc".to_vec(),
}],
remote_only: vec![],
};
let outcomes = sched.resolve(&div);
assert_eq!(outcomes.len(), 1);
assert!(matches!(&outcomes[0], Outcome::PeerUnavailable { .. }));
}
#[test]
fn evaluate_winner_when_one_dominates_others() {
let a = stamp_with_events(3);
let b = stamp_with_events(1);
let c = stamp_with_events(2);
let replicas = vec![
(Bytes::from_static(b"v_a"), a.clone()),
(Bytes::from_static(b"v_b"), b),
(Bytes::from_static(b"v_c"), c),
];
match RepairTask::evaluate(&replicas) {
RepairOutcome::Winner { value, clock } => {
assert_eq!(value, Bytes::from_static(b"v_a"));
assert_eq!(clock, a);
}
RepairOutcome::Siblings(s) => panic!("expected Winner, got Siblings({})", s.len()),
}
}
#[test]
fn evaluate_siblings_when_all_concurrent() {
let (a, bc) = Itc::seed().fork();
let (b, c) = bc.fork();
let mut a = a;
let mut b = b;
let mut c = c;
a.event();
b.event();
c.event();
let replicas = vec![
(Bytes::from_static(b"v_a"), a),
(Bytes::from_static(b"v_b"), b),
(Bytes::from_static(b"v_c"), c),
];
match RepairTask::evaluate(&replicas) {
RepairOutcome::Siblings(s) => {
assert_eq!(s.len(), 3);
}
RepairOutcome::Winner { .. } => panic!("expected Siblings(3), got Winner"),
}
}
#[test]
fn evaluate_siblings_excludes_dominated_entries() {
let (a_branch, c) = Itc::seed().fork();
let mut a = a_branch.clone();
a.event();
a.event();
let mut b = a_branch;
b.event();
let mut c = c;
c.event();
let replicas = vec![
(Bytes::from_static(b"v_a"), a),
(Bytes::from_static(b"v_b"), b),
(Bytes::from_static(b"v_c"), c),
];
match RepairTask::evaluate(&replicas) {
RepairOutcome::Siblings(s) => {
assert_eq!(s.len(), 2, "B is dominated by A and must be excluded");
let values: Vec<&Bytes> = s.iter().map(|(v, _)| v).collect();
assert!(values.contains(&&Bytes::from_static(b"v_a")));
assert!(values.contains(&&Bytes::from_static(b"v_c")));
assert!(!values.contains(&&Bytes::from_static(b"v_b")));
}
RepairOutcome::Winner { .. } => panic!("expected Siblings(2), got Winner"),
}
}
#[test]
fn evaluate_dedupes_equal_clocks() {
let a = stamp_with_events(2);
let replicas = vec![
(Bytes::from_static(b"v_a"), a.clone()),
(Bytes::from_static(b"v_a_dup"), a.clone()),
];
match RepairTask::evaluate(&replicas) {
RepairOutcome::Winner { value, clock } => {
assert_eq!(value, Bytes::from_static(b"v_a"));
assert_eq!(clock, a);
}
RepairOutcome::Siblings(s) => {
panic!("expected Winner after dedupe, got Siblings({})", s.len())
}
}
}
#[test]
fn resolve_with_warn_picks_lex_largest_on_siblings() {
let (a_branch, c) = Itc::seed().fork();
let mut a = a_branch.clone();
a.event();
let mut b = a_branch;
b.event();
let mut c = c;
c.event();
let _ = (a, b, c);
let (s1, s2) = forked_concurrent(1, 1);
let outcome = RepairOutcome::Siblings(vec![
(Bytes::from_static(b"alpha"), s1.clone()),
(Bytes::from_static(b"zulu"), s2.clone()),
(Bytes::from_static(b"mike"), s1),
]);
let (value, _) = outcome.resolve_with_warn(b"some-key");
assert_eq!(value, Bytes::from_static(b"zulu"));
}
#[test]
fn resolve_with_warn_passes_winner_through() {
let v = stamp_with_events(5);
let outcome = RepairOutcome::Winner {
value: Bytes::from_static(b"only"),
clock: v.clone(),
};
let (value, clock) = outcome.resolve_with_warn(b"k");
assert_eq!(value, Bytes::from_static(b"only"));
assert_eq!(clock, v);
}
#[test]
fn resolve_all_records_dispatched_metric() {
let (tx, _rx) = mpsc::channel::<RepairTask>(8);
let sink: Arc<dyn RepairSink> = Arc::new(MpscRepairSink::new(tx));
let order: Arc<dyn ClockOrder> = Arc::new(LexicographicOrder);
let metrics = Arc::new(AaeMetrics::new());
let sched =
RepairScheduler::new(11, sink, order).with_metrics(Arc::clone(&metrics), "dc1", "rA");
let divs = vec![
Divergence {
time_bucket: 0,
segment: 1,
local_only: vec![KeyEntry {
bucket: b"b".to_vec(),
key: b"k1".to_vec(),
vclock: b"vc".to_vec(),
}],
remote_only: vec![],
},
Divergence {
time_bucket: 0,
segment: 2,
local_only: vec![KeyEntry {
bucket: b"b".to_vec(),
key: b"k2".to_vec(),
vclock: b"vc".to_vec(),
}],
remote_only: vec![],
},
];
let outs = sched.resolve_all(&divs);
assert_eq!(outs.len(), 2);
let snap = metrics.snapshot();
assert_eq!(snap.repair_dispatched.len(), 1);
assert_eq!(snap.repair_dispatched[0].peer_idx, 11);
assert_eq!(snap.repair_dispatched[0].count, 2);
}
}