1#![expect(
2 clippy::unwrap_used,
3 reason = "ratchet MOIRAI-UNWRAP-1: pre-existing debt"
4)]
5
6use std::collections::VecDeque;
15use std::future::Future;
16use std::marker::Unpin;
17use std::pin::Pin;
18use std::sync::{Arc, Mutex};
19use std::task::{Context, Poll};
20
21use super::wait_queue::{WaitQueue, WaiterPoll};
22
23struct SharedState<T> {
24 buffer: VecDeque<T>,
25 capacity: usize,
26 sender_count: usize,
27 closed: bool,
28 send_waiters: WaitQueue<()>,
29 recv_waiters: WaitQueue<()>,
30}
31
32pub struct Sender<T> {
34 shared: Arc<Mutex<SharedState<T>>>,
35}
36
37impl<T> Clone for Sender<T> {
38 fn clone(&self) -> Self {
39 let mut shared = self.shared.lock().unwrap();
40 shared.sender_count += 1;
41 Sender {
42 shared: self.shared.clone(),
43 }
44 }
45}
46
47impl<T> Sender<T> {
48 pub fn send(&self, value: T) -> SendFuture<'_, T> {
53 SendFuture {
54 sender: self,
55 value: Some(value),
56 id: None,
57 }
58 }
59
60 pub fn try_send(&self, value: T) -> Result<(), T> {
67 let mut shared = self.shared.lock().unwrap();
68 if shared.closed {
69 return Err(value);
70 }
71 if shared.buffer.len() >= shared.capacity {
72 return Err(value);
73 }
74 shared.buffer.push_back(value);
75 let waker = shared.recv_waiters.grant_oldest(());
79 drop(shared);
80 if let Some(waker) = waker {
81 waker.wake();
82 }
83 Ok(())
84 }
85
86 pub fn is_closed(&self) -> bool {
88 self.shared.lock().unwrap().closed
89 }
90
91 pub fn sender_strong_count(&self) -> usize {
93 self.shared.lock().unwrap().sender_count
94 }
95}
96
97impl<T> Drop for Sender<T> {
98 fn drop(&mut self) {
99 let mut shared = self.shared.lock().unwrap();
100 shared.sender_count -= 1;
101 if shared.sender_count != 0 {
102 return;
103 }
104 shared.closed = true;
105 let wakers = shared.recv_waiters.grant_all(());
106 drop(shared);
107 for waker in wakers {
108 waker.wake();
109 }
110 }
111}
112
113pub struct SendFuture<'a, T> {
115 sender: &'a Sender<T>,
116 value: Option<T>,
117 id: Option<u64>,
118}
119
120impl<'a, T: Unpin> Future for SendFuture<'a, T> {
121 type Output = Result<(), T>;
122
123 fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
124 let this = self.get_mut();
125 let mut shared = this.sender.shared.lock().unwrap();
126
127 if shared.closed {
128 if let Some(id) = this.id.take() {
129 shared.send_waiters.deregister(id);
130 }
131 let value = this.value.take().unwrap();
132 return Poll::Ready(Err(value));
133 }
134
135 if shared.buffer.len() < shared.capacity {
136 if let Some(id) = this.id.take() {
137 shared.send_waiters.deregister(id);
138 }
139 shared.buffer.push_back(this.value.take().unwrap());
140 let waker = shared.recv_waiters.grant_oldest(());
141 drop(shared);
142 if let Some(waker) = waker {
143 waker.wake();
144 }
145 return Poll::Ready(Ok(()));
146 }
147
148 this.id = Some(match this.id {
152 Some(id) => match shared.send_waiters.poll_waiter(id, cx.waker()) {
153 WaiterPoll::Pending => id,
154 WaiterPoll::Granted(()) | WaiterPoll::NotRegistered => {
155 shared.send_waiters.register(cx.waker().clone())
156 }
157 },
158 None => shared.send_waiters.register(cx.waker().clone()),
159 });
160 Poll::Pending
161 }
162}
163
164impl<'a, T> Drop for SendFuture<'a, T> {
165 fn drop(&mut self) {
166 if let Some(id) = self.id
167 && let Ok(mut shared) = self.sender.shared.lock()
168 {
169 shared.send_waiters.deregister(id);
170 }
171 }
172}
173
174pub struct Receiver<T> {
176 shared: Arc<Mutex<SharedState<T>>>,
177}
178
179impl<T> Receiver<T> {
180 pub fn recv(&mut self) -> RecvFuture<'_, T> {
185 RecvFuture {
186 receiver: self,
187 id: None,
188 }
189 }
190
191 pub fn try_recv(&mut self) -> Option<T> {
193 let mut shared = self.shared.lock().unwrap();
194 let value = shared.buffer.pop_front();
195 let waker = if value.is_some() {
196 shared.send_waiters.grant_oldest(())
197 } else {
198 None
199 };
200 drop(shared);
201 if let Some(waker) = waker {
202 waker.wake();
203 }
204 value
205 }
206
207 pub fn close(&mut self) {
209 let mut shared = self.shared.lock().unwrap();
210 shared.closed = true;
211 let mut wakers = shared.send_waiters.grant_all(());
212 wakers.extend(shared.recv_waiters.grant_all(()));
213 drop(shared);
214 for waker in wakers {
215 waker.wake();
216 }
217 }
218}
219
220impl<T> Drop for Receiver<T> {
221 fn drop(&mut self) {
222 let mut shared = self.shared.lock().unwrap();
223 shared.closed = true;
224 let wakers = shared.send_waiters.grant_all(());
225 drop(shared);
226 for waker in wakers {
227 waker.wake();
228 }
229 }
230}
231
232pub struct RecvFuture<'a, T> {
234 receiver: &'a mut Receiver<T>,
235 id: Option<u64>,
236}
237
238impl<'a, T> Future for RecvFuture<'a, T> {
239 type Output = Result<T, ()>;
240
241 fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
242 let this = self.get_mut();
243 let mut shared = this.receiver.shared.lock().unwrap();
244 if let Some(value) = shared.buffer.pop_front() {
245 if let Some(id) = this.id.take() {
246 shared.recv_waiters.deregister(id);
247 }
248 let waker = shared.send_waiters.grant_oldest(());
249 drop(shared);
250 if let Some(waker) = waker {
251 waker.wake();
252 }
253 return Poll::Ready(Ok(value));
254 }
255 if shared.closed {
256 if let Some(id) = this.id.take() {
257 shared.recv_waiters.deregister(id);
258 }
259 return Poll::Ready(Err(()));
260 }
261
262 this.id = Some(match this.id {
263 Some(id) => match shared.recv_waiters.poll_waiter(id, cx.waker()) {
264 WaiterPoll::Pending => id,
265 WaiterPoll::Granted(()) | WaiterPoll::NotRegistered => {
266 shared.recv_waiters.register(cx.waker().clone())
267 }
268 },
269 None => shared.recv_waiters.register(cx.waker().clone()),
270 });
271 Poll::Pending
272 }
273}
274
275impl<'a, T> Drop for RecvFuture<'a, T> {
276 fn drop(&mut self) {
277 if let Some(id) = self.id
278 && let Ok(mut shared) = self.receiver.shared.lock()
279 {
280 shared.recv_waiters.deregister(id);
281 }
282 }
283}
284
285#[must_use]
287pub fn channel<T>(capacity: usize) -> (Sender<T>, Receiver<T>) {
288 let shared = Arc::new(Mutex::new(SharedState {
289 buffer: VecDeque::with_capacity(capacity),
290 capacity,
291 sender_count: 1,
292 closed: false,
293 send_waiters: WaitQueue::new(),
294 recv_waiters: WaitQueue::new(),
295 }));
296 (
297 Sender {
298 shared: shared.clone(),
299 },
300 Receiver { shared },
301 )
302}
303
304#[cfg(test)]
305mod tests {
306 use super::*;
307 use std::future::Future;
308 use std::pin::Pin;
309 use std::sync::{
310 Arc,
311 atomic::{AtomicUsize, Ordering},
312 };
313 use std::task::{Context, Poll, Wake, Waker};
314
315 fn poll_future<F: Future + Unpin>(future: &mut F) -> Poll<F::Output> {
316 let mut context = Context::from_waker(Waker::noop());
317 Pin::new(future).poll(&mut context)
318 }
319
320 fn poll_future_with_waker<F: Future + Unpin>(future: &mut F, waker: &Waker) -> Poll<F::Output> {
321 let mut context = Context::from_waker(waker);
322 Pin::new(future).poll(&mut context)
323 }
324
325 struct CountingWake(Arc<AtomicUsize>);
326
327 impl Wake for CountingWake {
328 fn wake(self: Arc<Self>) {
329 self.0.fetch_add(1, Ordering::Release);
330 }
331
332 fn wake_by_ref(self: &Arc<Self>) {
333 self.0.fetch_add(1, Ordering::Release);
334 }
335 }
336
337 #[test]
338 fn test_mpsc_send_recv() {
339 let (tx, mut rx) = channel(10);
340 tx.try_send(1).unwrap();
341 tx.try_send(2).unwrap();
342 tx.try_send(3).unwrap();
343 assert_eq!(rx.try_recv(), Some(1));
344 assert_eq!(rx.try_recv(), Some(2));
345 assert_eq!(rx.try_recv(), Some(3));
346 assert!(rx.try_recv().is_none());
347 }
348
349 #[test]
350 fn test_mpsc_closed_sender() {
351 let (tx, mut rx) = channel::<i32>(10);
352 tx.try_send(1).unwrap();
353 drop(tx);
354 assert_eq!(rx.try_recv(), Some(1));
355 assert!(rx.try_recv().is_none());
356 }
357
358 #[test]
359 fn test_mpsc_closed_receiver() {
360 let (tx, rx) = channel::<i32>(10);
361 drop(rx);
362 assert!(tx.try_send(1).is_err());
363 }
364
365 #[test]
366 fn test_mpsc_capacity() {
367 let (tx, mut rx) = channel(2);
368 assert!(tx.try_send(1).is_ok());
369 assert!(tx.try_send(2).is_ok());
370 assert!(tx.try_send(3).is_err());
371 let _ = rx.try_recv();
372 assert!(tx.try_send(3).is_ok());
373 }
374
375 #[test]
376 fn test_mpsc_sender_clone() {
377 let (tx1, mut rx) = channel(10);
378 let tx2 = tx1.clone();
379 tx1.try_send(1).unwrap();
380 tx2.try_send(2).unwrap();
381 drop(tx1);
382 drop(tx2);
383 assert_eq!(rx.try_recv(), Some(1));
384 assert_eq!(rx.try_recv(), Some(2));
385 assert!(rx.try_recv().is_none());
386 }
387
388 #[test]
389 fn test_mpsc_sender_strong_count() {
390 let (tx1, _) = channel::<i32>(10);
391 assert_eq!(tx1.sender_strong_count(), 1);
392 let tx2 = tx1.clone();
393 assert_eq!(tx1.sender_strong_count(), 2);
394 drop(tx2);
395 assert_eq!(tx1.sender_strong_count(), 1);
396 }
397
398 #[test]
399 fn test_mpsc_send_pending_then_recv() {
400 let (tx, mut rx) = channel(1);
401 tx.try_send(1).unwrap();
402 let mut send = tx.send(2);
403 assert!(matches!(poll_future(&mut send), Poll::Pending));
404 let _ = rx.try_recv();
405 assert!(matches!(poll_future(&mut send), Poll::Ready(Ok(()))));
406 }
407
408 #[test]
409 fn test_mpsc_async_recv_pending_then_send() {
410 let (tx, mut rx) = channel(1);
411 let mut recv = rx.recv();
412 assert!(matches!(poll_future(&mut recv), Poll::Pending));
413 tx.try_send(42).unwrap();
414 assert!(matches!(poll_future(&mut recv), Poll::Ready(Ok(42))));
415 }
416
417 #[test]
418 fn test_mpsc_send_future_dropped_cancels_waiter() {
419 let (tx, _rx) = channel(1);
420 tx.try_send(1).unwrap();
421 let mut send = tx.send(2);
423 assert!(matches!(poll_future(&mut send), Poll::Pending));
424 drop(send);
425 assert!(tx.shared.lock().unwrap().send_waiters.is_empty());
427 }
428
429 #[test]
430 fn test_mpsc_recv_future_dropped_cancels_waiter() {
431 let (tx, mut rx) = channel(1);
432 tx.try_send(1).unwrap();
433 let _ = rx.try_recv();
435 let mut recv = rx.recv();
436 assert!(matches!(poll_future(&mut recv), Poll::Pending));
437 drop(recv);
438 assert!(tx.shared.lock().unwrap().recv_waiters.is_empty());
440 }
441
442 #[test]
443 fn oldest_pending_sender_is_woken_first() {
444 let (tx, mut rx) = channel(1);
445 tx.try_send(1).expect("initial send must fill the channel");
446
447 let first_wakes = Arc::new(AtomicUsize::new(0));
448 let second_wakes = Arc::new(AtomicUsize::new(0));
449 let first_waker = Waker::from(Arc::new(CountingWake(Arc::clone(&first_wakes))));
450 let second_waker = Waker::from(Arc::new(CountingWake(Arc::clone(&second_wakes))));
451 let mut first = tx.send(2);
452 let mut second = tx.send(3);
453
454 assert!(poll_future_with_waker(&mut first, &first_waker).is_pending());
455 assert!(poll_future_with_waker(&mut second, &second_waker).is_pending());
456
457 assert_eq!(rx.try_recv(), Some(1));
458 assert_eq!(first_wakes.load(Ordering::Acquire), 1);
459 assert_eq!(second_wakes.load(Ordering::Acquire), 0);
460
461 assert!(matches!(
462 poll_future_with_waker(&mut first, &first_waker),
463 Poll::Ready(Ok(()))
464 ));
465 assert_eq!(rx.try_recv(), Some(2));
466 assert_eq!(second_wakes.load(Ordering::Acquire), 1);
467 assert!(matches!(
468 poll_future_with_waker(&mut second, &second_waker),
469 Poll::Ready(Ok(()))
470 ));
471 assert_eq!(rx.try_recv(), Some(3));
472 }
473}