1use std::cell::{Cell, RefCell};
3use std::task::{Context, Poll};
4use std::{collections::VecDeque, fmt, future::poll_fn, pin::Pin, rc::Rc, rc::Weak};
5
6use ntex_bytes::Bytes;
7
8use crate::{Stream, task::LocalWaker};
9
10const HIGH_WATERMARK: u32 = 32_768;
12const LOW_WATERMARK: u32 = HIGH_WATERMARK / 2;
14
15#[derive(Copy, Clone, Debug, PartialEq, Eq)]
17pub enum Status {
18 Eof,
20 Ready,
22 Dropped,
24}
25
26pub fn channel<E>() -> (Sender<E>, Receiver<E>) {
28 let inner = Rc::new(Inner::new(false));
29
30 (
31 Sender {
32 inner: Rc::downgrade(&inner),
33 },
34 Receiver { inner },
35 )
36}
37
38pub fn eof<E>() -> (Sender<E>, Receiver<E>) {
43 let inner = Rc::new(Inner::new(true));
44
45 (
46 Sender {
47 inner: Rc::downgrade(&inner),
48 },
49 Receiver { inner },
50 )
51}
52
53pub fn empty<E>(data: Option<Bytes>) -> Receiver<E> {
55 let rx = Receiver {
56 inner: Rc::new(Inner::new(true)),
57 };
58 if let Some(data) = data {
59 rx.put(data);
60 }
61 rx
62}
63
64#[derive(Debug)]
70pub struct Receiver<E> {
71 inner: Rc<Inner<E>>,
72}
73
74impl<E> Receiver<E> {
75 #[inline]
87 pub fn set_watermarks(&self, high: u32, low: u32) {
88 self.inner.set_watermarks(high, low);
89 }
90
91 #[inline]
95 #[deprecated(since = "4.2.0", note = "Use `Receiver::set_watermarks()` instead")]
96 pub fn max_buffer_size(&self, size: usize) {
97 let size = u32::try_from(size).unwrap_or(u32::MAX);
98 self.inner.set_watermarks(size, size / 2);
99 }
100
101 #[inline]
105 pub fn put(&self, data: Bytes) {
106 self.inner.unread_data(data);
107 }
108
109 #[inline]
110 pub fn is_eof(&self) -> bool {
114 self.inner.flags.get().contains(Flags::EOF)
115 }
116
117 #[inline]
118 pub async fn read(&self) -> Option<Result<Bytes, E>> {
120 poll_fn(|cx| self.poll_read(cx)).await
121 }
122
123 #[inline]
124 pub fn poll_read(&self, cx: &mut Context<'_>) -> Poll<Option<Result<Bytes, E>>> {
126 if let Some(data) = self.inner.get_data() {
127 Poll::Ready(Some(Ok(data)))
128 } else if let Some(err) = self.inner.err.take() {
129 self.inner.insert_flag(Flags::EOF);
130 Poll::Ready(Some(Err(err)))
131 } else if self.inner.flags.get().intersects(Flags::EOF | Flags::ERROR) {
132 Poll::Ready(None)
133 } else {
134 self.inner.recv_task.register(cx.waker());
135 Poll::Pending
136 }
137 }
138}
139
140impl<E> Stream for Receiver<E> {
141 type Item = Result<Bytes, E>;
142
143 fn poll_next(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
144 self.poll_read(cx)
145 }
146}
147
148impl<E> Drop for Receiver<E> {
149 fn drop(&mut self) {
150 self.inner.send_task.wake();
151 }
152}
153
154#[derive(Debug)]
159pub struct Sender<E> {
160 inner: Weak<Inner<E>>,
161}
162
163impl<E> Clone for Sender<E> {
164 fn clone(&self) -> Self {
165 Self {
166 inner: self.inner.clone(),
167 }
168 }
169}
170
171impl<E> Drop for Sender<E> {
172 fn drop(&mut self) {
173 if self.inner.weak_count() == 1
174 && let Some(shared) = self.inner.upgrade()
175 {
176 shared.insert_flag(Flags::EOF | Flags::SENDER_GONE);
177 shared.recv_task.wake();
179 }
180 }
181}
182
183impl<E> Sender<E> {
184 pub fn is_closed(&self) -> bool {
186 self.inner.strong_count() == 0
187 }
188
189 pub fn set_error(&self, err: E) {
191 if let Some(shared) = self.inner.upgrade() {
192 shared.set_error(err);
193 }
194 }
195
196 pub fn feed_eof(&self) {
198 if let Some(shared) = self.inner.upgrade() {
199 shared.feed_eof();
200 }
201 }
202
203 pub fn feed_data(&self, data: Bytes) {
207 if let Some(shared) = self.inner.upgrade() {
208 shared.feed_data(data);
209 }
210 }
211
212 pub async fn ready(&self) -> Status {
214 poll_fn(|cx| self.poll_ready(cx)).await
215 }
216
217 pub fn poll_ready(&self, cx: &mut Context<'_>) -> Poll<Status> {
222 if let Some(shared) = self.inner.upgrade() {
223 let flags = shared.flags.get();
224 if flags.intersects(Flags::SENDER_GONE | Flags::ERROR) {
225 Poll::Ready(Status::Dropped)
226 } else if flags.contains(Flags::EOF) {
227 Poll::Ready(Status::Eof)
228 } else if flags.contains(Flags::NEED_READ) {
229 Poll::Ready(Status::Ready)
230 } else {
231 shared.send_task.register(cx.waker());
232 Poll::Pending
233 }
234 } else {
235 Poll::Ready(Status::Dropped)
237 }
238 }
239}
240
241bitflags::bitflags! {
242 #[derive(Copy, Clone, Debug, Eq, PartialEq, Ord, PartialOrd, Hash)]
243 struct Flags: u8 {
244 const EOF = 0b0000_0001;
245 const ERROR = 0b0000_0010;
246 const NEED_READ = 0b0000_0100;
247 const SENDER_GONE = 0b0000_1000;
248 }
249}
250
251struct Inner<E> {
252 len: Cell<usize>,
253 flags: Cell<Flags>,
254 err: Cell<Option<E>>,
255 items: RefCell<VecDeque<Bytes>>,
256 recv_task: LocalWaker,
257 send_task: LocalWaker,
258 high_watermark: Cell<u32>,
259 low_watermark: Cell<u32>,
260}
261
262impl<E> Inner<E> {
263 fn new(eof: bool) -> Self {
264 let flags = if eof { Flags::EOF } else { Flags::NEED_READ };
265 Inner {
266 flags: Cell::new(flags),
267 len: Cell::new(0),
268 err: Cell::new(None),
269 items: RefCell::new(VecDeque::new()),
270 recv_task: LocalWaker::new(),
271 send_task: LocalWaker::new(),
272 high_watermark: Cell::new(HIGH_WATERMARK),
273 low_watermark: Cell::new(LOW_WATERMARK),
274 }
275 }
276
277 fn insert_flag(&self, f: Flags) {
278 let mut flags = self.flags.get();
279 flags.insert(f);
280 self.flags.set(flags);
281 }
282
283 fn remove_flag(&self, f: Flags) {
284 let mut flags = self.flags.get();
285 flags.remove(f);
286 self.flags.set(flags);
287 }
288
289 fn set_watermarks(&self, high: u32, low: u32) {
290 self.high_watermark.set(high);
291 self.low_watermark.set(low.min(high.saturating_sub(1)));
292
293 let flags = self.flags.get();
294 if flags.intersects(Flags::EOF | Flags::ERROR | Flags::SENDER_GONE) {
295 return;
296 }
297
298 if self.len.get() < high as usize {
299 if !flags.contains(Flags::NEED_READ) {
300 self.insert_flag(Flags::NEED_READ);
301 self.send_task.wake();
302 }
303 } else {
304 self.remove_flag(Flags::NEED_READ);
305 }
306 }
307
308 fn set_error(&self, err: E) {
309 self.err.set(Some(err));
310 self.insert_flag(Flags::ERROR);
311 self.recv_task.wake();
312 self.send_task.wake();
313 }
314
315 fn feed_eof(&self) {
316 self.insert_flag(Flags::EOF);
317 self.recv_task.wake();
318 self.send_task.wake();
319 }
320
321 fn feed_data(&self, data: Bytes) {
322 let len = self.len.get() + data.len();
323 self.len.set(len);
324 self.items.borrow_mut().push_back(data);
325 self.recv_task.wake();
326
327 if len >= self.high_watermark.get() as usize {
328 self.remove_flag(Flags::NEED_READ);
329 }
330 }
331
332 fn get_data(&self) -> Option<Bytes> {
333 self.items.borrow_mut().pop_front().inspect(|data| {
334 let len = self.len.get() - data.len();
335
336 self.len.set(len);
338 if len <= self.low_watermark.get() as usize {
339 self.insert_flag(Flags::NEED_READ);
340 self.send_task.wake();
341 }
342 })
343 }
344
345 fn unread_data(&self, data: Bytes) {
346 if !data.is_empty() {
347 self.len.set(self.len.get() + data.len());
348 self.items.borrow_mut().push_front(data);
349 }
350 }
351}
352
353impl<E> fmt::Debug for Inner<E> {
354 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
355 f.debug_struct("Inner")
356 .field("len", &self.len)
357 .field("flags", &self.flags)
358 .field("items", &self.items.borrow())
359 .field("high_watermark", &self.high_watermark)
360 .field("low_watermark", &self.low_watermark)
361 .field("recv_task", &self.recv_task)
362 .field("send_task", &self.send_task)
363 .finish()
364 }
365}
366
367#[cfg(test)]
368mod tests {
369 use super::*;
370 use crate::future::lazy;
371
372 #[ntex::test]
373 async fn test_sender_drop_wakes_receiver() {
374 let (tx, rx) = channel::<()>();
375 let handle = crate::spawn(async move { rx.read().await.is_none() });
376 crate::time::sleep(crate::time::Millis(10)).await;
377
378 drop(tx);
379 let res = crate::time::timeout(crate::time::Millis(1000), handle).await;
380 assert!(matches!(res, Ok(Ok(true))));
381 }
382
383 #[ntex::test]
384 async fn test_eof() {
385 let (tx, rx) = eof::<()>();
386 rx.set_watermarks(100, 50);
387 assert!(rx.read().await.is_none());
388 assert_eq!(tx.ready().await, Status::Eof);
389 }
390
391 #[ntex::test]
392 async fn test_closed() {
393 let (tx, rx) = channel::<()>();
395 assert!(!tx.is_closed());
396 drop(rx);
397 assert!(tx.is_closed());
398
399 let (tx, rx) = channel::<()>();
401 drop(tx);
402 assert_eq!(rx.read().await, None);
403 }
404
405 #[ntex::test]
406 async fn test_unread_data() {
407 let (_, payload) = channel::<()>();
408
409 payload.put(Bytes::from("data"));
410 assert_eq!(payload.inner.len.get(), 4);
411 assert_eq!(
412 Bytes::from("data"),
413 poll_fn(|cx| payload.poll_read(cx)).await.unwrap().unwrap()
414 );
415 }
416
417 #[ntex::test]
418 async fn buffer_size_updates_sender_readiness() {
419 let (tx, rx) = channel::<()>();
420 rx.set_watermarks(4, 2);
421 tx.feed_data(Bytes::from_static(b"data"));
422
423 assert!(lazy(|cx| tx.poll_ready(cx)).await.is_pending());
424 assert!(rx.inner.send_task.is_set());
425
426 rx.set_watermarks(5, 2);
427 assert!(!rx.inner.send_task.is_set());
428 assert_eq!(
429 lazy(|cx| tx.poll_ready(cx)).await,
430 Poll::Ready(Status::Ready)
431 );
432
433 rx.set_watermarks(4, 2);
434 assert!(lazy(|cx| tx.poll_ready(cx)).await.is_pending());
435 }
436
437 #[ntex::test]
438 async fn sender_resumes_at_low_watermark() {
439 let (tx, rx) = channel::<()>();
440 rx.set_watermarks(8, 4);
441 for _ in 0..4 {
442 tx.feed_data(Bytes::from_static(b"da"));
443 }
444 assert!(lazy(|cx| tx.poll_ready(cx)).await.is_pending());
445
446 assert!(rx.read().await.is_some());
448 assert!(rx.inner.send_task.is_set());
449 assert!(lazy(|cx| tx.poll_ready(cx)).await.is_pending());
450
451 assert!(rx.read().await.is_some());
453 assert!(!rx.inner.send_task.is_set());
454 assert_eq!(
455 lazy(|cx| tx.poll_ready(cx)).await,
456 Poll::Ready(Status::Ready)
457 );
458
459 let (tx, rx) = channel::<()>();
461 rx.set_watermarks(4, 10);
462 for _ in 0..3 {
463 tx.feed_data(Bytes::from_static(b"da"));
464 }
465 assert!(rx.read().await.is_some());
466 assert!(lazy(|cx| tx.poll_ready(cx)).await.is_pending());
467 assert!(rx.read().await.is_some());
468 assert_eq!(
469 lazy(|cx| tx.poll_ready(cx)).await,
470 Poll::Ready(Status::Ready)
471 );
472 }
473
474 #[ntex::test]
475 async fn custom_low_watermark() {
476 let (tx, rx) = channel::<()>();
477 rx.set_watermarks(8, 2);
478 for _ in 0..4 {
479 tx.feed_data(Bytes::from_static(b"da"));
480 }
481 assert!(rx.read().await.is_some());
482 assert!(rx.read().await.is_some());
483 assert!(lazy(|cx| tx.poll_ready(cx)).await.is_pending());
484 assert!(rx.read().await.is_some());
485 assert_eq!(
486 lazy(|cx| tx.poll_ready(cx)).await,
487 Poll::Ready(Status::Ready)
488 );
489 }
490
491 #[ntex::test]
492 #[allow(deprecated)]
493 async fn deprecated_max_buffer_size() {
494 let (_tx, rx) = channel::<()>();
495 rx.max_buffer_size(10);
496 assert_eq!(rx.inner.high_watermark.get(), 10);
497 assert_eq!(rx.inner.low_watermark.get(), 5);
498 }
499
500 #[ntex::test]
501 async fn test_sender_clone() {
502 let (sender, payload) = channel::<()>();
503 assert!(!payload.is_eof());
504 let sender2 = sender.clone();
505 assert!(!payload.is_eof());
506 drop(sender2);
507 assert!(!payload.is_eof());
508 drop(sender);
509 assert!(payload.is_eof());
510 }
511
512 #[ntex::test]
513 async fn test_ready_terminal_states() {
514 let (tx, _rx) = channel::<()>();
515 assert_eq!(tx.ready().await, Status::Ready);
516 tx.set_error(());
517 assert_eq!(tx.ready().await, Status::Dropped);
518
519 let (tx, _rx) = channel::<()>();
520 tx.feed_eof();
521 assert_eq!(tx.ready().await, Status::Eof);
522
523 let (tx, rx) = channel::<()>();
524 drop(rx);
525 assert_eq!(tx.ready().await, Status::Dropped);
526 }
527
528 #[ntex::test]
529 async fn test_empty() {
530 let rx = empty::<()>(None);
531 assert!(rx.is_eof());
532 assert!(rx.read().await.is_none());
533
534 let mut rx = empty::<()>(Some(Bytes::from_static(b"data")));
535 assert!(format!("{rx:?}").contains("high_watermark"));
536 assert_eq!(
537 crate::future::stream_recv(&mut rx).await.unwrap().unwrap(),
538 Bytes::from_static(b"data")
539 );
540 assert!(crate::future::stream_recv(&mut rx).await.is_none());
541 }
542
543 #[ntex::test]
544 async fn test_error() {
545 let (tx, rx) = channel::<&'static str>();
546 tx.feed_data(Bytes::from_static(b"data"));
547 tx.set_error("err");
548
549 assert_eq!(
551 rx.read().await.unwrap().unwrap(),
552 Bytes::from_static(b"data")
553 );
554 assert_eq!(rx.read().await.unwrap().unwrap_err(), "err");
555 assert!(rx.is_eof());
556 assert!(rx.read().await.is_none());
557 }
558}