1use crate::loom::cell::UnsafeCell;
2use crate::loom::future::AtomicWaker;
3use crate::loom::sync::atomic::AtomicUsize;
4use crate::loom::sync::Arc;
5use crate::runtime::park::CachedParkThread;
6use crate::sync::mpsc::error::TryRecvError;
7use crate::sync::mpsc::{bounded, list, unbounded};
8use crate::sync::notify::Notify;
9use crate::util::cacheline::CachePadded;
10
11use std::fmt;
12use std::panic;
13use std::process;
14use std::sync::atomic::Ordering::{AcqRel, Acquire, Relaxed, Release};
15use std::task::Poll::{Pending, Ready};
16use std::task::{ready, Context, Poll};
17
18pub(crate) struct Tx<T, S> {
20 inner: Arc<Chan<T, S>>,
21}
22
23impl<T, S: fmt::Debug> fmt::Debug for Tx<T, S> {
24 fn fmt(&self, fmt: &mut fmt::Formatter<'_>) -> fmt::Result {
25 fmt.debug_struct("Tx").field("inner", &self.inner).finish()
26 }
27}
28
29pub(crate) struct Rx<T, S: Semaphore> {
31 inner: Arc<Chan<T, S>>,
32}
33
34impl<T, S: Semaphore + fmt::Debug> fmt::Debug for Rx<T, S> {
35 fn fmt(&self, fmt: &mut fmt::Formatter<'_>) -> fmt::Result {
36 fmt.debug_struct("Rx").field("inner", &self.inner).finish()
37 }
38}
39
40pub(crate) trait Semaphore {
41 fn is_idle(&self) -> bool;
42
43 fn add_permit(&self);
44
45 fn add_permits(&self, n: usize);
46
47 fn close(&self);
48
49 fn is_closed(&self) -> bool;
50}
51
52pub(super) struct Chan<T, S> {
53 tx: CachePadded<list::Tx<T>>,
55
56 rx_waker: CachePadded<AtomicWaker>,
58
59 notify_rx_closed: Notify,
61
62 semaphore: S,
64
65 tx_count: AtomicUsize,
69
70 tx_weak_count: AtomicUsize,
72
73 rx_fields: UnsafeCell<RxFields<T>>,
75}
76
77impl<T, S> fmt::Debug for Chan<T, S>
78where
79 S: fmt::Debug,
80{
81 fn fmt(&self, fmt: &mut fmt::Formatter<'_>) -> fmt::Result {
82 fmt.debug_struct("Chan")
83 .field("tx", &*self.tx)
84 .field("semaphore", &self.semaphore)
85 .field("rx_waker", &*self.rx_waker)
86 .field("tx_count", &self.tx_count)
87 .field("rx_fields", &"...")
88 .finish()
89 }
90}
91
92struct RxFields<T> {
94 list: list::Rx<T>,
96
97 rx_closed: bool,
99}
100
101impl<T> fmt::Debug for RxFields<T> {
102 fn fmt(&self, fmt: &mut fmt::Formatter<'_>) -> fmt::Result {
103 fmt.debug_struct("RxFields")
104 .field("list", &self.list)
105 .field("rx_closed", &self.rx_closed)
106 .finish()
107 }
108}
109
110unsafe impl<T: Send, S: Send> Send for Chan<T, S> {}
111unsafe impl<T: Send, S: Sync> Sync for Chan<T, S> {}
112impl<T, S> panic::RefUnwindSafe for Chan<T, S> {}
113impl<T, S> panic::UnwindSafe for Chan<T, S> {}
114
115pub(crate) fn channel<T, S: Semaphore>(semaphore: S) -> (Tx<T, S>, Rx<T, S>) {
116 let (tx, rx) = list::channel();
117 channel_from_list(tx, rx, semaphore)
118}
119
120#[cfg(all(test, not(loom)))]
121pub(crate) fn channel_from_index<T, S: Semaphore>(
122 start_index: usize,
123 semaphore: S,
124) -> (Tx<T, S>, Rx<T, S>) {
125 let (tx, rx) = list::channel_from_index(start_index);
126 channel_from_list(tx, rx, semaphore)
127}
128
129fn channel_from_list<T, S: Semaphore>(
130 tx: list::Tx<T>,
131 rx: list::Rx<T>,
132 semaphore: S,
133) -> (Tx<T, S>, Rx<T, S>) {
134 let chan = Arc::new(Chan {
135 notify_rx_closed: Notify::new(),
136 tx: CachePadded::new(tx),
137 semaphore,
138 rx_waker: CachePadded::new(AtomicWaker::new()),
139 tx_count: AtomicUsize::new(1),
140 tx_weak_count: AtomicUsize::new(0),
141 rx_fields: UnsafeCell::new(RxFields {
142 list: rx,
143 rx_closed: false,
144 }),
145 });
146
147 (Tx::new(chan.clone()), Rx::new(chan))
148}
149
150impl<T, S> Tx<T, S> {
153 fn new(chan: Arc<Chan<T, S>>) -> Tx<T, S> {
154 Tx { inner: chan }
155 }
156
157 pub(super) fn strong_count(&self) -> usize {
158 self.inner.tx_count.load(Acquire)
159 }
160
161 pub(super) fn weak_count(&self) -> usize {
162 self.inner.tx_weak_count.load(Relaxed)
163 }
164
165 pub(super) fn downgrade(&self) -> Arc<Chan<T, S>> {
166 self.inner.increment_weak_count();
167
168 self.inner.clone()
169 }
170
171 pub(super) fn upgrade(chan: Arc<Chan<T, S>>) -> Option<Self> {
173 let mut tx_count = chan.tx_count.load(Acquire);
174
175 loop {
176 if tx_count == 0 {
177 return None;
179 }
180
181 match chan
182 .tx_count
183 .compare_exchange_weak(tx_count, tx_count + 1, AcqRel, Acquire)
184 {
185 Ok(_) => return Some(Tx { inner: chan }),
186 Err(prev_count) => tx_count = prev_count,
187 }
188 }
189 }
190
191 pub(super) fn semaphore(&self) -> &S {
192 &self.inner.semaphore
193 }
194
195 pub(crate) fn send(&self, value: T) {
197 self.inner.send(value);
198 }
199
200 pub(crate) fn wake_rx(&self) {
202 self.inner.rx_waker.wake();
203 }
204
205 pub(crate) fn same_channel(&self, other: &Self) -> bool {
207 Arc::ptr_eq(&self.inner, &other.inner)
208 }
209}
210
211impl<T, S: Semaphore> Tx<T, S> {
212 pub(crate) fn is_closed(&self) -> bool {
213 self.inner.semaphore.is_closed()
214 }
215
216 pub(crate) async fn closed(&self) {
217 let notified = self.inner.notify_rx_closed.notified();
221
222 if self.inner.semaphore.is_closed() {
223 return;
224 }
225 notified.await;
226 }
227}
228
229impl<T, S> Clone for Tx<T, S> {
230 fn clone(&self) -> Tx<T, S> {
231 self.inner.tx_count.fetch_add(1, Relaxed);
234
235 Tx {
236 inner: self.inner.clone(),
237 }
238 }
239}
240
241impl<T, S> Drop for Tx<T, S> {
242 fn drop(&mut self) {
243 if self.inner.tx_count.fetch_sub(1, AcqRel) != 1 {
244 return;
245 }
246
247 self.inner.tx.close();
249
250 self.wake_rx();
252 }
253}
254
255impl<T, S: Semaphore> Rx<T, S> {
258 fn new(chan: Arc<Chan<T, S>>) -> Rx<T, S> {
259 Rx { inner: chan }
260 }
261
262 pub(crate) fn close(&mut self) {
263 self.inner.rx_fields.with_mut(|rx_fields_ptr| {
264 let rx_fields = unsafe { &mut *rx_fields_ptr };
265
266 if rx_fields.rx_closed {
267 return;
268 }
269
270 rx_fields.rx_closed = true;
271 });
272
273 self.inner.semaphore.close();
274 self.inner.notify_rx_closed.notify_waiters();
275 }
276
277 pub(crate) fn is_closed(&self) -> bool {
278 self.inner.semaphore.is_closed() || self.inner.tx_count.load(Acquire) == 0
288 }
289
290 pub(crate) fn is_empty(&self) -> bool {
291 self.inner.rx_fields.with(|rx_fields_ptr| {
292 let rx_fields = unsafe { &*rx_fields_ptr };
293 rx_fields.list.is_empty(&self.inner.tx)
294 })
295 }
296
297 pub(crate) fn len(&self) -> usize {
298 self.inner.rx_fields.with(|rx_fields_ptr| {
299 let rx_fields = unsafe { &*rx_fields_ptr };
300 rx_fields.list.len(&self.inner.tx)
301 })
302 }
303
304 pub(crate) fn recv(&mut self, cx: &mut Context<'_>) -> Poll<Option<T>> {
306 use super::block::Read;
307
308 ready!(crate::trace::trace_leaf(cx));
309
310 let coop = ready!(crate::task::coop::poll_proceed(cx));
312
313 self.inner.rx_fields.with_mut(|rx_fields_ptr| {
314 let rx_fields = unsafe { &mut *rx_fields_ptr };
315
316 macro_rules! try_recv {
317 () => {
318 match rx_fields.list.pop(&self.inner.tx) {
319 Some(Read::Value(value)) => {
320 self.inner.semaphore.add_permit();
321 coop.made_progress();
322 return Ready(Some(value));
323 }
324 Some(Read::Closed) => {
325 debug_assert!(self.inner.semaphore.is_idle());
330 coop.made_progress();
331 return Ready(None);
332 }
333 None => {} }
335 };
336 }
337
338 try_recv!();
339
340 self.inner.rx_waker.register_by_ref(cx.waker());
341
342 try_recv!();
346
347 if rx_fields.rx_closed && self.inner.semaphore.is_idle() {
348 coop.made_progress();
349 Ready(None)
350 } else {
351 Pending
352 }
353 })
354 }
355
356 pub(crate) fn recv_many(
361 &mut self,
362 cx: &mut Context<'_>,
363 buffer: &mut Vec<T>,
364 limit: usize,
365 ) -> Poll<usize> {
366 use super::block::Read;
367
368 ready!(crate::trace::trace_leaf(cx));
369
370 let coop = ready!(crate::task::coop::poll_proceed(cx));
372
373 if limit == 0 {
374 coop.made_progress();
375 return Ready(0usize);
376 }
377
378 let mut remaining = limit;
379 let initial_length = buffer.len();
380
381 self.inner.rx_fields.with_mut(|rx_fields_ptr| {
382 let rx_fields = unsafe { &mut *rx_fields_ptr };
383 macro_rules! try_recv {
384 () => {
385 while remaining > 0 {
386 match rx_fields.list.pop(&self.inner.tx) {
387 Some(Read::Value(value)) => {
388 remaining -= 1;
389 buffer.push(value);
390 }
391
392 Some(Read::Closed) => {
393 let number_added = buffer.len() - initial_length;
394 if number_added > 0 {
395 self.inner.semaphore.add_permits(number_added);
396 }
397 debug_assert!(self.inner.semaphore.is_idle());
402 coop.made_progress();
403 return Ready(number_added);
404 }
405
406 None => {
407 break; }
409 }
410 }
411 let number_added = buffer.len() - initial_length;
412 if number_added > 0 {
413 self.inner.semaphore.add_permits(number_added);
414 coop.made_progress();
415 return Ready(number_added);
416 }
417 };
418 }
419
420 try_recv!();
421
422 self.inner.rx_waker.register_by_ref(cx.waker());
423
424 try_recv!();
428
429 if rx_fields.rx_closed && self.inner.semaphore.is_idle() {
430 debug_assert_eq!(buffer.len(), initial_length);
431 coop.made_progress();
432 Ready(0usize)
433 } else {
434 Pending
435 }
436 })
437 }
438
439 pub(crate) fn try_recv(&mut self) -> Result<T, TryRecvError> {
441 use super::list::TryPopResult;
442
443 self.inner.rx_fields.with_mut(|rx_fields_ptr| {
444 let rx_fields = unsafe { &mut *rx_fields_ptr };
445
446 macro_rules! try_recv {
447 () => {
448 match rx_fields.list.try_pop(&self.inner.tx) {
449 TryPopResult::Ok(value) => {
450 self.inner.semaphore.add_permit();
451 return Ok(value);
452 }
453 TryPopResult::Closed => return Err(TryRecvError::Disconnected),
454 TryPopResult::Empty
456 if rx_fields.rx_closed && self.inner.semaphore.is_idle() =>
457 {
458 return Err(TryRecvError::Disconnected)
459 }
460 TryPopResult::Empty => return Err(TryRecvError::Empty),
461 TryPopResult::Busy => {} }
463 };
464 }
465
466 try_recv!();
467
468 self.inner.rx_waker.wake();
476
477 let mut park = CachedParkThread::new();
479 let waker = park.waker().unwrap();
480 loop {
481 self.inner.rx_waker.register_by_ref(&waker);
482 try_recv!();
485 park.park();
486 }
487 })
488 }
489
490 pub(super) fn semaphore(&self) -> &S {
491 &self.inner.semaphore
492 }
493
494 pub(super) fn sender_strong_count(&self) -> usize {
495 self.inner.tx_count.load(Acquire)
496 }
497
498 pub(super) fn sender_weak_count(&self) -> usize {
499 self.inner.tx_weak_count.load(Relaxed)
500 }
501}
502
503impl<T, S: Semaphore> Drop for Rx<T, S> {
504 fn drop(&mut self) {
505 use super::block::Read::Value;
506
507 self.close();
508
509 self.inner.rx_fields.with_mut(|rx_fields_ptr| {
510 let rx_fields = unsafe { &mut *rx_fields_ptr };
511 struct Guard<'a, T, S: Semaphore> {
512 list: &'a mut list::Rx<T>,
513 tx: &'a list::Tx<T>,
514 sem: &'a S,
515 }
516
517 impl<'a, T, S: Semaphore> Guard<'a, T, S> {
518 fn drain(&mut self) {
519 while let Some(Value(_)) = self.list.pop(self.tx) {
521 self.sem.add_permit();
522 }
523 }
524 }
525
526 impl<'a, T, S: Semaphore> Drop for Guard<'a, T, S> {
527 fn drop(&mut self) {
528 self.drain();
529 }
530 }
531
532 let mut guard = Guard {
533 list: &mut rx_fields.list,
534 tx: &self.inner.tx,
535 sem: &self.inner.semaphore,
536 };
537
538 self.inner.rx_waker.take_waker();
542
543 guard.drain();
544 });
545 }
546}
547
548impl<T, S> Chan<T, S> {
551 fn send(&self, value: T) {
552 self.tx.push(value);
554
555 self.rx_waker.wake();
557 }
558
559 pub(super) fn decrement_weak_count(&self) {
560 self.tx_weak_count.fetch_sub(1, Relaxed);
561 }
562
563 pub(super) fn increment_weak_count(&self) {
564 self.tx_weak_count.fetch_add(1, Relaxed);
565 }
566
567 pub(super) fn strong_count(&self) -> usize {
568 self.tx_count.load(Acquire)
569 }
570
571 pub(super) fn weak_count(&self) -> usize {
572 self.tx_weak_count.load(Relaxed)
573 }
574}
575
576impl<T, S> Drop for Chan<T, S> {
577 fn drop(&mut self) {
578 use super::block::Read::Value;
579
580 self.rx_fields.with_mut(|rx_fields_ptr| {
583 let rx_fields = unsafe { &mut *rx_fields_ptr };
584
585 while let Some(Value(_)) = rx_fields.list.pop(&self.tx) {}
586 unsafe { rx_fields.list.free_blocks() };
587 });
588 }
589}
590
591impl Semaphore for bounded::Semaphore {
594 fn add_permit(&self) {
595 self.semaphore.release(1);
596 }
597
598 fn add_permits(&self, n: usize) {
599 self.semaphore.release(n)
600 }
601
602 fn is_idle(&self) -> bool {
603 self.semaphore.available_permits() == self.bound
604 }
605
606 fn close(&self) {
607 self.semaphore.close();
608 }
609
610 fn is_closed(&self) -> bool {
611 self.semaphore.is_closed()
612 }
613}
614
615impl Semaphore for unbounded::Semaphore {
618 fn add_permit(&self) {
619 let prev = self.0.fetch_sub(2, Release);
620
621 if prev >> 1 == 0 {
622 process::abort();
624 }
625 }
626
627 fn add_permits(&self, n: usize) {
628 let prev = self.0.fetch_sub(n << 1, Release);
629
630 if (prev >> 1) < n {
631 process::abort();
633 }
634 }
635
636 fn is_idle(&self) -> bool {
637 self.0.load(Acquire) >> 1 == 0
638 }
639
640 fn close(&self) {
641 self.0.fetch_or(1, Release);
642 }
643
644 fn is_closed(&self) -> bool {
645 self.0.load(Acquire) & 1 == 1
646 }
647}