use std::io;
use std::mem;
use std::panic::{AssertUnwindSafe, catch_unwind};
use std::sync::{Condvar, Mutex, MutexGuard, PoisonError};
use std::task::Waker;
use std::time::{Duration, Instant};
use crate::thread::ThreadStartError;
pub(super) const CONNECT_REPROBE_INTERVAL: Duration = Duration::from_millis(100);
struct Registry {
started: bool,
next_key: u64,
wakers: Vec<(u64, Waker)>,
}
static REGISTRY: Mutex<Registry> = Mutex::new(Registry {
started: false,
next_key: 0,
wakers: Vec::new(),
});
static REGISTERED: Condvar = Condvar::new();
fn registry() -> MutexGuard<'static, Registry> {
REGISTRY.lock().unwrap_or_else(PoisonError::into_inner)
}
pub(super) struct Registration {
key: Option<u64>,
}
impl Registration {
pub(super) const fn new() -> Self {
Self { key: None }
}
pub(super) fn arm(&mut self, waker: &Waker) -> io::Result<()> {
let mut registry = registry();
if !registry.started {
std::thread::Builder::new()
.name("moirai-connect-reprobe".to_owned())
.spawn(run)
.map_err(|source| ThreadStartError::new("connect re-probe", source))?;
registry.started = true;
}
if let Some(key) = self.key
&& let Some((_, current)) = registry.wakers.iter_mut().find(|(k, _)| *k == key)
{
if !current.will_wake(waker) {
let replaced = mem::replace(current, waker.clone());
drop(registry);
drop(replaced);
}
return Ok(());
}
let key = registry.next_key;
registry.next_key += 1;
registry.wakers.push((key, waker.clone()));
self.key = Some(key);
REGISTERED.notify_one();
Ok(())
}
}
impl Drop for Registration {
fn drop(&mut self) {
let Some(key) = self.key else {
return;
};
let removed = {
let mut registry = registry();
let position = registry.wakers.iter().position(|(k, _)| *k == key);
position.map(|index| registry.wakers.swap_remove(index))
};
drop(removed);
}
}
fn run() {
let mut registry = registry();
loop {
registry = REGISTERED
.wait_while(registry, |registry| registry.wakers.is_empty())
.unwrap_or_else(PoisonError::into_inner);
let due = Instant::now() + CONNECT_REPROBE_INTERVAL;
while let Some(remaining) = due.checked_duration_since(Instant::now())
&& !remaining.is_zero()
{
registry = REGISTERED
.wait_timeout(registry, remaining)
.unwrap_or_else(PoisonError::into_inner)
.0;
}
let wakers: Vec<Waker> = registry
.wakers
.iter()
.filter_map(|(_, waker)| catch_unwind(AssertUnwindSafe(|| waker.clone())).ok())
.collect();
drop(registry);
for waker in wakers {
let _contained = catch_unwind(AssertUnwindSafe(|| waker.wake()));
}
#[cfg(test)]
ticks::completed();
registry = self::registry();
}
}
#[cfg(test)]
pub(super) fn registered() -> usize {
registry().wakers.len()
}
#[cfg(test)]
pub(super) mod ticks {
use std::sync::{Condvar, Mutex, PoisonError};
use std::time::Duration;
static COMPLETED: Mutex<u64> = Mutex::new(0);
static ADVANCED: Condvar = Condvar::new();
pub(in crate::net::connect) fn completed() {
*COMPLETED.lock().unwrap_or_else(PoisonError::into_inner) += 1;
ADVANCED.notify_all();
}
pub(in crate::net::connect) fn count() -> u64 {
*COMPLETED.lock().unwrap_or_else(PoisonError::into_inner)
}
pub(in crate::net::connect) fn wait_for(target: u64, limit: Duration) -> u64 {
let completed = COMPLETED.lock().unwrap_or_else(PoisonError::into_inner);
*ADVANCED
.wait_timeout_while(completed, limit, |completed| *completed < target)
.unwrap_or_else(PoisonError::into_inner)
.0
}
}