use std::collections::{HashMap, HashSet};
use std::sync::{Arc, Mutex};
use std::thread;
use std::time::Duration;
use crossbeam::channel::{self, Receiver, Sender};
use rand::RngExt as _;
use uuid::Uuid;
use crate::error::{Error, Result};
use crate::raft::kv::{self, Command, Kv};
use crate::raft::{
Envelope, Index, Log, Message, Node, NodeID, Options, Request, RequestID, Response, Status,
TICK_INTERVAL,
};
use crate::storage::BitCask;
#[derive(Default)]
struct TransportInner {
partitions: HashSet<(NodeID, NodeID)>,
drop_rate: f64,
reorder: bool,
held: HashMap<NodeID, Envelope>,
online: HashMap<NodeID, bool>,
mailboxes: HashMap<NodeID, Sender<Envelope>>,
}
#[derive(Clone, Default)]
pub struct Transport {
inner: Arc<Mutex<TransportInner>>,
}
impl Transport {
fn register(&self, id: NodeID, tx: Sender<Envelope>) {
let mut g = self.inner.lock().expect("transport lock");
g.mailboxes.insert(id, tx);
g.online.insert(id, true);
}
fn set_online(&self, id: NodeID, online: bool) {
let mut g = self.inner.lock().expect("transport lock");
g.online.insert(id, online);
if !online {
g.held.remove(&id);
}
}
fn deliver(&self, msg: Envelope) {
let mut g = self.inner.lock().expect("transport lock");
let from = msg.from;
let to = msg.to;
if !g.online.get(&from).copied().unwrap_or(false) {
return;
}
if !g.online.get(&to).copied().unwrap_or(false) {
return;
}
if g.partitions.contains(&(from, to)) {
return;
}
if g.drop_rate > 0.0 && rand::rng().random::<f64>() < g.drop_rate {
return;
}
if g.reorder {
if let Some(prev) = g.held.remove(&to) {
if let Some(tx) = g.mailboxes.get(&to) {
let _ = tx.try_send(msg);
let _ = tx.try_send(prev);
}
return;
}
g.held.insert(to, msg);
return;
}
if let Some(tx) = g.mailboxes.get(&to) {
let _ = tx.try_send(msg);
}
}
pub fn partition(&self, a: NodeID, b: NodeID) {
let mut g = self.inner.lock().expect("transport lock");
g.partitions.insert((a, b));
g.partitions.insert((b, a));
}
pub fn partition_one_way(&self, from: NodeID, to: NodeID) {
let mut g = self.inner.lock().expect("transport lock");
g.partitions.insert((from, to));
}
pub fn partition_groups(&self, left: &[NodeID], right: &[NodeID]) {
let mut g = self.inner.lock().expect("transport lock");
for &a in left {
for &b in right {
g.partitions.insert((a, b));
g.partitions.insert((b, a));
}
}
}
pub fn heal_pair(&self, a: NodeID, b: NodeID) {
let mut g = self.inner.lock().expect("transport lock");
g.partitions.remove(&(a, b));
g.partitions.remove(&(b, a));
}
pub fn heal_all(&self) {
let mut g = self.inner.lock().expect("transport lock");
g.partitions.clear();
let held = std::mem::take(&mut g.held);
for (to, msg) in held {
if let Some(tx) = g.mailboxes.get(&to) {
let _ = tx.try_send(msg);
}
}
}
pub fn set_drop_rate(&self, rate: f64) {
let mut g = self.inner.lock().expect("transport lock");
g.drop_rate = rate.clamp(0.0, 1.0);
}
pub fn set_reorder(&self, on: bool) {
let mut g = self.inner.lock().expect("transport lock");
g.reorder = on;
if !on {
let held = std::mem::take(&mut g.held);
for (to, msg) in held {
if let Some(tx) = g.mailboxes.get(&to) {
let _ = tx.try_send(msg);
}
}
}
}
}
type RequestTx = Sender<(Request, Sender<Result<Response>>)>;
#[derive(Clone)]
pub struct Client {
request_txs: Arc<Mutex<HashMap<NodeID, RequestTx>>>,
preferred: NodeID,
attempts: u32,
per_attempt_timeout: Duration,
retry_sleep: Duration,
}
impl Client {
fn new(request_txs: HashMap<NodeID, RequestTx>, preferred: NodeID) -> Self {
Self {
request_txs: Arc::new(Mutex::new(request_txs)),
preferred,
attempts: 40,
per_attempt_timeout: Duration::from_millis(200),
retry_sleep: Duration::from_millis(50),
}
}
pub fn register_node(&self, id: NodeID, tx: RequestTx) {
self.request_txs.lock().expect("client lock").insert(id, tx);
}
pub fn unregister_node(&self, id: NodeID) {
self.request_txs.lock().expect("client lock").remove(&id);
}
pub fn preferred_hint(&mut self, id: NodeID) {
self.preferred = id;
}
pub fn request(&mut self, request: Request) -> Result<Response> {
let txs = self.request_txs.lock().expect("client lock").clone();
let mut order: Vec<NodeID> = txs.keys().copied().collect();
order.sort();
if let Some(pos) = order.iter().position(|&id| id == self.preferred) {
let id = order.remove(pos);
order.insert(0, id);
}
let mut last_err = Error::Abort;
for _ in 0..self.attempts {
for &node_id in &order {
let Some(tx) = txs.get(&node_id) else { continue };
let (resp_tx, resp_rx) = channel::bounded(1);
if tx.send((request.clone(), resp_tx)).is_err() {
continue;
}
match resp_rx.recv_timeout(self.per_attempt_timeout) {
Ok(Ok(resp)) => {
self.preferred = node_id;
return Ok(resp);
}
Ok(Err(Error::Abort)) => last_err = Error::Abort,
Ok(Err(e)) => return Err(e),
Err(_) => last_err = Error::IO("request timed out".into()),
}
}
thread::sleep(self.retry_sleep);
}
Err(last_err)
}
pub fn request_on(&mut self, node_id: NodeID, request: Request) -> Result<Response> {
let txs = self.request_txs.lock().expect("client lock").clone();
let Some(tx) = txs.get(&node_id).cloned() else {
return Err(Error::IO(format!("node {node_id} not registered")));
};
let mut last_err = Error::Abort;
for _ in 0..self.attempts {
let (resp_tx, resp_rx) = channel::bounded(1);
if tx.send((request.clone(), resp_tx)).is_err() {
return Err(Error::IO(format!("node {node_id} request channel closed")));
}
match resp_rx.recv_timeout(self.per_attempt_timeout) {
Ok(Ok(resp)) => {
self.preferred = node_id;
return Ok(resp);
}
Ok(Err(Error::Abort)) => last_err = Error::Abort,
Ok(Err(e)) => return Err(e),
Err(_) => last_err = Error::IO("request timed out".into()),
}
thread::sleep(self.retry_sleep);
}
Err(last_err)
}
pub fn put(&mut self, key: &str, value: &str) -> Result<Index> {
let req = Request::Write(kv::encode(&Command::Put {
key: key.into(),
value: value.into(),
}));
match self.request(req)? {
Response::Write(bytes) => match kv::decode::<kv::Response>(&bytes)? {
kv::Response::Put(index) => Ok(index),
other => Err(Error::InvalidData(format!("unexpected write response: {other:?}"))),
},
other => Err(Error::InvalidData(format!("expected Write, got {other:?}"))),
}
}
pub fn put_on(&mut self, node_id: NodeID, key: &str, value: &str) -> Result<Index> {
let req = Request::Write(kv::encode(&Command::Put {
key: key.into(),
value: value.into(),
}));
match self.request_on(node_id, req)? {
Response::Write(bytes) => match kv::decode::<kv::Response>(&bytes)? {
kv::Response::Put(index) => Ok(index),
other => Err(Error::InvalidData(format!("unexpected write response: {other:?}"))),
},
other => Err(Error::InvalidData(format!("expected Write, got {other:?}"))),
}
}
pub fn get(&mut self, key: &str) -> Result<Option<String>> {
let req = Request::Read(kv::encode(&Command::Get { key: key.into() }));
match self.request(req)? {
Response::Read(bytes) => match kv::decode::<kv::Response>(&bytes)? {
kv::Response::Get(v) => Ok(v),
other => Err(Error::InvalidData(format!("unexpected read response: {other:?}"))),
},
other => Err(Error::InvalidData(format!("expected Read, got {other:?}"))),
}
}
pub fn scan(&mut self) -> Result<std::collections::BTreeMap<String, String>> {
let req = Request::Read(kv::encode(&Command::Scan));
match self.request(req)? {
Response::Read(bytes) => match kv::decode::<kv::Response>(&bytes)? {
kv::Response::Scan(map) => Ok(map),
other => Err(Error::InvalidData(format!("unexpected scan response: {other:?}"))),
},
other => Err(Error::InvalidData(format!("expected Read, got {other:?}"))),
}
}
pub fn status(&mut self) -> Result<Status> {
match self.request(Request::Status)? {
Response::Status(s) => Ok(s),
other => Err(Error::InvalidData(format!("expected Status, got {other:?}"))),
}
}
pub fn status_on(&mut self, node_id: NodeID) -> Result<Status> {
match self.request_on(node_id, Request::Status)? {
Response::Status(s) => Ok(s),
other => Err(Error::InvalidData(format!("expected Status, got {other:?}"))),
}
}
pub fn change_membership(&mut self, voters: HashSet<NodeID>) -> Result<Index> {
match self.request(Request::ChangeMembership { voters })? {
Response::ChangeMembership { index } => Ok(index),
other => Err(Error::InvalidData(format!("expected ChangeMembership, got {other:?}"))),
}
}
}
pub fn wait_for_leader(client: &mut Client) -> Result<Status> {
let mut last_err = Error::Abort;
for _ in 0..100 {
match client.status() {
Ok(s) => return Ok(s),
Err(Error::Abort) => {
last_err = Error::Abort;
thread::sleep(Duration::from_millis(50));
}
Err(e) => return Err(e),
}
}
Err(last_err)
}
struct NodeControl {
stop_tx: Sender<()>,
}
pub struct Cluster {
opts: Options,
transport: Transport,
nodes: HashMap<NodeID, NodeControl>,
members: HashSet<NodeID>,
client: Client,
}
impl Cluster {
pub fn spawn(node_ids: &[NodeID]) -> Self {
Self::spawn_with_options(node_ids, test_options())
}
pub fn spawn_with_options(node_ids: &[NodeID], opts: Options) -> Self {
let transport = Transport::default();
let mut request_txs = HashMap::new();
let mut nodes = HashMap::new();
let members: HashSet<NodeID> = node_ids.iter().copied().collect();
for &id in node_ids {
let (control, req_tx) = spawn_node(id, &members, opts.clone(), transport.clone());
request_txs.insert(id, req_tx.clone());
nodes.insert(id, control);
}
let preferred = node_ids.first().copied().unwrap_or(1);
let client = Client::new(request_txs, preferred);
Self { opts, transport, nodes, members, client }
}
pub fn client(&self) -> Client {
self.client.clone()
}
pub fn transport(&self) -> Transport {
self.transport.clone()
}
pub fn members(&self) -> HashSet<NodeID> {
self.members.clone()
}
pub fn options(&self) -> &Options {
&self.opts
}
pub fn stop(&mut self, id: NodeID) {
if let Some(ctrl) = self.nodes.remove(&id) {
let _ = ctrl.stop_tx.send(());
self.client.unregister_node(id);
self.transport.set_online(id, false);
}
}
pub fn start(&mut self, id: NodeID) {
if self.nodes.contains_key(&id) {
return;
}
self.members.insert(id);
let (control, req_tx) =
spawn_node(id, &self.members, self.opts.clone(), self.transport.clone());
self.client.register_node(id, req_tx);
self.nodes.insert(id, control);
}
pub fn set_members(&mut self, members: HashSet<NodeID>) {
self.members = members;
}
pub fn partition(&self, a: NodeID, b: NodeID) {
self.transport.partition(a, b);
}
pub fn partition_groups(&self, left: &[NodeID], right: &[NodeID]) {
self.transport.partition_groups(left, right);
}
pub fn heal_all(&self) {
self.transport.heal_all();
}
pub fn set_drop_rate(&self, rate: f64) {
self.transport.set_drop_rate(rate);
}
pub fn set_reorder(&self, on: bool) {
self.transport.set_reorder(on);
}
pub fn is_running(&self, id: NodeID) -> bool {
self.nodes.contains_key(&id)
}
}
pub fn test_options() -> Options {
Options {
heartbeat_interval: 2,
election_timeout_range: 5..10,
max_append_entries: 100,
pre_vote: true,
check_quorum: true,
snapshot_threshold: 0,
}
}
fn spawn_node(
id: NodeID,
members: &HashSet<NodeID>,
opts: Options,
transport: Transport,
) -> (NodeControl, RequestTx) {
let peers: HashSet<NodeID> = members.iter().copied().filter(|&p| p != id).collect();
let (inbox_tx, inbox_rx) = channel::unbounded();
transport.register(id, inbox_tx);
let (request_tx, request_rx) = channel::unbounded();
let (stop_tx, stop_rx) = channel::bounded(1);
let (node_tx, node_rx) = channel::unbounded();
static SEQ: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
let seq = SEQ.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
let path = std::env::temp_dir().join(format!(
"raft-cluster-{}-{}-{}-{}.log",
std::process::id(),
id,
seq,
std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap()
.as_nanos()
));
let log = Log::new(Box::new(BitCask::new(path).expect("bitcask"))).expect("log");
let node = Node::new(id, peers, log, Kv::new(), node_tx, opts).expect("node");
let transport_out = transport.clone();
thread::spawn(move || {
run_node(node, inbox_rx, node_rx, request_rx, stop_rx, transport_out);
});
(NodeControl { stop_tx }, request_tx)
}
fn run_node(
mut node: Node,
peers_rx: Receiver<Envelope>,
node_rx: Receiver<Envelope>,
request_rx: Receiver<(Request, Sender<Result<Response>>)>,
stop_rx: Receiver<()>,
transport: Transport,
) {
let ticker = channel::tick(TICK_INTERVAL);
let mut response_txs: HashMap<RequestID, Sender<Result<Response>>> = HashMap::new();
let node_id = node.id();
loop {
crossbeam::select! {
recv(stop_rx) -> _ => break,
recv(ticker) -> _ => {
node = match node.tick() {
Ok(n) => n,
Err(_) => break,
};
}
recv(peers_rx) -> msg => {
let Ok(msg) = msg else { break };
node = match node.step(msg) {
Ok(n) => n,
Err(_) => break,
};
}
recv(node_rx) -> msg => {
let Ok(msg) = msg else { break };
if msg.to == node_id {
if let Message::ClientResponse { id, response } = msg.message {
if let Some(tx) = response_txs.remove(&id) {
let _ = tx.send(response);
}
}
continue;
}
transport.deliver(msg);
}
recv(request_rx) -> result => {
let Ok((request, response_tx)) = result else { break };
let id = Uuid::new_v4();
let msg = Envelope {
from: node.id(),
to: node.id(),
term: node.term(),
message: Message::ClientRequest { id, request },
};
response_txs.insert(id, response_tx);
node = match node.step(msg) {
Ok(n) => n,
Err(_) => break,
};
}
}
}
transport.set_online(node_id, false);
}