1use futures::Stream;
54use std::hash::Hash;
55use std::{collections::HashMap, sync::Arc};
56use tokio::sync::broadcast::error::RecvError;
57use tokio::sync::mpsc::{UnboundedReceiver, UnboundedSender, unbounded_channel};
58use tokio::sync::{RwLock, broadcast};
59
60pub trait Key: Hash + Eq + Clone + Send + Sync + 'static {}
63
64pub trait Value: Clone + Send + 'static {}
67
68impl<T: Clone + Send + 'static> Value for T {}
69impl<T: std::hash::Hash + Eq + Clone + Send + Sync + 'static> Key for T {}
70
71type Streams<K, V> = Arc<RwLock<HashMap<K, broadcast::Sender<V>>>>;
72
73pub struct KeyStream<K: Key, V: Value> {
91 broadcast_capacity: usize,
92 streams: Streams<K, V>,
93 sender: UnboundedSender<K>,
94 drop_keys_task: Option<tokio::task::JoinHandle<()>>,
95}
96
97pub struct KeySender<K: Key, V: Value> {
101 streams: Streams<K, V>,
102 broadcast_capacity: usize,
103 drop_notify: UnboundedSender<K>,
104}
105
106pub struct KeyReceiver<K: Key, V: Value> {
110 key: K,
111 receiver: broadcast::Receiver<V>,
112 on_drop: UnboundedSender<K>,
113}
114
115impl<K: Key, V: Value> KeyStream<K, V> {
116 pub fn new(broadcast_capacity: usize) -> Self {
118 let (sender, receiver) = unbounded_channel::<K>();
119 let streams = Arc::new(RwLock::new(HashMap::<K, broadcast::Sender<V>>::new()));
120 let on_drop_task = tokio::spawn(cleanup_keys(receiver, Arc::clone(&streams)));
121 Self {
122 broadcast_capacity,
123 streams,
124 sender,
125 drop_keys_task: Some(on_drop_task),
126 }
127 }
128
129 pub fn sender(&self) -> KeySender<K, V> {
131 KeySender::new(
132 Arc::clone(&self.streams),
133 self.broadcast_capacity,
134 self.sender.clone(),
135 )
136 }
137
138 pub async fn n_keys(&self) -> usize {
140 self.streams.read().await.len()
141 }
142
143 pub async fn keys_capacity(&self) -> usize {
145 self.streams.read().await.capacity()
146 }
147}
148
149impl<K: Key, V: Value> KeySender<K, V> {
150 fn new(streams: Streams<K, V>, broadcast_capacity: usize, sender: UnboundedSender<K>) -> Self {
151 Self {
152 streams,
153 broadcast_capacity,
154 drop_notify: sender,
155 }
156 }
157
158 pub async fn send(&self, key: &K, value: V) -> Result<usize, broadcast::error::SendError<V>> {
162 let streams = self.streams.read().await;
163 if let Some(sender) = streams.get(key) {
164 sender.send(value)
165 } else {
166 Ok(0)
167 }
168 }
169
170 pub async fn subscribe(&self, key: K) -> KeyReceiver<K, V> {
174 let streams = self.streams.read().await;
175 let inner = if let Some(sender) = streams.get(&key) {
176 sender.subscribe()
177 } else {
178 drop(streams);
179 let mut streams = self.streams.write().await;
180 let sender = streams.entry(key.clone()).or_insert_with(|| {
181 let (sender, _) = broadcast::channel(self.broadcast_capacity);
182 sender
183 });
184 sender.subscribe()
185 };
186 self.create_receiver(key, inner)
187 }
188
189 pub async fn n_keys(&self) -> usize {
191 self.streams.read().await.len()
192 }
193
194 pub async fn key_capacity(&self) -> usize {
196 self.streams.read().await.capacity()
197 }
198
199 fn create_receiver(&self, key: K, inner: broadcast::Receiver<V>) -> KeyReceiver<K, V> {
200 KeyReceiver {
201 key,
202 receiver: inner,
203 on_drop: self.drop_notify.clone(),
204 }
205 }
206}
207
208impl<K: Key, V: Value> KeyReceiver<K, V> {
209 pub async fn recv(&mut self) -> Result<V, broadcast::error::RecvError> {
212 self.receiver.recv().await
213 }
214
215 pub async fn try_recv(&mut self) -> Result<V, broadcast::error::TryRecvError> {
218 self.receiver.try_recv()
219 }
220
221 pub fn blocking_recv(&mut self) -> Result<V, broadcast::error::RecvError> {
225 self.receiver.blocking_recv()
226 }
227
228 pub fn to_async_stream(self) -> impl Stream<Item = Result<V, RecvError>> {
230 futures::stream::unfold(self, |mut receiver| async {
231 let item = receiver.recv().await;
232 Some((item, receiver))
233 })
234 }
235}
236
237impl<K: Key, V: Value> Clone for KeySender<K, V> {
238 fn clone(&self) -> Self {
239 Self {
240 streams: Arc::clone(&self.streams),
241 broadcast_capacity: self.broadcast_capacity,
242 drop_notify: self.drop_notify.clone(),
243 }
244 }
245}
246
247impl<K: Key, V: Value> Drop for KeyStream<K, V> {
248 fn drop(&mut self) {
249 if let Some(task) = self.drop_keys_task.take() {
250 task.abort();
251 }
252 }
253}
254
255impl<K: Key, V: Value> Drop for KeyReceiver<K, V> {
256 fn drop(&mut self) {
257 let _ = self.on_drop.send(self.key.clone());
258 }
259}
260
261async fn cleanup_keys<K: Key, V: Value>(
262 mut receiver: UnboundedReceiver<K>,
263 streams: Streams<K, V>,
264) {
265 while let Some(key) = receiver.recv().await {
266 let mut streams = streams.write().await;
267 if let Some(sender) = streams.get(&key)
268 && sender.receiver_count() == 0
269 {
270 streams.remove(&key);
271 optimize_dict_mem(&mut streams);
272 }
273 }
274}
275
276fn optimize_dict_mem<K: Key, V: Value>(streams: &mut HashMap<K, broadcast::Sender<V>>) {
277 let cap = streams.capacity() >> 1;
280 let len = streams.len();
281 if cap > 64 && len < cap {
282 streams.shrink_to(cap);
283 }
284}
285
286#[cfg(test)]
287mod tests {
288 use std::pin::Pin;
289 use tokio::task::JoinSet;
290 use tokio::time::{Duration, timeout};
291
292 use super::*;
293
294 #[tokio::test]
295 async fn test_recv() {
296 let key_stream = KeyStream::<String, String>::new(10);
297 let sender = key_stream.sender();
298 let mut receiver = sender.subscribe("1".to_string()).await;
299 assert_eq!(key_stream.n_keys().await, 1);
300 sender
301 .send(&"1".to_string(), "value".to_string())
302 .await
303 .unwrap();
304 assert_eq!(receiver.recv().await.unwrap(), "value".to_string());
305 }
306
307 #[tokio::test]
308 async fn test_no_receiver() {
309 let key_stream = KeyStream::<String, String>::new(10);
310 let sender = key_stream.sender();
311 let result = sender
312 .send(&"1".to_string(), "value".to_string())
313 .await
314 .unwrap();
315 assert_eq!(result, 0);
316 assert_eq!(key_stream.n_keys().await, 0);
317 }
318
319 #[tokio::test]
320 async fn test_send_after_drop() {
321 let key_stream = KeyStream::<String, String>::new(10);
322 let sender = key_stream.sender();
323 let receiver = sender.subscribe("1".to_string()).await;
324 assert_eq!(key_stream.n_keys().await, 1);
325 drop(receiver);
326 tokio::task::yield_now().await;
328 let result = sender
329 .send(&"1".to_string(), "value".to_string())
330 .await
331 .unwrap();
332 assert_eq!(result, 0);
333 assert_eq!(key_stream.n_keys().await, 0);
334 }
335
336 #[tokio::test]
337 async fn test_lagged() {
338 let key_stream = KeyStream::<String, String>::new(1);
339 let sender = key_stream.sender();
340 let mut receiver = sender.subscribe("1".to_string()).await;
341 let result = sender
342 .send(&"1".to_string(), "value1".to_string())
343 .await
344 .unwrap();
345 assert_eq!(result, 1);
346 let result = sender
347 .send(&"1".to_string(), "value2".to_string())
348 .await
349 .unwrap();
350 assert_eq!(result, 1);
351 assert_eq!(
352 receiver.recv().await,
353 Err(broadcast::error::RecvError::Lagged(1))
354 );
355 assert_eq!(receiver.recv().await.unwrap(), "value2".to_string());
356 }
357
358 #[tokio::test]
359 async fn test_recv_struct() {
360 #[derive(Clone, Debug)]
361 struct MyStruct {
362 field1: String,
363 field2: i32,
364 }
365 let key_stream = KeyStream::<String, MyStruct>::new(10);
366 let sender = key_stream.sender();
367 let mut receiver = sender.subscribe("1".to_string()).await;
368 assert_eq!(key_stream.n_keys().await, 1);
369 sender
370 .send(
371 &"1".to_string(),
372 MyStruct {
373 field1: "value".to_string(),
374 field2: 42,
375 },
376 )
377 .await
378 .unwrap();
379 let received = receiver.recv().await.unwrap();
380 assert_eq!(received.field1, "value".to_string());
381 assert_eq!(received.field2, 42);
382 }
383
384 #[tokio::test]
385 async fn test_recv_arc_struct() {
386 #[derive(Debug)]
387 struct MyStruct {
388 field1: String,
389 field2: i32,
390 }
391 let key_stream = KeyStream::<String, Arc<MyStruct>>::new(10);
392 let sender = key_stream.sender();
393 let mut receiver = sender.subscribe("1".to_string()).await;
394 assert_eq!(key_stream.n_keys().await, 1);
395 sender
396 .send(
397 &"1".to_string(),
398 Arc::new(MyStruct {
399 field1: "value".to_string(),
400 field2: 42,
401 }),
402 )
403 .await
404 .unwrap();
405 let received = receiver.recv().await.unwrap();
406
407 assert_eq!(received.field1, "value".to_string());
408 assert_eq!(received.field2, 42);
409 }
410
411 #[tokio::test]
412 async fn test_messages_broadcasted() {
413 let key_stream = KeyStream::<String, String>::new(10);
414 let sender = key_stream.sender();
415 let mut receiver1 = sender.subscribe("1".to_string()).await;
416 let mut receiver2 = sender.subscribe("1".to_string()).await;
417 assert_eq!(key_stream.n_keys().await, 1);
418 sender
419 .send(&"1".to_string(), "value".to_string())
420 .await
421 .unwrap();
422 assert_eq!(receiver1.recv().await.unwrap(), "value".to_string());
423 assert_eq!(receiver2.recv().await.unwrap(), "value".to_string());
424 }
425
426 #[tokio::test]
427 async fn test_key_filter() {
428 let key_stream = KeyStream::<String, String>::new(10);
429 let sender = key_stream.sender();
430 let mut receiver1 = sender.subscribe("1".to_string()).await;
431 let mut receiver2 = sender.subscribe("2".to_string()).await;
432 assert_eq!(key_stream.n_keys().await, 2);
433 sender
434 .send(&"1".to_string(), "value1".to_string())
435 .await
436 .unwrap();
437 sender
438 .send(&"2".to_string(), "value2".to_string())
439 .await
440 .unwrap();
441 assert_eq!(receiver1.recv().await.unwrap(), "value1".to_string());
442 assert_eq!(receiver2.recv().await.unwrap(), "value2".to_string());
443 assert_eq!(
444 receiver1.try_recv().await,
445 Err(broadcast::error::TryRecvError::Empty)
446 );
447 assert_eq!(
448 receiver2.try_recv().await,
449 Err(broadcast::error::TryRecvError::Empty)
450 );
451 }
452
453 #[tokio::test]
454 async fn test_key_drop() {
455 let key_stream = KeyStream::<String, String>::new(10);
456 let sender = key_stream.sender();
457 let receiver = sender.subscribe("1".to_string()).await;
458 assert_eq!(key_stream.n_keys().await, 1);
459 sender
460 .send(&"1".to_string(), "value".to_string())
461 .await
462 .unwrap();
463 drop(receiver);
464 tokio::task::yield_now().await;
466 assert_eq!(key_stream.n_keys().await, 0);
467 }
468
469 #[tokio::test]
470 async fn test_key_non_dropped_if_other_receiver_exists() {
471 let key_stream = KeyStream::<String, String>::new(10);
472 let sender = key_stream.sender();
473 let receiver1 = sender.subscribe("1".to_string()).await;
474 let _receiver2 = sender.subscribe("1".to_string()).await;
475 assert_eq!(key_stream.n_keys().await, 1);
476 sender
477 .send(&"1".to_string(), "value".to_string())
478 .await
479 .unwrap();
480 drop(receiver1);
481 tokio::task::yield_now().await;
483 assert_eq!(key_stream.n_keys().await, 1);
484 }
485
486 #[tokio::test]
487 async fn test_two_clients() {
488 let key_stream = KeyStream::<String, String>::new(10);
489 let sender1 = key_stream.sender();
490 let sender2 = key_stream.sender();
491 let mut receiver1 = sender1.subscribe("1".to_string()).await;
492 let mut receiver2 = sender2.subscribe("1".to_string()).await;
493
494 assert_eq!(key_stream.n_keys().await, 1);
495
496 sender1
497 .send(&"1".to_string(), "value".to_string())
498 .await
499 .unwrap();
500
501 assert_eq!(receiver1.recv().await.unwrap(), "value".to_string());
502 assert_eq!(receiver2.recv().await.unwrap(), "value".to_string());
503 assert_eq!(
504 receiver1.try_recv().await,
505 Err(broadcast::error::TryRecvError::Empty)
506 );
507 assert_eq!(
508 receiver2.try_recv().await,
509 Err(broadcast::error::TryRecvError::Empty)
510 );
511
512 drop(sender2);
513 tokio::task::yield_now().await;
515
516 assert_eq!(key_stream.n_keys().await, 1);
517
518 sender1
519 .send(&"1".to_string(), "value".to_string())
520 .await
521 .unwrap();
522
523 assert_eq!(receiver1.recv().await.unwrap(), "value".to_string());
524 }
525
526 #[tokio::test]
527 async fn test_reconnect() {
528 let key_stream = KeyStream::<String, String>::new(10);
529 let sender = key_stream.sender();
530 let mut receiver1 = sender.subscribe("1".to_string()).await;
531 assert_eq!(key_stream.n_keys().await, 1);
532 sender
533 .send(&"1".to_string(), "value".to_string())
534 .await
535 .unwrap();
536 assert_eq!(receiver1.recv().await.unwrap(), "value".to_string());
537 drop(receiver1);
538 tokio::task::yield_now().await;
540 assert_eq!(key_stream.n_keys().await, 0);
541 let mut receiver2 = sender.subscribe("1".to_string()).await;
542 assert_eq!(key_stream.n_keys().await, 1);
543 sender
544 .send(&"1".to_string(), "value2".to_string())
545 .await
546 .unwrap();
547 assert_eq!(receiver2.recv().await.unwrap(), "value2".to_string());
548 }
549
550 #[tokio::test]
551 async fn test_shrink_dict() {
552 let key_stream = KeyStream::<i32, String>::new(1);
553 let sender = key_stream.sender();
554 let mut set = JoinSet::new();
555 let mut subs = (0..500)
556 .map(|i| {
557 set.spawn({
558 let value = sender.clone();
559 async move { value.subscribe(i).await }
560 })
561 })
562 .collect::<Vec<_>>();
563 set.join_all().await;
564 assert_eq!(key_stream.n_keys().await, 500);
565 subs.drain(0..400);
566 tokio::task::yield_now().await;
568 assert!(key_stream.keys_capacity().await < 500);
569 }
570
571 #[tokio::test]
572 async fn test_drop_sender() {
573 let key_stream = KeyStream::<String, String>::new(10);
574 let sender = key_stream.sender();
575 let mut receiver = sender.subscribe("1".to_string()).await;
576 assert_eq!(key_stream.n_keys().await, 1);
577 sender
578 .send(&"1".to_string(), "value".to_string())
579 .await
580 .unwrap();
581 drop(sender);
582 tokio::task::yield_now().await;
583 assert_eq!(receiver.recv().await.unwrap(), "value".to_string());
584 assert_eq!(
585 receiver.try_recv().await,
586 Err(broadcast::error::TryRecvError::Empty)
587 );
588 }
589
590 #[tokio::test]
591 async fn test_drop_stream() {
592 let key_stream = KeyStream::<String, String>::new(10);
595 let sender = key_stream.sender();
596 drop(key_stream);
597 let mut receiver = sender.subscribe("1".to_string()).await;
598 sender
599 .send(&"1".to_string(), "value".to_string())
600 .await
601 .unwrap();
602 tokio::task::yield_now().await;
603 assert_eq!(receiver.recv().await.unwrap(), "value".to_string());
604 assert_eq!(
605 receiver.try_recv().await,
606 Err(broadcast::error::TryRecvError::Empty)
607 );
608 }
609
610 #[tokio::test]
611 async fn test_len_and_capacity() {
612 let key_stream = KeyStream::<String, String>::new(10);
613 let sender1 = key_stream.sender();
614 let sender2 = key_stream.sender();
615 assert_eq!(key_stream.n_keys().await, 0);
616 assert_eq!(sender1.n_keys().await, 0);
617 assert_eq!(sender2.n_keys().await, 0);
618 assert_eq!(
619 key_stream.keys_capacity().await,
620 sender1.key_capacity().await
621 );
622 assert_eq!(
623 key_stream.keys_capacity().await,
624 sender2.key_capacity().await
625 );
626 let _receiver1 = sender1.subscribe("1".to_string()).await;
627 let _receiver2 = sender2.subscribe("2".to_string()).await;
628 let _receiver3 = sender2.subscribe("3".to_string()).await;
629 assert_eq!(key_stream.n_keys().await, 3);
630 assert_eq!(sender1.n_keys().await, 3);
631 assert_eq!(sender2.n_keys().await, 3);
632 assert_eq!(
633 key_stream.keys_capacity().await,
634 sender1.key_capacity().await
635 );
636 assert_eq!(
637 key_stream.keys_capacity().await,
638 sender2.key_capacity().await
639 );
640 }
641
642 #[tokio::test]
643 async fn test_stream() {
644 use futures::StreamExt;
645 type WatchStream =
646 Pin<Box<dyn futures::Stream<Item = Result<String, RecvError>> + Send + 'static>>;
647 let key_stream = KeyStream::<String, String>::new(10);
648 let sender = key_stream.sender();
649 let receiver = sender.subscribe("1".to_string()).await;
650 sender
651 .send(&"1".to_string(), "value1".to_string())
652 .await
653 .unwrap();
654 let stream = receiver.to_async_stream();
655 let stream = Box::pin(stream) as WatchStream;
656 sender
657 .send(&"1".to_string(), "value2".to_string())
658 .await
659 .unwrap();
660 let msg = timeout(Duration::from_secs(1), stream.take(2).collect::<Vec<_>>()).await;
661 let msg = msg.expect("timeout");
662 assert_eq!(
663 msg,
664 vec![Ok("value1".to_string()), Ok("value2".to_string())]
665 );
666 }
667}