use std::collections::HashMap;
use std::fmt::Debug;
use std::hash::Hash;
use std::net::{IpAddr, SocketAddr};
use std::ops::RangeBounds;
use std::sync::Arc;
use std::time::{Duration, Instant};
use bincode::{DefaultOptions, Deserializer, Serializer};
use ipnet::IpNet;
use parking_lot::{MappedRwLockReadGuard, RwLock, RwLockReadGuard};
use rand::rngs::StdRng;
use rand::SeedableRng;
use serde::{de::DeserializeOwned, Deserialize, Serialize};
use tokio::net::UdpSocket;
use tokio::time::timeout;
use tracing::{debug, trace, warn};
use crate::auth;
use crate::bounds::Key;
use crate::clock::Timestamp;
use crate::fingerprint::Fingerprint;
use crate::gen_ip::{gen_ip, net_of};
use crate::proto;
use crate::reconcilable::ValueOnly;
use crate::reconcile_engine::{send_messages_to, send_to_retry, Message};
use crate::reconcile_store::Config;
use crate::HRTree;
const BUFFER_SIZE: usize = 65507;
const ACTIVITY_TIMEOUT: Duration = Duration::from_secs(1);
const PEER_EXPIRATION: Duration = Duration::from_secs(60);
type OnUpdateCallback<K, V> = Box<dyn Send + Sync + Fn(&K, &ValueOnly<V>)>;
type WireDated<V> = (Timestamp, Option<V>);
pub struct ReconcileMirror<K, V> {
tree: Arc<RwLock<HRTree<K, ValueOnly<V>>>>,
port: u16,
socket: Arc<UdpSocket>,
net: Arc<RwLock<IpNet>>,
rng: Arc<RwLock<StdRng>>,
peers: Arc<RwLock<HashMap<IpAddr, Instant>>>,
authenticator: auth::Authenticator,
on_update: Arc<RwLock<OnUpdateCallback<K, V>>>,
}
impl<K, V> Clone for ReconcileMirror<K, V> {
fn clone(&self) -> Self {
ReconcileMirror {
tree: self.tree.clone(),
port: self.port,
socket: self.socket.clone(),
net: self.net.clone(),
rng: self.rng.clone(),
peers: self.peers.clone(),
authenticator: self.authenticator.clone(),
on_update: self.on_update.clone(),
}
}
}
impl<K: Key, V: Clone + Debug + DeserializeOwned + Hash + Send + Serialize + Sync + 'static>
ReconcileMirror<K, V>
{
pub async fn new(config: Config) -> Self {
let socket = UdpSocket::bind(SocketAddr::new(config.listen_addr, config.port))
.await
.unwrap();
debug!(
"ReconcileMirror listening on: {}",
socket.local_addr().unwrap()
);
let authenticator = auth::Authenticator::new(config.cluster_key, config.encrypt);
if !authenticator.is_enabled() {
warn!(
"SECURITY: no cluster key set — the lightweight mirror accepts UNAUTHENTICATED \
datagrams. Set Config::with_cluster_key to match the dated cluster."
);
}
let nets: Vec<IpNet> = config.nets.iter().flatten().copied().collect();
let net = net_of(&nets, config.listen_addr)
.or_else(|| nets.first().copied())
.unwrap_or_else(|| "127.0.0.1/8".parse().unwrap());
ReconcileMirror {
tree: Arc::new(RwLock::new(HRTree::<K, ValueOnly<V>>::new())),
port: config.port,
socket: Arc::new(socket),
net: Arc::new(RwLock::new(net)),
rng: Arc::new(RwLock::new(StdRng::from_entropy())),
peers: Arc::new(RwLock::new(HashMap::new())),
authenticator,
on_update: Arc::new(RwLock::new(Box::new(|_, _| {}))),
}
}
pub fn with_seed(self, peer: IpAddr) -> Self {
self.peers.write().insert(peer, Instant::now());
self
}
pub fn set_net(&self, net: IpNet) {
*self.net.write() = net;
}
pub fn net(&self) -> IpNet {
*self.net.read()
}
pub fn add_on_update<F: Send + Sync + Fn(&K, &ValueOnly<V>) + 'static>(&self, on_update: F) {
*self.on_update.write() = Box::new(on_update);
}
pub fn get(&self, k: &K) -> Option<MappedRwLockReadGuard<'_, V>> {
let guard = self.tree.read();
RwLockReadGuard::try_map(guard, |tree| tree.get(k).and_then(|vo| vo.as_value())).ok()
}
pub fn contains_key(&self, k: &K) -> bool {
self.tree.read().get(k).is_some_and(|vo| !vo.is_tombstone())
}
pub fn len(&self) -> usize {
self.tree.read().len()
}
pub fn is_empty(&self) -> bool {
self.tree.read().is_empty()
}
pub fn fingerprint<R: RangeBounds<K>>(&self, range: R) -> Fingerprint {
self.tree.read().hash(&range)
}
fn get_peers(&self) -> Vec<IpAddr> {
let mut guard = self.peers.write();
guard.retain(|_, instant| instant.elapsed() < PEER_EXPIRATION);
guard.keys().cloned().collect()
}
fn integrate(&self, updates: Vec<(K, ValueOnly<V>)>) {
if updates.is_empty() {
return;
}
{
let hook = self.on_update.read();
for (k, vo) in &updates {
hook(k, vo);
}
}
let mut guard = self.tree.write();
for (k, vo) in updates {
guard.insert(k, vo);
}
}
pub async fn start_reconciliation(&self, send_buf: &mut Vec<u8>) {
let segments = proto::start_diff(&self.tree.read());
send_buf.clear();
for segment in segments {
Message::ValueComparisonItem::<K, WireDated<V>, ValueOnly<V>>(segment)
.serialize(&mut Serializer::new(&mut *send_buf, DefaultOptions::new()))
.unwrap();
}
let mut peers = self.get_peers();
let net = *self.net.read();
let addr = gen_ip(&mut *self.rng.write(), net);
peers.push(addr);
for peer in peers {
trace!("mirror start_diff {} bytes to {peer}", send_buf.len());
send_to_retry(
&self.socket,
&self.authenticator,
send_buf,
(peer, self.port),
)
.await
.unwrap();
}
}
async fn handle_messages(
&self,
payload: auth::Payload<'_>,
peer: SocketAddr,
send_buf: &mut Vec<u8>,
) {
let payload = payload.as_bytes();
trace!("mirror received {} bytes from {peer}", payload.len());
let mut value_in_comparison = Vec::new();
let mut value_updates: Vec<(K, ValueOnly<V>)> = Vec::new();
let mut deserializer = Deserializer::from_slice(payload, DefaultOptions::new());
loop {
match Message::<K, WireDated<V>, ValueOnly<V>>::deserialize(&mut deserializer) {
Err(ref kind) => {
if let bincode::ErrorKind::Io(err) = kind.as_ref() {
if err.kind() == std::io::ErrorKind::UnexpectedEof {
break;
}
}
warn!(
"mirror failed to deserialize datagram from {peer}, dropping it: {kind:?}"
);
break;
}
Ok(Message::ValueComparisonItem(segment)) => value_in_comparison.push(segment),
Ok(Message::ValueUpdate(update)) => value_updates.push(update),
Ok(Message::ComparisonItem(_)) | Ok(Message::Update(_)) | Ok(Message::Ack(_)) => {}
}
}
self.integrate(value_updates);
if !value_in_comparison.is_empty() {
debug!(
"mirror received {} value-only segments",
value_in_comparison.len()
);
let mut out_comparison = Vec::new();
let mut differences = Vec::new();
{
let guard = self.tree.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::<K, WireDated<V>, ValueOnly<V>>::ValueComparisonItem)
.collect();
send_messages_to(
&messages,
Arc::clone(&self.socket),
&self.authenticator,
&peer,
send_buf,
)
.await;
}
}
}
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 {
match timeout(ACTIVITY_TIMEOUT, self.socket.recv_from(&mut recv_buf)).await {
Err(_) => {
debug!("mirror: no recent activity; initiating value-only diff");
self.start_reconciliation(&mut send_buf).await;
}
Ok(Err(err)) => warn!("mirror network error in recv_from: {err}"),
Ok(Ok((size, peer))) => {
if peer.port() != self.port {
warn!(
"mirror received message from {peer}, but protocol port is {}",
self.port
);
}
if size == recv_buf.len() {
warn!("mirror buffer too small for message, discarded");
} else {
match self.authenticator.open(&recv_buf[..size]) {
Some(payload) => {
self.handle_messages(payload, peer, &mut send_buf).await;
self.peers.write().insert(peer.ip(), Instant::now());
}
None => trace!(
"mirror dropped datagram from {peer}: missing or invalid MAC"
),
}
}
}
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::reconcile_store::Config;
fn ephemeral_config() -> Config {
Config::default()
}
#[tokio::test]
async fn get_returns_integrated_value() {
let mirror = ReconcileMirror::<i32, String>::new(ephemeral_config()).await;
assert!(mirror.get(&1).is_none());
mirror.integrate(vec![(1, ValueOnly(Some("hello".to_string())))]);
assert_eq!(mirror.get(&1).as_deref(), Some(&"hello".to_string()));
assert!(mirror.contains_key(&1));
assert_eq!(mirror.len(), 1);
}
#[tokio::test]
async fn mirrors_tombstones() {
let mirror = ReconcileMirror::<i32, String>::new(ephemeral_config()).await;
mirror.integrate(vec![(1, ValueOnly(Some("v".to_string())))]);
assert_eq!(mirror.get(&1).as_deref(), Some(&"v".to_string()));
mirror.integrate(vec![(1, ValueOnly(None))]);
assert!(mirror.get(&1).is_none());
assert!(!mirror.contains_key(&1));
assert_eq!(mirror.len(), 1, "the tombstone is retained as an entry");
}
#[tokio::test]
async fn on_update_hook_fires() {
use std::sync::atomic::{AtomicUsize, Ordering};
let mirror = ReconcileMirror::<i32, i32>::new(ephemeral_config()).await;
let count = Arc::new(AtomicUsize::new(0));
let count2 = count.clone();
mirror.add_on_update(move |_, _| {
count2.fetch_add(1, Ordering::SeqCst);
});
mirror.integrate(vec![(1, ValueOnly(Some(10))), (2, ValueOnly(None))]);
assert_eq!(count.load(Ordering::SeqCst), 2);
}
#[tokio::test]
async fn value_fingerprint_is_timestamp_independent() {
let mirror = ReconcileMirror::<i32, String>::new(ephemeral_config()).await;
mirror.integrate(vec![
(1, ValueOnly(Some("a".to_string()))),
(2, ValueOnly(None)),
]);
let mut reference: HRTree<i32, ValueOnly<String>> = HRTree::new();
reference.insert(1, ValueOnly(Some("a".to_string())));
reference.insert(2, ValueOnly(None));
assert_eq!(mirror.fingerprint(..), reference.hash(&..));
}
#[test]
fn value_only_is_smaller_per_entry() {
let dated = std::mem::size_of::<(Timestamp, Option<u64>)>();
let light = std::mem::size_of::<ValueOnly<u64>>();
assert!(
light < dated,
"value-only entry ({light} B) should be smaller than dated entry ({dated} B)"
);
}
}