1use std::{fmt, future::poll_fn, mem, pin::Pin, task::Context, task::Poll};
2
3use crate::channel::bstream;
4use crate::http::{HeaderMap, error::PayloadError, h1, h2};
5use crate::util::{Bytes, Stream};
6
7pub type PayloadStream = Pin<Box<dyn Stream<Item = Result<Bytes, PayloadError>>>>;
9
10#[derive(Default)]
12pub enum Payload {
13 #[default]
15 None,
16 H1(h1::Payload),
18 H2(h2::Payload),
20 Stream(PayloadStream),
22}
23
24impl From<h1::Payload> for Payload {
25 fn from(v: h1::Payload) -> Self {
26 Payload::H1(v)
27 }
28}
29
30impl From<bstream::Receiver<PayloadError>> for Payload {
31 fn from(v: bstream::Receiver<PayloadError>) -> Self {
32 Payload::H1(v.into())
33 }
34}
35
36impl From<h2::Payload> for Payload {
37 fn from(v: h2::Payload) -> Self {
38 Payload::H2(v)
39 }
40}
41
42impl From<PayloadStream> for Payload {
43 fn from(pl: PayloadStream) -> Self {
44 Payload::Stream(pl)
45 }
46}
47
48impl fmt::Debug for Payload {
49 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
50 match self {
51 Payload::None => write!(f, "Payload::None"),
52 Payload::H1(pl) => write!(f, "Payload::H1({pl:?})"),
53 Payload::H2(pl) => write!(f, "Payload::H2({pl:?})"),
54 Payload::Stream(_) => write!(f, "Payload::Stream(..)"),
55 }
56 }
57}
58
59impl Payload {
60 #[must_use]
61 pub fn take(&mut self) -> Self {
63 mem::take(self)
64 }
65
66 #[must_use]
67 pub fn from_stream<S>(stream: S) -> Self
72 where
73 S: Stream<Item = Result<Bytes, PayloadError>> + 'static,
74 {
75 Payload::Stream(Box::pin(stream))
76 }
77
78 #[inline]
79 pub async fn recv(&mut self) -> Option<Result<Bytes, PayloadError>> {
81 poll_fn(|cx| self.poll_recv(cx)).await
82 }
83
84 #[inline]
85 pub fn poll_recv(&mut self, cx: &mut Context<'_>) -> Poll<Option<Result<Bytes, PayloadError>>> {
90 match self {
91 Payload::None => Poll::Ready(None),
92 Payload::H1(pl) => pl.poll_read(cx),
93 Payload::H2(pl) => pl.poll_read(cx),
94 Payload::Stream(pl) => Pin::new(pl).poll_next(cx),
95 }
96 }
97
98 pub fn trailers(&self) -> Option<HeaderMap> {
103 match self {
104 Payload::H1(pl) => pl.trailers(),
105 Payload::H2(pl) => pl.trailers(),
106 _ => None,
107 }
108 }
109}
110
111impl Stream for Payload {
112 type Item = Result<Bytes, PayloadError>;
113
114 #[inline]
115 fn poll_next(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
116 self.get_mut().poll_recv(cx)
117 }
118}
119
120#[cfg(test)]
121mod tests {
122 use super::*;
123
124 #[test]
125 fn payload_debug() {
126 assert!(format!("{:?}", Payload::None).contains("Payload::None"));
127 assert!(
128 format!(
129 "{:?}",
130 Payload::H1(crate::channel::bstream::channel().1.into())
131 )
132 .contains("Payload::H1")
133 );
134 assert!(
135 format!(
136 "{:?}",
137 Payload::Stream(Box::pin(crate::channel::bstream::channel().1))
138 )
139 .contains("Payload::Stream")
140 );
141
142 assert_eq!(
143 std::mem::size_of::<Payload>(),
144 std::mem::size_of::<Option<Payload>>()
145 );
146 }
147
148 #[crate::rt_test]
149 async fn payload_conversions() {
150 use crate::http::h1;
151
152 let (tx, rx) = crate::channel::bstream::channel();
153 let mut pl = Payload::from(h1::Payload::from(rx));
154 assert!(matches!(pl, Payload::H1(_)));
155 tx.feed_data(Bytes::from_static(b"data"));
156 tx.feed_eof();
157 assert_eq!(pl.recv().await.unwrap().unwrap(), "data");
158 assert!(pl.recv().await.is_none());
159 assert!(pl.trailers().is_none());
160
161 let mut pl = Payload::None;
162 assert!(pl.recv().await.is_none());
163 assert!(pl.trailers().is_none());
164
165 let (tx, rx) = crate::channel::bstream::channel();
166 let pl: PayloadStream = Box::pin(rx);
167 let mut pl = Payload::from(pl);
168 assert!(matches!(pl, Payload::Stream(_)));
169 tx.feed_eof();
170 assert!(pl.recv().await.is_none());
171 assert!(pl.trailers().is_none());
172 assert!(matches!(pl.take(), Payload::Stream(_)));
173 assert!(matches!(pl, Payload::None));
174 }
175}