Skip to main content

key_stream/
lib.rs

1//! Async key-based message streaming library.
2//!
3//! Enables sending messages to multiple receivers, grouped by keys, with automatic cleanup of unused keys.
4//! No messages are sent until a key is subscribed to, and keys are automatically removed when all receivers are dropped.
5//! Cleanup is performed by a background task, so task switching is required for timely removal.
6//! The memory usage of the keys map is optimized by shrinking it when many keys are removed.
7//!
8//! The main entry point is [`KeyStream`], created with [`KeyStream::new`].
9//! It manages the keys and contains a background task for cleaning up unused keys.
10//! The cleanup task is immediately notified when a key is dropped, rather than relying on periodic checks.
11//! Use [`KeyStream::sender`] to obtain a sender handle for sending messages and subscribing to keys.
12//! Use [`KeySender::send`] to send messages to a key, and [`KeySender::subscribe`] to subscribe to a key and receive a [`KeyReceiver`] for that key.
13//! [`KeyReceiver::recv`] can be used to receive messages for a key, waiting asynchronously until a message is available.
14//! [`KeyReceiver`] can also be converted into a stream of messages using [`KeyReceiver::to_async_stream`].
15//!
16//! # Examples
17//!
18//! ```
19//! use key_stream::KeyStream;
20//! use tokio;
21//!
22//! # #[tokio::main(flavor = "current_thread")]
23//! # async fn main() {
24//! let key_stream = KeyStream::<i32, String>::new(10);
25//! let sender = key_stream.sender();
26//! let mut receiver = sender.subscribe(1).await;
27//! sender.send(&1, "value".to_string()).await.unwrap();
28//! assert_eq!(receiver.recv().await.unwrap(), "value".to_string());
29//! # }
30//! ```
31//!
32//! ## Key cleanup example
33//!
34//! ```rust
35//! use key_stream::KeyStream;
36//! # #[tokio::main(flavor = "current_thread")]
37//! # async fn main() {
38//! let key_stream = KeyStream::<i32, String>::new(10);
39//! let sender = key_stream.sender();
40//! let receiver = sender.subscribe(1).await;
41//! assert_eq!(key_stream.n_keys().await, 1);
42//! drop(receiver);
43//! // give the key drop task a chance to run
44//! tokio::task::yield_now().await;
45//! let result = sender
46//!     .send(&1, "value".to_string())
47//!     .await
48//!     .unwrap();
49//! assert_eq!(result, 0);
50//! assert_eq!(key_stream.n_keys().await, 0);
51//! # }
52//! ```
53use 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
60/// Trait bound for types usable as keys in [`KeyStream`].
61/// Must be hashable, comparable, cloneable, thread-safe, and `'static`.
62pub trait Key: Hash + Eq + Clone + Send + Sync + 'static {}
63
64/// Trait bound for types usable as values in [`KeyStream`].
65/// Must be cloneable, thread-safe, and `'static`.
66pub 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
73/// The main entry point for key-based async message streaming.
74///
75/// Use [`KeyStream::new`] to create, then call [`KeyStream::sender`] to get a sender handle.
76///
77/// # Example
78/// ```
79/// use key_stream::KeyStream;
80/// use tokio;
81/// # #[tokio::main(flavor = "current_thread")]
82/// # async fn main() {
83/// let key_stream = KeyStream::<i32, String>::new(10);
84/// let sender = key_stream.sender();
85/// let mut receiver = sender.subscribe(1).await;
86/// sender.send(&1, "value".to_string()).await.unwrap();
87/// assert_eq!(receiver.recv().await.unwrap(), "value".to_string());
88/// # }
89/// ```
90pub 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
97/// Handle for sending and subscribing to messages by key.
98///
99/// Created via [`KeyStream::sender`].
100pub struct KeySender<K: Key, V: Value> {
101    streams: Streams<K, V>,
102    broadcast_capacity: usize,
103    drop_notify: UnboundedSender<K>,
104}
105
106/// A receiver for messages for a specific key.
107///
108/// Created via [`KeySender::subscribe`].
109pub 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    /// Create a new [`KeyStream`] with the given broadcast channel capacity per key.
117    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    /// Get a sender handle for publishing and subscribing to keys.
130    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    /// Get the number of keys currently tracked.
139    pub async fn n_keys(&self) -> usize {
140        self.streams.read().await.len()
141    }
142
143    /// Get the current capacity of the keys map.
144    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    /// Send a value to all receivers subscribed to the given key.
159    ///
160    /// Returns the number of receivers the message was sent to, or 0 if none.
161    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    /// Subscribe to messages for the given key.
171    ///
172    /// Returns a [`KeyReceiver`] for receiving messages.
173    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    /// Get the number of keys currently tracked.
190    pub async fn n_keys(&self) -> usize {
191        self.streams.read().await.len()
192    }
193
194    /// Get the current capacity of the keys map.
195    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    /// Receive the next message for this key, waiting asynchronously.
210    /// See, tokio::sync::broadcast::Receiver::recv for more details.
211    pub async fn recv(&mut self) -> Result<V, broadcast::error::RecvError> {
212        self.receiver.recv().await
213    }
214
215    /// Try to receive the next message for this key without waiting..
216    /// See, tokio::sync::broadcast::Receiver::try_recv for more details.
217    pub async fn try_recv(&mut self) -> Result<V, broadcast::error::TryRecvError> {
218        self.receiver.try_recv()
219    }
220
221    /// Receive the next message for this key, blocking the current thread.
222    /// Only use in synchronous contexts.
223    /// See, tokio::sync::broadcast::Receiver::blocking_recv for more details.
224    pub fn blocking_recv(&mut self) -> Result<V, broadcast::error::RecvError> {
225        self.receiver.blocking_recv()
226    }
227
228    /// Consume this receiver and convert it into a stream of messages for this key.
229    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    // If the number of keys is less than half the capacity and the capacity is big enough,
278    // shrink the capacity to save memory
279    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        // give the on_drop task a chance to run
327        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        // give the on_drop task a chance to run
465        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        // give the on_drop task a chance to run
482        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        // give the on_drop task a chance to run
514        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        // give the on_drop task a chance to run
539        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        // give the on_drop task a chance to run
567        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        // This test ensures that dropping the KeyStream doesn't cause any panics
593        // actual data may vary
594        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}