Skip to main content

ntex/http/h2/
payload.rs

1//! Payload stream
2use std::collections::VecDeque;
3use std::task::{Context, Poll, Waker};
4use std::{cell::Cell, cell::RefCell, fmt, future::poll_fn, pin::Pin, rc::Rc, rc::Weak};
5
6use ntex_h2::{self as h2};
7
8use crate::http::HeaderMap;
9use crate::util::{Bytes, Stream};
10use crate::{http::error::PayloadError, task::LocalWaker};
11
12bitflags::bitflags! {
13    #[derive(Copy, Clone, Debug, Eq, PartialEq, Ord, PartialOrd, Hash)]
14    struct Flags: u8 {
15        const EOF = 0b0000_0001;
16        const DROPPED = 0b0000_0010;
17        const ERROR = 0b0000_0100;
18    }
19}
20
21/// Buffered HTTP/2 payload stream.
22///
23/// Use [`Payload::read`] to receive the next chunk asynchronously, or consume
24/// the payload through its [`Stream`] implementation. Waiting tasks are woken
25/// when data, end-of-stream, or an error becomes available.
26///
27/// This type is not thread-safe and can also be used as a
28/// [`Response`](crate::http::Response) body stream.
29#[derive(Debug)]
30pub struct Payload {
31    inner: Rc<Inner>,
32}
33
34impl Payload {
35    /// Create payload stream.
36    ///
37    /// This method construct two objects responsible for bytes stream
38    /// generation.
39    ///
40    /// * `PayloadSender` - *Sender* side of the stream
41    ///
42    /// * `Payload` - *Receiver* side of the stream
43    pub fn create(cap: h2::Capacity) -> (PayloadSender, Payload) {
44        let shared = Rc::new(Inner::new(cap));
45
46        (
47            PayloadSender {
48                inner: Rc::downgrade(&shared),
49            },
50            Payload { inner: shared },
51        )
52    }
53
54    #[inline]
55    /// Receives the next payload chunk.
56    ///
57    /// Returns `None` after the complete payload has been received.
58    pub async fn read(&self) -> Option<Result<Bytes, PayloadError>> {
59        poll_fn(|cx| self.poll_read(cx)).await
60    }
61
62    #[inline]
63    /// Polls for the next payload chunk.
64    pub fn poll_read(&self, cx: &mut Context<'_>) -> Poll<Option<Result<Bytes, PayloadError>>> {
65        self.inner.readany(cx)
66    }
67
68    /// Returns the trailer fields received at the end of the payload.
69    ///
70    /// Trailers are available after the payload is complete, None is
71    /// returned if the payload is not complete or has no trailers.
72    pub fn trailers(&self) -> Option<HeaderMap> {
73        if self.inner.items.borrow().is_empty() {
74            self.inner.trailers.borrow().clone()
75        } else {
76            None
77        }
78    }
79}
80
81impl Drop for Payload {
82    fn drop(&mut self) {
83        self.inner.io_task.wake();
84        self.inner.insert_flags(Flags::DROPPED);
85        if let Some(f) = self.inner.on_drop.take() {
86            f();
87        }
88    }
89}
90
91impl Stream for Payload {
92    type Item = Result<Bytes, PayloadError>;
93
94    fn poll_next(
95        self: Pin<&mut Self>,
96        cx: &mut Context<'_>,
97    ) -> Poll<Option<Result<Bytes, PayloadError>>> {
98        self.inner.readany(cx)
99    }
100}
101
102#[derive(Debug)]
103/// Sender part of the payload stream
104pub struct PayloadSender {
105    inner: Weak<Inner>,
106}
107
108impl Drop for PayloadSender {
109    fn drop(&mut self) {
110        if let Some(shared) = self.inner.upgrade() {
111            drop(shared.on_drop.take());
112            shared.set_error(PayloadError::Incomplete(None));
113        }
114    }
115}
116
117impl PayloadSender {
118    /// Checks if the payload stream is dropped.
119    pub(crate) fn is_dropped(&self) -> bool {
120        self.inner.strong_count() == 0
121    }
122
123    /// Closes the payload stream with an error.
124    pub fn set_error(&self, err: PayloadError) {
125        if let Some(shared) = self.inner.upgrade() {
126            shared.set_error(err);
127        }
128    }
129
130    /// Sends the final payload chunk and closes the stream.
131    pub fn feed_eof(&self, data: Bytes, cap: Option<h2::Capacity>) {
132        if let Some(shared) = self.inner.upgrade() {
133            shared.feed_eof(data, cap);
134        }
135    }
136
137    /// Sends the trailer fields and closes the stream.
138    pub fn feed_trailers(&self, trailers: HeaderMap) {
139        if let Some(shared) = self.inner.upgrade() {
140            shared.feed_trailers(trailers);
141        }
142    }
143
144    /// Sends a payload chunk and updates the HTTP/2 flow-control capacity.
145    pub fn feed_data(&self, data: Bytes, cap: h2::Capacity) {
146        if let Some(shared) = self.inner.upgrade() {
147            shared.feed_data(data, cap);
148        }
149    }
150
151    /// Registers a callback that runs if the payload is dropped while the sender is alive.
152    pub(crate) fn on_drop(&self, f: impl FnOnce() + 'static) {
153        if let Some(shared) = self.inner.upgrade() {
154            shared.on_drop.set(Some(Box::new(f)));
155        }
156    }
157
158    pub(crate) fn on_cancel(&self, w: &Waker) -> Poll<()> {
159        if let Some(shared) = self.inner.upgrade() {
160            if shared.flags.get().contains(Flags::DROPPED) {
161                Poll::Ready(())
162            } else {
163                shared.io_task.register(w);
164                Poll::Pending
165            }
166        } else {
167            Poll::Ready(())
168        }
169    }
170}
171
172struct Inner {
173    flags: Cell<Flags>,
174    cap: Cell<Option<h2::Capacity>>,
175    err: Cell<Option<PayloadError>>,
176    items: RefCell<VecDeque<Bytes>>,
177    trailers: RefCell<Option<HeaderMap>>,
178    task: LocalWaker,
179    io_task: LocalWaker,
180    on_drop: Cell<Option<Box<dyn FnOnce()>>>,
181}
182
183impl Inner {
184    fn new(cap: h2::Capacity) -> Self {
185        Inner {
186            cap: Cell::new(Some(cap)),
187            flags: Cell::new(Flags::empty()),
188            err: Cell::new(None),
189            items: RefCell::new(VecDeque::new()),
190            trailers: RefCell::new(None),
191            task: LocalWaker::new(),
192            io_task: LocalWaker::new(),
193            on_drop: Cell::new(None),
194        }
195    }
196
197    fn insert_flags(&self, f: Flags) {
198        let mut flags = self.flags.get();
199        flags.insert(f);
200        self.flags.set(flags);
201    }
202
203    fn set_error(&self, err: PayloadError) {
204        // the first error is kept, a finished payload is not failed
205        if !self.flags.get().intersects(Flags::EOF | Flags::ERROR) {
206            self.insert_flags(Flags::ERROR);
207            self.err.set(Some(err));
208            self.task.wake();
209        }
210    }
211
212    fn feed_eof(&self, data: Bytes, cap: Option<h2::Capacity>) {
213        if let Some(cap) = cap {
214            self.cap.set(Some(self.cap.take().unwrap() + cap));
215        }
216        self.insert_flags(Flags::EOF);
217        if !data.is_empty() {
218            self.items.borrow_mut().push_back(data);
219        }
220        self.task.wake();
221    }
222
223    fn feed_trailers(&self, trailers: HeaderMap) {
224        if !self.flags.get().intersects(Flags::EOF | Flags::ERROR) {
225            *self.trailers.borrow_mut() = Some(trailers);
226            self.insert_flags(Flags::EOF);
227            self.task.wake();
228        }
229    }
230
231    fn feed_data(&self, data: Bytes, cap: h2::Capacity) {
232        self.cap.set(Some(self.cap.take().unwrap() + cap));
233        // empty DATA frames are not flow controlled, queueing them is unbounded
234        if !data.is_empty() {
235            self.items.borrow_mut().push_back(data);
236            self.task.wake();
237        }
238    }
239
240    fn readany(&self, cx: &mut Context<'_>) -> Poll<Option<Result<Bytes, PayloadError>>> {
241        if let Some(data) = self.items.borrow_mut().pop_front() {
242            let cap = self.cap.take().unwrap();
243            cap.consume(data.len() as u32);
244            self.cap.set(Some(cap));
245            Poll::Ready(Some(Ok(data)))
246        } else if let Some(err) = self.err.take() {
247            // the payload ends after an error
248            self.insert_flags(Flags::EOF);
249            Poll::Ready(Some(Err(err)))
250        } else if self.flags.get().contains(Flags::EOF) {
251            Poll::Ready(None)
252        } else {
253            self.task.register(cx.waker());
254            Poll::Pending
255        }
256    }
257}
258
259impl fmt::Debug for Inner {
260    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
261        let cap = self.cap.take().unwrap();
262        let err = self.err.take();
263        let result = f
264            .debug_struct("Inner")
265            .field("flags", &self.flags.get())
266            .field("capacity", &cap)
267            .field("error", &err)
268            .field("items", &self.items.borrow())
269            .field("trailers", &self.trailers.borrow())
270            .finish();
271
272        self.cap.set(Some(cap));
273        self.err.set(err);
274        result
275    }
276}
277
278#[cfg(test)]
279mod tests {
280    use std::sync::{Arc, atomic::AtomicUsize, atomic::Ordering};
281    use std::task::Wake;
282
283    use ntex_h2::client::SimpleClient;
284
285    use super::*;
286    use crate::http::{HeaderMap, Method};
287    use crate::io::{Io, IoBoxed, testing::IoTest};
288    use crate::{SharedCfg, time::Millis, time::sleep, util::ByteString};
289
290    struct Counter(AtomicUsize);
291
292    impl Wake for Counter {
293        fn wake(self: Arc<Self>) {
294            self.0.fetch_add(1, Ordering::SeqCst);
295        }
296    }
297
298    /// A pending read does not wake the payload io task, only dropping the payload does.
299    #[crate::rt_test]
300    async fn test_pending_read_does_not_wake_io_task() {
301        let (io, server) = IoTest::create();
302        io.remote_buffer_cap(64 * 1024);
303        let client = SimpleClient::new(
304            IoBoxed::from(Io::new(io, SharedCfg::default())),
305            false,
306            ByteString::from_static("localhost"),
307        );
308        server.write([0, 0, 0, 4, 0, 0, 0, 0, 0]);
309        sleep(Millis(50)).await;
310        let (snd, _rcv) = client
311            .send(Method::GET, "/".into(), HeaderMap::default(), true)
312            .await
313            .unwrap();
314
315        let (sender, payload) = Payload::create(snd.stream().empty_capacity());
316        let io_task = Arc::new(Counter(AtomicUsize::new(0)));
317        let io_waker = io_task.clone().into();
318        assert!(sender.on_cancel(&io_waker).is_pending());
319
320        let reader = Arc::new(Counter(AtomicUsize::new(0)));
321        let reader_waker = reader.clone().into();
322        let mut cx = Context::from_waker(&reader_waker);
323        for _ in 0..3 {
324            assert!(payload.poll_read(&mut cx).is_pending());
325        }
326        assert_eq!(io_task.0.load(Ordering::SeqCst), 0);
327
328        drop(payload);
329        assert_eq!(io_task.0.load(Ordering::SeqCst), 1);
330        assert!(sender.on_cancel(&io_waker).is_ready());
331    }
332
333    /// A ready read does not register the reader waker, new data does not wake
334    /// a reader that is not waiting.
335    #[crate::rt_test]
336    async fn test_ready_read_does_not_register_waker() {
337        let (io, server) = IoTest::create();
338        io.remote_buffer_cap(64 * 1024);
339        let client = SimpleClient::new(
340            IoBoxed::from(Io::new(io, SharedCfg::default())),
341            false,
342            ByteString::from_static("localhost"),
343        );
344        server.write([0, 0, 0, 4, 0, 0, 0, 0, 0]);
345        sleep(Millis(50)).await;
346        let (snd, rcv) = client
347            .send(Method::GET, "/".into(), HeaderMap::default(), true)
348            .await
349            .unwrap();
350
351        // `200` response and two DATA frames
352        server.write([0, 0, 1, 1, 4, 0, 0, 0, 1, 0x88]);
353        server.write([0, 0, 2, 0, 0, 0, 0, 0, 1, b'a', b'b']);
354        server.write([0, 0, 2, 0, 0, 0, 0, 0, 1, b'c', b'd']);
355        let _ = rcv.recv().await.unwrap();
356        let mut chunks = Vec::new();
357        for _ in 0..2 {
358            let msg = rcv.recv().await.unwrap();
359            let h2::MessageKind::Data(data, cap) = msg.kind else {
360                panic!("unexpected message: {msg:?}")
361            };
362            chunks.push((data, cap));
363        }
364
365        let (sender, payload) = Payload::create(snd.stream().empty_capacity());
366        let reader = Arc::new(Counter(AtomicUsize::new(0)));
367        let reader_waker = reader.clone().into();
368        let mut cx = Context::from_waker(&reader_waker);
369
370        let (data, cap) = chunks.remove(0);
371        sender.feed_data(data, cap);
372        assert!(matches!(payload.poll_read(&mut cx), Poll::Ready(Some(Ok(ref d))) if d == "ab"));
373
374        let (data, cap) = chunks.remove(0);
375        sender.feed_data(data, cap);
376        assert_eq!(reader.0.load(Ordering::SeqCst), 0);
377        assert!(matches!(payload.poll_read(&mut cx), Poll::Ready(Some(Ok(ref d))) if d == "cd"));
378    }
379
380    /// Trailers are available after all queued data is read.
381    #[crate::rt_test]
382    async fn test_trailers_after_data() {
383        let (io, server) = IoTest::create();
384        io.remote_buffer_cap(64 * 1024);
385        let client = SimpleClient::new(
386            IoBoxed::from(Io::new(io, SharedCfg::default())),
387            false,
388            ByteString::from_static("localhost"),
389        );
390        server.write([0, 0, 0, 4, 0, 0, 0, 0, 0]);
391        sleep(Millis(50)).await;
392        let (snd, rcv) = client
393            .send(Method::GET, "/".into(), HeaderMap::default(), true)
394            .await
395            .unwrap();
396
397        // `200` response and a DATA frame
398        server.write([0, 0, 1, 1, 4, 0, 0, 0, 1, 0x88]);
399        server.write([0, 0, 2, 0, 0, 0, 0, 0, 1, b'a', b'b']);
400        let _ = rcv.recv().await.unwrap();
401        let msg = rcv.recv().await.unwrap();
402        let h2::MessageKind::Data(data, cap) = msg.kind else {
403            panic!("unexpected message: {msg:?}")
404        };
405
406        let (sender, payload) = Payload::create(snd.stream().empty_capacity());
407        sender.feed_data(data, cap);
408        let mut trailers = HeaderMap::default();
409        trailers.insert(
410            crate::http::header::HeaderName::from_static("x-trailer"),
411            crate::http::header::HeaderValue::from_static("1"),
412        );
413        sender.feed_trailers(trailers);
414        // ignored, the payload is complete
415        sender.set_error(PayloadError::Incomplete(None));
416        assert!(payload.trailers().is_none());
417
418        assert_eq!(payload.read().await.unwrap().unwrap(), "ab");
419        assert_eq!(payload.trailers().unwrap().get("x-trailer").unwrap(), "1");
420        assert!(payload.read().await.is_none());
421        assert!(payload.trailers().is_some());
422    }
423
424    #[crate::rt_test]
425    async fn test_debug_and_stream() {
426        use crate::util::stream_recv;
427
428        let (io, server) = IoTest::create();
429        io.remote_buffer_cap(64 * 1024);
430        let client = SimpleClient::new(
431            IoBoxed::from(Io::new(io, SharedCfg::default())),
432            false,
433            ByteString::from_static("localhost"),
434        );
435        server.write([0, 0, 0, 4, 0, 0, 0, 0, 0]);
436        sleep(Millis(50)).await;
437        let (snd, _rcv) = client
438            .send(Method::GET, "/".into(), HeaderMap::default(), true)
439            .await
440            .unwrap();
441
442        let (sender, mut payload) = Payload::create(snd.stream().empty_capacity());
443        let s = format!("{payload:?}");
444        assert!(s.contains("Inner") && s.contains("capacity"), "{s}");
445
446        let mut trailers = HeaderMap::default();
447        trailers.insert(
448            crate::http::header::HeaderName::from_static("x-trailer"),
449            crate::http::header::HeaderValue::from_static("1"),
450        );
451        sender.feed_trailers(trailers);
452        assert!(stream_recv(&mut payload).await.is_none());
453        assert_eq!(payload.trailers().unwrap().len(), 1);
454    }
455}