use std::collections::HashMap;
use std::sync::{Arc, Mutex};
use std::time::Duration;
use tokio::sync::Notify;
use uuid::Uuid;
use crate::error::DurableError;
#[derive(Debug, Default)]
pub(crate) struct NotifyRegistry {
inner: Mutex<HashMap<Uuid, Arc<Notify>>>,
}
impl NotifyRegistry {
pub(crate) fn register(&self, key: Uuid, cap: Option<usize>) -> Option<Arc<Notify>> {
let mut guard = self
.inner
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
if let Some(existing) = guard.get(&key) {
return Some(existing.clone());
}
if let Some(cap) = cap
&& guard.len() >= cap
{
return None;
}
let notify = Arc::new(Notify::new());
guard.insert(key, notify.clone());
Some(notify)
}
pub(crate) fn wake(&self, key: Uuid) {
let notify = {
let mut guard = self
.inner
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
guard.remove(&key)
};
if let Some(notify) = notify {
notify.notify_waiters();
}
}
pub(crate) fn cancel(&self, key: Uuid) {
let mut guard = self
.inner
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
guard.remove(&key);
}
#[cfg(test)]
pub(crate) fn parked(&self) -> usize {
self.inner
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.len()
}
}
pub(crate) async fn wait_on_notify_or_poll<T, FBefore, FutBefore, FAfter, FutAfter>(
registry: &NotifyRegistry,
key: Uuid,
cap: Option<usize>,
mut wait_duration: impl FnMut() -> Duration,
mut check_before: FBefore,
mut check_after: FAfter,
) -> Result<T, DurableError>
where
FBefore: FnMut() -> FutBefore,
FutBefore: Future<Output = Result<Option<T>, DurableError>>,
FAfter: FnMut() -> FutAfter,
FutAfter: Future<Output = Result<Option<T>, DurableError>>,
{
loop {
if let Some(value) = check_before().await? {
return Ok(value);
}
let wait = wait_duration();
match registry.register(key, cap) {
Some(notify) => {
let notified = notify.notified();
tokio::pin!(notified);
notified.as_mut().enable();
if let Some(value) = check_after().await? {
registry.cancel(key);
return Ok(value);
}
tokio::select! {
() = notified => {}
() = tokio::time::sleep(wait) => {}
}
}
None => tokio::time::sleep(wait).await,
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::time::Duration;
#[test]
fn register_is_idempotent_per_key() {
let reg = NotifyRegistry::default();
let key = Uuid::now_v7();
let a = reg.register(key, None).unwrap();
let b = reg.register(key, None).unwrap();
assert!(Arc::ptr_eq(&a, &b), "the same key shares one Notify");
assert_eq!(reg.parked(), 1);
}
#[test]
fn cap_declines_new_keys_but_admits_existing() {
let reg = NotifyRegistry::default();
let first = Uuid::now_v7();
assert!(reg.register(first, Some(1)).is_some());
assert!(reg.register(Uuid::now_v7(), Some(1)).is_none());
assert!(reg.register(first, Some(1)).is_some());
}
#[tokio::test]
async fn wake_releases_a_parked_waiter() {
let reg = Arc::new(NotifyRegistry::default());
let key = Uuid::now_v7();
let notify = reg.register(key, None).unwrap();
let parked = {
let notify = notify.clone();
tokio::spawn(async move { notify.notified().await })
};
tokio::time::sleep(Duration::from_millis(20)).await;
reg.wake(key);
tokio::time::timeout(Duration::from_secs(1), parked)
.await
.expect("waiter is woken")
.expect("task joins");
assert_eq!(reg.parked(), 0, "wake drops the registration");
}
}