use std::collections::HashSet;
use std::hash::Hash;
use std::net::SocketAddr;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::Arc;
use std::time::{Duration, Instant};
use parking_lot::RwLock;
use serde::Serialize;
use tokio::time::sleep;
use tracing::{debug, error, instrument, trace, warn};
use crate::bounds::{Key, Value};
use crate::clock::Timestamp;
use crate::entry::{Entry, State};
use crate::observability;
use crate::transport::Transport;
use gossip::auth;
use gossip::replay;
use super::{Message, Replica, BUFFER_SIZE, MAX_SENDTO_RETRIES};
impl<K: Key + Hash, V: Value> Replica<K, V> {
pub(super) fn send_ports(&self) -> SendPorts<'_, dyn Transport<Addr = SocketAddr>> {
SendPorts {
transport: &*self.transport,
authenticator: &self.authenticator,
sender_counter: &self.sender_counter,
}
}
pub(super) fn try_claim_dump_slot(
&self,
peer: SocketAddr,
) -> Option<(BulkInFlightGuard, BulkDumpCountGuard)> {
if !self.bulk_in_flight.write().insert(peer) {
return None;
}
let budget = self.max_concurrent_bulk_dumps;
let claimed = self
.bulk_dumps_in_flight
.fetch_update(Ordering::AcqRel, Ordering::Acquire, |n| {
if n < budget {
Some(n + 1)
} else {
None
}
})
.is_ok();
if !claimed {
self.bulk_in_flight.write().remove(&peer);
trace!("skipped bulk dump to {peer}: global dump budget ({budget}) exhausted");
return None;
}
Some((
BulkInFlightGuard {
set: Arc::clone(&self.bulk_in_flight),
peer,
},
BulkDumpCountGuard {
counter: Arc::clone(&self.bulk_dumps_in_flight),
},
))
}
pub(super) fn spawn_paced_send(
&self,
messages: Vec<Message<K, Entry<Timestamp, V>, State<V>>>,
peer: SocketAddr,
peer_guard: BulkInFlightGuard,
global_guard: BulkDumpCountGuard,
) {
let transport = Arc::clone(&self.transport);
let authenticator = self.authenticator.clone();
let sender_counter = Arc::clone(&self.sender_counter);
let rate = self.bulk_send_rate;
tokio::spawn(async move {
let _peer_guard = peer_guard;
let _global_guard = global_guard;
let ports = SendPorts {
transport: &*transport,
authenticator: &authenticator,
sender_counter: &sender_counter,
};
let mut send_buf = Vec::new();
send_messages_paced(&messages, &ports, &peer, &mut send_buf, rate).await;
});
}
}
pub(crate) struct SendPorts<'a, T: ?Sized> {
pub(crate) transport: &'a T,
pub(crate) authenticator: &'a auth::Authenticator,
pub(crate) sender_counter: &'a replay::SenderCounter,
}
pub(crate) async fn send_to_retry<T: Transport<Addr = SocketAddr> + ?Sized>(
transport: &T,
authenticator: &auth::Authenticator,
sender_counter: &replay::SenderCounter,
buf: &[u8],
target: SocketAddr,
) -> std::io::Result<usize> {
let seq = sender_counter.next_seq();
let stamp = sender_counter.next_stamp();
let wire = authenticator.seal(seq, stamp, buf);
let mut res = Ok(0);
for _ in 0..MAX_SENDTO_RETRIES {
res = transport.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, V, P, T>(
messages: &[Message<K, V, P>],
ports: &SendPorts<'_, T>,
peer: &SocketAddr,
send_buf: &mut Vec<u8>,
) where
K: Serialize,
V: Serialize,
P: Serialize,
T: Transport<Addr = SocketAddr> + ?Sized,
{
send_messages_paced(messages, ports, 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, V, P, T>(
messages: &[Message<K, V, P>],
ports: &SendPorts<'_, T>,
peer: &SocketAddr,
send_buf: &mut Vec<u8>,
rate: Option<usize>,
) where
K: Serialize,
V: Serialize,
P: Serialize,
T: Transport<Addr = SocketAddr> + ?Sized,
{
debug!("sending {} messages to {peer}", messages.len());
let max_payload = BUFFER_SIZE - ports.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();
gossip::bincode::encode(message, send_buf)
.expect("serializing a protocol Message into an in-memory buffer cannot fail");
let this_message_len = send_buf.len() - last_size;
if send_buf.len() > max_payload {
if last_size > 0 {
trace!("sending {} bytes to {peer}", last_size);
if let Err(err) = send_to_retry(
ports.transport,
ports.authenticator,
ports.sender_counter,
&send_buf[..last_size],
*peer,
)
.await
{
warn!("failed to send datagram to {peer}: {err}; continuing");
} else {
trace!("sent {} bytes to {peer}", last_size);
}
sent_bytes += last_size;
pace(rate, start, sent_bytes).await;
}
send_buf.drain(..last_size);
if this_message_len > max_payload {
error!(
"dropping oversized message to {peer}: encodes to {this_message_len} bytes, \
exceeding the {max_payload}-byte datagram budget; this key will never \
converge on this peer until a smaller value is written"
);
observability::record_value_oversized();
send_buf.clear();
}
}
}
if !send_buf.is_empty() {
trace!("sending last {} bytes to {peer}", send_buf.len());
if let Err(err) = send_to_retry(
ports.transport,
ports.authenticator,
ports.sender_counter,
send_buf,
*peer,
)
.await
{
warn!("failed to send final datagram to {peer}: {err}; continuing");
} else {
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;
}
}
pub(super) struct BulkInFlightGuard {
set: Arc<RwLock<HashSet<SocketAddr>>>,
peer: SocketAddr,
}
impl Drop for BulkInFlightGuard {
fn drop(&mut self) {
self.set.write().remove(&self.peer);
}
}
pub(super) struct BulkDumpCountGuard {
counter: Arc<AtomicUsize>,
}
impl Drop for BulkDumpCountGuard {
fn drop(&mut self) {
self.counter.fetch_sub(1, Ordering::Release);
}
}