use std::collections::{BTreeMap, BTreeSet};
use std::time::Duration;
use openraft::Raft;
use openraft::async_runtime::watch::WatchReceiver;
use openraft::storage::RaftStateMachine;
use tsoracle_driver_openraft::TypeConfig;
type NodeId = u64;
const HANDOFF_WAIT: Duration = Duration::from_secs(5);
pub(crate) fn pick_handoff_target(
me: NodeId,
voters: &BTreeSet<NodeId>,
matched: &BTreeMap<NodeId, u64>,
) -> Option<NodeId> {
voters
.iter()
.copied()
.filter(|id| *id != me)
.max_by_key(|id| matched.get(id).copied().unwrap_or(0))
}
pub(crate) async fn graceful_leader_handoff<SM>(raft: &Raft<TypeConfig, SM>, me: NodeId)
where
SM: RaftStateMachine<TypeConfig> + Send + Sync + 'static,
{
if raft.current_leader().await != Some(me) {
return;
}
let (voters, matched) = {
let metrics = raft.metrics().borrow_watched().clone();
let voters: BTreeSet<NodeId> = metrics.membership_config.voter_ids().collect();
let matched: BTreeMap<NodeId, u64> = metrics
.replication
.unwrap_or_default()
.into_iter()
.filter_map(|(id, log)| log.map(|log_id| (id, log_id.index)))
.collect();
(voters, matched)
};
let Some(target) = pick_handoff_target(me, &voters, &matched) else {
return;
};
tracing::info!(target, "transferring leadership before drain");
if let Err(error) = raft.trigger().transfer_leader(target).await {
tracing::warn!(
?error,
"leadership transfer failed; draining without handoff"
);
return;
}
let moved = tokio::time::timeout(HANDOFF_WAIT, async {
while raft.current_leader().await == Some(me) {
tokio::time::sleep(Duration::from_millis(50)).await;
}
})
.await;
if moved.is_err() {
tracing::warn!("leadership did not move within the handoff window; draining anyway");
} else {
tracing::info!(target, "leadership handed off; draining");
}
}
#[cfg(test)]
mod tests {
use super::*;
fn voters(ids: &[NodeId]) -> BTreeSet<NodeId> {
ids.iter().copied().collect()
}
#[test]
fn picks_the_most_caught_up_follower() {
let v = voters(&[1, 2, 3]);
let matched = BTreeMap::from([(2, 40), (3, 70)]);
assert_eq!(pick_handoff_target(1, &v, &matched), Some(3));
}
#[test]
fn excludes_self_even_when_most_caught_up() {
let v = voters(&[1, 2, 3]);
let matched = BTreeMap::from([(1, 100), (2, 40), (3, 30)]);
assert_eq!(pick_handoff_target(1, &v, &matched), Some(2));
}
#[test]
fn treats_missing_replication_as_zero() {
let v = voters(&[1, 2, 3]);
let matched = BTreeMap::new();
assert!(matches!(pick_handoff_target(1, &v, &matched), Some(2 | 3)));
}
#[test]
fn single_voter_cluster_has_no_target() {
let v = voters(&[1]);
assert_eq!(pick_handoff_target(1, &v, &BTreeMap::new()), None);
}
#[test]
fn ignores_non_voter_progress() {
let v = voters(&[1, 2]);
let matched = BTreeMap::from([(2, 10), (9, 999)]);
assert_eq!(pick_handoff_target(1, &v, &matched), Some(2));
}
}