1use std::future::{Future, IntoFuture};
10use std::ops::Drop;
11use std::pin::Pin;
12use std::sync::{Arc, Mutex, OnceLock, Weak};
13use std::task::{Context, Poll, Waker, ready};
14
15use slotmap_careful::DenseSlotMap;
16
17slotmap_careful::new_key_type! { struct WakerKey; }
18
19#[derive(Debug)]
21pub struct Sender<T> {
22 shared: Weak<Shared<T>>,
24}
25
26#[derive(Clone, Debug)]
52pub struct Receiver<T> {
53 shared: Arc<Shared<T>>,
55}
56
57#[derive(Debug)]
80struct Shared<T> {
81 msg: OnceLock<Result<T, SenderDropped>>,
83 wakers: Mutex<Result<DenseSlotMap<WakerKey, Waker>, WakersAlreadyWoken>>,
89}
90
91#[derive(Debug)]
96pub struct BorrowedReceiverFuture<'a, T> {
97 shared: &'a Shared<T>,
99 waker_key: Option<WakerKey>,
101}
102
103#[derive(Debug)]
114pub struct ReceiverFuture<T> {
115 shared: Arc<Shared<T>>,
117 waker_key: Option<WakerKey>,
119}
120
121#[derive(Copy, Clone, Debug)]
128struct WakersAlreadyWoken;
129
130#[derive(Copy, Clone, Debug, thiserror::Error)]
132#[error("the message was already set")]
133struct MessageAlreadySet;
134
135#[derive(Copy, Clone, Debug, PartialEq, Eq, thiserror::Error)]
137#[error("the sender was dropped")]
138#[allow(clippy::exhaustive_structs)]
139pub struct SenderDropped;
140
141#[derive(Copy, Clone, Debug, PartialEq, Eq, thiserror::Error)]
143#[error("all the receivers were dropped")]
144#[allow(clippy::exhaustive_structs)]
145pub struct AllReceiversDropped;
146
147pub fn channel<T>() -> (Sender<T>, Receiver<T>) {
161 let shared = Arc::new(Shared {
162 msg: OnceLock::new(),
163 wakers: Mutex::new(Ok(DenseSlotMap::with_key())),
164 });
165
166 let sender = Sender {
167 shared: Arc::downgrade(&shared),
168 };
169
170 let receiver = Receiver { shared };
171
172 (sender, receiver)
173}
174
175impl<T> Sender<T> {
176 pub fn send(self, msg: T) {
180 Self::send_and_wake(&self.shared, Ok(msg))
182 .expect("could not set the message");
186 }
187
188 fn send_and_wake(
194 shared: &Weak<Shared<T>>,
195 msg: Result<T, SenderDropped>,
196 ) -> Result<(), MessageAlreadySet> {
197 let Some(shared) = shared.upgrade() else {
202 return Ok(());
204 };
205
206 shared.msg.set(msg).or(Err(MessageAlreadySet))?;
208
209 let mut wakers = {
210 let mut wakers = shared.wakers.lock().expect("poisoned");
211 std::mem::replace(&mut *wakers, Err(WakersAlreadyWoken))
220 .expect("wakers were taken more than once")
221 };
222
223 for (_key, waker) in wakers.drain() {
233 waker.wake();
234 }
235
236 Ok(())
237 }
238
239 pub fn is_cancelled(&self) -> bool {
249 self.shared.strong_count() == 0
250 }
251
252 pub fn subscribe(&self) -> Result<Receiver<T>, AllReceiversDropped> {
257 Ok(Receiver {
258 shared: self.shared.upgrade().ok_or(AllReceiversDropped)?,
259 })
260 }
261}
262
263impl<T> Drop for Sender<T> {
264 fn drop(&mut self) {
265 let _ = Self::send_and_wake(&self.shared, Err(SenderDropped));
269 }
270}
271
272impl<T> Receiver<T> {
273 pub fn borrowed(&self) -> BorrowedReceiverFuture<'_, T> {
280 BorrowedReceiverFuture {
281 shared: &self.shared,
282 waker_key: None,
283 }
284 }
285
286 pub fn is_ready(&self) -> bool {
290 self.shared.msg.get().is_some()
291 }
292}
293
294impl<T: Clone> IntoFuture for Receiver<T> {
295 type Output = Result<T, SenderDropped>;
296 type IntoFuture = ReceiverFuture<T>;
297
298 fn into_future(self) -> Self::IntoFuture {
300 ReceiverFuture {
301 shared: self.shared,
302 waker_key: None,
303 }
304 }
305}
306
307impl<'a, T> Future for BorrowedReceiverFuture<'a, T> {
308 type Output = Result<&'a T, SenderDropped>;
309
310 fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
311 let self_ = self.get_mut();
312 receiver_fut_poll(self_.shared, &mut self_.waker_key, cx.waker())
313 }
314}
315
316impl<T> Drop for BorrowedReceiverFuture<'_, T> {
317 fn drop(&mut self) {
318 receiver_fut_drop(self.shared, &mut self.waker_key);
319 }
320}
321
322impl<T: Clone> Future for ReceiverFuture<T> {
323 type Output = Result<T, SenderDropped>;
324
325 fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
326 let self_ = self.get_mut();
327 let poll = receiver_fut_poll(&self_.shared, &mut self_.waker_key, cx.waker());
328 Poll::Ready(ready!(poll)).map_ok(Clone::clone)
329 }
330}
331
332impl<T> Drop for ReceiverFuture<T> {
333 fn drop(&mut self) {
334 receiver_fut_drop(&self.shared, &mut self.waker_key);
335 }
336}
337
338fn receiver_fut_poll<'a, T>(
340 shared: &'a Shared<T>,
341 waker_key: &mut Option<WakerKey>,
342 new_waker: &Waker,
343) -> Poll<Result<&'a T, SenderDropped>> {
344 if let Some(msg) = shared.msg.get() {
346 return Poll::Ready(msg.as_ref().or(Err(SenderDropped)));
347 }
348
349 let mut wakers = shared.wakers.lock().expect("poisoned");
350
351 if let Some(msg) = shared.msg.get() {
353 return Poll::Ready(msg.as_ref().or(Err(SenderDropped)));
354 }
355
356 let wakers = wakers.as_mut().expect("wakers were already woken");
360
361 match waker_key {
362 Some(waker_key) => {
364 let waker = wakers
366 .get_mut(*waker_key)
367 .expect("waker key is missing from map");
370 waker.clone_from(new_waker);
371 }
372 None => {
374 let new_key = wakers.insert(new_waker.clone());
376 *waker_key = Some(new_key);
377 }
378 }
379
380 Poll::Pending
381}
382
383fn receiver_fut_drop<T>(shared: &Shared<T>, waker_key: &mut Option<WakerKey>) {
385 if let Some(waker_key) = waker_key.take() {
386 let mut wakers = shared.wakers.lock().expect("poisoned");
387 if let Ok(wakers) = wakers.as_mut() {
388 let waker = wakers.remove(waker_key);
389 debug_assert!(waker.is_some(), "the waker key was not found");
392 }
393 }
394}
395
396#[cfg(test)]
397mod test {
398 #![allow(clippy::unwrap_used)]
399
400 use super::*;
401
402 use futures::future::FutureExt;
403 use tor_rtcompat::SpawnExt;
404
405 impl<T> Shared<T> {
406 fn count_wakers(&self) -> usize {
408 self.wakers
409 .lock()
410 .expect("poisoned")
411 .as_ref()
412 .map(|x| x.len())
413 .unwrap_or(0)
414 }
415 }
416
417 #[test]
418 fn standard_usage() {
419 tor_rtmock::MockRuntime::test_with_various(|_rt| async move {
420 let (tx, rx) = channel();
421 tx.send(0_u8);
422 assert_eq!(rx.borrowed().await, Ok(&0));
423
424 let (tx, rx) = channel();
425 tx.send(0_u8);
426 assert_eq!(rx.await, Ok(0));
427 });
428 }
429
430 #[test]
431 fn immediate_drop() {
432 let _ = channel::<()>();
433
434 let (tx, rx) = channel::<()>();
435 drop(tx);
436 drop(rx);
437
438 let (tx, rx) = channel::<()>();
439 drop(rx);
440 drop(tx);
441 }
442
443 #[test]
444 fn drop_sender() {
445 tor_rtmock::MockRuntime::test_with_various(|_rt| async move {
446 let (tx, rx_1) = channel::<u8>();
447
448 let rx_2 = rx_1.clone();
449 drop(tx);
450 let rx_3 = rx_1.clone();
451 assert_eq!(rx_1.borrowed().await, Err(SenderDropped));
452 assert_eq!(rx_2.borrowed().await, Err(SenderDropped));
453 assert_eq!(rx_3.borrowed().await, Err(SenderDropped));
454 });
455 }
456
457 #[test]
458 fn clone_before_send() {
459 tor_rtmock::MockRuntime::test_with_various(|_rt| async move {
460 let (tx, rx_1) = channel();
461
462 let rx_2 = rx_1.clone();
463 tx.send(0_u8);
464 assert_eq!(rx_1.borrowed().await, Ok(&0));
465 assert_eq!(rx_2.borrowed().await, Ok(&0));
466 });
467 }
468
469 #[test]
470 fn clone_after_send() {
471 tor_rtmock::MockRuntime::test_with_various(|_rt| async move {
472 let (tx, rx_1) = channel();
473
474 tx.send(0_u8);
475 let rx_2 = rx_1.clone();
476 assert_eq!(rx_1.borrowed().await, Ok(&0));
477 assert_eq!(rx_2.borrowed().await, Ok(&0));
478 });
479 }
480
481 #[test]
482 fn clone_after_borrowed() {
483 tor_rtmock::MockRuntime::test_with_various(|_rt| async move {
484 let (tx, rx_1) = channel();
485
486 tx.send(0_u8);
487 assert_eq!(rx_1.borrowed().await, Ok(&0));
488 let rx_2 = rx_1.clone();
489 assert_eq!(rx_2.borrowed().await, Ok(&0));
490 });
491 }
492
493 #[test]
494 fn drop_one_receiver() {
495 tor_rtmock::MockRuntime::test_with_various(|_rt| async move {
496 let (tx, rx_1) = channel();
497
498 let rx_2 = rx_1.clone();
499 drop(rx_1);
500 tx.send(0_u8);
501 assert_eq!(rx_2.borrowed().await, Ok(&0));
502 });
503 }
504
505 #[test]
506 fn drop_all_receivers() {
507 let (tx, rx_1) = channel();
508
509 let rx_2 = rx_1.clone();
510 drop(rx_1);
511 drop(rx_2);
512 tx.send(0_u8);
513 }
514
515 #[test]
516 fn drop_fut() {
517 let (_tx, rx) = channel::<u8>();
518 let fut = rx.borrowed();
519 assert_eq!(rx.shared.count_wakers(), 0);
520 drop(fut);
521 assert_eq!(rx.shared.count_wakers(), 0);
522
523 let (tx, rx) = channel();
525 tx.send(0_u8);
526 let fut = rx.borrowed();
527 assert_eq!(rx.shared.count_wakers(), 0);
528 drop(fut);
529 assert_eq!(rx.shared.count_wakers(), 0);
530
531 let (_tx, rx) = channel::<u8>();
533 let mut fut = Box::pin(rx.borrowed());
534 assert_eq!(rx.shared.count_wakers(), 0);
535 assert_eq!(fut.as_mut().now_or_never(), None);
536 assert_eq!(rx.shared.count_wakers(), 1);
537 drop(fut);
538 assert_eq!(rx.shared.count_wakers(), 0);
539
540 let (tx, rx) = channel();
542 let mut fut = Box::pin(rx.borrowed());
543 assert_eq!(rx.shared.count_wakers(), 0);
544 assert_eq!(fut.as_mut().now_or_never(), None);
545 assert_eq!(rx.shared.count_wakers(), 1);
546 tx.send(0_u8);
547 assert_eq!(rx.shared.count_wakers(), 0);
548 drop(fut);
549 }
550
551 #[test]
552 fn drop_owned_fut() {
553 let (_tx, rx) = channel::<u8>();
554 let fut = rx.clone().into_future();
555 assert_eq!(rx.shared.count_wakers(), 0);
556 drop(fut);
557 assert_eq!(rx.shared.count_wakers(), 0);
558
559 let (tx, rx) = channel();
561 tx.send(0_u8);
562 let fut = rx.clone().into_future();
563 assert_eq!(rx.shared.count_wakers(), 0);
564 drop(fut);
565 assert_eq!(rx.shared.count_wakers(), 0);
566
567 let (_tx, rx) = channel::<u8>();
569 let mut fut = Box::pin(rx.clone().into_future());
570 assert_eq!(rx.shared.count_wakers(), 0);
571 assert_eq!(fut.as_mut().now_or_never(), None);
572 assert_eq!(rx.shared.count_wakers(), 1);
573 drop(fut);
574 assert_eq!(rx.shared.count_wakers(), 0);
575
576 let (tx, rx) = channel();
578 let mut fut = Box::pin(rx.clone().into_future());
579 assert_eq!(rx.shared.count_wakers(), 0);
580 assert_eq!(fut.as_mut().now_or_never(), None);
581 assert_eq!(rx.shared.count_wakers(), 1);
582 tx.send(0_u8);
583 assert_eq!(rx.shared.count_wakers(), 0);
584 drop(fut);
585 }
586
587 #[test]
588 fn is_ready_after_send() {
589 let (tx, rx_1) = channel();
590 assert!(!rx_1.is_ready());
591 let rx_2 = rx_1.clone();
592 assert!(!rx_2.is_ready());
593
594 tx.send(0_u8);
595
596 assert!(rx_1.is_ready());
597 assert!(rx_2.is_ready());
598
599 let rx_3 = rx_1.clone();
600 assert!(rx_3.is_ready());
601 }
602
603 #[test]
604 fn is_ready_after_drop() {
605 let (tx, rx_1) = channel::<u8>();
606 assert!(!rx_1.is_ready());
607 let rx_2 = rx_1.clone();
608 assert!(!rx_2.is_ready());
609
610 drop(tx);
611
612 assert!(rx_1.is_ready());
613 assert!(rx_2.is_ready());
614
615 let rx_3 = rx_1.clone();
616 assert!(rx_3.is_ready());
617 }
618
619 #[test]
620 fn is_cancelled() {
621 let (tx, rx) = channel::<u8>();
622 assert!(!tx.is_cancelled());
623 drop(rx);
624 assert!(tx.is_cancelled());
625
626 let (tx, rx_1) = channel::<u8>();
627 assert!(!tx.is_cancelled());
628 let rx_2 = rx_1.clone();
629 drop(rx_1);
630 assert!(!tx.is_cancelled());
631 drop(rx_2);
632 assert!(tx.is_cancelled());
633 }
634
635 #[test]
636 fn recv_in_task() {
637 tor_rtmock::MockRuntime::test_with_various(|rt| async move {
638 let (tx, rx) = channel();
639
640 let join = rt
641 .spawn_with_handle(async move {
642 assert_eq!(rx.borrowed().await, Ok(&0));
643 assert_eq!(rx.await, Ok(0));
644 })
645 .unwrap();
646
647 tx.send(0_u8);
648
649 join.await;
650 });
651 }
652
653 #[test]
654 fn recv_multiple_in_task() {
655 tor_rtmock::MockRuntime::test_with_various(|rt| async move {
656 let (tx, rx) = channel();
657 let rx_1 = rx.clone();
658 let rx_2 = rx.clone();
659
660 let join_1 = rt
661 .spawn_with_handle(async move {
662 assert_eq!(rx_1.borrowed().await, Ok(&0));
663 })
664 .unwrap();
665 let join_2 = rt
666 .spawn_with_handle(async move {
667 assert_eq!(rx_2.await, Ok(0));
668 })
669 .unwrap();
670
671 tx.send(0_u8);
672
673 join_1.await;
674 join_2.await;
675 assert_eq!(rx.borrowed().await, Ok(&0));
676 });
677 }
678
679 #[test]
680 fn recv_multiple_times() {
681 tor_rtmock::MockRuntime::test_with_various(|_rt| async move {
682 let (tx, rx) = channel();
683 let rx_subscribed = tx.subscribe().unwrap();
684
685 tx.send(0_u8);
686 assert_eq!(rx.borrowed().await, Ok(&0));
687 assert_eq!(rx.borrowed().await, Ok(&0));
688 assert_eq!(rx.clone().await, Ok(0));
689 assert_eq!(rx.await, Ok(0));
690 assert_eq!(rx_subscribed.await, Ok(0));
691 });
692 }
693
694 #[test]
695 fn stress() {
696 tor_rtmock::MockRuntime::test_with_various(|rt| async move {
710 let (tx, rx) = channel();
711
712 rt.spawn(async move {
713 for _ in 0..20 {
716 tor_rtcompat::task::yield_now().await;
717 }
718 tx.send(0_u8);
719 })
720 .unwrap();
721
722 let mut joins = vec![];
723 for _ in 0..100 {
724 let rx_clone = rx.clone();
725 let join = rt
726 .spawn_with_handle(async move { rx_clone.borrowed().await.cloned() })
727 .unwrap();
728 joins.push(join);
729 tor_rtcompat::task::yield_now().await;
731 }
732
733 for join in joins {
734 assert!(matches!(join.await, Ok(0)));
735 }
736 });
737 }
738}