ntex_util/channel/
mpsc.rs1use std::collections::VecDeque;
3use std::future::poll_fn;
4use std::{fmt, panic::UnwindSafe, pin::Pin, task::Context, task::Poll};
5
6use futures_core::{FusedStream, Stream};
7
8use super::cell::Cell;
9use crate::task::LocalWaker;
10
11pub fn channel<T>() -> (Sender<T>, Receiver<T>) {
13 let shared = Cell::new(Shared {
14 has_receiver: true,
15 buffer: VecDeque::new(),
16 blocked_recv: LocalWaker::new(),
17 closed: LocalWaker::new(),
18 });
19 let sender = Sender {
20 shared: shared.clone(),
21 };
22 let receiver = Receiver { shared };
23 (sender, receiver)
24}
25
26#[derive(Debug)]
27struct Shared<T> {
28 buffer: VecDeque<T>,
29 blocked_recv: LocalWaker,
30 closed: LocalWaker,
31 has_receiver: bool,
32}
33
34impl<T> Shared<T> {
35 fn close(&mut self) {
36 self.has_receiver = false;
37 self.blocked_recv.wake();
38 self.closed.wake();
39 }
40}
41
42#[derive(Debug)]
46pub struct Sender<T> {
47 shared: Cell<Shared<T>>,
48}
49
50impl<T> Unpin for Sender<T> {}
51
52impl<T> Sender<T> {
53 pub fn send(&self, item: T) -> Result<(), SendError<T>> {
55 let shared = self.shared.get_mut();
56 if !shared.has_receiver {
57 return Err(SendError(item)); }
59 shared.buffer.push_back(item);
60 shared.blocked_recv.wake();
61 Ok(())
62 }
63
64 pub fn close(&self) {
70 self.shared.get_mut().close();
71 }
72
73 pub fn is_closed(&self) -> bool {
75 self.shared.strong_count() == 1 || !self.shared.get_ref().has_receiver
76 }
77
78 pub fn poll_closed(&self, cx: &mut Context<'_>) -> Poll<()> {
83 let shared = self.shared.get_mut();
84 if shared.has_receiver {
85 shared.closed.register(cx.waker());
86 Poll::Pending
87 } else {
88 Poll::Ready(())
89 }
90 }
91
92 pub async fn closed(&self) {
96 poll_fn(|cx| self.poll_closed(cx)).await;
97 }
98}
99
100impl<T> Clone for Sender<T> {
101 fn clone(&self) -> Self {
102 Sender {
103 shared: self.shared.clone(),
104 }
105 }
106}
107
108impl<T> Drop for Sender<T> {
109 fn drop(&mut self) {
110 let count = self.shared.strong_count();
111 let shared = self.shared.get_mut();
112
113 if shared.has_receiver && count == 2 {
115 shared.blocked_recv.wake();
117 }
118 }
119}
120
121#[derive(Debug)]
125pub struct Receiver<T> {
126 shared: Cell<Shared<T>>,
127}
128
129impl<T> Receiver<T> {
130 pub fn sender(&self) -> Sender<T> {
132 Sender {
133 shared: self.shared.clone(),
134 }
135 }
136
137 pub fn close(&self) {
142 self.shared.get_mut().close();
143 }
144
145 pub fn is_closed(&self) -> bool {
147 self.shared.strong_count() == 1 || !self.shared.get_ref().has_receiver
148 }
149
150 pub async fn recv(&self) -> Option<T> {
155 poll_fn(|cx| self.poll_recv(cx)).await
156 }
157
158 pub fn poll_recv(&self, cx: &mut Context<'_>) -> Poll<Option<T>> {
163 let shared = self.shared.get_mut();
164
165 if let Some(msg) = shared.buffer.pop_front() {
166 Poll::Ready(Some(msg))
167 } else if shared.has_receiver {
168 shared.blocked_recv.register(cx.waker());
169 if self.shared.strong_count() == 1 {
170 Poll::Ready(None)
173 } else {
174 Poll::Pending
175 }
176 } else {
177 Poll::Ready(None)
178 }
179 }
180}
181
182impl<T> Unpin for Receiver<T> {}
183
184impl<T> Stream for Receiver<T> {
185 type Item = T;
186
187 fn poll_next(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
188 self.poll_recv(cx)
189 }
190}
191
192impl<T> FusedStream for Receiver<T> {
193 fn is_terminated(&self) -> bool {
196 self.is_closed() && self.shared.get_ref().buffer.is_empty()
197 }
198}
199
200impl<T> UnwindSafe for Receiver<T> {}
201
202impl<T> Drop for Receiver<T> {
203 fn drop(&mut self) {
204 let shared = self.shared.get_mut();
205 shared.has_receiver = false;
206 let buffer = std::mem::take(&mut shared.buffer);
207 shared.closed.wake();
208
209 drop(buffer);
211 }
212}
213
214pub struct SendError<T>(T);
217
218impl<T> std::error::Error for SendError<T> {}
219
220impl<T> fmt::Debug for SendError<T> {
221 fn fmt(&self, fmt: &mut fmt::Formatter<'_>) -> fmt::Result {
222 fmt.debug_tuple("SendError").field(&"...").finish()
223 }
224}
225
226impl<T> fmt::Display for SendError<T> {
227 fn fmt(&self, fmt: &mut fmt::Formatter<'_>) -> fmt::Result {
228 write!(fmt, "send failed because receiver is gone")
229 }
230}
231
232impl<T> SendError<T> {
233 pub fn into_inner(self) -> T {
235 self.0
236 }
237}
238
239#[cfg(test)]
240mod tests {
241 use super::*;
242 use crate::{future::lazy, future::stream_recv};
243
244 #[ntex::test]
245 async fn test_mpsc() {
246 let (tx, mut rx) = channel();
247 assert!(format!("{tx:?}").contains("Sender"));
248 assert!(format!("{rx:?}").contains("Receiver"));
249
250 tx.send("test").unwrap();
251 assert_eq!(stream_recv(&mut rx).await.unwrap(), "test");
252
253 let tx2 = tx.clone();
254 tx2.send("test2").unwrap();
255 assert_eq!(stream_recv(&mut rx).await.unwrap(), "test2");
256
257 assert_eq!(
258 lazy(|cx| Pin::new(&mut rx).poll_next(cx)).await,
259 Poll::Pending
260 );
261 drop(tx2);
262 assert_eq!(
263 lazy(|cx| Pin::new(&mut rx).poll_next(cx)).await,
264 Poll::Pending
265 );
266 drop(tx);
267
268 let (tx, mut rx) = channel::<String>();
269 tx.close();
270 assert_eq!(stream_recv(&mut rx).await, None);
271
272 let (tx, rx) = channel();
273 tx.send("test").unwrap();
274 drop(rx);
275 assert!(tx.send("test").is_err());
276
277 let (tx, _) = channel();
278 let tx2 = tx.clone();
279 tx.close();
280 assert!(tx.send("test").is_err());
281 assert!(tx2.send("test").is_err());
282
283 let err = SendError("test");
284 assert!(format!("{err:?}").contains("SendError"));
285 assert!(format!("{err}").contains("send failed because receiver is gone"));
286 assert_eq!(err.into_inner(), "test");
287 }
288
289 #[ntex::test]
290 async fn test_close() {
291 let (tx, rx) = channel::<()>();
292 assert!(!tx.is_closed());
293 assert!(!rx.is_closed());
294 assert!(!rx.is_terminated());
295
296 tx.close();
297 assert!(tx.is_closed());
298 assert!(rx.is_closed());
299 assert!(rx.is_terminated());
300
301 let (tx, rx) = channel::<()>();
302 assert_eq!(lazy(|cx| rx.poll_recv(cx)).await, Poll::Pending);
303 assert!(rx.shared.get_ref().blocked_recv.is_set());
304 rx.close();
305 assert!(tx.is_closed());
306 assert!(!rx.shared.get_ref().blocked_recv.is_set());
307 assert_eq!(lazy(|cx| rx.poll_recv(cx)).await, Poll::Ready(None));
308
309 let (tx, rx) = channel::<()>();
310 drop(tx);
311 assert!(rx.is_closed());
312 assert!(rx.is_terminated());
313 let _tx = rx.sender();
314 assert!(!rx.is_closed());
315 assert!(!rx.is_terminated());
316 }
317
318 #[ntex::test]
319 async fn test_poll_closed() {
320 let (tx, rx) = channel::<()>();
321 assert_eq!(lazy(|cx| tx.poll_closed(cx)).await, Poll::Pending);
322 assert!(tx.shared.get_ref().closed.is_set());
323 drop(rx);
324 assert!(!tx.shared.get_ref().closed.is_set());
325 assert_eq!(lazy(|cx| tx.poll_closed(cx)).await, Poll::Ready(()));
326 tx.closed().await;
327
328 let (tx, rx) = channel::<()>();
329 assert_eq!(lazy(|cx| tx.poll_closed(cx)).await, Poll::Pending);
330 rx.close();
331 assert!(!tx.shared.get_ref().closed.is_set());
332 tx.closed().await;
333
334 let (tx, _rx) = channel::<()>();
335 let tx2 = tx.clone();
336 assert_eq!(lazy(|cx| tx2.poll_closed(cx)).await, Poll::Pending);
337 tx.close();
338 assert!(!tx.shared.get_ref().closed.is_set());
339 tx2.closed().await;
340 }
341
342 #[test]
343 fn test_drop_receiver_reentrant_send() {
344 use std::rc::Rc;
345
346 struct Msg(Option<Sender<Msg>>, Rc<std::cell::Cell<usize>>);
347
348 impl Drop for Msg {
349 fn drop(&mut self) {
350 self.1.set(self.1.get() + 1);
351 if let Some(tx) = self.0.take() {
352 assert!(tx.send(Msg(None, self.1.clone())).is_err());
353 }
354 }
355 }
356
357 let drops = Rc::new(std::cell::Cell::new(0));
358 let (tx, rx) = channel();
359 for _ in 0..4 {
360 tx.send(Msg(Some(tx.clone()), drops.clone())).unwrap();
361 }
362 drop(rx);
363 assert_eq!(drops.get(), 8);
364 assert!(tx.is_closed());
365 }
366
367 #[ntex::test]
368 async fn test_fused_drains_buffer() {
369 let (tx, mut rx) = channel();
370 tx.send(1).unwrap();
371 tx.send(2).unwrap();
372 drop(tx);
373
374 assert!(!rx.is_terminated());
375 assert_eq!(stream_recv(&mut rx).await, Some(1));
376 assert_eq!(stream_recv(&mut rx).await, Some(2));
377 assert!(rx.is_terminated());
378 assert_eq!(stream_recv(&mut rx).await, None);
379
380 let (tx, mut rx) = channel();
381 tx.send(1).unwrap();
382 rx.close();
383 assert!(!rx.is_terminated());
384 assert_eq!(stream_recv(&mut rx).await, Some(1));
385 assert!(rx.is_terminated());
386 }
387
388 #[ntex::test]
389 async fn test_mpsc_recv() {
390 let (tx, rx) = channel();
391 tx.send(1).unwrap();
392 assert_eq!(rx.recv().await, Some(1));
393 drop(tx);
394 assert_eq!(rx.recv().await, None);
395 }
396}