use std::collections::VecDeque;
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::{Arc, Mutex, RwLock};
use std::thread::ThreadId;
use dashmap::DashMap;
use crate::atom::Atom;
#[derive(Copy, Clone, Debug, Eq, PartialEq)]
pub enum ConnectionDownReason {
PeerClosed,
ReadError,
WriteError,
WriteTimeout,
ManualDisconnect,
HeartbeatTimeout,
ControlOverflow,
}
#[derive(Copy, Clone, Debug, Eq, PartialEq)]
pub struct ConnectionDownEvent {
pub node: Atom,
pub reason: ConnectionDownReason,
}
type ConnectionDownCallback = dyn Fn(ConnectionDownEvent) + Send + Sync + 'static;
#[derive(Clone, Default)]
pub struct ConnectionDownHook {
callback: Arc<RwLock<Option<Arc<ConnectionDownCallback>>>>,
}
impl ConnectionDownHook {
#[must_use]
pub fn new() -> Self {
Self::default()
}
pub fn register<F>(&self, callback: F)
where
F: Fn(ConnectionDownEvent) + Send + Sync + 'static,
{
let mut slot = self
.callback
.write()
.unwrap_or_else(|error| error.into_inner());
*slot = Some(Arc::new(callback));
}
pub fn unregister(&self) {
let mut slot = self
.callback
.write()
.unwrap_or_else(|error| error.into_inner());
*slot = None;
}
#[must_use]
pub fn is_registered(&self) -> bool {
self.callback
.read()
.unwrap_or_else(|error| error.into_inner())
.is_some()
}
pub(crate) fn invoke(&self, event: ConnectionDownEvent) {
let callback = self
.callback
.read()
.unwrap_or_else(|error| error.into_inner())
.clone();
if let Some(callback) = callback {
callback(event);
}
}
}
#[derive(Copy, Clone, Debug, Eq, PartialEq, Ord, PartialOrd, Hash)]
pub struct ConnectionGeneration(u64);
impl ConnectionGeneration {
#[must_use]
pub const fn get(self) -> u64 {
self.0
}
#[must_use]
pub const fn from_raw(raw: u64) -> Self {
Self(raw)
}
}
#[derive(Copy, Clone, Debug, Eq, PartialEq)]
#[non_exhaustive]
pub struct NodeUp {
pub node: Atom,
pub generation: ConnectionGeneration,
pub peer_creation: u32,
}
#[derive(Copy, Clone, Debug, Eq, PartialEq)]
#[non_exhaustive]
pub struct NodeDown {
pub node: Atom,
pub generation: ConnectionGeneration,
pub reason: ConnectionDownReason,
}
#[derive(Copy, Clone, Debug, Eq, PartialEq)]
#[non_exhaustive]
pub enum ConnectionEvent {
Up(NodeUp),
Down(NodeDown),
}
impl ConnectionEvent {
#[must_use]
pub fn up(node: Atom, generation: ConnectionGeneration, peer_creation: u32) -> Self {
Self::Up(NodeUp {
node,
generation,
peer_creation,
})
}
#[must_use]
pub fn down(
node: Atom,
generation: ConnectionGeneration,
reason: ConnectionDownReason,
) -> Self {
Self::Down(NodeDown {
node,
generation,
reason,
})
}
#[must_use]
pub fn node(&self) -> Atom {
match self {
Self::Up(up) => up.node,
Self::Down(down) => down.node,
}
}
#[must_use]
pub fn generation(&self) -> ConnectionGeneration {
match self {
Self::Up(up) => up.generation,
Self::Down(down) => down.generation,
}
}
#[must_use]
pub fn down_reason(&self) -> Option<ConnectionDownReason> {
match self {
Self::Up(_) => None,
Self::Down(down) => Some(down.reason),
}
}
}
#[derive(Copy, Clone, Debug, Eq, PartialEq, Hash)]
pub struct SubscriberId(u64);
type ConnectionEventCallback = dyn Fn(ConnectionEvent) + Send + Sync + 'static;
#[derive(Default)]
pub(crate) struct ConnectionEventHub {
subscribers: RwLock<Vec<(SubscriberId, Arc<ConnectionEventCallback>)>>,
next_subscriber_id: AtomicU64,
legacy_down: ConnectionDownHook,
generations: DashMap<Atom, u64>,
queue: Mutex<VecDeque<ConnectionEvent>>,
dispatch_gate: Mutex<()>,
dispatch_owner: Mutex<Option<ThreadId>>,
}
impl ConnectionEventHub {
pub(crate) fn new() -> Self {
Self::default()
}
pub(crate) fn subscribe<F>(&self, callback: F) -> SubscriberId
where
F: Fn(ConnectionEvent) + Send + Sync + 'static,
{
let id = SubscriberId(self.next_subscriber_id.fetch_add(1, Ordering::Relaxed));
let mut subscribers = self
.subscribers
.write()
.unwrap_or_else(|error| error.into_inner());
subscribers.push((id, Arc::new(callback)));
id
}
pub(crate) fn unsubscribe(&self, id: SubscriberId) -> bool {
let mut subscribers = self
.subscribers
.write()
.unwrap_or_else(|error| error.into_inner());
let before = subscribers.len();
subscribers.retain(|(subscriber_id, _)| *subscriber_id != id);
subscribers.len() != before
}
pub(crate) fn legacy_down_hook(&self) -> ConnectionDownHook {
self.legacy_down.clone()
}
pub(crate) fn next_generation(&self, node: Atom) -> ConnectionGeneration {
let mut entry = self.generations.entry(node).or_insert(0);
*entry += 1;
ConnectionGeneration(*entry)
}
pub(crate) fn last_generation(&self, node: Atom) -> Option<ConnectionGeneration> {
self.generations
.get(&node)
.map(|entry| ConnectionGeneration(*entry.value()))
}
pub(crate) fn enqueue(&self, event: ConnectionEvent) {
let mut queue = self.queue.lock().unwrap_or_else(|error| error.into_inner());
queue.push_back(event);
}
pub(crate) fn dispatch(&self) {
let me = std::thread::current().id();
{
let owner = self
.dispatch_owner
.lock()
.unwrap_or_else(|error| error.into_inner());
if *owner == Some(me) {
return;
}
} let _gate = self
.dispatch_gate
.lock()
.unwrap_or_else(|error| error.into_inner());
let _owner = OwnerGuard::set(&self.dispatch_owner, me);
self.drain_queue();
}
pub(crate) fn subscribe_with_snapshot<F>(
&self,
callback: F,
live_peers: impl Fn() -> Vec<NodeUp>,
) -> SubscriberId
where
F: Fn(ConnectionEvent) + Send + Sync + 'static,
{
let me = std::thread::current().id();
{
let owner = self
.dispatch_owner
.lock()
.unwrap_or_else(|error| error.into_inner());
if *owner == Some(me) {
return self.subscribe(callback);
}
} let _gate = self
.dispatch_gate
.lock()
.unwrap_or_else(|error| error.into_inner());
let _owner = OwnerGuard::set(&self.dispatch_owner, me);
let rows = loop {
self.drain_queue();
let rows = live_peers();
let queue_empty = self
.queue
.lock()
.unwrap_or_else(|error| error.into_inner())
.is_empty();
if queue_empty {
break rows;
}
};
let callback: Arc<ConnectionEventCallback> = Arc::new(callback);
for row in rows {
callback(ConnectionEvent::Up(row));
}
let id = SubscriberId(self.next_subscriber_id.fetch_add(1, Ordering::Relaxed));
self.subscribers
.write()
.unwrap_or_else(|error| error.into_inner())
.push((id, Arc::clone(&callback)));
self.drain_queue();
id
}
fn drain_queue(&self) {
loop {
let event = {
let mut queue = self.queue.lock().unwrap_or_else(|error| error.into_inner());
queue.pop_front()
}; let Some(event) = event else { break };
let snapshot: Vec<Arc<ConnectionEventCallback>> = self
.subscribers
.read()
.unwrap_or_else(|error| error.into_inner())
.iter()
.map(|(_, callback)| Arc::clone(callback))
.collect();
for callback in snapshot {
callback(event);
}
if let ConnectionEvent::Down(down) = event {
self.legacy_down.invoke(ConnectionDownEvent {
node: down.node,
reason: down.reason,
});
}
}
}
}
struct OwnerGuard<'a> {
owner: &'a Mutex<Option<ThreadId>>,
}
impl<'a> OwnerGuard<'a> {
fn set(owner: &'a Mutex<Option<ThreadId>>, holder: ThreadId) -> Self {
*owner.lock().unwrap_or_else(|error| error.into_inner()) = Some(holder);
Self { owner }
}
}
impl Drop for OwnerGuard<'_> {
fn drop(&mut self) {
*self.owner.lock().unwrap_or_else(|error| error.into_inner()) = None;
}
}