use std::collections::HashMap;
use std::net::SocketAddr;
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::mpsc::{self, Receiver, RecvTimeoutError, Sender};
use std::sync::{Arc, Mutex};
use std::thread;
use std::time::Duration;
use beamr::atom::{Atom, AtomTable};
use beamr::distribution::ConnectionManager;
use beamr::distribution::resolver::{NodeResolver, ResolveError, ResolveFuture};
use dashmap::DashMap;
use tokio::runtime::{Builder, Handle, Runtime};
use crate::api::kv::{KvKey, KvValue};
use crate::branch::ShardId;
use crate::sync::SyncNodeId;
use crate::sync::ballot::{Ballot, Stamp};
use crate::sync::consistency::{
CasVote, ConsistencyError, QuorumOutcome, RejectKind, StrongConsistency, quorum_size,
wait_for_cas_quorum_from_receiver,
};
use crate::sync::membership::WriteMembership;
use crate::tree::Hash;
use super::protocol::{
AckOutcome, BatchWriteAck, BatchWriteEntry, BatchWriteProposal, Nack, Prepare, Promise,
PushResponse, RejectReason, ShardSyncRequest, SyncError, SyncMessage, WriteAck, WriteId,
WriteProposal, encode_beamr_sync_frame, register_beamr_sync_handler,
send_sync_message_via_beamr,
};
type WriteRegistry = Arc<DashMap<WriteId, Sender<CasVote<SyncNodeId>>>>;
type ElectionRegistry = Arc<DashMap<ShardId, Sender<ElectionVote>>>;
type CatchUpRegistry = Arc<DashMap<ShardId, Sender<PushResponse>>>;
#[derive(Debug, Clone)]
pub enum ElectionVote {
Promised(Promise),
Nacked(Nack),
}
#[derive(Debug, Clone)]
pub struct ElectionOutcome {
pub ballot: Ballot,
pub promises: Vec<Promise>,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum ElectionError {
Lost { highest_seen: Ballot },
Timeout {
required: usize,
promised_votes: usize,
},
Transport(String),
}
impl std::fmt::Display for ElectionError {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::Lost { highest_seen } => write!(
formatter,
"election lost: a higher ballot {highest_seen:?} was promised elsewhere"
),
Self::Timeout {
required,
promised_votes,
} => write!(
formatter,
"election timed out: required {required} promises, collected {promised_votes}"
),
Self::Transport(message) => write!(formatter, "election transport error: {message}"),
}
}
}
impl std::error::Error for ElectionError {}
const DEFAULT_COOKIE: &str = "haematite-distribution-cookie";
#[derive(Default)]
struct EndpointResolver {
nodes: Mutex<HashMap<String, SocketAddr>>,
}
impl EndpointResolver {
fn insert(&self, name: &str, addr: SocketAddr) {
self.nodes
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.insert(name.to_owned(), addr);
}
}
impl NodeResolver for EndpointResolver {
fn resolve<'a>(&'a self, name: &'a str) -> ResolveFuture<'a> {
let result = self
.nodes
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.get(name)
.copied()
.ok_or(ResolveError::NotFound);
Box::pin(async move { result })
}
}
pub type InboundSync = Result<SyncMessage, SyncError>;
pub struct DistributionEndpoint {
atom_table: Arc<AtomTable>,
resolver: Arc<EndpointResolver>,
manager: ConnectionManager,
accept: AcceptGuard,
runtime: Option<Arc<Runtime>>,
inbound: Mutex<Receiver<InboundSync>>,
local_name: String,
local_addr: SocketAddr,
local_creation: u32,
write_counter: AtomicU64,
registry: WriteRegistry,
elections: ElectionRegistry,
catch_ups: CatchUpRegistry,
}
impl std::fmt::Debug for DistributionEndpoint {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
formatter
.debug_struct("DistributionEndpoint")
.field("local_name", &self.local_name)
.field("local_addr", &self.local_addr)
.finish_non_exhaustive()
}
}
pub struct ProposeWrite {
pub key: KvKey,
pub expected: Option<Hash>,
pub value: KvValue,
pub ttl: Option<Duration>,
}
impl DistributionEndpoint {
pub fn bind(
local_name: impl Into<String>,
listen_addr: SocketAddr,
local_creation: u32,
cookie: Option<&str>,
) -> Result<Self, SyncError> {
ensure_outside_runtime()?;
let local_name = local_name.into();
let runtime = Arc::new(
Builder::new_multi_thread()
.worker_threads(2)
.enable_all()
.build()
.map_err(|_error| SyncError::TransportRuntimeUnavailable)?,
);
let atom_table = Arc::new(AtomTable::with_common_atoms());
let resolver = Arc::new(EndpointResolver::default());
let manager = ConnectionManager::new(
Arc::clone(&atom_table),
Arc::clone(&resolver) as Arc<dyn NodeResolver + Send + Sync>,
cookie.unwrap_or(DEFAULT_COOKIE),
local_name.clone(),
local_creation,
);
let accept = runtime
.block_on(manager.listen(listen_addr))
.map_err(|error: std::io::Error| SyncError::TransportBind(error.to_string()))?;
let local_addr = accept.local_addr();
resolver.insert(&local_name, local_addr);
let (tx, inbound) = mpsc::channel::<InboundSync>();
let registry: WriteRegistry = Arc::new(DashMap::new());
let elections: ElectionRegistry = Arc::new(DashMap::new());
let catch_ups: CatchUpRegistry = Arc::new(DashMap::new());
register_inbound_drain(
&manager,
tx,
Arc::clone(®istry),
Arc::clone(&elections),
Arc::clone(&catch_ups),
local_creation,
);
Ok(Self {
atom_table,
resolver,
manager,
accept: AcceptGuard::new(accept),
runtime: Some(runtime),
inbound: Mutex::new(inbound),
local_name,
local_addr,
local_creation,
write_counter: AtomicU64::new(0),
registry,
elections,
catch_ups,
})
}
#[must_use]
pub fn local_name(&self) -> &str {
&self.local_name
}
#[must_use]
pub const fn local_addr(&self) -> SocketAddr {
self.local_addr
}
#[must_use]
pub const fn local_creation(&self) -> u32 {
self.local_creation
}
pub fn add_peer(&self, name: &str, addr: SocketAddr) {
self.resolver.insert(name, addr);
}
#[must_use]
pub fn peer_atom(&self, name: &str) -> Atom {
self.atom_table.intern(name)
}
pub fn connect(&self, peer_name: &str) -> Result<(), SyncError> {
ensure_outside_runtime()?;
let manager = self.manager.clone();
let peer_name = peer_name.to_owned();
self.runtime()?
.block_on(async move { manager.connect(&peer_name).await })
.map(drop)
.map_err(|_error| SyncError::TransportConnectFailed)
}
#[must_use]
pub fn connected_nodes(&self) -> Vec<Atom> {
self.manager.connected_nodes()
}
#[must_use]
pub fn is_connected(&self, peer_name: &str) -> bool {
self.manager
.get_connection(self.atom_table.intern(peer_name))
.is_some()
}
pub fn send(&self, remote: Atom, message: &SyncMessage) -> Result<(), SyncError> {
ensure_outside_runtime()?;
let handle = self.runtime()?.handle().clone();
send_sync_message_via_beamr(&self.manager, remote, message, |connection, frame| {
handle.block_on(async move {
connection
.write_raw(&frame)
.await
.map_err(|_error| SyncError::TransportWrite)
})
})
}
pub fn send_to(&self, peer_name: &str, message: &SyncMessage) -> Result<(), SyncError> {
self.send(self.atom_table.intern(peer_name), message)
}
pub fn propose_write(
&self,
write: ProposeWrite,
shard_id: ShardId,
epoch: Ballot,
membership: &WriteMembership,
timeout: Duration,
) -> Result<QuorumOutcome<SyncNodeId>, ConsistencyError> {
self.propose_write_stamped(
write,
shard_id,
Stamp::new(epoch, 0),
false,
membership,
timeout,
)
}
pub fn propose_write_stamped(
&self,
write: ProposeWrite,
shard_id: ShardId,
stamp: Stamp,
tombstone: bool,
membership: &WriteMembership,
timeout: Duration,
) -> Result<QuorumOutcome<SyncNodeId>, ConsistencyError> {
let ProposeWrite {
key,
expected,
value,
ttl,
} = write;
if Handle::try_current().is_ok() {
return Err(ConsistencyError::TransportUnavailable);
}
let write_id = WriteId {
origin: SyncNodeId::new(self.local_name.clone()),
origin_creation: self.local_creation,
counter: self.write_counter.fetch_add(1, Ordering::Relaxed),
};
let (vote_tx, vote_rx) = mpsc::channel::<CasVote<SyncNodeId>>();
self.registry.insert(write_id.clone(), vote_tx);
let _guard = RegistryGuard {
registry: &self.registry,
write_id: write_id.clone(),
};
let handle = self
.runtime()
.map_err(|_error| ConsistencyError::TransportUnavailable)?
.handle()
.clone();
let proposal = WriteProposal {
write_id,
shard_id,
key,
expected,
value,
ttl,
epoch: stamp.epoch,
seq: stamp.seq,
tombstone,
};
let frame = encode_beamr_sync_frame(&SyncMessage::WriteProposal(proposal))
.map_err(|_error| ConsistencyError::TransportUnavailable)?;
let frame = Arc::new(frame);
for target in &membership.send_targets {
let manager = self.manager.clone();
let remote = self.atom_table.intern(target.as_str());
let frame = Arc::clone(&frame);
handle.spawn(async move {
match manager.get_connection(remote) {
Some(connection) => {
if let Err(error) = connection.write_raw(frame.as_slice()).await {
log::warn!("write proposal send failed: {error}");
}
}
None => log::warn!("write proposal send target unreachable"),
}
});
}
let strong = StrongConsistency::new(membership.total_nodes, timeout);
wait_for_cas_quorum_from_receiver(strong, &vote_rx)
}
pub fn propose_batch_stamped(
&self,
shard_id: ShardId,
entries: Vec<BatchWriteEntry>,
stamp: Stamp,
membership: &WriteMembership,
timeout: Duration,
) -> Result<QuorumOutcome<SyncNodeId>, ConsistencyError> {
if Handle::try_current().is_ok() {
return Err(ConsistencyError::TransportUnavailable);
}
let write_id = WriteId {
origin: SyncNodeId::new(self.local_name.clone()),
origin_creation: self.local_creation,
counter: self.write_counter.fetch_add(1, Ordering::Relaxed),
};
let (vote_tx, vote_rx) = mpsc::channel::<CasVote<SyncNodeId>>();
self.registry.insert(write_id.clone(), vote_tx);
let _guard = RegistryGuard {
registry: &self.registry,
write_id: write_id.clone(),
};
let handle = self
.runtime()
.map_err(|_error| ConsistencyError::TransportUnavailable)?
.handle()
.clone();
let proposal = BatchWriteProposal {
write_id,
shard_id,
entries,
stamp,
};
let frame = encode_beamr_sync_frame(&SyncMessage::BatchWriteProposal(proposal))
.map_err(|_error| ConsistencyError::TransportUnavailable)?;
let frame = Arc::new(frame);
for target in &membership.send_targets {
let manager = self.manager.clone();
let remote = self.atom_table.intern(target.as_str());
let frame = Arc::clone(&frame);
handle.spawn(async move {
match manager.get_connection(remote) {
Some(connection) => {
if let Err(error) = connection.write_raw(frame.as_slice()).await {
log::warn!("batch write proposal send failed: {error}");
}
}
None => log::warn!("batch write proposal send target unreachable"),
}
});
}
let strong = StrongConsistency::new(membership.total_nodes, timeout);
wait_for_cas_quorum_from_receiver(strong, &vote_rx)
}
pub fn run_prepare_round(
&self,
shard_id: ShardId,
ballot: &Ballot,
self_promise: Promise,
membership: &WriteMembership,
timeout: Duration,
) -> Result<Vec<Promise>, ElectionError> {
if Handle::try_current().is_ok() {
return Err(ElectionError::Transport(
"acquire_shard blocked from inside the distribution runtime".to_owned(),
));
}
let required = quorum_size(membership.total_nodes)
.map_err(|error| ElectionError::Transport(error.to_string()))?;
let (vote_tx, vote_rx) = mpsc::channel::<ElectionVote>();
self.elections.insert(shard_id, vote_tx);
let _guard = ElectionGuard {
elections: &self.elections,
shard_id,
};
let handle = self
.runtime()
.map_err(|error| ElectionError::Transport(error.to_string()))?
.handle()
.clone();
let frame = encode_beamr_sync_frame(&SyncMessage::Prepare(Prepare {
shard_id,
ballot: ballot.clone(),
}))
.map_err(|error| ElectionError::Transport(error.to_string()))?;
let frame = Arc::new(frame);
for target in &membership.send_targets {
let manager = self.manager.clone();
let remote = self.atom_table.intern(target.as_str());
let frame = Arc::clone(&frame);
handle.spawn(async move {
match manager.get_connection(remote) {
Some(connection) => {
if let Err(error) = connection.write_raw(frame.as_slice()).await {
log::warn!("prepare send failed: {error}");
}
}
None => log::warn!("prepare send target unreachable"),
}
});
}
collect_prepare_votes(ballot, required, self_promise, &vote_rx, timeout)
}
pub fn run_catch_up_round(
&self,
shard_id: ShardId,
source: &SyncNodeId,
from_root: Option<Hash>,
timeout: Duration,
) -> Result<PushResponse, SyncError> {
if Handle::try_current().is_ok() {
return Err(SyncError::TransportBlockingFromAsync);
}
let (tx, rx) = mpsc::channel::<PushResponse>();
self.catch_ups.insert(shard_id, tx);
let _guard = CatchUpGuard {
catch_ups: &self.catch_ups,
shard_id,
};
let handle = self.runtime()?.handle().clone();
let request = ShardSyncRequest::new(
shard_id,
SyncNodeId::new(self.local_name.clone()),
from_root,
);
let frame = encode_beamr_sync_frame(&SyncMessage::ShardSyncRequest(request))?;
let frame = Arc::new(frame);
let manager = self.manager.clone();
let remote = self.atom_table.intern(source.as_str());
let frame_for_send = Arc::clone(&frame);
handle.spawn(async move {
match manager.get_connection(remote) {
Some(connection) => {
if let Err(error) = connection.write_raw(frame_for_send.as_slice()).await {
log::warn!("catch-up request send failed: {error}");
}
}
None => log::warn!("catch-up request source unreachable"),
}
});
match rx.recv_timeout(timeout) {
Ok(response) => Ok(response),
Err(RecvTimeoutError::Timeout | RecvTimeoutError::Disconnected) => {
Err(SyncError::TransportDrainDisconnected)
}
}
}
pub fn send_message_fire_and_forget(
&self,
target: &SyncNodeId,
message: &SyncMessage,
) -> Result<(), SyncError> {
let frame = encode_beamr_sync_frame(message)?;
let frame = Arc::new(frame);
let handle = self.runtime()?.handle().clone();
let manager = self.manager.clone();
let remote = self.atom_table.intern(target.as_str());
handle.spawn(async move {
match manager.get_connection(remote) {
Some(connection) => {
if let Err(error) = connection.write_raw(frame.as_slice()).await {
log::warn!("catch-up response send failed: {error}");
}
}
None => log::warn!("catch-up response target unreachable"),
}
});
Ok(())
}
pub fn recv_inbound(&self, timeout: Duration) -> Result<Option<InboundSync>, SyncError> {
let inbound = self
.inbound
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
match inbound.recv_timeout(timeout) {
Ok(message) => Ok(Some(message)),
Err(RecvTimeoutError::Timeout) => Ok(None),
Err(RecvTimeoutError::Disconnected) => Err(SyncError::TransportDrainDisconnected),
}
}
pub fn runtime_handle(&self) -> Result<Handle, SyncError> {
Ok(self.runtime()?.handle().clone())
}
fn runtime(&self) -> Result<&Arc<Runtime>, SyncError> {
self.runtime
.as_ref()
.ok_or(SyncError::TransportRuntimeUnavailable)
}
}
fn ensure_outside_runtime() -> Result<(), SyncError> {
if Handle::try_current().is_ok() {
return Err(SyncError::TransportBlockingFromAsync);
}
Ok(())
}
impl Drop for DistributionEndpoint {
fn drop(&mut self) {
self.accept.shutdown();
if let Some(runtime) = self.runtime.take() {
thread::spawn(move || drop(runtime));
}
}
}
struct AcceptGuard {
handle: beamr::distribution::connection::AcceptHandle,
}
impl AcceptGuard {
const fn new(handle: beamr::distribution::connection::AcceptHandle) -> Self {
Self { handle }
}
fn shutdown(&self) {
self.handle.shutdown();
}
}
struct RegistryGuard<'registry> {
registry: &'registry WriteRegistry,
write_id: WriteId,
}
impl Drop for RegistryGuard<'_> {
fn drop(&mut self) {
self.registry.remove(&self.write_id);
}
}
struct ElectionGuard<'registry> {
elections: &'registry ElectionRegistry,
shard_id: ShardId,
}
impl Drop for ElectionGuard<'_> {
fn drop(&mut self) {
self.elections.remove(&self.shard_id);
}
}
struct CatchUpGuard<'registry> {
catch_ups: &'registry CatchUpRegistry,
shard_id: ShardId,
}
impl Drop for CatchUpGuard<'_> {
fn drop(&mut self) {
self.catch_ups.remove(&self.shard_id);
}
}
fn collect_prepare_votes(
ballot: &Ballot,
required: usize,
self_promise: Promise,
receiver: &Receiver<ElectionVote>,
timeout: Duration,
) -> Result<Vec<Promise>, ElectionError> {
use std::collections::HashSet;
use std::time::Instant;
let mut promised_nodes: HashSet<SyncNodeId> = HashSet::new();
promised_nodes.insert(self_promise.promiser.clone());
let mut promises = vec![self_promise];
let mut highest_seen = ballot.clone();
if promises.len() >= required {
return Ok(promises);
}
let deadline = Instant::now() + timeout;
loop {
let Some(remaining) = deadline.checked_duration_since(Instant::now()) else {
return Err(finish_loss(required, promises.len(), ballot, &highest_seen));
};
match receiver.recv_timeout(remaining) {
Ok(ElectionVote::Promised(promise)) => {
if &promise.ballot != ballot {
if promise.ballot > highest_seen {
highest_seen = promise.ballot.clone();
}
continue;
}
if promised_nodes.insert(promise.promiser.clone()) {
promises.push(promise);
if promises.len() >= required {
return Ok(promises);
}
}
}
Ok(ElectionVote::Nacked(nack)) => {
if nack.promised > highest_seen {
highest_seen = nack.promised;
}
}
Err(RecvTimeoutError::Timeout | RecvTimeoutError::Disconnected) => {
return Err(finish_loss(required, promises.len(), ballot, &highest_seen));
}
}
}
}
fn finish_loss(
required: usize,
promised_votes: usize,
own_ballot: &Ballot,
highest_seen: &Ballot,
) -> ElectionError {
if highest_seen > own_ballot {
ElectionError::Lost {
highest_seen: highest_seen.clone(),
}
} else {
ElectionError::Timeout {
required,
promised_votes,
}
}
}
fn register_inbound_drain(
manager: &ConnectionManager,
sender: Sender<InboundSync>,
registry: WriteRegistry,
elections: ElectionRegistry,
catch_ups: CatchUpRegistry,
local_creation: u32,
) {
register_beamr_sync_handler(manager, move |decoded| {
match decoded {
Ok(SyncMessage::WriteAck(ack)) => route_write_ack(®istry, local_creation, &ack),
Ok(SyncMessage::BatchWriteAck(ack)) => {
route_batch_write_ack(®istry, local_creation, &ack);
}
Ok(SyncMessage::Promise(promise)) => {
route_election_vote(
&elections,
promise.shard_id,
ElectionVote::Promised(promise),
);
}
Ok(SyncMessage::Nack(nack)) => {
route_election_vote(&elections, nack.shard_id, ElectionVote::Nacked(nack));
}
Ok(SyncMessage::PushResponse(response)) => {
route_catch_up_response(&catch_ups, response);
}
other => {
let _ = sender.send(other);
}
}
});
}
fn route_election_vote(elections: &ElectionRegistry, shard_id: ShardId, vote: ElectionVote) {
let Some(sender) = elections.get(&shard_id) else {
return;
};
let _ = sender.send(vote);
}
fn route_catch_up_response(catch_ups: &CatchUpRegistry, response: PushResponse) {
let Some(sender) = catch_ups.get(&response.shard_id) else {
return;
};
let _ = sender.send(response);
}
fn route_write_ack(registry: &WriteRegistry, local_creation: u32, ack: &WriteAck) {
route_ack_outcome(
registry,
local_creation,
&ack.write_id,
&ack.acker,
ack.outcome,
);
}
fn route_batch_write_ack(registry: &WriteRegistry, local_creation: u32, ack: &BatchWriteAck) {
route_ack_outcome(
registry,
local_creation,
&ack.write_id,
&ack.acker,
ack.outcome,
);
}
fn route_ack_outcome(
registry: &WriteRegistry,
local_creation: u32,
write_id: &WriteId,
acker: &SyncNodeId,
outcome: AckOutcome,
) {
if write_id.origin_creation != local_creation {
return;
}
let Some(sender) = registry.get(write_id) else {
return;
};
let vote = match outcome {
AckOutcome::Applied => CasVote::Accept(acker.clone()),
AckOutcome::Rejected(RejectReason::Fenced) => {
CasVote::Reject(acker.clone(), RejectKind::EpochFence)
}
AckOutcome::Rejected(RejectReason::CasMismatch) => {
CasVote::Reject(acker.clone(), RejectKind::CasMismatch)
}
AckOutcome::Rejected(RejectReason::ApplyError) => CasVote::Fault(acker.clone()),
};
let _ = sender.send(vote);
}
#[cfg(test)]
#[path = "endpoint_route_ack_tests.rs"]
mod route_ack_tests;