use std::time::{Duration, Instant};
use dynomite::events::{PeerId, TokenRange};
use dynomite::hashkit::DynToken;
use gen_fsm::{Action, EventType, FsmHandler, TimeoutKind, Transition};
use throttle_core::{SystemClock, Throttle};
pub const DEFAULT_CHUNK_SIZE: u64 = 1024;
pub const DEFAULT_MAX_IN_FLIGHT: u64 = 4;
pub const DEFAULT_CHUNKS_PER_SEC: u64 = 100;
pub const NEGOTIATING_STATE_TIMEOUT: Duration = Duration::from_secs(10);
pub const SENDING_EVENT_TIMEOUT: Duration = Duration::from_secs(30);
pub const FLUSHING_STATE_TIMEOUT: Duration = Duration::from_mins(1);
pub const FINALIZING_STATE_TIMEOUT: Duration = Duration::from_secs(30);
#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
pub enum State {
Init,
Negotiating,
Sending,
Flushing,
Finalizing,
Failed,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct TokenCursor {
pub start: DynToken,
pub end: DynToken,
pub keys_drained: u64,
pub total_keys: u64,
}
impl TokenCursor {
#[must_use]
pub fn new(range: &TokenRange, total_keys: u64) -> Self {
Self {
start: range.start().clone(),
end: range.end().clone(),
keys_drained: 0,
total_keys,
}
}
#[must_use]
pub const fn is_drained(&self) -> bool {
self.keys_drained >= self.total_keys
}
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct SendRequest {
pub src_peer: PeerId,
pub dst_peer: PeerId,
pub token_range: TokenRange,
pub total_keys: u64,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct Chunk {
pub chunk_id: u64,
pub keys: u64,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum Event {
SendRequestReceived(SendRequest),
NegotiationAck {
accepted: bool,
max_chunk_size: u64,
},
NextChunkBuilt(Chunk),
ChunkAcked {
chunk_id: u64,
},
BatchDone,
BatchAcked,
FinalizeAcked,
PeerError(String),
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum HandoffOutcome {
Completed {
keys_transferred: u64,
duration: Duration,
},
Failed {
reason: String,
partial_count: u64,
last_state: State,
},
}
pub struct HandoffHandler {
src_peer: PeerId,
dst_peer: PeerId,
token_range: TokenRange,
chunk_size: u64,
max_in_flight: u64,
throttle: Throttle<SystemClock>,
cursor: TokenCursor,
sent_chunks: u64,
acked_chunks: u64,
seen_acks: Vec<u64>,
started_at: Instant,
last_error: Option<String>,
last_state: State,
}
impl HandoffHandler {
#[must_use]
pub fn new(
src_peer: PeerId,
dst_peer: PeerId,
token_range: TokenRange,
total_keys: u64,
) -> Self {
Self::with_settings(
src_peer,
dst_peer,
token_range,
total_keys,
DEFAULT_CHUNK_SIZE,
DEFAULT_MAX_IN_FLIGHT,
DEFAULT_CHUNKS_PER_SEC,
)
}
#[must_use]
pub fn with_settings(
src_peer: PeerId,
dst_peer: PeerId,
token_range: TokenRange,
total_keys: u64,
chunk_size: u64,
max_in_flight: u64,
chunks_per_sec: u64,
) -> Self {
let chunk_size = chunk_size.max(1);
let max_in_flight = max_in_flight.max(1);
let burst = chunks_per_sec.max(1);
let cursor = TokenCursor::new(&token_range, total_keys);
Self {
src_peer,
dst_peer,
token_range,
chunk_size,
max_in_flight,
throttle: Throttle::new(burst, chunks_per_sec),
cursor,
sent_chunks: 0,
acked_chunks: 0,
seen_acks: Vec::new(),
started_at: Instant::now(),
last_error: None,
last_state: State::Init,
}
}
#[must_use]
pub const fn src_peer(&self) -> PeerId {
self.src_peer
}
#[must_use]
pub const fn dst_peer(&self) -> PeerId {
self.dst_peer
}
#[must_use]
pub const fn token_range(&self) -> &TokenRange {
&self.token_range
}
#[must_use]
pub const fn chunk_size(&self) -> u64 {
self.chunk_size
}
#[must_use]
pub const fn max_in_flight(&self) -> u64 {
self.max_in_flight
}
#[must_use]
pub const fn sent_chunks(&self) -> u64 {
self.sent_chunks
}
#[must_use]
pub const fn acked_chunks(&self) -> u64 {
self.acked_chunks
}
#[must_use]
pub const fn cursor(&self) -> &TokenCursor {
&self.cursor
}
#[must_use]
pub fn has_in_flight_capacity(&self) -> bool {
self.sent_chunks.saturating_sub(self.acked_chunks) < self.max_in_flight
}
pub fn try_admit_chunk(&self) -> bool {
self.throttle.try_acquire(1)
}
#[must_use]
pub fn acked_keys(&self) -> u64 {
let raw = self.acked_chunks.saturating_mul(self.chunk_size);
let total = self.cursor.total_keys;
if total == 0 {
raw
} else {
raw.min(total)
}
}
#[must_use]
pub fn elapsed(&self) -> Duration {
self.started_at.elapsed()
}
#[must_use]
pub const fn last_state(&self) -> State {
self.last_state
}
fn record_state(&mut self, state: State) {
self.last_state = state;
}
}
impl FsmHandler for HandoffHandler {
type State = State;
type Event = Event;
type Reply = ();
type Stop = HandoffOutcome;
fn initial(&self) -> State {
State::Init
}
fn on_enter(&mut self, state: State) -> Transition<Self> {
if state != State::Failed {
self.record_state(state);
}
match state {
State::Init => Transition::Keep(vec![]),
State::Negotiating => {
Transition::Keep(vec![Action::set_state_timeout(NEGOTIATING_STATE_TIMEOUT)])
}
State::Sending => {
Transition::Keep(vec![Action::set_event_timeout(SENDING_EVENT_TIMEOUT)])
}
State::Flushing => {
Transition::Keep(vec![Action::set_state_timeout(FLUSHING_STATE_TIMEOUT)])
}
State::Finalizing => {
Transition::Keep(vec![Action::set_state_timeout(FINALIZING_STATE_TIMEOUT)])
}
State::Failed => Transition::Stop(HandoffOutcome::Failed {
reason: self
.last_error
.clone()
.unwrap_or_else(|| "handoff failed".to_string()),
partial_count: self.acked_keys(),
last_state: self.last_state,
}),
}
}
fn handle(&mut self, state: State, _et: EventType, ev: Event) -> Transition<Self> {
self.record_state(state);
if let Event::PeerError(msg) = ev.clone() {
self.last_error = Some(msg);
return Transition::Next(State::Failed, vec![]);
}
match (state, ev) {
(State::Init, Event::SendRequestReceived(req)) => {
if req.src_peer != self.src_peer
|| req.dst_peer != self.dst_peer
|| req.token_range != self.token_range
{
self.last_error = Some(format!(
"send request mismatch: got src={} dst={} expected src={} dst={}",
req.src_peer, req.dst_peer, self.src_peer, self.dst_peer,
));
return Transition::Next(State::Failed, vec![]);
}
if req.total_keys != self.cursor.total_keys {
self.cursor.total_keys = req.total_keys;
}
Transition::Next(State::Negotiating, vec![])
}
(
State::Negotiating,
Event::NegotiationAck {
accepted,
max_chunk_size,
},
) => {
if !accepted {
self.last_error = Some("negotiation rejected".to_string());
return Transition::Next(State::Failed, vec![]);
}
if max_chunk_size > 0 && max_chunk_size < self.chunk_size {
self.chunk_size = max_chunk_size;
}
Transition::Next(State::Sending, vec![])
}
(State::Sending, Event::NextChunkBuilt(chunk)) => {
self.sent_chunks = self.sent_chunks.saturating_add(1);
self.cursor.keys_drained = self.cursor.keys_drained.saturating_add(chunk.keys);
Transition::Keep(vec![Action::set_event_timeout(SENDING_EVENT_TIMEOUT)])
}
(State::Sending, Event::ChunkAcked { chunk_id }) => {
if chunk_id >= self.sent_chunks {
self.last_error = Some(format!(
"ack for unknown chunk_id={chunk_id} (sent={})",
self.sent_chunks
));
return Transition::Next(State::Failed, vec![]);
}
if self.seen_acks.binary_search(&chunk_id).is_ok() {
return Transition::Keep(vec![Action::set_event_timeout(
SENDING_EVENT_TIMEOUT,
)]);
}
let pos = self
.seen_acks
.binary_search(&chunk_id)
.unwrap_or_else(|p| p);
self.seen_acks.insert(pos, chunk_id);
self.acked_chunks = self.acked_chunks.saturating_add(1);
Transition::Keep(vec![Action::set_event_timeout(SENDING_EVENT_TIMEOUT)])
}
(State::Sending, Event::BatchDone) => Transition::Next(State::Flushing, vec![]),
(State::Flushing, Event::BatchAcked) => Transition::Next(State::Finalizing, vec![]),
(State::Flushing, Event::ChunkAcked { chunk_id }) => {
if chunk_id < self.sent_chunks && self.seen_acks.binary_search(&chunk_id).is_err() {
let pos = self
.seen_acks
.binary_search(&chunk_id)
.unwrap_or_else(|p| p);
self.seen_acks.insert(pos, chunk_id);
self.acked_chunks = self.acked_chunks.saturating_add(1);
}
Transition::Keep(vec![])
}
(State::Finalizing, Event::FinalizeAcked) => {
let outcome = HandoffOutcome::Completed {
keys_transferred: self.acked_keys(),
duration: self.elapsed(),
};
Transition::Stop(outcome)
}
_ => Transition::Keep(vec![]),
}
}
fn on_timeout(&mut self, state: State, kind: TimeoutKind) -> Transition<Self> {
let label = match kind {
TimeoutKind::State => "state timeout",
TimeoutKind::Event => "event timeout",
TimeoutKind::Generic(name) => name,
};
self.last_error = Some(format!("{label} in {state:?}"));
self.record_state(state);
Transition::Next(State::Failed, vec![])
}
}
#[cfg(test)]
mod tests {
use super::*;
fn range() -> TokenRange {
TokenRange::new(DynToken::from_u32(0), DynToken::from_u32(8192))
}
fn handler(total: u64) -> HandoffHandler {
HandoffHandler::with_settings(7, 11, range(), total, 64, 4, 1_000_000)
}
fn assert_set_state(transition: &Transition<HandoffHandler>, 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 assert_set_event(transition: &Transition<HandoffHandler>, expected: Duration) {
match transition {
Transition::Keep(actions) | Transition::Next(_, actions) => {
let found = actions
.iter()
.any(|a| matches!(a, Action::SetEventTimeout(d) if *d == expected));
assert!(
found,
"expected SetEventTimeout({expected:?}); actions = {actions:?}"
);
}
Transition::Stop(_) => panic!("expected Keep/Next, got Stop"),
}
}
#[test]
fn negotiating_entry_arms_state_timeout() {
let mut h = handler(8);
let t = h.on_enter(State::Negotiating);
assert_set_state(&t, NEGOTIATING_STATE_TIMEOUT);
}
#[test]
fn sending_entry_arms_event_timeout() {
let mut h = handler(8);
let t = h.on_enter(State::Sending);
assert_set_event(&t, SENDING_EVENT_TIMEOUT);
}
#[test]
fn flushing_entry_arms_state_timeout() {
let mut h = handler(8);
let t = h.on_enter(State::Flushing);
assert_set_state(&t, FLUSHING_STATE_TIMEOUT);
}
#[test]
fn finalizing_entry_arms_state_timeout() {
let mut h = handler(8);
let t = h.on_enter(State::Finalizing);
assert_set_state(&t, FINALIZING_STATE_TIMEOUT);
}
#[test]
fn negotiation_accepted_clamps_chunk_size() {
let mut h = handler(1024);
let t = h.handle(
State::Negotiating,
EventType::Cast,
Event::NegotiationAck {
accepted: true,
max_chunk_size: 16,
},
);
match t {
Transition::Next(State::Sending, _) => {}
other => panic!("expected Next(Sending), got {other:?}"),
}
assert_eq!(h.chunk_size(), 16);
}
#[test]
fn double_ack_does_not_double_count() {
let mut h = handler(64);
let _ = h.handle(
State::Sending,
EventType::Cast,
Event::NextChunkBuilt(Chunk {
chunk_id: 0,
keys: 64,
}),
);
let _ = h.handle(
State::Sending,
EventType::Cast,
Event::ChunkAcked { chunk_id: 0 },
);
let _ = h.handle(
State::Sending,
EventType::Cast,
Event::ChunkAcked { chunk_id: 0 },
);
assert_eq!(h.acked_chunks(), 1);
}
#[test]
fn ack_for_unknown_chunk_advances_to_failed() {
let mut h = handler(64);
let t = h.handle(
State::Sending,
EventType::Cast,
Event::ChunkAcked { chunk_id: 99 },
);
match t {
Transition::Next(State::Failed, _) => {}
other => panic!("expected Next(Failed), got {other:?}"),
}
}
}