use dashmap::{
DashMap,
mapref::entry::Entry,
};
use flume::{bounded, Receiver, Sender, TrySendError};
use std::sync::Arc;
use thiserror::Error;
pub use flume;
#[non_exhaustive]
#[derive(Error, Debug, PartialEq, Eq)]
pub enum Error {
#[error("Channel does not exist")]
Nonexistent,
#[error("Channel is full")]
Full,
#[error("Channel disconnected")]
Disconnected,
#[error("Channel already exists")]
AlreadyExists(String),
}
#[derive(Debug, Clone)]
pub struct ChannelMap<T> {
channels: Arc<DashMap<String, Sender<T>>>,
buffer: usize,
}
impl<T> ChannelMap<T> {
#[must_use]
pub fn new() -> Self {
Self::new_with_buffer(100)
}
#[must_use]
pub fn new_with_buffer(buffer: usize) -> Self {
Self {
channels: Arc::new(DashMap::new()),
buffer,
}
}
pub fn add(&self, name: &str) -> Result<Receiver<T>, Error> {
self.add_with_buffer(name, self.buffer)
}
pub fn add_with_buffer(&self, name: &str, buffer: usize) -> Result<Receiver<T>, Error> {
let (sender, receiver) = bounded(buffer);
match self.channels.entry(name.to_owned()) {
Entry::Occupied(_) => return Err(Error::AlreadyExists(name.to_owned())),
Entry::Vacant(entry) => entry.insert(sender),
};
Ok(receiver)
}
#[must_use]
pub fn len(&self) -> usize {
self.channels.len()
}
#[must_use]
pub fn is_empty(&self) -> bool {
self.channels.is_empty()
}
#[must_use]
pub fn contains(&self, name: &str) -> bool {
self.channels.contains_key(name)
}
pub fn iter(&self) -> impl Iterator<Item = Sender<T>> {
self.channels.iter().map(|x| x.value().clone())
}
pub fn keys(&self) -> impl Iterator<Item = String> {
self.channels.iter().map(|x| x.key().clone())
}
pub fn clear(&self) {
self.channels.clear();
}
#[allow(clippy::must_use_candidate)]
pub fn remove(&self, name: &str) -> bool {
self.channels.remove(name).is_some()
}
#[must_use]
pub fn get(&self, name: &str) -> Option<Sender<T>> {
Some(self.channels.get(name)?.value().clone())
}
pub fn send(&self, name: &str, msg: T) -> Result<(), Error> {
let tx = self.get(name).ok_or(Error::Nonexistent)?;
if tx.send(msg).is_err() {
self.remove(name);
return Err(Error::Disconnected);
}
Ok(())
}
pub async fn send_async(&self, name: &str, msg: T) -> Result<(), Error> {
let tx = self.get(name).ok_or(Error::Nonexistent)?;
if tx.send_async(msg).await.is_err() {
self.remove(name);
return Err(Error::Disconnected);
}
Ok(())
}
pub fn try_send(&self, name: &str, msg: T) -> Result<(), Error> {
let tx = self.get(name).ok_or(Error::Nonexistent)?;
match tx.try_send(msg) {
Err(TrySendError::Full(_)) => Err(Error::Full),
Err(TrySendError::Disconnected(_)) => {
self.remove(name);
Err(Error::Disconnected)
},
Ok(()) => Ok(())
}
}
#[must_use]
pub fn get_inner(&self) -> Arc<DashMap<String, Sender<T>>> {
self.channels.clone()
}
}
impl<T> Default for ChannelMap<T> {
fn default() -> Self {
Self::new()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn send_channel() {
let channelmap = ChannelMap::new();
let receiver = channelmap.add("foo").unwrap();
assert!(channelmap.contains("foo"));
channelmap.send("foo", "bar").unwrap();
assert_eq!(receiver.recv(), Ok("bar"));
}
#[tokio::test]
async fn send_channel_async() {
let channelmap = ChannelMap::new();
let receiver = channelmap.add("foo").unwrap();
channelmap.send_async("foo", "bar").await.unwrap();
assert_eq!(receiver.recv_async().await, Ok("bar"));
}
#[test]
fn remove_channel() {
let channelmap: ChannelMap<&str> = ChannelMap::new();
let _receiver = channelmap.add("foo").unwrap();
channelmap.remove("foo");
assert_eq!(channelmap.send("foo", "bar"), Err(Error::Nonexistent));
}
#[test]
fn send_channel_threads() {
let channelmap = ChannelMap::new();
let channelmap_clone = channelmap.clone();
let receiver: Receiver<&str> = channelmap_clone.add("foo").unwrap();
let thread = std::thread::spawn(move || {
assert_eq!(receiver.recv(), Ok("bar"));
});
std::thread::spawn(move || {
assert!(channelmap.send("foo", "bar").is_ok());
});
assert!(thread.join().is_ok()); }
#[test]
fn test_length() {
let channelmap: ChannelMap<&str> = ChannelMap::new();
assert!(channelmap.is_empty());
let _receiver = channelmap.add("foo").unwrap();
assert_eq!(channelmap.len(), 1);
channelmap.remove("foo");
assert!(channelmap.is_empty());
assert_eq!(channelmap.len(), 0);
}
#[test]
fn already_exists() {
let channelmap: ChannelMap<&str> = ChannelMap::new();
let _receiver = channelmap.add("foo").unwrap();
assert_eq!(channelmap.add("foo").unwrap_err(), Error::AlreadyExists("foo".to_owned()));
}
#[test]
fn dropped_rx() {
let channelmap = ChannelMap::new();
{
let _receiver = channelmap.add("foo").unwrap();
}
assert_eq!(channelmap.send("foo", "bar").unwrap_err(), Error::Disconnected);
}
}