use std::collections::hash_map::DefaultHasher;
use std::collections::{HashMap, HashSet};
use std::fmt::Debug;
use std::hash::{Hash, Hasher};
use std::net::{IpAddr, SocketAddr};
use std::ops::RangeBounds;
use std::sync::atomic::{AtomicU32, AtomicUsize, Ordering};
use std::sync::Arc;
use std::time::{Duration, Instant};
use bincode::{DefaultOptions, Deserializer, Serializer};
use ipnet::IpNet;
use parking_lot::RwLock;
use rand::rngs::StdRng;
use rand::seq::SliceRandom;
use rand::SeedableRng;
use serde::{Deserialize, Serialize};
use socket2::SockRef;
use tokio::net::{ToSocketAddrs, UdpSocket};
use tokio::time::{sleep, timeout};
use tracing::{debug, error, info, instrument, trace, warn};
use crate::auth;
use crate::bounds::{Key, Value};
use crate::clock::{Clock, HlcClock, Timestamp, Timestamped};
use crate::discovery::{Discovery, RandomProbe};
use crate::fingerprint::Fingerprint;
use crate::gen_ip::{host_net, net_of};
use crate::observability;
use crate::proto::{self, HashSegment};
use crate::reconcilable::{MaybeTombstone, Projectable, Reconcilable};
use crate::reconcile_store::{Config, MAX_NETS};
use crate::HRTree;
const BUFFER_SIZE: usize = 65507;
const PEER_EXPIRATION: Duration = Duration::from_secs(60);
const MAX_SENDTO_RETRIES: u32 = 4;
type PreInsertCallback<K, V> = Box<dyn Send + Sync + Fn(&K, &V)>;
pub(crate) fn version_hash<V: Hash>(value: &V) -> u64 {
let mut hasher = DefaultHasher::new();
value.hash(&mut hasher);
hasher.finish()
}
fn derive_local_net(nets: &[IpNet], listen_addr: IpAddr) -> IpNet {
net_of(nets, listen_addr).unwrap_or_else(|| {
warn!(
"listen address {listen_addr} is contained in none of the configured networks \
{nets:?}; cannot identify the local network — treating only this node as local, so \
every peer is remote and reconciled on the throttled cross-network cadence. Declare \
the network containing {listen_addr} via Config::with_net or ReconcileStore::add_net.",
);
host_net(listen_addr)
})
}
fn set_socket_buffers(socket: &UdpSocket, config: &Config) {
let sock = SockRef::from(socket);
if let Some(size) = config.recv_buffer_size {
match sock.set_recv_buffer_size(size) {
Ok(()) => match sock.recv_buffer_size() {
Ok(actual) => debug!(
"gossip socket SO_RCVBUF: requested {size} B, OS granted {actual} B \
(raise net.core.rmem_max if a larger buffer is needed)"
),
Err(e) => debug!("could not read back SO_RCVBUF: {e}"),
},
Err(e) => warn!("failed to set gossip socket SO_RCVBUF to {size} B: {e}"),
}
}
if let Some(size) = config.send_buffer_size {
match sock.set_send_buffer_size(size) {
Ok(()) => match sock.send_buffer_size() {
Ok(actual) => debug!(
"gossip socket SO_SNDBUF: requested {size} B, OS granted {actual} B \
(raise net.core.wmem_max if a larger buffer is needed)"
),
Err(e) => debug!("could not read back SO_SNDBUF: {e}"),
},
Err(e) => warn!("failed to set gossip socket SO_SNDBUF to {size} B: {e}"),
}
}
}
pub(crate) struct ReconcileEngine<K, V: Projectable> {
inner: Arc<Inner<K, V>>,
}
pub(crate) struct Inner<K, V: Projectable> {
pub(crate) map: Arc<RwLock<HRTree<K, V>>>,
pub(crate) projection: Arc<RwLock<HRTree<K, V::Projected>>>,
port: u16,
socket: Arc<UdpSocket>,
nets: Arc<RwLock<Vec<IpNet>>>,
local_net: Arc<RwLock<IpNet>>,
listen_addr: IpAddr,
remote_interval: Arc<AtomicU32>,
remote_fanout: Arc<AtomicUsize>,
reconcile_interval: Arc<RwLock<Duration>>,
bulk_send_rate: Option<usize>,
bulk_in_flight: Arc<RwLock<HashSet<SocketAddr>>>,
round: Arc<AtomicU32>,
rng: Arc<RwLock<StdRng>>,
probe: Arc<dyn Discovery>,
pub(crate) peers: Arc<RwLock<HashMap<IpAddr, Instant>>>,
pub(crate) pre_insert: Arc<RwLock<PreInsertCallback<K, V>>>,
authenticator: auth::Authenticator,
pub(crate) members: Arc<RwLock<HashSet<IpAddr>>>,
pub(crate) tombstone_acks: Arc<RwLock<HashMap<K, HashMap<IpAddr, u64>>>>,
clock: Arc<dyn Clock>,
}
impl<K, V: Projectable> Clone for ReconcileEngine<K, V> {
fn clone(&self) -> Self {
ReconcileEngine {
inner: Arc::clone(&self.inner),
}
}
}
impl<K, V: Projectable> std::ops::Deref for ReconcileEngine<K, V> {
type Target = Inner<K, V>;
fn deref(&self) -> &Self::Target {
&self.inner
}
}
#[derive(Clone, Debug, Deserialize, Serialize)]
pub(crate) enum Message<K: Serialize, V: Serialize, P: Serialize> {
ComparisonItem(HashSegment<K>),
Update((K, V)),
Ack((K, u64)),
ValueComparisonItem(HashSegment<K>),
ValueUpdate((K, P)),
}
impl<K: Key, V: Value + MaybeTombstone + Projectable + Reconcilable + Timestamped>
ReconcileEngine<K, V>
{
pub async fn new(config: Config) -> Self {
let node_id = config.node_id.unwrap_or_else(rand::random);
let clock: Arc<dyn Clock> = Arc::new(HlcClock::new(node_id));
Self::build(config, clock).await
}
#[cfg(test)]
pub(crate) async fn new_with_clock(config: Config, clock: Arc<dyn Clock>) -> Self {
Self::build(config, clock).await
}
async fn build(config: Config, clock: Arc<dyn Clock>) -> Self {
let socket = UdpSocket::bind(SocketAddr::new(config.listen_addr, config.port))
.await
.unwrap();
info!("Listening on: {}", socket.local_addr().unwrap());
set_socket_buffers(&socket, &config);
let authenticator = auth::Authenticator::new(config.cluster_key, config.encrypt);
if authenticator.is_encrypted() {
debug!("per-datagram authenticated encryption (XChaCha20-Poly1305) ENABLED");
} else if authenticator.is_enabled() {
debug!("per-datagram MAC authentication ENABLED");
} else {
warn!(
"SECURITY: no cluster key set — UDP reconciliation is UNAUTHENTICATED. Any host \
that can send UDP to this port can forge updates and poison the cluster via \
last-write-wins. Set Config::with_cluster_key on every node, or restrict the \
network to a trusted underlay. See REVIEW.md F3."
);
}
let map = HRTree::<K, V>::new();
let projection = HRTree::<K, V::Projected>::new();
let mut nets: Vec<IpNet> = config.nets.iter().flatten().copied().collect();
if nets.is_empty() {
nets.push("127.0.0.1/8".parse().unwrap());
}
let local_net = derive_local_net(&nets, config.listen_addr);
let nets = Arc::new(RwLock::new(nets));
let rng = Arc::new(RwLock::new(StdRng::from_entropy()));
let probe: Arc<dyn Discovery> =
Arc::new(RandomProbe::new(Arc::clone(&nets), Arc::clone(&rng)));
ReconcileEngine {
inner: Arc::new(Inner {
map: Arc::new(RwLock::new(map)),
projection: Arc::new(RwLock::new(projection)),
port: config.port,
socket: Arc::new(socket),
nets,
local_net: Arc::new(RwLock::new(local_net)),
listen_addr: config.listen_addr,
remote_interval: Arc::new(AtomicU32::new(config.remote_interval)),
remote_fanout: Arc::new(AtomicUsize::new(config.remote_fanout)),
reconcile_interval: Arc::new(RwLock::new(config.reconcile_interval)),
bulk_send_rate: config.bulk_send_rate,
bulk_in_flight: Arc::new(RwLock::new(HashSet::new())),
round: Arc::new(AtomicU32::new(0)),
rng,
probe,
peers: Arc::new(RwLock::new(HashMap::new())),
pre_insert: Arc::new(RwLock::new(Box::new(|_, _| {}))),
authenticator,
members: Arc::new(RwLock::new(HashSet::new())),
tombstone_acks: Arc::new(RwLock::new(HashMap::new())),
clock,
}),
}
}
pub fn fingerprint<R: RangeBounds<K>>(&self, range: R) -> Fingerprint {
self.map.read().hash(&range)
}
pub fn value_fingerprint<R: RangeBounds<K>>(&self, range: R) -> Fingerprint {
self.projection.read().hash(&range)
}
fn map_insert(&self, guard: &mut HRTree<K, V>, key: K, value: V) -> Option<V> {
self.projection.write().insert(key.clone(), value.project());
guard.insert(key, value)
}
pub(crate) fn gc_remove(&self, key: &K) -> Option<V> {
let mut guard = self.map.write();
self.projection.write().remove(key);
guard.remove(key)
}
pub fn clock_now(&self) -> Timestamp {
self.clock.now()
}
pub(crate) fn set_nets(&self, nets: &[IpNet]) {
let nets = nets.to_vec();
let local = derive_local_net(&nets, self.listen_addr);
*self.local_net.write() = local;
*self.nets.write() = nets;
}
pub(crate) fn add_net(&self, net: IpNet) -> bool {
let mut guard = self.nets.write();
if guard.contains(&net) {
return true;
}
if guard.len() >= MAX_NETS {
warn!("cannot add network {net}: already at the maximum of {MAX_NETS} networks");
return false;
}
guard.push(net);
*self.local_net.write() = derive_local_net(&guard, self.listen_addr);
true
}
pub(crate) fn remove_net(&self, net: IpNet) -> bool {
let mut guard = self.nets.write();
let before = guard.len();
guard.retain(|n| *n != net);
let removed = guard.len() != before;
if removed {
*self.local_net.write() = derive_local_net(&guard, self.listen_addr);
}
removed
}
pub(crate) fn nets(&self) -> Vec<IpNet> {
self.nets.read().clone()
}
pub(crate) fn local_net(&self) -> IpNet {
*self.local_net.read()
}
pub(crate) fn set_remote_interval(&self, interval: u32) {
self.remote_interval.store(interval, Ordering::Relaxed);
}
pub(crate) fn set_remote_fanout(&self, fanout: usize) {
self.remote_fanout.store(fanout, Ordering::Relaxed);
}
pub(crate) fn set_reconcile_interval(&self, interval: Duration) {
*self.reconcile_interval.write() = interval;
}
fn get_peers(&self) -> Vec<IpAddr> {
let mut guard = self.peers.write();
guard.retain(|_, instant| instant.elapsed() < PEER_EXPIRATION);
guard.keys().cloned().collect()
}
pub fn just_insert(&self, key: K, value: V) -> Option<V> {
(self.pre_insert.read())(&key, &value);
if value.is_tombstone() {
observability::record_remove();
} else {
observability::record_insert();
}
let mut guard = self.map.write();
self.map_insert(&mut guard, key, value)
}
fn broadcast(&self, messages: Vec<Message<K, V, V::Projected>>) {
let peers = self.get_peers();
let port = self.port;
let socket = Arc::clone(&self.socket);
let authenticator = self.authenticator.clone();
tokio::spawn(async move {
let mut send_buf = Vec::new();
for addr in peers {
let peer = SocketAddr::new(addr, port);
send_messages_to(
&messages,
Arc::clone(&socket),
&authenticator,
&peer,
&mut send_buf,
)
.await;
}
});
}
fn spawn_paced_send(&self, messages: Vec<Message<K, V, V::Projected>>, peer: SocketAddr) {
if !self.bulk_in_flight.write().insert(peer) {
return;
}
let guard = BulkInFlightGuard {
set: Arc::clone(&self.bulk_in_flight),
peer,
};
let socket = Arc::clone(&self.socket);
let authenticator = self.authenticator.clone();
let rate = self.bulk_send_rate;
tokio::spawn(async move {
let _guard = guard;
let mut send_buf = Vec::new();
send_messages_paced(
&messages,
socket,
&authenticator,
&peer,
&mut send_buf,
rate,
)
.await;
});
}
pub fn insert(&self, key: K, value: V) -> Option<V> {
let ret = self.just_insert(key.clone(), value.clone());
self.broadcast(vec![Message::Update::<K, V, V::Projected>((key, value))]);
ret
}
pub(crate) fn broadcast_update(&self, key: K, value: V) {
self.broadcast(vec![Message::Update::<K, V, V::Projected>((key, value))]);
}
pub fn just_insert_bulk(&self, key_values: &[(K, V)]) {
for (key, value) in key_values {
(self.pre_insert.read())(key, value);
if value.is_tombstone() {
observability::record_remove();
} else {
observability::record_insert();
}
}
let mut guard = self.map.write();
for (key, value) in key_values {
self.map_insert(&mut guard, key.clone(), value.clone());
}
}
pub fn insert_bulk(&self, key_values: &[(K, V)]) {
self.just_insert_bulk(key_values);
let messages: Vec<_> = key_values
.iter()
.map(|kv| Message::Update::<K, V, V::Projected>(kv.clone()))
.collect();
self.broadcast(messages);
}
#[instrument(name = "reconcile.run", skip_all, fields(port = self.port))]
pub async fn run(self) {
let mut recv_buf = [0; BUFFER_SIZE + 1];
let mut send_buf = Vec::new();
self.start_reconciliation(&mut send_buf).await;
loop {
let recv_timeout = *self.reconcile_interval.read();
match timeout(recv_timeout, self.socket.recv_from(&mut recv_buf)).await {
Err(_) => {
debug!("no recent activity; initiating diff protocol");
self.start_reconciliation(&mut send_buf).await;
}
Ok(Err(err)) => {
warn!("network error in recv_from: {err}");
observability::record_datagram_dropped("recv_error");
}
Ok(Ok((size, peer))) => {
observability::record_bytes_received(size);
if peer.port() != self.port {
warn!(
"received message from {peer}, but protocol port is {}",
self.port
);
}
if size == recv_buf.len() {
warn!("Buffer too small for message, discarded");
observability::record_datagram_dropped("too_large");
} else {
match self.authenticator.open(&recv_buf[..size]) {
Some(payload) => {
let spoke_dated =
self.handle_messages(payload, peer, &mut send_buf).await;
if spoke_dated {
let addr = peer.ip();
self.peers.write().insert(addr, Instant::now());
self.members.write().insert(addr);
}
}
None => {
trace!("dropped datagram from {peer}: missing or invalid MAC");
observability::record_datagram_dropped("bad_mac");
}
}
}
}
}
}
}
#[instrument(name = "reconcile.round", skip_all)]
pub async fn start_reconciliation(&self, send_buf: &mut Vec<u8>) {
let timer = observability::timer();
observability::record_reconcile_round();
let segments = {
let guard = self.map.read();
proto::start_diff(&guard)
};
send_buf.clear();
for segment in segments {
Message::ComparisonItem::<K, V, V::Projected>(segment)
.serialize(&mut Serializer::new(&mut *send_buf, DefaultOptions::new()))
.expect("serializing a ComparisonItem into an in-memory buffer cannot fail");
}
let nets = self.nets.read().clone();
let local = *self.local_net.read();
let remote_interval = self.remote_interval.load(Ordering::Relaxed).max(1);
let remote_fanout = self.remote_fanout.load(Ordering::Relaxed);
let round = self.round.fetch_add(1, Ordering::Relaxed);
let do_remote = round.is_multiple_of(remote_interval);
let known = self.get_peers();
let mut targets: HashSet<IpAddr> = HashSet::new();
targets.extend(self.probe.discover().await.unwrap_or_default());
for &addr in &known {
if local.contains(&addr) {
targets.insert(addr);
}
}
if do_remote {
let remote_nets: Vec<IpNet> = nets.iter().copied().filter(|&n| n != local).collect();
let mut buckets: HashMap<Option<usize>, Vec<IpAddr>> = HashMap::new();
for &addr in &known {
if local.contains(&addr) {
continue; }
let bucket = remote_nets.iter().position(|n| n.contains(&addr));
buckets.entry(bucket).or_default().push(addr);
}
let mut rng = self.rng.write();
for (_, mut peers) in buckets {
peers.shuffle(&mut *rng);
targets.extend(peers.into_iter().take(remote_fanout));
}
}
for peer in targets {
trace!("start_diff {} bytes to {peer}", send_buf.len());
send_to_retry(
&self.socket,
&self.authenticator,
send_buf,
(peer, self.port),
)
.await
.unwrap();
}
observability::record_round_duration(timer);
}
#[instrument(name = "reconcile.handle", skip_all, fields(peer = %peer))]
async fn handle_messages(
&self,
payload: auth::Payload<'_>,
peer: SocketAddr,
send_buf: &mut Vec<u8>,
) -> bool {
let timer = observability::timer();
let payload = payload.as_bytes();
trace!("received {} bytes from {peer}", payload.len());
let mut in_comparison = Vec::new();
let mut updates: Vec<(K, V)> = Vec::new();
let mut acks: Vec<(K, u64)> = Vec::new();
let mut value_in_comparison = Vec::new();
let mut deserializer = Deserializer::from_slice(payload, DefaultOptions::new());
loop {
match Message::<K, V, V::Projected>::deserialize(&mut deserializer) {
Err(ref kind) => {
if let bincode::ErrorKind::Io(err) = kind.as_ref() {
if err.kind() == std::io::ErrorKind::UnexpectedEof {
break;
}
}
warn!("failed to deserialize datagram from {peer}, dropping it: {kind:?}");
observability::record_datagram_dropped("malformed");
break;
}
Ok(Message::ComparisonItem(segment)) => in_comparison.push(segment),
Ok(Message::Update(update)) => updates.push(update),
Ok(Message::Ack(ack)) => acks.push(ack),
Ok(Message::ValueComparisonItem(segment)) => value_in_comparison.push(segment),
Ok(Message::ValueUpdate(_)) => {}
}
}
let spoke_dated = !in_comparison.is_empty() || !updates.is_empty() || !acks.is_empty();
if !acks.is_empty() {
let peer_ip = peer.ip();
let mut reciprocal_acks = Vec::new();
{
let map_guard = self.map.read();
let mut guard = self.tombstone_acks.write();
for (key, version) in acks {
let entry = guard.entry(key.clone()).or_default();
let already = entry.get(&peer_ip) == Some(&version);
entry.insert(peer_ip, version);
if !already {
let holds_same_tombstone = map_guard
.get(&key)
.is_some_and(|v| v.is_tombstone() && version_hash(v) == version);
if holds_same_tombstone {
reciprocal_acks
.push(Message::Ack::<K, V, V::Projected>((key, version)));
}
}
}
}
if !reciprocal_acks.is_empty() {
send_messages_to(
&reciprocal_acks,
Arc::clone(&self.socket),
&self.authenticator,
&peer,
send_buf,
)
.await;
}
}
if !in_comparison.is_empty() {
debug!("received {} segments", in_comparison.len());
let mut differences = Vec::new();
let mut out_comparison = Vec::new();
{
let guard = self.map.read();
proto::diff_round(&guard, in_comparison, &mut out_comparison, &mut differences);
}
if !out_comparison.is_empty() {
debug!("returning {} segments", out_comparison.len());
trace!("segments: {out_comparison:?}");
let messages: Vec<_> = out_comparison
.into_iter()
.map(Message::ComparisonItem::<K, V, V::Projected>)
.collect();
send_messages_to(
&messages,
Arc::clone(&self.socket),
&self.authenticator,
&peer,
send_buf,
)
.await;
}
if !differences.is_empty() {
debug!("returning {} diff_ranges", differences.len());
trace!("diff_ranges: {differences:?}");
let updates: Vec<Message<K, V, V::Projected>> = {
let guard = self.map.read();
let mut updates = Vec::new();
for range in differences {
for (k, v) in guard.get_range(&range) {
updates.push(Message::Update((k.clone(), v.clone())));
}
}
updates
};
if !updates.is_empty() {
self.spawn_paced_send(updates, peer);
}
}
}
if !updates.is_empty() {
debug!("received {} updates", updates.len());
observability::record_updates_received(updates.len());
let mut acks_to_send = Vec::new();
let mut to_apply: Vec<(K, V)> = Vec::new();
{
let guard = self.map.read();
for (k, remote_v) in updates {
self.clock.observe(remote_v.timestamp());
match guard.get(&k) {
Some(local_v) => {
let merged_v = local_v.reconcile(&remote_v);
if merged_v != *local_v {
to_apply.push((k, merged_v));
} else if local_v.is_tombstone() {
acks_to_send.push(Message::Ack::<K, V, V::Projected>((
k,
version_hash(local_v),
)));
}
}
None => to_apply.push((k, remote_v)),
}
}
}
for (k, v) in &to_apply {
(self.pre_insert.read())(k, v);
}
if !to_apply.is_empty() {
let mut guard = self.map.write();
for (k, v) in to_apply {
let merged_v = match guard.get(&k) {
Some(local_v) => local_v.reconcile(&v),
None => v,
};
let version = merged_v.is_tombstone().then(|| version_hash(&merged_v));
self.map_insert(&mut guard, k.clone(), merged_v);
if let Some(version) = version {
acks_to_send.push(Message::Ack::<K, V, V::Projected>((k, version)));
}
}
}
if !acks_to_send.is_empty() {
send_messages_to(
&acks_to_send,
Arc::clone(&self.socket),
&self.authenticator,
&peer,
send_buf,
)
.await;
}
}
if !value_in_comparison.is_empty() {
debug!("received {} value-only segments", value_in_comparison.len());
let mut differences = Vec::new();
let mut out_comparison = Vec::new();
{
let guard = self.projection.read();
proto::diff_round(
&guard,
value_in_comparison,
&mut out_comparison,
&mut differences,
);
}
if !out_comparison.is_empty() {
let messages: Vec<_> = out_comparison
.into_iter()
.map(Message::ValueComparisonItem::<K, V, V::Projected>)
.collect();
send_messages_to(
&messages,
Arc::clone(&self.socket),
&self.authenticator,
&peer,
send_buf,
)
.await;
}
if !differences.is_empty() {
let updates: Vec<Message<K, V, V::Projected>> = {
let guard = self.projection.read();
let mut updates = Vec::new();
for range in differences {
for (k, p) in guard.get_range(&range) {
updates.push(Message::ValueUpdate((k.clone(), p.clone())));
}
}
updates
};
if !updates.is_empty() {
self.spawn_paced_send(updates, peer);
}
}
}
observability::record_handle_duration(timer);
spoke_dated
}
pub(crate) fn is_tombstone_stable(&self, key: &K, version: u64) -> bool {
let members = self.members.read();
if members.is_empty() {
return true;
}
let acks = self.tombstone_acks.read();
let Some(key_acks) = acks.get(key) else {
return false;
};
members
.iter()
.all(|peer| key_acks.get(peer) == Some(&version))
}
pub(crate) fn forget_tombstone(&self, key: &K) {
self.tombstone_acks.write().remove(key);
}
pub(crate) fn decommission_peer(&self, peer: IpAddr) {
self.members.write().remove(&peer);
for key_acks in self.tombstone_acks.write().values_mut() {
key_acks.remove(&peer);
}
}
pub(crate) fn seed_peer(&self, peer: IpAddr) {
self.peers.write().insert(peer, Instant::now());
}
pub(crate) fn members_snapshot(&self) -> HashSet<IpAddr> {
self.members.read().clone()
}
pub(crate) fn listen_addr(&self) -> IpAddr {
self.listen_addr
}
}
pub(crate) async fn send_to_retry<A: ToSocketAddrs>(
socket: &UdpSocket,
authenticator: &auth::Authenticator,
buf: &[u8],
target: A,
) -> std::io::Result<usize> {
let framed = authenticator.seal(buf);
let wire: &[u8] = framed.as_deref().unwrap_or(buf);
let mut res = Ok(0);
for _ in 0..MAX_SENDTO_RETRIES {
res = socket.send_to(wire, &target).await;
if res.is_ok() {
break;
}
tokio::time::sleep(Duration::from_millis(1)).await;
}
match &res {
Ok(sent) => observability::record_bytes_sent(*sent),
Err(err) => {
error!("send_to failed after {MAX_SENDTO_RETRIES} retries: {err}");
observability::record_send_failure();
}
}
res
}
pub(crate) async fn send_messages_to<K: Serialize, V: Serialize, P: Serialize>(
messages: &[Message<K, V, P>],
socket: Arc<UdpSocket>,
authenticator: &auth::Authenticator,
peer: &SocketAddr,
send_buf: &mut Vec<u8>,
) {
send_messages_paced(messages, socket, authenticator, peer, send_buf, None).await
}
#[instrument(name = "reconcile.send", skip_all, fields(peer = %peer, count = messages.len()))]
pub(crate) async fn send_messages_paced<K: Serialize, V: Serialize, P: Serialize>(
messages: &[Message<K, V, P>],
socket: Arc<UdpSocket>,
authenticator: &auth::Authenticator,
peer: &SocketAddr,
send_buf: &mut Vec<u8>,
rate: Option<usize>,
) {
debug!("sending {} messages to {peer}", messages.len());
let max_payload = BUFFER_SIZE - authenticator.overhead();
send_buf.clear();
let start = Instant::now();
let mut sent_bytes: usize = 0;
for message in messages {
let last_size = send_buf.len();
message
.serialize(&mut Serializer::new(&mut *send_buf, DefaultOptions::new()))
.expect("serializing a protocol Message into an in-memory buffer cannot fail");
if send_buf.len() > max_payload {
trace!("sending {} bytes to {peer}", last_size);
send_to_retry(&socket, authenticator, &send_buf[..last_size], peer)
.await
.unwrap();
trace!("sent {} bytes to {peer}", last_size);
send_buf.drain(..last_size);
sent_bytes += last_size;
pace(rate, start, sent_bytes).await;
}
}
trace!("sending last {} bytes to {peer}", send_buf.len());
send_to_retry(&socket, authenticator, send_buf, peer)
.await
.unwrap();
trace!("sent last {} bytes to {peer}", send_buf.len());
}
async fn pace(rate: Option<usize>, start: Instant, sent_bytes: usize) {
let Some(rate) = rate.filter(|&r| r > 0) else {
return;
};
let expected = Duration::from_secs_f64(sent_bytes as f64 / rate as f64);
if let Some(delay) = expected.checked_sub(start.elapsed()) {
sleep(delay).await;
}
}
struct BulkInFlightGuard {
set: Arc<RwLock<HashSet<SocketAddr>>>,
peer: SocketAddr,
}
impl Drop for BulkInFlightGuard {
fn drop(&mut self) {
self.set.write().remove(&self.peer);
}
}
#[cfg(test)]
mod deadlock_regressions {
use crate::auth;
use crate::clock::Timestamp;
use crate::reconcilable::ValueOnly;
use crate::reconcile_engine::ReconcileEngine;
use crate::{reconcile_store::Config, ReconcileStore};
use bincode::{DefaultOptions, Serializer};
use serde::Serialize;
use std::net::SocketAddr;
use std::sync::{
atomic::{AtomicBool, Ordering},
mpsc, Arc,
};
use std::time::Duration;
use super::Message;
#[tokio::test(flavor = "multi_thread")]
async fn pre_insert_hook_can_call_insert_again_without_deadlock() {
let config = Config::default()
.with_port(8080)
.with_listen_addr("127.0.0.44".parse().unwrap());
let svc = ReconcileStore::new(config).await;
svc.insert_bulk(&[(1, 10_u8)]);
let flag = Arc::new(AtomicBool::new(false));
let flag2 = flag.clone();
let hook_svc = svc.clone();
let once = Arc::new(AtomicBool::new(false));
let guard = once.clone();
svc.add_pre_insert(move |&k, &v| {
if !guard.swap(true, Ordering::SeqCst) {
let _ = hook_svc.just_insert(k + 100, v.1.unwrap_or_default() + 100);
}
flag2.store(true, Ordering::SeqCst);
});
let _ = svc.just_insert(42, 99);
assert!(
flag.load(Ordering::SeqCst),
"The pre-insert hook never ran to completion (likely deadlocked)"
);
}
fn update_message_bytes(key: i32, value: (Timestamp, Option<u8>)) -> Vec<u8> {
let message = Message::Update::<i32, (Timestamp, Option<u8>), ValueOnly<u8>>((key, value));
let mut buf = Vec::new();
message
.serialize(&mut Serializer::new(&mut buf, DefaultOptions::new()))
.unwrap();
buf
}
#[test]
fn pre_insert_hook_can_call_insert_again_from_network_path_without_deadlock() {
let (tx, rx) = mpsc::channel();
std::thread::spawn(move || {
let rt = tokio::runtime::Builder::new_current_thread()
.enable_all()
.build()
.unwrap();
let reinserted = rt.block_on(async {
let config = Config::default()
.with_port(8083)
.with_listen_addr("127.0.0.50".parse().unwrap());
let engine = ReconcileEngine::<i32, (Timestamp, Option<u8>)>::new(config).await;
let hook_engine = engine.clone();
let once = Arc::new(AtomicBool::new(false));
let guard = once.clone();
*engine.pre_insert.write() =
Box::new(move |&k: &i32, v: &(Timestamp, Option<u8>)| {
if !guard.swap(true, Ordering::SeqCst) {
let inner = (
Timestamp::new(u64::MAX, 1, 0),
Some(v.1.unwrap_or_default() + 100),
);
let _ = hook_engine.just_insert(k + 100, inner);
}
});
let bytes = update_message_bytes(42, (Timestamp::new(u64::MAX, 0, 0), Some(99)));
let payload = auth::Authenticator::new(None, false)
.open(&bytes)
.expect("unauthenticated mode clears any datagram");
let peer: SocketAddr = "127.0.0.51:8083".parse().unwrap();
let mut send_buf = Vec::new();
engine.handle_messages(payload, peer, &mut send_buf).await;
let map_guard = engine.map.read();
let reinserted = map_guard.get(&142).and_then(|(_, v)| *v);
drop(map_guard);
reinserted
});
let _ = tx.send(reinserted);
});
match rx.recv_timeout(Duration::from_secs(5)) {
Ok(reinserted) => assert_eq!(
reinserted,
Some(199),
"the re-entrant insert from the network-path hook did not take effect"
),
Err(_) => panic!(
"the network-path pre-insert hook deadlocked (it ran under the map write lock, so \
its re-entrant insert could not re-acquire the lock)"
),
}
}
}
#[cfg(test)]
mod auth_attack {
use std::time::Duration;
use bincode::{DefaultOptions, Serializer};
use serde::Serialize;
use tokio::net::UdpSocket;
use super::Message;
use crate::clock::Timestamp;
use crate::{auth, reconcile_store::Config, ReconcileStore};
fn forged_update() -> Vec<u8> {
let far_future = Timestamp::new(u64::MAX, 0, 0);
let message = Message::Update::<
i32,
(Timestamp, Option<String>),
crate::reconcilable::ValueOnly<String>,
>((0, (far_future, Some("evil".to_string()))));
let mut buf = Vec::new();
message
.serialize(&mut Serializer::new(&mut buf, DefaultOptions::new()))
.unwrap();
buf
}
#[tokio::test(flavor = "multi_thread")]
async fn forged_datagram_is_ignored() {
let key = [0x42u8; auth::KEY_LEN];
let port = 8082;
let victim_addr = "127.0.0.48";
let config = Config::default()
.with_port(port)
.with_listen_addr(victim_addr.parse().unwrap())
.with_cluster_key(key);
let store = ReconcileStore::<i32, String>::new(config).await;
store.just_insert(0, "legit".to_string());
let task = tokio::spawn(store.clone().run());
let attacker = UdpSocket::bind("127.0.0.49:0").await.unwrap();
let target = format!("{victim_addr}:{port}");
let forged = forged_update();
attacker.send_to(&forged, &target).await.unwrap();
let wrong_key_sealed = auth::Authenticator::new(Some([0x99u8; auth::KEY_LEN]), false)
.seal(&forged)
.expect("enabled");
attacker.send_to(&wrong_key_sealed, &target).await.unwrap();
tokio::time::sleep(Duration::from_millis(200)).await;
assert_eq!(store.get(&0).as_deref(), Some(&"legit".to_string()));
task.abort();
}
}
#[cfg(test)]
mod causal_stability {
use std::net::IpAddr;
use crate::clock::Timestamp;
use crate::reconcile_engine::{version_hash, ReconcileEngine};
use crate::reconcile_store::Config;
type Tombstoned = (Timestamp, Option<i32>);
async fn engine(addr: &str) -> ReconcileEngine<i32, Tombstoned> {
let config = Config::default()
.with_port(8080)
.with_listen_addr(addr.parse().unwrap());
ReconcileEngine::new(config).await
}
#[tokio::test]
async fn tombstone_not_stable_until_all_members_ack() {
let eng = engine("127.0.0.60").await;
let peer_a: IpAddr = "127.0.0.61".parse().unwrap();
let peer_b: IpAddr = "127.0.0.62".parse().unwrap();
let key = 7;
let tombstone: Tombstoned = (Timestamp::new(1, 0, 0), None);
let version = version_hash(&tombstone);
assert!(eng.is_tombstone_stable(&key, version));
eng.members.write().insert(peer_a);
eng.members.write().insert(peer_b);
assert!(!eng.is_tombstone_stable(&key, version));
eng.tombstone_acks
.write()
.entry(key)
.or_default()
.insert(peer_a, version);
assert!(!eng.is_tombstone_stable(&key, version));
eng.tombstone_acks
.write()
.entry(key)
.or_default()
.insert(peer_b, version.wrapping_add(1));
assert!(!eng.is_tombstone_stable(&key, version));
eng.tombstone_acks
.write()
.entry(key)
.or_default()
.insert(peer_b, version);
assert!(eng.is_tombstone_stable(&key, version));
}
#[tokio::test]
async fn decommission_releases_a_silent_peer() {
let eng = engine("127.0.0.63").await;
let live: IpAddr = "127.0.0.64".parse().unwrap();
let gone: IpAddr = "127.0.0.65".parse().unwrap();
let key = 9;
let tombstone: Tombstoned = (Timestamp::new(1, 0, 0), None);
let version = version_hash(&tombstone);
eng.members.write().insert(live);
eng.members.write().insert(gone);
eng.tombstone_acks
.write()
.entry(key)
.or_default()
.insert(live, version);
assert!(!eng.is_tombstone_stable(&key, version));
eng.decommission_peer(gone);
assert!(eng.is_tombstone_stable(&key, version));
eng.forget_tombstone(&key);
assert!(eng.tombstone_acks.read().get(&key).is_none());
}
}
#[cfg(test)]
mod clock_port {
use std::sync::Arc;
use crate::clock::{ManualClock, Timestamp};
use crate::reconcile_engine::ReconcileEngine;
use crate::reconcile_store::Config;
#[tokio::test]
async fn engine_mints_through_the_injected_clock() {
let config = Config::default()
.with_port(8080)
.with_listen_addr("127.0.0.70".parse().unwrap());
let clock = Arc::new(ManualClock::new(42));
let eng: ReconcileEngine<i32, (Timestamp, Option<i32>)> =
ReconcileEngine::new_with_clock(config, clock).await;
assert_eq!(eng.clock_now(), Timestamp::new(0, 1, 42));
assert_eq!(eng.clock_now(), Timestamp::new(0, 2, 42));
}
}
#[cfg(test)]
mod pacing {
use std::net::SocketAddr;
use std::sync::Arc;
use std::time::{Duration, Instant};
use tokio::net::UdpSocket;
use crate::auth::Authenticator;
use crate::reconcilable::ValueOnly;
use crate::reconcile_engine::{send_messages_paced, Message};
type Msg = Message<u64, Vec<u8>, ValueOnly<u8>>;
fn bulk_updates(n: u64, value_len: usize) -> Vec<Msg> {
(0..n)
.map(|k| Message::Update((k, vec![0u8; value_len])))
.collect()
}
async fn time_send(messages: &[Msg], rate: Option<usize>) -> Duration {
let socket = Arc::new(UdpSocket::bind("127.0.0.1:0").await.unwrap());
let authenticator = Authenticator::new(None, false);
let peer: SocketAddr = "127.0.0.1:9".parse().unwrap(); let mut send_buf = Vec::new();
let start = Instant::now();
send_messages_paced(messages, socket, &authenticator, &peer, &mut send_buf, rate).await;
start.elapsed()
}
#[tokio::test]
async fn bulk_send_rate_meters_the_transfer() {
let messages = bulk_updates(256, 1024);
let unpaced = time_send(&messages, None).await;
assert!(
unpaced < Duration::from_millis(200),
"unpaced send should be near-instant, took {unpaced:?}"
);
let paced = time_send(&messages, Some(512 * 1024)).await;
assert!(
paced >= Duration::from_millis(300),
"paced send should be metered to ~0.5 s, took {paced:?}"
);
}
#[tokio::test]
async fn zero_or_none_rate_does_not_pace() {
let messages = bulk_updates(256, 1024);
assert!(time_send(&messages, None).await < Duration::from_millis(200));
assert!(time_send(&messages, Some(0)).await < Duration::from_millis(200));
}
}
#[cfg(test)]
mod socket_buffers {
use socket2::SockRef;
use crate::clock::Timestamp;
use crate::reconcile_engine::ReconcileEngine;
use crate::reconcile_store::Config;
type Tombstoned = (Timestamp, Option<i32>);
async fn engine(addr: &str, config: Config) -> ReconcileEngine<i32, Tombstoned> {
ReconcileEngine::new(config.with_listen_addr(addr.parse().unwrap())).await
}
#[tokio::test]
async fn recv_buffer_size_is_configurable() {
let big = engine("127.0.0.90", Config::default()).await;
let small = engine(
"127.0.0.91",
Config::default().with_recv_buffer_size(8 * 1024),
)
.await;
let big_buf = SockRef::from(&*big.socket).recv_buffer_size().unwrap();
let small_buf = SockRef::from(&*small.socket).recv_buffer_size().unwrap();
assert!(
big_buf > small_buf,
"the multi-MiB default ({big_buf} B) should exceed an explicitly tiny buffer \
({small_buf} B)"
);
}
#[tokio::test]
async fn recv_buffer_size_none_leaves_os_default() {
let config = Config {
recv_buffer_size: None,
..Config::default()
};
let eng = engine("127.0.0.92", config).await;
let buf = SockRef::from(&*eng.socket).recv_buffer_size().unwrap();
assert!(buf > 0, "a socket always has a positive receive buffer");
}
}