use std::collections::BTreeSet;
use std::sync::atomic::{AtomicBool, AtomicU64, Ordering};
use std::sync::{Arc, Mutex, RwLock, Weak};
use std::time::Duration;
use tokio::sync::watch;
use unb_client::{EndpointSet, TransportKind};
use unb_core::{NodeIdentity, RetirementReason, RouteDelta, RouteSnapshot};
use unb_runtime::{CancellationToken, Wire};
use crate::connect::EndpointDialer;
use crate::node::Node;
const INITIAL_BACKOFF_MS: u64 = 100;
const MAX_BACKOFF_MS: u64 = 5_000;
const STABLE_HEALTH: Duration = Duration::from_secs(30);
#[derive(Clone, Debug, PartialEq, Eq, thiserror::Error)]
pub enum ConnectError {
#[error("no supported endpoint in set")]
NoSupportedEndpoint,
#[error("dial failed for {transport:?}: {message}")]
Dial {
transport: TransportKind,
message: String,
},
#[error("dial timed out for {transport:?}")]
DialTimedOut { transport: TransportKind },
#[error("peer establishment failed: {message}")]
Establishment { message: String },
#[error("peer identity mismatch: expected {expected:?}, got {actual:?}")]
IdentityMismatch {
expected: String,
actual: Option<String>,
},
#[error("the owning node is shut down")]
NodeShutdown,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum DisconnectReason {
ExplicitDisconnect,
NodeShutdown,
SessionRetired { reason: RetirementReason },
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum ConnectionStatus {
Connecting,
Connected,
Disconnected { reason: DisconnectReason },
}
#[derive(Clone)]
struct SessionBinding {
session_id: String,
wire: Arc<Wire>,
}
#[derive(Clone)]
struct MaintenanceTask {
generation: u64,
cancellation: CancellationToken,
}
pub(crate) struct ReadinessWait {
status: watch::Receiver<ConnectionStatus>,
wait_for_change: bool,
}
impl ReadinessWait {
pub(crate) async fn wait(mut self) -> Result<(), DisconnectReason> {
if self.wait_for_change {
self.status
.changed()
.await
.map_err(|_| DisconnectReason::NodeShutdown)?;
}
loop {
let status = self.status.borrow_and_update().clone();
match status {
ConnectionStatus::Connected => return Ok(()),
ConnectionStatus::Disconnected { reason } => return Err(reason),
ConnectionStatus::Connecting => {}
}
if self.status.changed().await.is_err() {
return Err(DisconnectReason::NodeShutdown);
}
}
}
}
#[derive(Clone)]
pub struct PeerConnection {
node: Weak<Node>,
peer: Arc<str>,
endpoints: Arc<RwLock<EndpointSet>>,
dialer: Arc<RwLock<Option<Arc<dyn EndpointDialer>>>>,
identity: Arc<RwLock<NodeIdentity>>,
session: Arc<Mutex<Option<SessionBinding>>>,
status: watch::Sender<ConnectionStatus>,
generation: Arc<AtomicU64>,
terminal: Arc<AtomicBool>,
maintenance: Arc<Mutex<Option<MaintenanceTask>>>,
stability: Arc<Mutex<Option<CancellationToken>>>,
backoff_ms: Arc<AtomicU64>,
destinations: Arc<RwLock<BTreeSet<String>>>,
}
impl PeerConnection {
pub(crate) fn new(
node: Weak<Node>,
identity: NodeIdentity,
endpoints: EndpointSet,
session_id: String,
wire: Arc<Wire>,
dialer: Option<Arc<dyn EndpointDialer>>,
) -> PeerConnection {
Self::from_binding(node, identity, endpoints, session_id, wire, dialer)
}
pub(crate) fn passive(
node: Weak<Node>,
identity: NodeIdentity,
session_id: String,
wire: Arc<Wire>,
) -> PeerConnection {
Self::from_binding(node, identity, EndpointSet::new(), session_id, wire, None)
}
fn from_binding(
node: Weak<Node>,
identity: NodeIdentity,
endpoints: EndpointSet,
session_id: String,
wire: Arc<Wire>,
dialer: Option<Arc<dyn EndpointDialer>>,
) -> PeerConnection {
let (status, _) = watch::channel(ConnectionStatus::Connected);
PeerConnection {
node,
peer: Arc::from(identity.node_id.as_str()),
endpoints: Arc::new(RwLock::new(endpoints)),
dialer: Arc::new(RwLock::new(dialer)),
identity: Arc::new(RwLock::new(identity)),
session: Arc::new(Mutex::new(Some(SessionBinding { session_id, wire }))),
status,
generation: Arc::new(AtomicU64::new(0)),
terminal: Arc::new(AtomicBool::new(false)),
maintenance: Arc::new(Mutex::new(None)),
stability: Arc::new(Mutex::new(None)),
backoff_ms: Arc::new(AtomicU64::new(INITIAL_BACKOFF_MS)),
destinations: Arc::new(RwLock::new(BTreeSet::new())),
}
}
pub fn peer(&self) -> &str {
&self.peer
}
pub fn status(&self) -> ConnectionStatus {
self.status.borrow().clone()
}
pub fn changed(&self) -> impl std::future::Future<Output = ConnectionStatus> + Send + 'static {
let mut receiver = self.status.subscribe();
async move {
if receiver.changed().await.is_err() {
return receiver.borrow().clone();
}
receiver.borrow_and_update().clone()
}
}
pub fn disconnect(&self) {
let binding = self.terminate(DisconnectReason::ExplicitDisconnect);
if let Some(binding) = binding {
binding.wire.shutdown();
}
}
pub(crate) fn owner(&self) -> Option<Arc<Node>> {
self.node.upgrade()
}
pub(crate) fn readiness_wait(&self) -> Option<ReadinessWait> {
let status = self.status.subscribe();
let wait_for_change = match status.borrow().clone() {
ConnectionStatus::Connected => true,
ConnectionStatus::Connecting => false,
ConnectionStatus::Disconnected { .. } => return None,
};
Some(ReadinessWait {
status,
wait_for_change,
})
}
pub(crate) fn carried_destination(&self, destination: &str) -> bool {
self.destinations
.read()
.unwrap_or_else(|poisoned| poisoned.into_inner())
.contains(destination)
}
pub(crate) fn replace_destinations(&self, snapshot: &RouteSnapshot) {
*self
.destinations
.write()
.unwrap_or_else(|poisoned| poisoned.into_inner()) = snapshot
.routes
.iter()
.map(|route| route.destination.clone())
.collect();
}
pub(crate) fn apply_destination_delta(&self, delta: &RouteDelta) {
let mut destinations = self
.destinations
.write()
.unwrap_or_else(|poisoned| poisoned.into_inner());
for withdrawal in &delta.withdraw {
destinations.remove(&withdrawal.destination);
}
destinations.extend(delta.upsert.iter().map(|route| route.destination.clone()));
}
pub(crate) fn is_terminal(&self) -> bool {
self.terminal.load(Ordering::Acquire)
}
pub(crate) fn endpoints(&self) -> EndpointSet {
self.endpoints
.read()
.unwrap_or_else(|poisoned| poisoned.into_inner())
.clone()
}
pub(crate) fn replace_endpoints(&self, endpoints: EndpointSet) {
*self
.endpoints
.write()
.unwrap_or_else(|poisoned| poisoned.into_inner()) = endpoints;
}
pub(crate) fn replace_dialer(&self, dialer: Option<Arc<dyn EndpointDialer>>) {
*self
.dialer
.write()
.unwrap_or_else(|poisoned| poisoned.into_inner()) = dialer;
}
fn dialer(&self) -> Option<Arc<dyn EndpointDialer>> {
self.dialer
.read()
.unwrap_or_else(|poisoned| poisoned.into_inner())
.clone()
}
fn can_dial(&self) -> bool {
!self
.endpoints
.read()
.unwrap_or_else(|poisoned| poisoned.into_inner())
.is_empty()
}
pub(crate) fn bind(&self, identity: NodeIdentity, session_id: String, wire: Arc<Wire>) -> bool {
let mut maintenance = self
.maintenance
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner());
if self.is_terminal() {
return false;
}
self.generation.fetch_add(1, Ordering::AcqRel);
if let Some(task) = maintenance.take() {
task.cancellation.cancel();
}
self.install(identity, session_id, wire)
}
fn install(&self, identity: NodeIdentity, session_id: String, wire: Arc<Wire>) -> bool {
let mut session = self
.session
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner());
if wire.is_closed() {
return false;
}
*session = Some(SessionBinding {
session_id: session_id.clone(),
wire,
});
drop(session);
*self
.identity
.write()
.unwrap_or_else(|poisoned| poisoned.into_inner()) = identity;
self.publish(ConnectionStatus::Connected);
self.arm_stability(session_id);
true
}
pub(crate) fn retire(&self, session_id: &str, reason: RetirementReason) {
let should_start = {
let _maintenance = self
.maintenance
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner());
let mut session = self
.session
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner());
let removed = if session
.as_ref()
.is_some_and(|binding| binding.session_id == session_id)
{
session.take()
} else {
None
};
drop(session);
if removed.is_none() || self.is_terminal() {
return;
}
self.cancel_stability();
if self.can_dial() {
self.publish(ConnectionStatus::Connecting);
true
} else {
self.publish(ConnectionStatus::Disconnected {
reason: DisconnectReason::SessionRetired { reason },
});
false
}
};
if should_start {
self.start_maintenance();
}
}
pub(crate) fn node_shutdown(&self) {
let binding = self.terminate(DisconnectReason::NodeShutdown);
if let Some(binding) = binding {
binding.wire.shutdown();
}
}
pub(crate) fn publish(&self, status: ConnectionStatus) {
if *self.status.borrow() != status {
self.status.send_replace(status);
}
}
fn start_maintenance(&self) {
if self.is_terminal() || !self.can_dial() {
return;
}
let task = {
let mut maintenance = self
.maintenance
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner());
if maintenance.is_some() || self.status() != ConnectionStatus::Connecting {
return;
}
let task = MaintenanceTask {
generation: self.generation.fetch_add(1, Ordering::AcqRel) + 1,
cancellation: CancellationToken::new(),
};
*maintenance = Some(task.clone());
task
};
let connection = self.clone();
unb_runtime::RuntimeHandle::current().spawn(async move {
connection.run_maintenance(task).await;
});
}
async fn run_maintenance(&self, task: MaintenanceTask) {
loop {
if !self.maintenance_is_current(&task) {
return;
}
let Some(node) = self.owner() else {
self.node_shutdown();
return;
};
if node.cancellation().is_cancelled() {
self.node_shutdown();
return;
}
let result = node
.reconnect_peer(self.peer(), &self.endpoints(), self.dialer())
.await;
if !self.maintenance_is_current(&task) {
if let Ok(candidate) = result {
candidate.candidate_wire.shutdown();
}
return;
}
match result {
Ok(candidate) => {
let installed = self.install_maintenance_candidate(
&task,
candidate.identity,
candidate.selected.session_id,
candidate.selected.wire,
);
if !installed {
candidate.candidate_wire.shutdown();
}
return;
}
Err(_) => {
let base = self
.backoff_ms
.load(Ordering::Acquire)
.clamp(INITIAL_BACKOFF_MS, MAX_BACKOFF_MS);
self.backoff_ms.store(
base.saturating_mul(2).min(MAX_BACKOFF_MS),
Ordering::Release,
);
let delay = jittered_delay(base);
tokio::select! {
biased;
() = task.cancellation.cancelled() => return,
() = node.cancellation().cancelled() => {
self.node_shutdown();
return;
}
() = n0_future::time::sleep(delay) => {}
}
}
}
}
}
fn maintenance_is_current(&self, task: &MaintenanceTask) -> bool {
!self.is_terminal()
&& !task.cancellation.is_cancelled()
&& self.generation.load(Ordering::Acquire) == task.generation
&& self
.maintenance
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner())
.as_ref()
.is_some_and(|current| current.generation == task.generation)
}
fn install_maintenance_candidate(
&self,
task: &MaintenanceTask,
identity: NodeIdentity,
session_id: String,
wire: Arc<Wire>,
) -> bool {
let mut maintenance = self
.maintenance
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner());
if self.is_terminal()
|| task.cancellation.is_cancelled()
|| self.generation.load(Ordering::Acquire) != task.generation
|| !maintenance
.as_ref()
.is_some_and(|current| current.generation == task.generation)
{
return false;
}
if !self.install(identity, session_id, wire) {
return false;
}
maintenance.take();
true
}
fn arm_stability(&self, session_id: String) {
self.cancel_stability();
let cancellation = CancellationToken::new();
*self
.stability
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner()) = Some(cancellation.clone());
let connection = self.clone();
unb_runtime::RuntimeHandle::current().spawn(async move {
tokio::select! {
biased;
() = cancellation.cancelled() => return,
() = n0_future::time::sleep(STABLE_HEALTH) => {}
}
if connection.is_terminal() || connection.status() != ConnectionStatus::Connected {
return;
}
let exact = connection
.session
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner())
.as_ref()
.is_some_and(|binding| binding.session_id == session_id);
if exact {
connection
.backoff_ms
.store(INITIAL_BACKOFF_MS, Ordering::Release);
}
});
}
fn cancel_stability(&self) {
if let Some(cancellation) = self
.stability
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner())
.take()
{
cancellation.cancel();
}
}
fn terminate(&self, reason: DisconnectReason) -> Option<SessionBinding> {
let mut maintenance = self
.maintenance
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner());
if self.terminal.swap(true, Ordering::AcqRel) {
return None;
}
self.generation.fetch_add(1, Ordering::AcqRel);
if let Some(task) = maintenance.take() {
task.cancellation.cancel();
}
self.cancel_stability();
let binding = self
.session
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner())
.take();
self.publish(ConnectionStatus::Disconnected { reason });
binding
}
}
fn jittered_delay(base_ms: u64) -> Duration {
let spread = base_ms / 5;
let minimum = base_ms.saturating_sub(spread);
let maximum = base_ms.saturating_add(spread);
Duration::from_millis(fastrand::u64(minimum..=maximum))
}
#[cfg(test)]
mod tests {
use super::*;
fn test_connection() -> PeerConnection {
let node = Node::builder("local")
.insecure_accept_declared_peer_identities()
.build()
.unwrap();
let (pipe, _remote) = unb_client::pair();
PeerConnection::passive(
Arc::downgrade(&node),
NodeIdentity {
node_id: "peer".into(),
instance_id: "peer-instance".into(),
epoch: 1,
proof: serde_json::Value::Null,
},
"session-1".into(),
Arc::new(Wire::open(pipe)),
)
}
#[test]
fn jitter_stays_within_twenty_percent_at_every_backoff_edge() {
for base in [
INITIAL_BACKOFF_MS,
200,
400,
800,
1_600,
3_200,
MAX_BACKOFF_MS,
] {
for _ in 0..128 {
let delay = jittered_delay(base).as_millis() as u64;
assert!(delay >= base - base / 5);
assert!(delay <= base + base / 5);
}
}
}
#[tokio::test(start_paused = true)]
async fn backoff_resets_only_after_thirty_seconds_on_the_exact_healthy_binding() {
let connection = test_connection();
connection.backoff_ms.store(800, Ordering::Release);
connection.arm_stability("session-1".into());
tokio::task::yield_now().await;
tokio::time::advance(Duration::from_secs(29)).await;
tokio::task::yield_now().await;
assert_eq!(connection.backoff_ms.load(Ordering::Acquire), 800);
tokio::time::advance(Duration::from_secs(1)).await;
tokio::task::yield_now().await;
assert_eq!(
connection.backoff_ms.load(Ordering::Acquire),
INITIAL_BACKOFF_MS
);
}
#[tokio::test(start_paused = true)]
async fn a_retired_brief_binding_does_not_reset_accumulated_backoff() {
let connection = test_connection();
connection.backoff_ms.store(800, Ordering::Release);
connection.arm_stability("session-1".into());
tokio::task::yield_now().await;
connection.retire("session-1", RetirementReason::TransportFailed);
tokio::time::advance(STABLE_HEALTH).await;
tokio::task::yield_now().await;
assert_eq!(connection.backoff_ms.load(Ordering::Acquire), 800);
}
#[tokio::test(start_paused = true)]
async fn maintenance_election_keeps_exactly_one_supervisor_generation() {
let connection = test_connection();
connection.replace_endpoints(EndpointSet::from(crate::Endpoint {
kind: crate::TransportKind::WebSocket,
address: "ws://127.0.0.1:9".into(),
cert_hash: None,
}));
connection.publish(ConnectionStatus::Connecting);
connection.start_maintenance();
let generation = connection
.maintenance
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner())
.as_ref()
.expect("one maintenance supervisor")
.generation;
connection.start_maintenance();
assert_eq!(
connection
.maintenance
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner())
.as_ref()
.expect("the original supervisor remains elected")
.generation,
generation
);
connection.disconnect();
tokio::task::yield_now().await;
assert!(connection
.maintenance
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner())
.is_none());
}
#[tokio::test]
async fn a_connected_recovery_hint_waits_for_the_next_recovery_cycle() {
let connection = test_connection();
let waiter = connection.readiness_wait().unwrap();
let waiting = tokio::spawn(waiter.wait());
tokio::task::yield_now().await;
assert!(!waiting.is_finished());
connection.publish(ConnectionStatus::Connecting);
connection.publish(ConnectionStatus::Connected);
assert!(waiting.await.unwrap().is_ok());
connection.disconnect();
}
#[tokio::test]
async fn stale_session_retirement_cannot_replace_or_retire_the_current_binding() {
let connection = test_connection();
let (pipe, _remote) = unb_client::pair();
assert!(connection.bind(
NodeIdentity {
node_id: "peer".into(),
instance_id: "peer-instance-2".into(),
epoch: 2,
proof: serde_json::Value::Null,
},
"session-2".into(),
Arc::new(Wire::open(pipe)),
));
connection.retire("session-1", RetirementReason::TransportFailed);
assert_eq!(connection.status(), ConnectionStatus::Connected);
assert_eq!(
connection
.session
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner())
.as_ref()
.map(|binding| binding.session_id.as_str()),
Some("session-2")
);
connection.disconnect();
}
#[tokio::test]
async fn terminal_or_closed_bindings_cannot_be_installed() {
let connection = test_connection();
connection.disconnect();
let (late_pipe, _late_remote) = unb_client::pair();
assert!(!connection.bind(
NodeIdentity {
node_id: "peer".into(),
instance_id: "peer-instance-2".into(),
epoch: 2,
proof: serde_json::Value::Null,
},
"session-2".into(),
Arc::new(Wire::open(late_pipe)),
));
assert_eq!(
connection.status(),
ConnectionStatus::Disconnected {
reason: DisconnectReason::ExplicitDisconnect,
}
);
let connection = test_connection();
let (closed_pipe, _closed_remote) = unb_client::pair();
let closed = Arc::new(Wire::open(closed_pipe));
closed.shutdown();
assert!(!connection.bind(
NodeIdentity {
node_id: "peer".into(),
instance_id: "peer-instance-2".into(),
epoch: 2,
proof: serde_json::Value::Null,
},
"session-2".into(),
closed,
));
connection.disconnect();
}
}