1use 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#[derive(Debug)]
43pub struct Payload(pub crate::http::Payload);
44
45impl Payload {
46 #[inline]
47 pub fn into_inner(self) -> crate::http::Payload {
49 self.0
50 }
51
52 #[inline]
53 pub async fn recv(&mut self) -> Option<Result<Bytes, error::PayloadError>> {
55 self.0.recv().await
56 }
57
58 #[inline]
59 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
79impl<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
120impl<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
169impl<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 cfg.check_mimetype(req)?;
216
217 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#[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 pub fn new(limit: usize) -> Self {
249 PayloadConfig {
250 limit,
251 ..Default::default()
252 }
253 }
254
255 #[must_use]
256 pub fn limit(mut self, limit: usize) -> Self {
260 self.limit = limit;
261 self
262 }
263
264 #[must_use]
265 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 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 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
333struct 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 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 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 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 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 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 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}