1use std::{fmt, future::Future, ops, pin::Pin, sync::Arc, task::Context, task::Poll};
3
4use serde::{Serialize, de::DeserializeOwned};
5
6#[cfg(feature = "compress")]
7use crate::http::encoding::Decoder;
8use crate::http::header::CONTENT_LENGTH;
9use crate::http::{HttpMessage, Payload, Response, StatusCode, error::PayloadError};
10use crate::util::BoxFuture;
11use crate::web::error::{JsonError, JsonPayloadError, WebResponseError};
12use crate::web::{FromRequest, HttpRequest, Responder, State};
13
14pub struct Json<T>(pub T);
70
71impl<T> Json<T> {
72 pub fn into_inner(self) -> T {
74 self.0
75 }
76}
77
78impl<T> ops::Deref for Json<T> {
79 type Target = T;
80
81 fn deref(&self) -> &T {
82 &self.0
83 }
84}
85
86impl<T> ops::DerefMut for Json<T> {
87 fn deref_mut(&mut self) -> &mut T {
88 &mut self.0
89 }
90}
91
92impl<T> fmt::Debug for Json<T>
93where
94 T: fmt::Debug,
95{
96 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
97 f.debug_tuple("Json").field(&self.0).finish()
98 }
99}
100
101impl<T> fmt::Display for Json<T>
102where
103 T: fmt::Display,
104{
105 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
106 fmt::Display::fmt(&self.0, f)
107 }
108}
109
110impl<St, T: Serialize> Responder<St> for Json<T>
111where
112 St: State,
113 JsonError: WebResponseError<St, St::Error>,
114{
115 async fn respond_to(self, st: &St, _: &HttpRequest) -> Response {
116 let body = match crate::http::helpers::json_body(&self.0) {
117 Ok(body) => body,
118 Err(e) => return e.error_response(st),
119 };
120
121 Response::builder(StatusCode::OK)
122 .content_type("application/json")
123 .body(body)
124 }
125}
126
127impl<St, T> FromRequest<St> for Json<T>
159where
160 St: State,
161 T: DeserializeOwned + 'static,
162{
163 type Error = JsonPayloadError;
164
165 async fn from_request(
166 _: &St,
167 req: &HttpRequest,
168 payload: &mut Payload,
169 ) -> Result<Self, Self::Error> {
170 let req2 = req.clone();
171 let (limit, ctype) = req
172 .app_state::<JsonConfig>()
173 .map_or((32768, None), |c| (c.limit, c.content_type.as_ref()));
174
175 match JsonBody::new(req, payload, ctype).limit(limit).await {
176 Err(e) => {
177 log::debug!(
178 "Failed to deserialize Json from payload. \
179 Request path: {}",
180 req2.path()
181 );
182 Err(e)
183 }
184 Ok(data) => Ok(Json(data)),
185 }
186 }
187}
188
189#[derive(Clone)]
223pub struct JsonConfig {
224 limit: usize,
225 content_type: Option<Arc<dyn Fn(mime::Mime) -> bool + Send + Sync>>,
226}
227
228impl JsonConfig {
229 #[must_use]
230 pub fn limit(mut self, limit: usize) -> Self {
234 self.limit = limit;
235 self
236 }
237
238 #[must_use]
239 pub fn content_type<F>(mut self, predicate: F) -> Self
244 where
245 F: Fn(mime::Mime) -> bool + Send + Sync + 'static,
246 {
247 self.content_type = Some(Arc::new(predicate));
248 self
249 }
250}
251
252impl Default for JsonConfig {
253 fn default() -> Self {
254 JsonConfig {
255 limit: 32768,
256 content_type: None,
257 }
258 }
259}
260
261impl fmt::Debug for JsonConfig {
262 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
263 f.debug_struct("JsonConfig")
264 .field("limit", &self.limit)
265 .field(
266 "content_type",
267 &self
268 .content_type
269 .as_ref()
270 .map(|_| "Arc<dyn Fn(mime::Mime) -> bool + Send + Sync>"),
271 )
272 .finish()
273 }
274}
275
276struct JsonBody<U> {
285 limit: usize,
286 length: Option<usize>,
287 #[cfg(feature = "compress")]
288 stream: Option<Decoder<Payload>>,
289 #[cfg(not(feature = "compress"))]
290 stream: Option<Payload>,
291 err: Option<JsonPayloadError>,
292 fut: Option<BoxFuture<'static, Result<U, JsonPayloadError>>>,
293}
294
295impl<U> JsonBody<U>
296where
297 U: DeserializeOwned + 'static,
298{
299 fn new(
301 req: &HttpRequest,
302 payload: &mut Payload,
303 ctype: Option<&Arc<dyn Fn(mime::Mime) -> bool + Send + Sync>>,
304 ) -> Self {
305 let json = if let Ok(Some(mime)) = req.mime_type() {
307 mime.subtype() == mime::JSON
308 || mime.suffix() == Some(mime::JSON)
309 || ctype.as_ref().is_some_and(|predicate| predicate(mime))
310 } else {
311 false
312 };
313
314 if !json {
315 return JsonBody {
316 limit: 262_144,
317 length: None,
318 stream: None,
319 fut: None,
320 err: Some(JsonPayloadError::ContentType),
321 };
322 }
323
324 let len = match req.headers().get(&CONTENT_LENGTH).map(|l| {
325 l.to_str()
326 .ok()
327 .and_then(|s| s.parse::<usize>().ok())
328 .ok_or(PayloadError::UnknownLength)
329 }) {
330 None => None,
331 Some(Ok(len)) => Some(len),
332 Some(Err(e)) => {
333 return JsonBody {
334 limit: 262_144,
335 length: None,
336 stream: None,
337 fut: None,
338 err: Some(JsonPayloadError::Payload(e)),
339 };
340 }
341 };
342
343 #[cfg(feature = "compress")]
344 let payload = Decoder::from_headers(payload.take(), req.headers());
345 #[cfg(not(feature = "compress"))]
346 let payload = payload.take();
347
348 JsonBody {
349 limit: 262_144,
350 length: len,
351 stream: Some(payload),
352 fut: None,
353 err: None,
354 }
355 }
356
357 fn limit(mut self, limit: usize) -> Self {
359 self.limit = limit;
360 self
361 }
362}
363
364impl<U> Future for JsonBody<U>
365where
366 U: DeserializeOwned + 'static,
367{
368 type Output = Result<U, JsonPayloadError>;
369
370 fn poll(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
371 if let Some(ref mut fut) = self.fut {
372 return Pin::new(fut).poll(cx);
373 }
374
375 if let Some(err) = self.err.take() {
376 return Poll::Ready(Err(err));
377 }
378
379 let (limit, length) = (self.limit, self.length);
380 if let Some(len) = length
381 && len > limit
382 {
383 return Poll::Ready(Err(JsonPayloadError::Overflow));
384 }
385 let mut stream = self.stream.take().unwrap();
386
387 self.fut = Some(Box::pin(async move {
388 let body = super::read_body(&mut stream, limit, length, |_| JsonPayloadError::Overflow)
389 .await?;
390 Ok(serde_json::from_slice::<U>(&body)?)
391 }));
392
393 self.poll(cx)
394 }
395}
396
397#[cfg(test)]
398mod tests {
399 use super::*;
400 use crate::http::header;
401 use crate::util::Bytes;
402 use crate::web::test::{TestRequest, from_request, respond_to};
403
404 #[derive(serde::Serialize, serde::Deserialize, PartialEq, Debug, thiserror::Error)]
405 #[error("MyObject({name})")]
406 struct MyObject {
407 name: String,
408 }
409
410 fn json_eq(err: &JsonPayloadError, other: &JsonPayloadError) -> bool {
411 if let JsonPayloadError::Overflow = err
412 && let JsonPayloadError::Overflow = other
413 {
414 return true;
415 } else if let JsonPayloadError::ContentType = err
416 && let JsonPayloadError::ContentType = other
417 {
418 return true;
419 }
420 false
421 }
422
423 #[test]
424 fn test_json() {
425 let mut j = Json(MyObject {
426 name: "test2".to_string(),
427 });
428 assert_eq!(j.name, "test2");
429 j.name = "test".to_string();
430 assert_eq!(j.name, "test");
431 assert!(format!("{j:?}").contains("Json"));
432 assert!(format!("{j}").contains("test"));
433
434 let cfg = JsonConfig::default().content_type(|mime: mime::Mime| {
435 mime.type_() == mime::TEXT && mime.subtype() == mime::PLAIN
436 });
437 assert!(format!("{cfg:?}").contains("JsonConfig"));
438 }
439
440 #[crate::rt_test]
441 async fn test_responder() {
442 let req = TestRequest::default().to_http_request();
443
444 let j = Json(MyObject {
445 name: "test".to_string(),
446 });
447 let resp = respond_to(j, &req).await;
448 assert_eq!(resp.status(), StatusCode::OK);
449 assert_eq!(
450 resp.headers().get(header::CONTENT_TYPE).unwrap(),
451 header::HeaderValue::from_static("application/json")
452 );
453
454 assert_eq!(resp.get_body_ref(), b"{\"name\":\"test\"}");
455 }
456
457 #[crate::rt_test]
458 async fn test_responder_serialize_error() {
459 struct Invalid;
460
461 impl Serialize for Invalid {
462 fn serialize<S: serde::Serializer>(&self, _: S) -> Result<S::Ok, S::Error> {
463 Err(serde::ser::Error::custom("invalid"))
464 }
465 }
466
467 let req = TestRequest::default().to_http_request();
468 let resp = respond_to(Json(Invalid), &req).await;
469 assert_eq!(resp.status(), StatusCode::INTERNAL_SERVER_ERROR);
470 }
471
472 #[crate::rt_test]
473 async fn test_extract() {
474 let (req, mut pl, ()) = TestRequest::default()
475 .header(
476 header::CONTENT_TYPE,
477 header::HeaderValue::from_static("application/json"),
478 )
479 .header(
480 header::CONTENT_LENGTH,
481 header::HeaderValue::from_static("16"),
482 )
483 .payload(Bytes::from_static(b"{\"name\": \"test\"}"))
484 .to_http_parts();
485
486 let s = from_request::<_, Json<MyObject>>(&(), &req, &mut pl)
487 .await
488 .unwrap();
489 assert_eq!(s.name, "test");
490 assert_eq!(
491 s.into_inner(),
492 MyObject {
493 name: "test".to_string()
494 }
495 );
496
497 let (req, mut pl, ()) = TestRequest::default()
498 .header(
499 header::CONTENT_TYPE,
500 header::HeaderValue::from_static("application/json"),
501 )
502 .header(
503 header::CONTENT_LENGTH,
504 header::HeaderValue::from_static("16"),
505 )
506 .payload(Bytes::from_static(b"{\"name\": \"test\"}"))
507 .app_state(JsonConfig::default().limit(10))
508 .to_http_parts();
509
510 let s = from_request::<_, Json<MyObject>>(&(), &req, &mut pl).await;
511 assert!(
512 format!("{}", s.err().unwrap()).contains("Json payload size is bigger than allowed")
513 );
514
515 let (req, mut pl, ()) = TestRequest::default()
516 .header(
517 header::CONTENT_TYPE,
518 header::HeaderValue::from_static("application/json"),
519 )
520 .header(
521 header::CONTENT_LENGTH,
522 header::HeaderValue::from_static("16"),
523 )
524 .payload(Bytes::from_static(b"--name-: -test--"))
525 .to_http_parts();
526 let s = from_request::<_, Json<serde_json::Value>>(&(), &req, &mut pl).await;
527 assert!(format!("{:?}", s.err().unwrap()).contains("Deserialize(Error("));
528 }
529
530 #[crate::rt_test]
531 async fn test_json_body() {
532 let (req, mut pl, ()) = TestRequest::default().to_http_parts();
533 let json = JsonBody::<MyObject>::new(&req, &mut pl, None).await;
534 assert!(json_eq(
535 &json.err().unwrap(),
536 &JsonPayloadError::ContentType
537 ));
538
539 let (req, mut pl, ()) = TestRequest::default()
540 .header(
541 header::CONTENT_TYPE,
542 header::HeaderValue::from_static("application/text"),
543 )
544 .to_http_parts();
545 let json = JsonBody::<MyObject>::new(&req, &mut pl, None).await;
546 assert!(json_eq(
547 &json.err().unwrap(),
548 &JsonPayloadError::ContentType
549 ));
550
551 let (req, mut pl, ()) = TestRequest::default()
552 .header(
553 header::CONTENT_TYPE,
554 header::HeaderValue::from_static("application/json"),
555 )
556 .header(
557 header::CONTENT_LENGTH,
558 header::HeaderValue::from_static("10000"),
559 )
560 .to_http_parts();
561
562 let json = JsonBody::<MyObject>::new(&req, &mut pl, None)
563 .limit(100)
564 .await;
565 assert!(json_eq(&json.err().unwrap(), &JsonPayloadError::Overflow));
566
567 let (req, mut pl, ()) = TestRequest::default()
568 .header(
569 header::CONTENT_TYPE,
570 header::HeaderValue::from_static("application/json"),
571 )
572 .header(
573 header::CONTENT_LENGTH,
574 header::HeaderValue::from_static("16"),
575 )
576 .payload(Bytes::from_static(b"{\"name\": \"test\"}"))
577 .to_http_parts();
578
579 let json = JsonBody::<MyObject>::new(&req, &mut pl, None).await;
580 assert_eq!(
581 json.ok().unwrap(),
582 MyObject {
583 name: "test".to_owned()
584 }
585 );
586
587 let (req, mut pl, ()) = TestRequest::default()
588 .header(
589 header::CONTENT_TYPE,
590 header::HeaderValue::from_static("application/json"),
591 )
592 .header(
593 header::CONTENT_LENGTH,
594 header::HeaderValue::from_static("16x"),
595 )
596 .payload(Bytes::from_static(b"{\"name\": \"test\"}"))
597 .to_http_parts();
598
599 let json = JsonBody::<MyObject>::new(&req, &mut pl, None).await;
600 assert!(matches!(
601 json.err().unwrap(),
602 JsonPayloadError::Payload(PayloadError::UnknownLength)
603 ));
604 }
605
606 #[crate::rt_test]
607 async fn test_with_json_and_bad_content_type() {
608 let (req, mut pl, ()) = TestRequest::with_header(
609 header::CONTENT_TYPE,
610 header::HeaderValue::from_static("text/plain"),
611 )
612 .header(
613 header::CONTENT_LENGTH,
614 header::HeaderValue::from_static("16"),
615 )
616 .payload(Bytes::from_static(b"{\"name\": \"test\"}"))
617 .app_state(JsonConfig::default().limit(4096))
618 .to_http_parts();
619
620 let s = from_request::<_, Json<MyObject>>(&(), &req, &mut pl).await;
621 assert!(s.is_err());
622 }
623
624 #[crate::rt_test]
625 async fn test_with_json_and_good_custom_content_type() {
626 let (req, mut pl, ()) = TestRequest::with_header(
627 header::CONTENT_TYPE,
628 header::HeaderValue::from_static("text/plain"),
629 )
630 .header(
631 header::CONTENT_LENGTH,
632 header::HeaderValue::from_static("16"),
633 )
634 .payload(Bytes::from_static(b"{\"name\": \"test\"}"))
635 .app_state(JsonConfig::default().content_type(|mime: mime::Mime| {
636 mime.type_() == mime::TEXT && mime.subtype() == mime::PLAIN
637 }))
638 .to_http_parts();
639
640 let s = from_request::<_, Json<MyObject>>(&(), &req, &mut pl).await;
641 assert!(s.is_ok());
642 }
643
644 #[crate::rt_test]
645 async fn test_with_json_and_bad_custom_content_type() {
646 let (req, mut pl, ()) = TestRequest::with_header(
647 header::CONTENT_TYPE,
648 header::HeaderValue::from_static("text/html"),
649 )
650 .header(
651 header::CONTENT_LENGTH,
652 header::HeaderValue::from_static("16"),
653 )
654 .payload(Bytes::from_static(b"{\"name\": \"test\"}"))
655 .app_state(JsonConfig::default().content_type(|mime: mime::Mime| {
656 mime.type_() == mime::TEXT && mime.subtype() == mime::PLAIN
657 }))
658 .to_http_parts();
659
660 let s = from_request::<_, Json<MyObject>>(&(), &req, &mut pl).await;
661 assert!(s.is_err());
662 }
663}