Skip to main content

ntex/web/types/
payload.rs

1//! Payload/Bytes/String extractors
2use std::{
3    borrow::Cow, convert::Infallible, fmt, future::Future, pin::Pin, str, sync::Arc, task::Context,
4    task::Poll,
5};
6
7use encoding_rs::UTF_8;
8use mime::Mime;
9
10use crate::http::{HttpMessage, error, header};
11use crate::util::{BoxFuture, Bytes, Stream};
12use crate::web::{FromRequest, HttpRequest, State, error::PayloadError};
13
14/// Payload extractor returns request's payload stream.
15///
16/// ## Example
17///
18/// ```rust
19/// use std::future::Future;
20/// use ntex::util::{BytesMut, Stream};
21/// use ntex::web::{self, error, App, HttpResponse};
22///
23/// /// extract binary data from request
24/// async fn index(mut body: web::types::Payload) -> Result<HttpResponse, error::PayloadError>
25/// {
26///     let mut bytes = BytesMut::new();
27///     while let Some(item) = ntex::util::stream_recv(&mut body).await {
28///         bytes.extend_from_slice(&item?);
29///     }
30///
31///     format!("Body {:?}!", bytes);
32///     Ok(HttpResponse::Ok().build())
33/// }
34///
35/// fn main() {
36///     let app = App::default().service(
37///         web::resource("/index.html").route(
38///             web::get().to(index))
39///     );
40/// }
41/// ```
42#[derive(Debug)]
43pub struct Payload(pub crate::http::Payload);
44
45impl Payload {
46    #[inline]
47    /// Deconstruct to a inner value
48    pub fn into_inner(self) -> crate::http::Payload {
49        self.0
50    }
51
52    #[inline]
53    /// Attempt to pull out the next value of this payload.
54    pub async fn recv(&mut self) -> Option<Result<Bytes, error::PayloadError>> {
55        self.0.recv().await
56    }
57
58    #[inline]
59    /// Attempt to pull out the next value of this payload, registering
60    /// the current task for wakeup if the value is not yet available,
61    /// and returning None if the payload is exhausted.
62    pub fn poll_recv(
63        &mut self,
64        cx: &mut Context<'_>,
65    ) -> Poll<Option<Result<Bytes, error::PayloadError>>> {
66        self.0.poll_recv(cx)
67    }
68}
69
70impl Stream for Payload {
71    type Item = Result<Bytes, error::PayloadError>;
72
73    #[inline]
74    fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
75        self.poll_recv(cx)
76    }
77}
78
79/// Get request's payload stream
80///
81/// ## Example
82///
83/// ```rust
84/// use std::future::Future;
85/// use ntex::util::{BytesMut, Stream};
86/// use ntex::web::{self, error, App, HttpResponse};
87///
88/// /// extract binary data from request
89/// async fn index(mut body: web::types::Payload) -> Result<HttpResponse, error::PayloadError>
90/// {
91///     let mut bytes = BytesMut::new();
92///     while let Some(item) = ntex::util::stream_recv(&mut body).await {
93///         bytes.extend_from_slice(&item?);
94///     }
95///
96///     format!("Body {:?}!", bytes);
97///     Ok(HttpResponse::Ok().build())
98/// }
99///
100/// fn main() {
101///     let app = App::default().service(
102///         web::resource("/index.html").route(
103///             web::get().to(index))
104///     );
105/// }
106/// ```
107impl<St: State> FromRequest<St> for Payload {
108    type Error = Infallible;
109
110    #[inline]
111    async fn from_request(
112        _: &St,
113        _: &HttpRequest,
114        payload: &mut crate::http::Payload,
115    ) -> Result<Payload, Self::Error> {
116        Ok(Payload(payload.take()))
117    }
118}
119
120/// Request binary data from a request's payload.
121///
122/// Loads request's payload and construct Bytes instance.
123///
124/// [`PayloadConfig`] allows to configure
125/// extraction process.
126///
127/// ## Example
128///
129/// ```rust
130/// use ntex::{web, util::Bytes};
131///
132/// /// extract binary data from request
133/// async fn index(body: Bytes) -> String {
134///     format!("Body {:?}!", body)
135/// }
136///
137/// fn main() {
138///     let app = web::App::default().service(
139///         web::resource("/index.html").route(
140///             web::get().to(index))
141///     );
142/// }
143/// ```
144impl<St: State> FromRequest<St> for Bytes {
145    type Error = PayloadError;
146
147    async fn from_request(
148        _: &St,
149        req: &HttpRequest,
150        payload: &mut crate::http::Payload,
151    ) -> Result<Bytes, Self::Error> {
152        let tmp;
153        let cfg = if let Some(cfg) = req.app_state::<PayloadConfig>() {
154            cfg
155        } else {
156            tmp = PayloadConfig::default();
157            &tmp
158        };
159
160        if let Err(e) = cfg.check_mimetype(req) {
161            Err(e)
162        } else {
163            let limit = cfg.limit;
164            HttpMessageBody::new(req, payload).limit(limit).await
165        }
166    }
167}
168
169/// Extract text information from a request's body.
170///
171/// Text extractor automatically decode body according to the request's charset.
172///
173/// [`PayloadConfig`] allows to configure
174/// extraction process.
175///
176/// ## Example
177///
178/// ```rust
179/// use ntex::web::{self, App, FromRequest, WebAppConfig};
180///
181/// /// extract text data from request
182/// async fn index(text: String) -> String {
183///     format!("Body {}!", text)
184/// }
185///
186/// fn main() {
187///     let cfg = WebAppConfig::new()
188///         .set_state(web::types::PayloadConfig::new(4096)); // <- limit size of the payload
189///
190///     let app = App::default()
191///         .with_config(cfg)
192///         .service(
193///             web::resource("/index.html")
194///                 .route(web::get().to(index))  // <- register handler with extractor params
195///     );
196/// }
197/// ```
198impl<St: State> FromRequest<St> for String {
199    type Error = PayloadError;
200
201    async fn from_request(
202        _: &St,
203        req: &HttpRequest,
204        payload: &mut crate::http::Payload,
205    ) -> Result<String, Self::Error> {
206        let tmp;
207        let cfg = if let Some(cfg) = req.app_state::<PayloadConfig>() {
208            cfg
209        } else {
210            tmp = PayloadConfig::default();
211            &tmp
212        };
213
214        // check content-type
215        cfg.check_mimetype(req)?;
216
217        // check charset
218        let encoding = match req.encoding() {
219            Ok(enc) => enc,
220            Err(e) => return Err(PayloadError::from(e)),
221        };
222        let limit = cfg.limit;
223        let body = HttpMessageBody::new(req, payload).limit(limit).await?;
224
225        if encoding == UTF_8 {
226            Ok(str::from_utf8(body.as_ref())
227                .map_err(|_| PayloadError::Decoding)?
228                .to_owned())
229        } else {
230            Ok(encoding
231                .decode_without_bom_handling_and_without_replacement(&body)
232                .map(Cow::into_owned)
233                .ok_or(PayloadError::Decoding)?)
234        }
235    }
236}
237
238/// Payload configuration for request's payload.
239#[derive(Clone)]
240pub struct PayloadConfig {
241    limit: usize,
242    content_type: Option<Arc<dyn Fn(Mime) -> bool + Send + Sync>>,
243}
244
245impl PayloadConfig {
246    #[must_use]
247    /// Create `PayloadConfig` instance and set max size of payload.
248    pub fn new(limit: usize) -> Self {
249        PayloadConfig {
250            limit,
251            ..Default::default()
252        }
253    }
254
255    #[must_use]
256    /// Change max size of payload.
257    ///
258    /// By default max size is 256Kb.
259    pub fn limit(mut self, limit: usize) -> Self {
260        self.limit = limit;
261        self
262    }
263
264    #[must_use]
265    /// Set predicate for allowed content types.
266    ///
267    /// The predicate receives the request's parsed `Content-Type`. A request
268    /// without a `Content-Type` header, or one the predicate rejects, fails
269    /// with a content-type error. By default the content type is not checked.
270    pub fn content_type<F>(mut self, predicate: F) -> Self
271    where
272        F: Fn(Mime) -> bool + Send + Sync + 'static,
273    {
274        self.content_type = Some(Arc::new(predicate));
275        self
276    }
277
278    #[must_use]
279    /// Set required mime-type of the request.
280    ///
281    /// This is a shorthand for [`content_type`](Self::content_type) with a
282    /// predicate that accepts only `mt`. The comparison includes parameters,
283    /// such as `charset`. By default mime type is not enforced.
284    pub fn mimetype(self, mt: Mime) -> Self {
285        self.content_type(move |req_mt| req_mt == mt)
286    }
287
288    fn check_mimetype(&self, req: &HttpRequest) -> Result<(), PayloadError> {
289        // check content-type
290        if let Some(ref predicate) = self.content_type {
291            match req.mime_type() {
292                Ok(Some(req_mt)) => {
293                    if !predicate(req_mt) {
294                        return Err(PayloadError::from(error::ContentTypeError::Unexpected));
295                    }
296                }
297                Ok(None) => {
298                    return Err(PayloadError::from(error::ContentTypeError::Expected));
299                }
300                Err(err) => {
301                    return Err(err.into());
302                }
303            }
304        }
305        Ok(())
306    }
307}
308
309impl Default for PayloadConfig {
310    fn default() -> Self {
311        PayloadConfig {
312            limit: 262_144,
313            content_type: None,
314        }
315    }
316}
317
318impl fmt::Debug for PayloadConfig {
319    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
320        f.debug_struct("PayloadConfig")
321            .field("limit", &self.limit)
322            .field(
323                "content_type",
324                &self
325                    .content_type
326                    .as_ref()
327                    .map(|_| "Arc<dyn Fn(mime::Mime) -> bool + Send + Sync>"),
328            )
329            .finish()
330    }
331}
332
333/// Future that resolves to a complete http message body.
334///
335/// Load http message body.
336///
337/// By default only 256Kb payload reads to a memory, then
338/// `PayloadError::Overflow` is returned. Use `HttpMessageBody::limit()`
339/// method to change upper limit.
340struct HttpMessageBody {
341    limit: usize,
342    length: Option<usize>,
343    #[cfg(feature = "compress")]
344    stream: Option<crate::http::encoding::Decoder<crate::http::Payload>>,
345    #[cfg(not(feature = "compress"))]
346    stream: Option<crate::http::Payload>,
347    err: Option<PayloadError>,
348    fut: Option<BoxFuture<'static, Result<Bytes, PayloadError>>>,
349}
350
351impl HttpMessageBody {
352    /// Create `MessageBody` for request.
353    fn new(req: &HttpRequest, payload: &mut crate::http::Payload) -> HttpMessageBody {
354        let mut len = None;
355        if let Some(l) = req.headers().get(&header::CONTENT_LENGTH) {
356            if let Ok(s) = l.to_str() {
357                if let Ok(l) = s.parse::<usize>() {
358                    len = Some(l);
359                } else {
360                    return Self::err(PayloadError::Payload(error::PayloadError::UnknownLength));
361                }
362            } else {
363                return Self::err(PayloadError::Payload(error::PayloadError::UnknownLength));
364            }
365        }
366
367        #[cfg(feature = "compress")]
368        let stream = Some(crate::http::encoding::Decoder::from_headers(
369            payload.take(),
370            req.headers(),
371        ));
372        #[cfg(not(feature = "compress"))]
373        let stream = Some(payload.take());
374
375        HttpMessageBody {
376            stream,
377            limit: 262_144,
378            length: len,
379            fut: None,
380            err: None,
381        }
382    }
383
384    /// Change max size of payload. By default max size is 256Kb
385    fn limit(mut self, limit: usize) -> Self {
386        self.limit = limit;
387        self
388    }
389
390    fn err(e: PayloadError) -> Self {
391        HttpMessageBody {
392            stream: None,
393            limit: 262_144,
394            fut: None,
395            err: Some(e),
396            length: None,
397        }
398    }
399}
400
401impl Future for HttpMessageBody {
402    type Output = Result<Bytes, PayloadError>;
403
404    fn poll(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
405        if let Some(ref mut fut) = self.fut {
406            return Pin::new(fut).poll(cx);
407        }
408
409        if let Some(err) = self.err.take() {
410            return Poll::Ready(Err(err));
411        }
412
413        if let Some(len) = self.length
414            && len > self.limit
415        {
416            return Poll::Ready(Err(PayloadError::from(error::PayloadError::Overflow)));
417        }
418
419        // future
420        let (limit, length) = (self.limit, self.length);
421        let mut stream = self.stream.take().unwrap();
422        self.fut = Some(Box::pin(async move {
423            super::read_body(&mut stream, limit, length, |_| {
424                PayloadError::from(error::PayloadError::Overflow)
425            })
426            .await
427        }));
428        self.poll(cx)
429    }
430}
431
432#[cfg(test)]
433mod tests {
434    use super::*;
435    use crate::web::test::{TestRequest, from_request};
436
437    #[crate::rt_test]
438    async fn test_payload_config() {
439        let req = TestRequest::default().to_http_request();
440        let cfg = PayloadConfig::default()
441            .limit(5)
442            .mimetype(mime::APPLICATION_JSON);
443        assert!(cfg.check_mimetype(&req).is_err());
444
445        let req =
446            TestRequest::with_header(header::CONTENT_TYPE, "application/x-www-form-urlencoded")
447                .to_http_request();
448        assert!(cfg.check_mimetype(&req).is_err());
449
450        let req =
451            TestRequest::with_header(header::CONTENT_TYPE, "application/json").to_http_request();
452        assert!(cfg.check_mimetype(&req).is_ok());
453
454        let cfg = PayloadConfig::default()
455            .content_type(|mt| mt.type_() == mime::TEXT && mt.subtype() == mime::PLAIN);
456        let req =
457            TestRequest::with_header(header::CONTENT_TYPE, "application/json").to_http_request();
458        assert!(cfg.check_mimetype(&req).is_err());
459
460        let req = TestRequest::with_header(header::CONTENT_TYPE, "text/plain; charset=utf-8")
461            .to_http_request();
462        assert!(cfg.check_mimetype(&req).is_ok());
463        assert!(format!("{cfg:?}").contains("PayloadConfig"));
464    }
465
466    #[crate::rt_test]
467    async fn test_payload() {
468        let (req, mut pl, ()) = TestRequest::with_header(header::CONTENT_LENGTH, "11")
469            .payload(Bytes::from_static(b"hello=world"))
470            .to_http_parts();
471
472        let mut s = from_request::<_, Payload>(&(), &req, &mut pl)
473            .await
474            .unwrap();
475        let b = crate::util::stream_recv(&mut s).await.unwrap().unwrap();
476        assert_eq!(b, Bytes::from_static(b"hello=world"));
477    }
478
479    #[crate::rt_test]
480    async fn test_payload_recv() {
481        let (req, mut pl, ()) = TestRequest::with_header(header::CONTENT_LENGTH, "11")
482            .payload(Bytes::from_static(b"hello=world"))
483            .to_http_parts();
484
485        let mut s = from_request::<_, Payload>(&(), &req, &mut pl)
486            .await
487            .unwrap();
488        let b = s.recv().await.unwrap().unwrap();
489        assert_eq!(b, Bytes::from_static(b"hello=world"));
490    }
491
492    #[crate::rt_test]
493    async fn test_bytes() {
494        let (req, mut pl, ()) = TestRequest::with_header(header::CONTENT_LENGTH, "11")
495            .payload(Bytes::from_static(b"hello=world"))
496            .to_http_parts();
497
498        let s = from_request::<_, Bytes>(&(), &req, &mut pl).await.unwrap();
499        assert_eq!(s, Bytes::from_static(b"hello=world"));
500
501        let (req, mut pl, ()) = TestRequest::with_header(header::CONTENT_LENGTH, "11")
502            .payload(Bytes::from_static(b"hello=world"))
503            .app_state(PayloadConfig::default().mimetype(mime::APPLICATION_JSON))
504            .to_http_parts();
505        assert!(from_request::<_, Bytes>(&(), &req, &mut pl).await.is_err());
506    }
507
508    #[crate::rt_test]
509    async fn test_string() {
510        let (req, mut pl, ()) = TestRequest::with_header(header::CONTENT_LENGTH, "11")
511            .payload(Bytes::from_static(b"hello=world"))
512            .to_http_parts();
513
514        let s = from_request::<_, String>(&(), &req, &mut pl).await.unwrap();
515        assert_eq!(s, "hello=world");
516
517        let (req, mut pl, ()) = TestRequest::with_header(header::CONTENT_LENGTH, "11")
518            .header(header::CONTENT_TYPE, "text/plain; charset=cp1251")
519            .payload(Bytes::from_static(b"hello=world"))
520            .to_http_parts();
521        let s = from_request::<_, String>(&(), &req, &mut pl).await.unwrap();
522        assert_eq!(s, "hello=world");
523
524        let (req, mut pl, ()) = TestRequest::with_header(header::CONTENT_LENGTH, "11")
525            .payload(Bytes::from_static(b"hello=world"))
526            .app_state(PayloadConfig::default().mimetype(mime::APPLICATION_JSON))
527            .to_http_parts();
528        assert!(from_request::<_, String>(&(), &req, &mut pl).await.is_err());
529    }
530
531    #[crate::rt_test]
532    async fn test_message_body() {
533        let (req, mut pl, ()) = TestRequest::with_header(header::CONTENT_LENGTH, "xxxx")
534            .to_srv_request()
535            .into_parts();
536        let res = HttpMessageBody::new(&req, &mut pl).await;
537        match res.err().unwrap() {
538            PayloadError::Payload(error::PayloadError::UnknownLength) => (),
539            _ => unreachable!("error"),
540        }
541
542        let (req, mut pl, ()) = TestRequest::with_header(header::CONTENT_LENGTH, "1000000")
543            .to_srv_request()
544            .into_parts();
545        let res = HttpMessageBody::new(&req, &mut pl).await;
546        match res.err().unwrap() {
547            PayloadError::Payload(error::PayloadError::Overflow) => (),
548            _ => unreachable!("error"),
549        }
550
551        let (req, mut pl, ()) = TestRequest::default()
552            .payload(Bytes::from_static(b"test"))
553            .to_http_parts();
554        let res = HttpMessageBody::new(&req, &mut pl).await;
555        assert_eq!(res.ok().unwrap(), Bytes::from_static(b"test"));
556
557        let (req, mut pl, ()) = TestRequest::default()
558            .payload(Bytes::from_static(b"11111111111111"))
559            .to_http_parts();
560        let res = HttpMessageBody::new(&req, &mut pl).limit(5).await;
561        match res.err().unwrap() {
562            PayloadError::Payload(error::PayloadError::Overflow) => (),
563            _ => unreachable!("error"),
564        }
565    }
566
567    #[crate::rt_test]
568    async fn test_payload_errors() {
569        let cfg = PayloadConfig::new(5);
570        assert_eq!(cfg.limit, 5);
571
572        // invalid charset
573        let (req, mut pl, ()) = TestRequest::with_header(header::CONTENT_LENGTH, "11")
574            .header(header::CONTENT_TYPE, "text/plain; charset=unknown")
575            .payload(Bytes::from_static(b"hello=world"))
576            .to_http_parts();
577        assert!(from_request::<_, String>(&(), &req, &mut pl).await.is_err());
578
579        // invalid mime type
580        let cfg = PayloadConfig::default().mimetype(mime::APPLICATION_JSON);
581        let req = TestRequest::with_header(header::CONTENT_TYPE, "invalid").to_http_request();
582        assert!(matches!(
583            cfg.check_mimetype(&req),
584            Err(PayloadError::ContentType(_))
585        ));
586
587        // non-ascii content-length
588        let (req, mut pl, ()) = TestRequest::with_header(
589            header::CONTENT_LENGTH,
590            header::HeaderValue::from_bytes(b"1\xff").unwrap(),
591        )
592        .to_http_parts();
593        let res = HttpMessageBody::new(&req, &mut pl).await;
594        assert!(matches!(
595            res,
596            Err(PayloadError::Payload(error::PayloadError::UnknownLength))
597        ));
598    }
599}