Skip to main content

ntex_util/channel/
bstream.rs

1//! A buffered stream of byte chunks.
2use 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
10/// Default high watermark, 32 KiB
11const HIGH_WATERMARK: u32 = 32_768;
12/// Default low watermark, 16 KiB
13const LOW_WATERMARK: u32 = HIGH_WATERMARK / 2;
14
15/// Indicates the current status of a byte stream.
16#[derive(Copy, Clone, Debug, PartialEq, Eq)]
17pub enum Status {
18    /// End of stream reached.
19    Eof,
20    /// Stream is ready to accept more bytes.
21    Ready,
22    /// The receiver side has been dropped.
23    Dropped,
24}
25
26/// Creates a byte stream and returns its sender and receiver.
27pub 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
38/// Creates a byte stream that starts at EOF.
39///
40/// Data may still be added through the returned sender. The receiver yields
41/// that buffered data first and then completes.
42pub 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
53/// Creates a receiver that contains optional data and is already at EOF.
54pub 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/// A buffered stream of byte chunks.
65///
66/// The receiver yields chunks in insertion order. Its configured buffer size
67/// is a cooperative backpressure threshold for the sender, not a hard memory
68/// limit.
69#[derive(Debug)]
70pub struct Receiver<E> {
71    inner: Rc<Inner<E>>,
72}
73
74impl<E> Receiver<E> {
75    /// Sets the sender backpressure watermarks.
76    ///
77    /// Once buffered data reaches `high` bytes, [`Sender::poll_ready`] stops
78    /// reporting [`Status::Ready`] until the receiver drains the buffer to
79    /// `low` bytes or less, so the sender is not woken for every consumed
80    /// chunk. `low` is capped below `high`.
81    ///
82    /// Sending does not enforce the watermarks, so producers must cooperate
83    /// by waiting for readiness. Changing the watermarks immediately updates
84    /// and, when needed, wakes sender readiness: the sender is ready if
85    /// buffered data is below `high`. The defaults are 32 KiB and 16 KiB.
86    #[inline]
87    pub fn set_watermarks(&self, high: u32, low: u32) {
88        self.inner.set_watermarks(high, low);
89    }
90
91    /// Sets the sender backpressure threshold.
92    ///
93    /// Sets the high watermark to `size` and the low watermark to half of it.
94    #[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    /// Puts previously read data back at the front of the stream.
102    ///
103    /// This may grow the buffer past its readiness threshold.
104    #[inline]
105    pub fn put(&self, data: Bytes) {
106        self.inner.unread_data(data);
107    }
108
109    #[inline]
110    /// Returns `true` once EOF has been marked.
111    ///
112    /// Buffered chunks may still be available after this returns `true`.
113    pub fn is_eof(&self) -> bool {
114        self.inner.flags.get().contains(Flags::EOF)
115    }
116
117    #[inline]
118    /// Waits for and returns the next chunk, stream error, or EOF.
119    pub async fn read(&self) -> Option<Result<Bytes, E>> {
120        poll_fn(|cx| self.poll_read(cx)).await
121    }
122
123    #[inline]
124    /// Polls for the next chunk, stream error, or EOF.
125    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/// Sender side of the byte stream.
155///
156/// Clones share one readiness registration. If several clones poll readiness,
157/// the most recently registered task is the one that will be woken.
158#[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            // a pending read completes with EOF
178            shared.recv_task.wake();
179        }
180    }
181}
182
183impl<E> Sender<E> {
184    /// Returns `true` if the receiver has been dropped.
185    pub fn is_closed(&self) -> bool {
186        self.inner.strong_count() == 0
187    }
188
189    /// Stores a terminal stream error and wakes both sides.
190    pub fn set_error(&self, err: E) {
191        if let Some(shared) = self.inner.upgrade() {
192            shared.set_error(err);
193        }
194    }
195
196    /// Marks the stream as EOF and wakes both sides.
197    pub fn feed_eof(&self) {
198        if let Some(shared) = self.inner.upgrade() {
199            shared.feed_eof();
200        }
201    }
202
203    /// Adds a chunk to the stream.
204    ///
205    /// This method does not enforce the configured backpressure threshold.
206    pub fn feed_data(&self, data: Bytes) {
207        if let Some(shared) = self.inner.upgrade() {
208            shared.feed_data(data);
209        }
210    }
211
212    /// Waits until the stream needs more data or reaches a terminal state.
213    pub async fn ready(&self) -> Status {
214        poll_fn(|cx| self.poll_ready(cx)).await
215    }
216
217    /// Polls until the stream needs more data or reaches a terminal state.
218    ///
219    /// Terminal states take precedence: [`Status::Dropped`] is returned once
220    /// the receiver is gone or an error was set, [`Status::Eof`] after EOF.
221    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            // receiver is gone
236            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            // wake up sender once the buffer is drained to the low watermark
337            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        // drop receiver
394        let (tx, rx) = channel::<()>();
395        assert!(!tx.is_closed());
396        drop(rx);
397        assert!(tx.is_closed());
398
399        // drop sender
400        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        // above the low watermark
447        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        // at the low watermark
452        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        // low watermark is capped below the high watermark
460        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        // buffered data is returned before the error, then eof
550        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}