use std::collections::HashMap;
use std::sync::Mutex;
use tokio::sync::{Notify, mpsc};
use weida_core::Error;
pub struct NameRegistry<T> {
max_name_bytes: usize,
entries: Mutex<HashMap<String, mpsc::UnboundedSender<T>>>,
bound: Notify,
}
impl<T> NameRegistry<T> {
pub fn new(max_name_bytes: usize) -> NameRegistry<T> {
NameRegistry {
max_name_bytes,
entries: Mutex::new(HashMap::new()),
bound: Notify::new(),
}
}
pub const fn max_name_bytes(&self) -> usize {
self.max_name_bytes
}
pub fn validate(&self, name: &str) -> Result<(), Error> {
if name.is_empty() || name.len() > self.max_name_bytes {
return Err(Error::InvalidAddress(format!(
"name must be 1..={} bytes: {name:?}",
self.max_name_bytes
)));
}
if name.bytes().any(|b| b < 0x20) {
return Err(Error::InvalidAddress(format!(
"invalid byte in name: {name:?}"
)));
}
Ok(())
}
pub fn bind(&self, name: &str) -> Result<mpsc::UnboundedReceiver<T>, Error> {
self.validate(name)?;
let mut entries = self.entries.lock().expect("name registry poisoned");
if entries.contains_key(name) {
return Err(Error::AlreadyRegistered);
}
let (tx, rx) = mpsc::unbounded_channel();
entries.insert(name.to_owned(), tx);
drop(entries);
self.bound.notify_waiters();
Ok(rx)
}
pub fn unbind(&self, name: &str) {
self.entries
.lock()
.expect("name registry poisoned")
.remove(name);
}
pub fn lookup(&self, name: &str) -> Option<mpsc::UnboundedSender<T>> {
self.entries
.lock()
.expect("name registry poisoned")
.get(name)
.cloned()
}
pub async fn wait_bound(&self, name: &str) {
loop {
let bound = self.bound.notified();
tokio::pin!(bound);
bound.as_mut().enable();
if self.lookup(name).is_some() {
return;
}
bound.await;
}
}
}
impl<T> std::fmt::Debug for NameRegistry<T> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("NameRegistry")
.field("max_name_bytes", &self.max_name_bytes)
.finish_non_exhaustive()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn a_name_is_bounded_by_the_registry_budget() {
let registry: NameRegistry<u8> = NameRegistry::new(8);
assert_eq!(registry.max_name_bytes(), 8);
assert!(registry.validate("orders").is_ok());
assert!(registry.validate(&"x".repeat(8)).is_ok());
assert!(registry.validate(&"x".repeat(9)).is_err());
assert!(registry.validate("").is_err());
assert!(registry.validate("has\ncontrol").is_err());
}
#[test]
fn one_owner_per_name() {
let registry: NameRegistry<u8> = NameRegistry::new(16);
let first = registry.bind("orders").expect("first bind");
assert!(matches!(
registry.bind("orders"),
Err(Error::AlreadyRegistered)
));
drop(first);
assert!(matches!(
registry.bind("orders"),
Err(Error::AlreadyRegistered)
));
registry.unbind("orders");
assert!(registry.bind("orders").is_ok());
}
#[test]
fn a_dial_reaches_the_owner_and_an_unbound_name_reaches_nobody() {
let registry: NameRegistry<u8> = NameRegistry::new(16);
assert!(registry.lookup("orders").is_none());
let mut incoming = registry.bind("orders").expect("bind");
registry
.lookup("orders")
.expect("bound name has a sender")
.send(9)
.expect("the owner is still listening");
assert_eq!(incoming.try_recv().expect("delivered"), 9);
registry.unbind("orders");
assert!(registry.lookup("orders").is_none());
}
#[test]
fn an_oversized_name_cannot_be_bound() {
let registry: NameRegistry<u8> = NameRegistry::new(4);
let err = registry.bind("toolong").unwrap_err();
assert!(matches!(err, Error::InvalidAddress(_)), "{err:?}");
}
}