use std::sync::Arc;
use std::sync::atomic::{AtomicUsize, Ordering};
use tokio::sync::Notify;
use crate::realtime::error::RealtimeError;
#[derive(Debug)]
struct Shared {
live: AtomicUsize,
drain: Notify,
}
#[derive(Clone)]
pub struct Registry {
shared: Arc<Shared>,
}
#[derive(Debug)]
pub struct ConnectionGuard {
shared: Arc<Shared>,
}
impl Registry {
#[must_use]
pub fn new() -> Self {
Self {
shared: Arc::new(Shared {
live: AtomicUsize::new(0),
drain: Notify::new(),
}),
}
}
#[must_use]
pub fn live_count(&self) -> usize {
self.shared.live.load(Ordering::Relaxed)
}
pub fn acquire(&self, max: usize) -> Result<ConnectionGuard, RealtimeError> {
loop {
let current = self.shared.live.load(Ordering::Relaxed);
if current >= max {
return Err(RealtimeError::ConnectionLimit);
}
if self
.shared
.live
.compare_exchange(current, current + 1, Ordering::Relaxed, Ordering::Relaxed)
.is_ok()
{
return Ok(ConnectionGuard {
shared: Arc::clone(&self.shared),
});
}
}
}
pub async fn drain(&self, bound: std::time::Duration) -> Result<(), RealtimeError> {
self.shared.drain.notify_waiters();
let result = tokio::time::timeout(bound, async {
loop {
if self.shared.live.load(Ordering::Relaxed) == 0 {
return;
}
tokio::task::yield_now().await;
}
})
.await;
match result {
Ok(()) => Ok(()),
Err(_) => Err(RealtimeError::Shutdown {
remaining: self.shared.live.load(Ordering::Relaxed),
}),
}
}
#[must_use]
pub fn drain_signal(&self) -> &Notify {
&self.shared.drain
}
}
impl Default for Registry {
fn default() -> Self {
Self::new()
}
}
impl std::fmt::Debug for Registry {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("Registry")
.field("live", &self.live_count())
.finish()
}
}
impl ConnectionGuard {
#[must_use]
pub fn live_count(&self) -> usize {
self.shared.live.load(Ordering::Relaxed)
}
}
impl Drop for ConnectionGuard {
fn drop(&mut self) {
loop {
let current = self.shared.live.load(Ordering::Relaxed);
if current == 0 {
break;
}
if self
.shared
.live
.compare_exchange(current, current - 1, Ordering::Relaxed, Ordering::Relaxed)
.is_ok()
{
self.shared.drain.notify_one();
break;
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn acquire_and_drop_track_count() {
let reg = Registry::new();
assert_eq!(reg.live_count(), 0);
let g = reg.acquire(8).expect("under cap");
assert_eq!(reg.live_count(), 1);
assert_eq!(g.live_count(), 1);
drop(g);
assert_eq!(reg.live_count(), 0);
}
#[test]
fn acquire_enforces_cap() {
let reg = Registry::new();
let _g1 = reg.acquire(2).expect("first");
let _g2 = reg.acquire(2).expect("second");
let outcome = reg.acquire(2);
assert!(
matches!(outcome, Err(RealtimeError::ConnectionLimit)),
"{outcome:?}"
);
drop(_g2);
let _g3 = reg.acquire(2).expect("slot freed");
}
#[tokio::test]
async fn drain_completes_when_connections_drop() {
let reg = Registry::new();
let _g = reg.acquire(8).expect("under cap");
assert_eq!(reg.live_count(), 1);
let reg2 = reg.clone();
let drain_task =
tokio::spawn(async move { reg2.drain(std::time::Duration::from_secs(5)).await });
tokio::task::yield_now().await;
drop(_g);
let result = tokio::time::timeout(std::time::Duration::from_secs(5), drain_task)
.await
.expect("drain task did not hang")
.expect("drain task did not panic");
assert!(result.is_ok(), "{result:?}");
assert_eq!(reg.live_count(), 0);
}
#[tokio::test]
async fn drain_times_out_with_remaining_count() {
let reg = Registry::new();
let _g = reg.acquire(8).expect("under cap");
let result = reg.drain(std::time::Duration::from_millis(50)).await;
assert!(
matches!(result, Err(RealtimeError::Shutdown { remaining: 1 })),
"{result:?}"
);
}
#[test]
fn registry_is_send_sync_clone_for_appstate() {
fn assert_send_sync_clone<T: Send + Sync + Clone>() {}
assert_send_sync_clone::<Registry>();
}
}