1use 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#[derive(Debug)]
30pub struct Payload {
31 inner: Rc<Inner>,
32}
33
34impl Payload {
35 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 pub async fn read(&self) -> Option<Result<Bytes, PayloadError>> {
59 poll_fn(|cx| self.poll_read(cx)).await
60 }
61
62 #[inline]
63 pub fn poll_read(&self, cx: &mut Context<'_>) -> Poll<Option<Result<Bytes, PayloadError>>> {
65 self.inner.readany(cx)
66 }
67
68 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)]
103pub 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 pub(crate) fn is_dropped(&self) -> bool {
120 self.inner.strong_count() == 0
121 }
122
123 pub fn set_error(&self, err: PayloadError) {
125 if let Some(shared) = self.inner.upgrade() {
126 shared.set_error(err);
127 }
128 }
129
130 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 pub fn feed_trailers(&self, trailers: HeaderMap) {
139 if let Some(shared) = self.inner.upgrade() {
140 shared.feed_trailers(trailers);
141 }
142 }
143
144 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 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 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 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 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 #[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 #[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 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 #[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 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 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}