use std::sync::Arc;
use std::time::Duration;
use dynomite::events::{ClusterEvent, EventManager, PeerId, TokenRange};
use gen_fsm::{
Action, DriverError, EventType, FsmDriver, FsmHandler, StopReason, TimeoutKind, Transition,
};
use tokio::sync::mpsc;
use tokio::task::JoinHandle;
use crate::aae::exchange::{Divergence, ExchangeError, PeerView};
use crate::aae::tictac::{KeyEntry, Tree};
pub const DEFAULT_STATE_TIMEOUT: Duration = Duration::from_secs(30);
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum State {
Init,
Compare,
Repair,
Finalize,
Failed,
}
#[derive(Debug)]
pub enum Event {
PeerRootsReceived(Vec<(u32, u64)>),
SegmentDiffsComputed(Vec<Divergence>),
RepairProgress(usize),
AllRepaired,
PeerError(String),
}
#[derive(Debug, Clone)]
pub enum ExchangeOutcome {
Completed {
repaired: usize,
divergences: usize,
},
Failed {
reason: String,
last_state: State,
},
}
#[derive(Debug, Clone)]
pub enum FsmRequest {
FetchPeerRoots,
ComputeDivergences {
peer_roots: Vec<(u32, u64)>,
},
ApplyRepairs {
divergences: Vec<Divergence>,
},
Finalize,
}
pub struct ExchangeHandler {
peer_idx: PeerId,
partition: TokenRange,
local_roots: Vec<(u32, u64)>,
peer_roots: Vec<(u32, u64)>,
divergences: Vec<Divergence>,
repaired: usize,
last_error: Option<String>,
state_timeout: Duration,
request_tx: mpsc::UnboundedSender<FsmRequest>,
events: Option<Arc<EventManager>>,
local_tree: Arc<Tree>,
}
impl ExchangeHandler {
#[must_use]
pub fn new(
local_tree: Arc<Tree>,
peer_idx: PeerId,
partition: TokenRange,
request_tx: mpsc::UnboundedSender<FsmRequest>,
) -> Self {
Self {
peer_idx,
partition,
local_roots: Vec::new(),
peer_roots: Vec::new(),
divergences: Vec::new(),
repaired: 0,
last_error: None,
state_timeout: DEFAULT_STATE_TIMEOUT,
request_tx,
events: None,
local_tree,
}
}
#[must_use]
pub fn with_state_timeout(mut self, timeout: Duration) -> Self {
self.state_timeout = timeout;
self
}
#[must_use]
pub fn with_events(mut self, events: Arc<EventManager>) -> Self {
self.events = Some(events);
self
}
#[must_use]
pub fn partition(&self) -> &TokenRange {
&self.partition
}
#[must_use]
pub const fn peer_idx(&self) -> PeerId {
self.peer_idx
}
#[must_use]
pub const fn repaired(&self) -> usize {
self.repaired
}
#[must_use]
pub fn divergences(&self) -> &[Divergence] {
&self.divergences
}
#[must_use]
pub fn local_tree(&self) -> &Tree {
self.local_tree.as_ref()
}
fn dispatch(&mut self, req: FsmRequest) -> bool {
if self.request_tx.send(req).is_ok() {
true
} else {
self.last_error
.get_or_insert_with(|| "io worker channel closed".to_string());
false
}
}
}
impl FsmHandler for ExchangeHandler {
type State = State;
type Event = Event;
type Reply = ();
type Stop = ExchangeOutcome;
fn initial(&self) -> State {
State::Init
}
fn on_enter(&mut self, state: State) -> Transition<Self> {
match state {
State::Init => {
self.local_roots = self.local_tree.roots();
if let Some(ev) = self.events.as_ref() {
ev.publish(ClusterEvent::AaeExchangeStarted {
with_peer: self.peer_idx,
partition: self.partition.clone(),
ts: std::time::SystemTime::now(),
});
}
if !self.dispatch(FsmRequest::FetchPeerRoots) {
return Transition::Next(State::Failed, vec![]);
}
Transition::Keep(vec![Action::set_state_timeout(self.state_timeout)])
}
State::Compare => {
if !self.dispatch(FsmRequest::ComputeDivergences {
peer_roots: self.peer_roots.clone(),
}) {
return Transition::Next(State::Failed, vec![]);
}
Transition::Keep(vec![Action::set_state_timeout(self.state_timeout)])
}
State::Repair => {
if self.divergences.is_empty() {
return Transition::Next(State::Finalize, vec![]);
}
if !self.dispatch(FsmRequest::ApplyRepairs {
divergences: self.divergences.clone(),
}) {
return Transition::Next(State::Failed, vec![]);
}
Transition::Keep(vec![Action::set_state_timeout(self.state_timeout)])
}
State::Finalize => {
if let Some(ev) = self.events.as_ref() {
let repaired = u64::try_from(self.repaired).unwrap_or(u64::MAX);
ev.publish(ClusterEvent::AaeExchangeCompleted {
with_peer: self.peer_idx,
partition: self.partition.clone(),
repaired,
ts: std::time::SystemTime::now(),
});
}
let _ = self.dispatch(FsmRequest::Finalize);
Transition::Stop(ExchangeOutcome::Completed {
repaired: self.repaired,
divergences: self.divergences.len(),
})
}
State::Failed => Transition::Stop(ExchangeOutcome::Failed {
reason: self
.last_error
.clone()
.unwrap_or_else(|| "exchange failed".to_string()),
last_state: State::Failed,
}),
}
}
fn handle(&mut self, state: State, _et: EventType, ev: Event) -> Transition<Self> {
if let Event::PeerError(msg) = &ev {
self.last_error = Some(msg.clone());
return Transition::Next(State::Failed, vec![]);
}
match (state, ev) {
(State::Init, Event::PeerRootsReceived(roots)) => {
self.peer_roots = roots;
if self.peer_roots == self.local_roots {
Transition::Next(State::Finalize, vec![])
} else {
Transition::Next(State::Compare, vec![])
}
}
(State::Compare, Event::SegmentDiffsComputed(diffs)) => {
let empty = diffs.is_empty();
self.divergences = diffs;
if empty {
Transition::Next(State::Finalize, vec![])
} else {
Transition::Next(State::Repair, vec![])
}
}
(State::Repair, Event::RepairProgress(n)) => {
self.repaired = self.repaired.saturating_add(n);
Transition::Keep(vec![])
}
(State::Repair, Event::AllRepaired) => Transition::Next(State::Finalize, vec![]),
_ => Transition::Keep(vec![]),
}
}
fn on_timeout(&mut self, state: State, kind: TimeoutKind) -> Transition<Self> {
let _ = kind;
self.last_error
.get_or_insert_with(|| format!("state timeout in {state:?}"));
Transition::Next(State::Failed, vec![])
}
}
pub trait PeerViewAsync: Send + Sync + 'static {
fn roots(
&self,
) -> impl std::future::Future<Output = Result<Vec<(u32, u64)>, ExchangeError>> + Send;
fn segments(
&self,
time_bucket: u32,
) -> impl std::future::Future<Output = Result<Vec<(u32, u64)>, ExchangeError>> + Send;
fn keys_in_segment(
&self,
time_bucket: u32,
segment: u32,
) -> impl std::future::Future<Output = Result<Vec<KeyEntry>, ExchangeError>> + Send;
}
pub struct SyncPeerViewAdapter<V> {
inner: V,
}
impl<V> SyncPeerViewAdapter<V> {
#[must_use]
pub const fn new(inner: V) -> Self {
Self { inner }
}
}
impl<V> PeerViewAsync for SyncPeerViewAdapter<V>
where
V: PeerView + Send + Sync + 'static,
{
async fn roots(&self) -> Result<Vec<(u32, u64)>, ExchangeError> {
self.inner.roots()
}
async fn segments(&self, time_bucket: u32) -> Result<Vec<(u32, u64)>, ExchangeError> {
self.inner.segments(time_bucket)
}
async fn keys_in_segment(
&self,
time_bucket: u32,
segment: u32,
) -> Result<Vec<KeyEntry>, ExchangeError> {
self.inner.keys_in_segment(time_bucket, segment)
}
}
#[must_use]
pub fn spawn_exchange<P>(
handler: ExchangeHandler,
request_rx: mpsc::UnboundedReceiver<FsmRequest>,
peer: Arc<P>,
) -> (FsmDriver<ExchangeHandler>, JoinHandle<()>)
where
P: PeerViewAsync,
{
let local_tree = Arc::clone(&handler.local_tree);
let driver = FsmDriver::start(handler);
let driver_for_io = driver.clone();
let io = tokio::spawn(io_loop(driver_for_io, request_rx, peer, local_tree));
(driver, io)
}
pub async fn exchange_with_peer<P>(
local_tree: Arc<Tree>,
peer: Arc<P>,
peer_idx: PeerId,
partition: TokenRange,
) -> Result<ExchangeOutcome, ExchangeError>
where
P: PeerViewAsync,
{
let (req_tx, req_rx) = mpsc::unbounded_channel::<FsmRequest>();
let handler = ExchangeHandler::new(local_tree, peer_idx, partition, req_tx);
let (driver, io) = spawn_exchange(handler, req_rx, peer);
let stop = driver
.join()
.await
.map_err(|e: DriverError| ExchangeError::BadPayload(format!("fsm driver: {e}")))?;
io.abort();
let _ = io.await;
match stop {
StopReason::Handler(outcome) => Ok(outcome),
StopReason::Closed => Err(ExchangeError::BadPayload(
"fsm driver closed before outcome".to_string(),
)),
}
}
async fn io_loop<P>(
driver: FsmDriver<ExchangeHandler>,
mut rx: mpsc::UnboundedReceiver<FsmRequest>,
peer: Arc<P>,
local_tree: Arc<Tree>,
) where
P: PeerViewAsync,
{
while let Some(req) = rx.recv().await {
match req {
FsmRequest::FetchPeerRoots => match peer.roots().await {
Ok(roots) => driver.cast(Event::PeerRootsReceived(roots)).await,
Err(e) => driver.cast(Event::PeerError(e.to_string())).await,
},
FsmRequest::ComputeDivergences { peer_roots } => {
match compute_divergences(local_tree.as_ref(), peer.as_ref(), &peer_roots).await {
Ok(divs) => driver.cast(Event::SegmentDiffsComputed(divs)).await,
Err(e) => driver.cast(Event::PeerError(e.to_string())).await,
}
}
FsmRequest::ApplyRepairs { divergences } => {
let count = divergences.len();
for _ in 0..count {
driver.cast(Event::RepairProgress(1)).await;
}
driver.cast(Event::AllRepaired).await;
}
FsmRequest::Finalize => {
break;
}
}
}
}
async fn compute_divergences<P>(
local: &Tree,
peer: &P,
peer_roots: &[(u32, u64)],
) -> Result<Vec<Divergence>, ExchangeError>
where
P: PeerViewAsync,
{
let local_roots = local.roots();
let dr = Tree::diverging_time_buckets(&local_roots, peer_roots);
let mut out = Vec::new();
for tb in dr {
let local_segs = local.segments(tb)?;
let peer_segs = peer.segments(tb).await?;
let ds = Tree::diverging_segments(&local_segs, &peer_segs);
for seg in ds {
let local_keys = local.keys_in_segment(tb, seg)?;
let peer_keys = peer.keys_in_segment(tb, seg).await?;
let (local_only, remote_only) = symmetric_difference(&local_keys, &peer_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)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::aae::exchange::LocalPeerView;
use crate::aae::tictac::TreeShape;
use dynomite::hashkit::DynToken;
use std::time::Duration;
fn shape() -> TreeShape {
TreeShape {
n_time_buckets: 4,
n_segments: 32,
time_window_seconds: 60,
}
}
fn partition() -> TokenRange {
TokenRange::new(DynToken::from_u32(0), DynToken::from_u32(1024))
}
fn handler(tree: Tree) -> (ExchangeHandler, mpsc::UnboundedReceiver<FsmRequest>) {
let (tx, rx) = mpsc::unbounded_channel::<FsmRequest>();
(ExchangeHandler::new(Arc::new(tree), 7, partition(), tx), rx)
}
fn assert_state_timeout(transition: &Transition<ExchangeHandler>, expected: Duration) {
match transition {
Transition::Keep(actions) | Transition::Next(_, actions) => {
let found = actions
.iter()
.any(|a| matches!(a, Action::SetStateTimeout(d) if *d == expected));
assert!(
found,
"expected SetStateTimeout({expected:?}); actions = {actions:?}"
);
}
Transition::Stop(_) => panic!("expected Keep/Next, got Stop"),
}
}
fn drain(rx: &mut mpsc::UnboundedReceiver<FsmRequest>) -> Vec<FsmRequest> {
let mut out = Vec::new();
while let Ok(req) = rx.try_recv() {
out.push(req);
}
out
}
#[test]
fn exchange_init_state_sets_30s_timeout() {
let tree = Tree::new(shape());
let (mut h, mut rx) = handler(tree);
let t = h.on_enter(State::Init);
assert_state_timeout(&t, DEFAULT_STATE_TIMEOUT);
let reqs = drain(&mut rx);
assert!(matches!(reqs.as_slice(), [FsmRequest::FetchPeerRoots]));
}
#[test]
fn peer_root_match_skips_to_finalize() {
let mut tree = Tree::new(shape());
tree.insert(b"users", b"alice", b"vc1", 0);
let (mut h, _rx) = handler(tree);
let _ = h.on_enter(State::Init);
let local = h.local_roots.clone();
let t = h.handle(
State::Init,
EventType::Cast,
Event::PeerRootsReceived(local),
);
match t {
Transition::Next(State::Finalize, _) => {}
other => panic!("expected Next(Finalize), got {other:?}"),
}
}
#[test]
fn peer_root_mismatch_advances_to_compare() {
let mut tree = Tree::new(shape());
tree.insert(b"users", b"alice", b"vc1", 0);
let (mut h, _rx) = handler(tree);
let _ = h.on_enter(State::Init);
let mut diff = h.local_roots.clone();
if diff.is_empty() {
diff.push((0, 0xdead_beef));
} else {
diff[0].1 ^= 0x1;
}
let t = h.handle(State::Init, EventType::Cast, Event::PeerRootsReceived(diff));
match t {
Transition::Next(State::Compare, _) => {}
other => panic!("expected Next(Compare), got {other:?}"),
}
}
#[test]
fn compare_timeout_transitions_to_failed() {
let tree = Tree::new(shape());
let (mut h, _rx) = handler(tree);
let t = h.on_timeout(State::Compare, TimeoutKind::State);
match t {
Transition::Next(State::Failed, _) => {}
other => panic!("expected Next(Failed), got {other:?}"),
}
let stop = h.on_enter(State::Failed);
match stop {
Transition::Stop(ExchangeOutcome::Failed { reason, .. }) => {
assert!(
reason.contains("timeout"),
"expected reason to mention timeout; got {reason:?}"
);
}
other => panic!("expected Stop(Failed), got {other:?}"),
}
}
#[test]
fn repair_progress_increments_counter() {
let tree = Tree::new(shape());
let (mut h, _rx) = handler(tree);
let _ = h.handle(State::Repair, EventType::Cast, Event::RepairProgress(1));
let _ = h.handle(State::Repair, EventType::Cast, Event::RepairProgress(2));
assert_eq!(h.repaired(), 3);
}
#[test]
fn all_repaired_advances_to_finalize() {
let tree = Tree::new(shape());
let (mut h, _rx) = handler(tree);
let t = h.handle(State::Repair, EventType::Cast, Event::AllRepaired);
match t {
Transition::Next(State::Finalize, _) => {}
other => panic!("expected Next(Finalize), got {other:?}"),
}
}
#[test]
fn peer_error_at_any_state_transitions_to_failed() {
for state in [State::Init, State::Compare, State::Repair, State::Finalize] {
let tree = Tree::new(shape());
let (mut h, _rx) = handler(tree);
let t = h.handle(
state,
EventType::Cast,
Event::PeerError(format!("boom in {state:?}")),
);
match t {
Transition::Next(State::Failed, _) => {}
other => panic!("expected Next(Failed) from {state:?}, got {other:?}"),
}
let stop = h.on_enter(State::Failed);
match stop {
Transition::Stop(ExchangeOutcome::Failed { reason, .. }) => {
assert!(reason.contains("boom"));
}
other => panic!("expected Stop(Failed), got {other:?}"),
}
}
}
#[test]
fn empty_diffs_skip_repair_and_go_to_finalize() {
let tree = Tree::new(shape());
let (mut h, _rx) = handler(tree);
let t = h.handle(
State::Compare,
EventType::Cast,
Event::SegmentDiffsComputed(vec![]),
);
match t {
Transition::Next(State::Finalize, _) => {}
other => panic!("expected Next(Finalize), got {other:?}"),
}
}
#[test]
fn diffs_found_advances_to_repair_and_dispatches_request() {
let tree = Tree::new(shape());
let (mut h, mut rx) = handler(tree);
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 t = h.handle(
State::Compare,
EventType::Cast,
Event::SegmentDiffsComputed(vec![div.clone()]),
);
match t {
Transition::Next(State::Repair, _) => {}
other => panic!("expected Next(Repair), got {other:?}"),
}
let t2 = h.on_enter(State::Repair);
assert_state_timeout(&t2, DEFAULT_STATE_TIMEOUT);
let reqs = drain(&mut rx);
assert_eq!(reqs.len(), 1, "expected one ApplyRepairs request");
assert!(matches!(
&reqs[0],
FsmRequest::ApplyRepairs { divergences } if divergences == &vec![div.clone()]
));
}
#[test]
fn finalize_entry_stops_with_completed_outcome() {
let tree = Tree::new(shape());
let (mut h, _rx) = handler(tree);
h.repaired = 4;
let _ = h.handle(
State::Compare,
EventType::Cast,
Event::SegmentDiffsComputed(vec![Divergence {
time_bucket: 0,
segment: 1,
local_only: vec![],
remote_only: vec![],
}]),
);
let stop = h.on_enter(State::Finalize);
match stop {
Transition::Stop(ExchangeOutcome::Completed {
repaired,
divergences,
}) => {
assert_eq!(repaired, 4);
assert_eq!(divergences, 1);
}
other => panic!("expected Stop(Completed), got {other:?}"),
}
}
#[tokio::test]
async fn end_to_end_exchange_via_fsm_finds_diff() {
let mut a = Tree::new(shape());
let mut b = Tree::new(shape());
for i in 0..50u32 {
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 b = Arc::new(b);
let view = SyncPeerViewAdapter::new(BorrowedPeerView {
tree: Arc::clone(&b),
});
let peer = Arc::new(view);
let local = Arc::new(a);
let outcome = exchange_with_peer(local, peer, 1, partition())
.await
.unwrap();
match outcome {
ExchangeOutcome::Completed {
repaired,
divergences,
} => {
assert!(divergences >= 1, "expected at least one divergence");
assert_eq!(repaired, divergences, "io_loop posts one progress per div");
}
ExchangeOutcome::Failed { reason, last_state } => {
panic!("unexpected Failed({last_state:?}, {reason})");
}
}
}
struct BorrowedPeerView {
tree: Arc<Tree>,
}
impl PeerView for BorrowedPeerView {
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)
}
}
#[tokio::test]
async fn end_to_end_with_failing_peer_yields_failed_outcome() {
struct Boom;
impl PeerView for Boom {
fn roots(&self) -> Result<Vec<(u32, u64)>, ExchangeError> {
Err(ExchangeError::BadPayload("synthetic peer down".into()))
}
fn segments(&self, _: u32) -> Result<Vec<(u32, u64)>, ExchangeError> {
Err(ExchangeError::BadPayload("synthetic peer down".into()))
}
fn keys_in_segment(&self, _: u32, _: u32) -> Result<Vec<KeyEntry>, ExchangeError> {
Err(ExchangeError::BadPayload("synthetic peer down".into()))
}
}
let local = Arc::new(Tree::new(shape()));
let peer = Arc::new(SyncPeerViewAdapter::new(Boom));
let outcome = exchange_with_peer(local, peer, 9, partition())
.await
.unwrap();
match outcome {
ExchangeOutcome::Failed { reason, .. } => {
assert!(reason.contains("synthetic peer down"));
}
ExchangeOutcome::Completed { .. } => panic!("expected Failed"),
}
}
fn takes_peer_view(_: &impl PeerView) {}
#[test]
fn local_peer_view_is_peer_view() {
let t = Tree::new(shape());
takes_peer_view(&LocalPeerView::new(&t));
}
}