use futures::Stream;
use std::hash::Hash;
use std::{collections::HashMap, sync::Arc};
use tokio::sync::broadcast::error::RecvError;
use tokio::sync::mpsc::{UnboundedReceiver, UnboundedSender, unbounded_channel};
use tokio::sync::{RwLock, broadcast};
pub trait Key: Hash + Eq + Clone + Send + Sync + 'static {}
pub trait Value: Clone + Send + 'static {}
impl<T: Clone + Send + 'static> Value for T {}
impl<T: std::hash::Hash + Eq + Clone + Send + Sync + 'static> Key for T {}
type Streams<K, V> = Arc<RwLock<HashMap<K, broadcast::Sender<V>>>>;
pub struct KeyStream<K: Key, V: Value> {
broadcast_capacity: usize,
streams: Streams<K, V>,
sender: UnboundedSender<K>,
drop_keys_task: Option<tokio::task::JoinHandle<()>>,
}
pub struct KeySender<K: Key, V: Value> {
streams: Streams<K, V>,
broadcast_capacity: usize,
drop_notify: UnboundedSender<K>,
}
pub struct KeyReceiver<K: Key, V: Value> {
key: K,
receiver: broadcast::Receiver<V>,
on_drop: UnboundedSender<K>,
}
impl<K: Key, V: Value> KeyStream<K, V> {
pub fn new(broadcast_capacity: usize) -> Self {
let (sender, receiver) = unbounded_channel::<K>();
let streams = Arc::new(RwLock::new(HashMap::<K, broadcast::Sender<V>>::new()));
let on_drop_task = tokio::spawn(cleanup_keys(receiver, Arc::clone(&streams)));
Self {
broadcast_capacity,
streams,
sender,
drop_keys_task: Some(on_drop_task),
}
}
pub fn sender(&self) -> KeySender<K, V> {
KeySender::new(
Arc::clone(&self.streams),
self.broadcast_capacity,
self.sender.clone(),
)
}
pub async fn n_keys(&self) -> usize {
self.streams.read().await.len()
}
pub async fn key_capacity(&self) -> usize {
self.streams.read().await.capacity()
}
}
impl<K: Key, V: Value> KeySender<K, V> {
fn new(streams: Streams<K, V>, broadcast_capacity: usize, sender: UnboundedSender<K>) -> Self {
Self {
streams,
broadcast_capacity,
drop_notify: sender,
}
}
pub async fn send(&self, key: &K, value: V) -> Result<usize, broadcast::error::SendError<V>> {
let streams = self.streams.read().await;
if let Some(sender) = streams.get(key) {
sender.send(value)
} else {
Ok(0)
}
}
pub async fn subscribe(&self, key: K) -> KeyReceiver<K, V> {
let streams = self.streams.read().await;
let inner = if let Some(sender) = streams.get(&key) {
sender.subscribe()
} else {
drop(streams);
let mut streams = self.streams.write().await;
let sender = streams.entry(key.clone()).or_insert_with(|| {
let (sender, _) = broadcast::channel(self.broadcast_capacity);
sender
});
sender.subscribe()
};
self.create_receiver(key, inner)
}
pub async fn n_keys(&self) -> usize {
self.streams.read().await.len()
}
pub async fn key_capacity(&self) -> usize {
self.streams.read().await.capacity()
}
fn create_receiver(&self, key: K, inner: broadcast::Receiver<V>) -> KeyReceiver<K, V> {
KeyReceiver {
key,
receiver: inner,
on_drop: self.drop_notify.clone(),
}
}
}
impl<K: Key, V: Value> KeyReceiver<K, V> {
pub async fn recv(&mut self) -> Result<V, broadcast::error::RecvError> {
self.receiver.recv().await
}
pub async fn try_recv(&mut self) -> Result<V, broadcast::error::TryRecvError> {
self.receiver.try_recv()
}
pub fn blocking_recv(&mut self) -> Result<V, broadcast::error::RecvError> {
self.receiver.blocking_recv()
}
pub fn to_async_stream(self) -> impl Stream<Item = Result<V, RecvError>> {
futures::stream::unfold(self, |mut receiver| async {
let item = receiver.recv().await;
Some((item, receiver))
})
}
}
impl<K: Key, V: Value> Clone for KeySender<K, V> {
fn clone(&self) -> Self {
Self {
streams: Arc::clone(&self.streams),
broadcast_capacity: self.broadcast_capacity,
drop_notify: self.drop_notify.clone(),
}
}
}
impl<K: Key, V: Value> Drop for KeyStream<K, V> {
fn drop(&mut self) {
if let Some(task) = self.drop_keys_task.take() {
task.abort();
}
}
}
impl<K: Key, V: Value> Drop for KeyReceiver<K, V> {
fn drop(&mut self) {
let _ = self.on_drop.send(self.key.clone());
}
}
async fn cleanup_keys<K: Key, V: Value>(
mut receiver: UnboundedReceiver<K>,
streams: Streams<K, V>,
) {
while let Some(key) = receiver.recv().await {
let mut streams = streams.write().await;
if let Some(sender) = streams.get(&key)
&& sender.receiver_count() == 0
{
streams.remove(&key);
optimize_dict_mem(&mut streams);
}
}
}
fn optimize_dict_mem<K: Key, V: Value>(streams: &mut HashMap<K, broadcast::Sender<V>>) {
let cap = streams.capacity() >> 1;
let len = streams.len();
if cap > 64 && len < cap {
streams.shrink_to(cap);
}
}
#[cfg(test)]
mod tests {
use std::pin::Pin;
use tokio::task::JoinSet;
use tokio::time::{Duration, timeout};
use super::*;
#[tokio::test]
async fn test_recv() {
let key_stream = KeyStream::<String, String>::new(10);
let sender = key_stream.sender();
let mut receiver = sender.subscribe("1".to_string()).await;
assert_eq!(key_stream.n_keys().await, 1);
sender
.send(&"1".to_string(), "value".to_string())
.await
.unwrap();
assert_eq!(receiver.recv().await.unwrap(), "value".to_string());
}
#[tokio::test]
async fn test_no_receiver() {
let key_stream = KeyStream::<String, String>::new(10);
let sender = key_stream.sender();
let result = sender
.send(&"1".to_string(), "value".to_string())
.await
.unwrap();
assert_eq!(result, 0);
assert_eq!(key_stream.n_keys().await, 0);
}
#[tokio::test]
async fn test_send_after_drop() {
let key_stream = KeyStream::<String, String>::new(10);
let sender = key_stream.sender();
let receiver = sender.subscribe("1".to_string()).await;
assert_eq!(key_stream.n_keys().await, 1);
drop(receiver);
tokio::task::yield_now().await;
let result = sender
.send(&"1".to_string(), "value".to_string())
.await
.unwrap();
assert_eq!(result, 0);
assert_eq!(key_stream.n_keys().await, 0);
}
#[tokio::test]
async fn test_lagged() {
let key_stream = KeyStream::<String, String>::new(1);
let sender = key_stream.sender();
let mut receiver = sender.subscribe("1".to_string()).await;
let result = sender
.send(&"1".to_string(), "value1".to_string())
.await
.unwrap();
assert_eq!(result, 1);
let result = sender
.send(&"1".to_string(), "value2".to_string())
.await
.unwrap();
assert_eq!(result, 1);
assert_eq!(
receiver.recv().await,
Err(broadcast::error::RecvError::Lagged(1))
);
assert_eq!(receiver.recv().await.unwrap(), "value2".to_string());
}
#[tokio::test]
async fn test_recv_struct() {
#[derive(Clone, Debug)]
struct MyStruct {
field1: String,
field2: i32,
}
let key_stream = KeyStream::<String, MyStruct>::new(10);
let sender = key_stream.sender();
let mut receiver = sender.subscribe("1".to_string()).await;
assert_eq!(key_stream.n_keys().await, 1);
sender
.send(
&"1".to_string(),
MyStruct {
field1: "value".to_string(),
field2: 42,
},
)
.await
.unwrap();
let received = receiver.recv().await.unwrap();
assert_eq!(received.field1, "value".to_string());
assert_eq!(received.field2, 42);
}
#[tokio::test]
async fn test_recv_arc_struct() {
#[derive(Debug)]
struct MyStruct {
field1: String,
field2: i32,
}
let key_stream = KeyStream::<String, Arc<MyStruct>>::new(10);
let sender = key_stream.sender();
let mut receiver = sender.subscribe("1".to_string()).await;
assert_eq!(key_stream.n_keys().await, 1);
sender
.send(
&"1".to_string(),
Arc::new(MyStruct {
field1: "value".to_string(),
field2: 42,
}),
)
.await
.unwrap();
let received = receiver.recv().await.unwrap();
assert_eq!(received.field1, "value".to_string());
assert_eq!(received.field2, 42);
}
#[tokio::test]
async fn test_messages_broadcasted() {
let key_stream = KeyStream::<String, String>::new(10);
let sender = key_stream.sender();
let mut receiver1 = sender.subscribe("1".to_string()).await;
let mut receiver2 = sender.subscribe("1".to_string()).await;
assert_eq!(key_stream.n_keys().await, 1);
sender
.send(&"1".to_string(), "value".to_string())
.await
.unwrap();
assert_eq!(receiver1.recv().await.unwrap(), "value".to_string());
assert_eq!(receiver2.recv().await.unwrap(), "value".to_string());
}
#[tokio::test]
async fn test_key_filter() {
let key_stream = KeyStream::<String, String>::new(10);
let sender = key_stream.sender();
let mut receiver1 = sender.subscribe("1".to_string()).await;
let mut receiver2 = sender.subscribe("2".to_string()).await;
assert_eq!(key_stream.n_keys().await, 2);
sender
.send(&"1".to_string(), "value1".to_string())
.await
.unwrap();
sender
.send(&"2".to_string(), "value2".to_string())
.await
.unwrap();
assert_eq!(receiver1.recv().await.unwrap(), "value1".to_string());
assert_eq!(receiver2.recv().await.unwrap(), "value2".to_string());
assert_eq!(
receiver1.try_recv().await,
Err(broadcast::error::TryRecvError::Empty)
);
assert_eq!(
receiver2.try_recv().await,
Err(broadcast::error::TryRecvError::Empty)
);
}
#[tokio::test]
async fn test_key_drop() {
let key_stream = KeyStream::<String, String>::new(10);
let sender = key_stream.sender();
let receiver = sender.subscribe("1".to_string()).await;
assert_eq!(key_stream.n_keys().await, 1);
sender
.send(&"1".to_string(), "value".to_string())
.await
.unwrap();
drop(receiver);
tokio::task::yield_now().await;
assert_eq!(key_stream.n_keys().await, 0);
}
#[tokio::test]
async fn test_key_non_dropped_if_other_receiver_exists() {
let key_stream = KeyStream::<String, String>::new(10);
let sender = key_stream.sender();
let receiver1 = sender.subscribe("1".to_string()).await;
let _receiver2 = sender.subscribe("1".to_string()).await;
assert_eq!(key_stream.n_keys().await, 1);
sender
.send(&"1".to_string(), "value".to_string())
.await
.unwrap();
drop(receiver1);
tokio::task::yield_now().await;
assert_eq!(key_stream.n_keys().await, 1);
}
#[tokio::test]
async fn test_two_clients() {
let key_stream = KeyStream::<String, String>::new(10);
let sender1 = key_stream.sender();
let sender2 = key_stream.sender();
let mut receiver1 = sender1.subscribe("1".to_string()).await;
let mut receiver2 = sender2.subscribe("1".to_string()).await;
assert_eq!(key_stream.n_keys().await, 1);
sender1
.send(&"1".to_string(), "value".to_string())
.await
.unwrap();
assert_eq!(receiver1.recv().await.unwrap(), "value".to_string());
assert_eq!(receiver2.recv().await.unwrap(), "value".to_string());
assert_eq!(
receiver1.try_recv().await,
Err(broadcast::error::TryRecvError::Empty)
);
assert_eq!(
receiver2.try_recv().await,
Err(broadcast::error::TryRecvError::Empty)
);
drop(sender2);
tokio::task::yield_now().await;
assert_eq!(key_stream.n_keys().await, 1);
sender1
.send(&"1".to_string(), "value".to_string())
.await
.unwrap();
assert_eq!(receiver1.recv().await.unwrap(), "value".to_string());
}
#[tokio::test]
async fn test_reconnect() {
let key_stream = KeyStream::<String, String>::new(10);
let sender = key_stream.sender();
let mut receiver1 = sender.subscribe("1".to_string()).await;
assert_eq!(key_stream.n_keys().await, 1);
sender
.send(&"1".to_string(), "value".to_string())
.await
.unwrap();
assert_eq!(receiver1.recv().await.unwrap(), "value".to_string());
drop(receiver1);
tokio::task::yield_now().await;
assert_eq!(key_stream.n_keys().await, 0);
let mut receiver2 = sender.subscribe("1".to_string()).await;
assert_eq!(key_stream.n_keys().await, 1);
sender
.send(&"1".to_string(), "value2".to_string())
.await
.unwrap();
assert_eq!(receiver2.recv().await.unwrap(), "value2".to_string());
}
#[tokio::test]
async fn test_resubscribe_before_cleanup_runs() {
let key_stream = KeyStream::<String, String>::new(10);
let sender = key_stream.sender();
let receiver1 = sender.subscribe("1".to_string()).await;
assert_eq!(key_stream.n_keys().await, 1);
drop(receiver1);
let mut receiver2 = sender.subscribe("1".to_string()).await;
assert_eq!(key_stream.n_keys().await, 1);
tokio::task::yield_now().await;
assert_eq!(key_stream.n_keys().await, 1);
let result = sender
.send(&"1".to_string(), "value".to_string())
.await
.unwrap();
assert_eq!(result, 1);
assert_eq!(receiver2.recv().await.unwrap(), "value".to_string());
}
#[tokio::test]
async fn test_shrink_dict() {
let key_stream = KeyStream::<i32, String>::new(1);
let sender = key_stream.sender();
let mut set = JoinSet::new();
let mut subs = (0..500)
.map(|i| {
set.spawn({
let value = sender.clone();
async move { value.subscribe(i).await }
})
})
.collect::<Vec<_>>();
set.join_all().await;
assert_eq!(key_stream.n_keys().await, 500);
subs.drain(0..400);
tokio::task::yield_now().await;
assert!(key_stream.key_capacity().await < 500);
}
#[tokio::test]
async fn test_drop_sender() {
let key_stream = KeyStream::<String, String>::new(10);
let sender = key_stream.sender();
let mut receiver = sender.subscribe("1".to_string()).await;
assert_eq!(key_stream.n_keys().await, 1);
sender
.send(&"1".to_string(), "value".to_string())
.await
.unwrap();
drop(sender);
tokio::task::yield_now().await;
assert_eq!(receiver.recv().await.unwrap(), "value".to_string());
assert_eq!(
receiver.try_recv().await,
Err(broadcast::error::TryRecvError::Empty)
);
}
#[tokio::test]
async fn test_drop_stream() {
let key_stream = KeyStream::<String, String>::new(10);
let sender = key_stream.sender();
drop(key_stream);
let mut receiver = sender.subscribe("1".to_string()).await;
sender
.send(&"1".to_string(), "value".to_string())
.await
.unwrap();
tokio::task::yield_now().await;
assert_eq!(receiver.recv().await.unwrap(), "value".to_string());
assert_eq!(
receiver.try_recv().await,
Err(broadcast::error::TryRecvError::Empty)
);
}
#[tokio::test]
async fn test_len_and_capacity() {
let key_stream = KeyStream::<String, String>::new(10);
let sender1 = key_stream.sender();
let sender2 = key_stream.sender();
assert_eq!(key_stream.n_keys().await, 0);
assert_eq!(sender1.n_keys().await, 0);
assert_eq!(sender2.n_keys().await, 0);
assert_eq!(
key_stream.key_capacity().await,
sender1.key_capacity().await
);
assert_eq!(
key_stream.key_capacity().await,
sender2.key_capacity().await
);
let _receiver1 = sender1.subscribe("1".to_string()).await;
let _receiver2 = sender2.subscribe("2".to_string()).await;
let _receiver3 = sender2.subscribe("3".to_string()).await;
assert_eq!(key_stream.n_keys().await, 3);
assert_eq!(sender1.n_keys().await, 3);
assert_eq!(sender2.n_keys().await, 3);
assert_eq!(
key_stream.key_capacity().await,
sender1.key_capacity().await
);
assert_eq!(
key_stream.key_capacity().await,
sender2.key_capacity().await
);
}
#[tokio::test]
async fn test_stream() {
use futures::StreamExt;
type WatchStream =
Pin<Box<dyn futures::Stream<Item = Result<String, RecvError>> + Send + 'static>>;
let key_stream = KeyStream::<String, String>::new(10);
let sender = key_stream.sender();
let receiver = sender.subscribe("1".to_string()).await;
sender
.send(&"1".to_string(), "value1".to_string())
.await
.unwrap();
let stream = receiver.to_async_stream();
let stream = Box::pin(stream) as WatchStream;
sender
.send(&"1".to_string(), "value2".to_string())
.await
.unwrap();
let msg = timeout(Duration::from_secs(1), stream.take(2).collect::<Vec<_>>()).await;
let msg = msg.expect("timeout");
assert_eq!(
msg,
vec![Ok("value1".to_string()), Ok("value2".to_string())]
);
}
}