1use std::cell::{Cell, Ref, RefMut};
2use std::task::{Context, Poll, ready};
3use std::{fmt, future::Future, marker::PhantomData, pin::Pin};
4
5use serde::de::DeserializeOwned;
6
7#[cfg(feature = "cookie")]
8use coo_kie::{Cookie, ParseError as CookieParseError};
9
10use crate::Cfg;
11use crate::error::Error;
12use crate::http::header::{AsName, CONTENT_LENGTH, HeaderValue};
13use crate::http::{HeaderMap, HttpMessage, Payload, ResponseHead, StatusCode, Version};
14use crate::http::{error::PayloadError, helpers::take_trimmed};
15use crate::time::{Deadline, Millis};
16use crate::util::{Bytes, BytesMut, Extensions, Stream};
17
18use super::error::{ClientPayloadError, JsonPayloadError};
19use super::{ClientConfig, ServiceResponse};
20
21pub struct ClientResponse {
23 pub(crate) head: ResponseHead,
24 pub(crate) payload: Cell<Option<Payload>>,
25 pub(crate) config: Cfg<ClientConfig>,
26}
27
28impl HttpMessage for ClientResponse {
29 fn message_headers(&self) -> &HeaderMap {
30 &self.head.headers
31 }
32
33 fn message_extensions(&self) -> Ref<'_, Extensions> {
34 self.head.extensions()
35 }
36
37 fn message_extensions_mut(&self) -> RefMut<'_, Extensions> {
38 self.head.extensions_mut()
39 }
40
41 #[cfg(feature = "cookie")]
42 fn cookies(&self) -> Result<Ref<'_, Vec<Cookie<'static>>>, CookieParseError> {
44 use crate::http::header::SET_COOKIE;
45
46 struct Cookies(Vec<Cookie<'static>>);
47
48 if self.message_extensions().get::<Cookies>().is_none() {
49 let mut cookies = Vec::new();
50 for hdr in self.message_headers().get_all(&SET_COOKIE) {
51 let s = std::str::from_utf8(hdr.as_bytes()).map_err(CookieParseError::from)?;
52 cookies.push(Cookie::parse_encoded(s)?.into_owned());
53 }
54 self.message_extensions_mut().insert(Cookies(cookies));
55 }
56 Ok(Ref::map(self.message_extensions(), |ext| {
57 &ext.get::<Cookies>().unwrap().0
58 }))
59 }
60}
61
62impl ClientResponse {
63 #[doc(hidden)]
65 pub fn new(head: ResponseHead, payload: Payload, config: Cfg<ClientConfig>) -> Self {
66 ClientResponse {
67 head,
68 config,
69 payload: Cell::new(Some(payload)),
70 }
71 }
72
73 #[cfg(feature = "ws")]
74 pub(crate) fn with_empty_payload(head: ResponseHead, config: Cfg<ClientConfig>) -> Self {
75 ClientResponse::new(head, Payload::None, config)
76 }
77
78 #[inline]
79 pub(crate) fn head(&self) -> &ResponseHead {
80 &self.head
81 }
82
83 #[inline]
84 pub(crate) fn head_mut(&mut self) -> &mut ResponseHead {
85 &mut self.head
86 }
87
88 #[inline]
90 pub fn version(&self) -> Version {
91 self.head().version
92 }
93
94 #[inline]
96 pub fn status(&self) -> StatusCode {
97 self.head().status
98 }
99
100 #[inline]
101 pub fn header<N: AsName>(&self, name: N) -> Option<&HeaderValue> {
103 self.head().headers.get(name)
104 }
105
106 #[inline]
107 pub fn headers(&self) -> &HeaderMap {
109 &self.head().headers
110 }
111
112 #[inline]
113 pub fn headers_mut(&mut self) -> &mut HeaderMap {
115 &mut self.head_mut().headers
116 }
117
118 pub fn set_payload(&self, payload: Payload) {
122 self.payload.set(Some(payload));
123 }
124
125 #[must_use]
126 pub fn take_payload(&self) -> Payload {
130 if let Some(pl) = self.payload.take() {
131 pl
132 } else {
133 Payload::None
134 }
135 }
136
137 #[inline]
139 pub fn extensions(&self) -> Ref<'_, Extensions> {
140 self.head().extensions()
141 }
142
143 #[inline]
145 pub fn extensions_mut(&self) -> RefMut<'_, Extensions> {
146 self.head().extensions_mut()
147 }
148}
149
150impl ClientResponse {
151 pub fn body(&self) -> MessageBody {
153 MessageBody::new(self)
154 }
155
156 pub fn json<T: DeserializeOwned>(&self) -> JsonBody<T> {
164 JsonBody::new(self)
165 }
166}
167
168impl Stream for ClientResponse {
169 type Item = Result<Bytes, Error<ClientPayloadError>>;
170
171 fn poll_next(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
172 if let Some(mut pl) = self.payload.take() {
173 let result = Pin::new(&mut pl).poll_next(cx);
174 self.payload.set(Some(pl));
175 Poll::Ready(
176 ready!(result).map(|item| item.map_err(|e| Error::from(ClientPayloadError(e)))),
177 )
178 } else {
179 Poll::Ready(None)
180 }
181 }
182}
183
184impl From<ServiceResponse> for ClientResponse {
185 fn from(res: ServiceResponse) -> Self {
186 Self {
187 head: res.head,
188 payload: Cell::new(Some(res.payload)),
189 config: res.config,
190 }
191 }
192}
193
194impl fmt::Debug for ClientResponse {
195 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
196 writeln!(f, "\nClientResponse {:?} {}", self.version(), self.status())?;
197 writeln!(f, " headers:")?;
198 for (key, val) in self.headers() {
199 writeln!(f, " {key:?}: {val:?}")?;
200 }
201 Ok(())
202 }
203}
204
205#[derive(Debug)]
206pub struct MessageBody {
208 length: Option<usize>,
209 err: Option<Error<ClientPayloadError>>,
210 fut: Option<ReadBody>,
211 config: Cfg<ClientConfig>,
212}
213
214impl MessageBody {
215 pub fn new(res: &ClientResponse) -> MessageBody {
220 let config = res.config.clone();
221
222 let len = match content_length(res) {
223 Ok(len) => len,
224 Err(e) => return Self::err(Error::from(e).with_service(config.service()), config),
225 };
226
227 MessageBody {
228 config,
229 length: len,
230 err: None,
231 fut: Some(ReadBody::new(
232 res.take_payload(),
233 res.config.response_payload_limit(),
234 res.config.response_payload_timeout(),
235 res.config.clone(),
236 )),
237 }
238 }
239
240 #[must_use]
241 pub fn limit(mut self, limit: usize) -> Self {
245 if let Some(ref mut fut) = self.fut {
246 fut.limit = limit;
247 }
248 self
249 }
250
251 #[must_use]
252 pub fn timeout<T: Into<Millis>>(mut self, to: T) -> Self {
256 if let Some(ref mut fut) = self.fut {
257 fut.timeout.reset(to.into());
258 }
259 self
260 }
261
262 fn err(e: Error<ClientPayloadError>, config: Cfg<ClientConfig>) -> Self {
263 MessageBody {
264 config,
265 fut: None,
266 err: Some(e),
267 length: None,
268 }
269 }
270}
271
272impl Future for MessageBody {
273 type Output = Result<Bytes, Error<ClientPayloadError>>;
274
275 fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
276 let this = self.get_mut();
277
278 if let Some(err) = this.err.take() {
279 return Poll::Ready(Err(err));
280 }
281
282 if let Some(len) = this.length.take() {
283 let limit = this.fut.as_ref().unwrap().limit;
284 if limit > 0 && len > limit {
285 return Poll::Ready(Err(Error::from(ClientPayloadError(PayloadError::Overflow))
286 .with_service(this.config.service())));
287 }
288 }
289
290 Pin::new(&mut *this.fut.as_mut().unwrap())
291 .poll(cx)
292 .map_err(|e| e.map(ClientPayloadError::from))
293 }
294}
295
296#[derive(Debug)]
297pub struct JsonBody<U> {
305 length: Option<usize>,
306 err: Option<Error<JsonPayloadError>>,
307 fut: Option<ReadBody>,
308 config: Cfg<ClientConfig>,
309 _t: PhantomData<U>,
310}
311
312impl<U> JsonBody<U>
313where
314 U: DeserializeOwned,
315{
316 #[must_use]
317 pub fn new(res: &ClientResponse) -> Self {
322 let config = res.config.clone();
323
324 let json = if let Ok(Some(mime)) = res.mime_type() {
326 mime.subtype() == mime::JSON || mime.suffix() == Some(mime::JSON)
327 } else {
328 false
329 };
330 if !json {
331 let err =
332 Some(Error::from(JsonPayloadError::ContentType).with_service(config.service()));
333 return JsonBody {
334 err,
335 config,
336 length: None,
337 fut: None,
338 _t: PhantomData,
339 };
340 }
341
342 let len = match content_length(res) {
343 Ok(len) => len,
344 Err(e) => {
345 return JsonBody {
346 err: Some(
347 Error::from(JsonPayloadError::Payload(e)).with_service(config.service()),
348 ),
349 config,
350 length: None,
351 fut: None,
352 _t: PhantomData,
353 };
354 }
355 };
356
357 JsonBody {
358 config,
359 length: len,
360 err: None,
361 fut: Some(ReadBody::new(
362 res.take_payload(),
363 res.config.response_payload_limit(),
364 res.config.response_payload_timeout(),
365 res.config.clone(),
366 )),
367 _t: PhantomData,
368 }
369 }
370
371 #[must_use]
372 pub fn limit(mut self, limit: usize) -> Self {
376 if let Some(ref mut fut) = self.fut {
377 fut.limit = limit;
378 }
379 self
380 }
381
382 #[must_use]
383 pub fn timeout<T: Into<Millis>>(mut self, to: T) -> Self {
387 if let Some(ref mut fut) = self.fut {
388 fut.timeout.reset(to.into());
389 }
390 self
391 }
392}
393
394impl<U> Unpin for JsonBody<U> where U: DeserializeOwned {}
395
396impl<U> Future for JsonBody<U>
397where
398 U: DeserializeOwned,
399{
400 type Output = Result<U, Error<JsonPayloadError>>;
401
402 fn poll(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
403 if let Some(err) = self.err.take() {
404 return Poll::Ready(Err(err));
405 }
406
407 if let Some(len) = self.length.take() {
408 let limit = self.fut.as_ref().unwrap().limit;
409 if limit > 0 && len > limit {
410 return Poll::Ready(Err(Error::from(JsonPayloadError::Payload(
411 ClientPayloadError(PayloadError::Overflow),
412 ))
413 .with_service(self.config.service())));
414 }
415 }
416
417 let this = self.get_mut();
418 let body = match Pin::new(&mut *this.fut.as_mut().unwrap()).poll(cx) {
419 Poll::Ready(result) => result.map_err(|e| e.map(JsonPayloadError::from))?,
420 Poll::Pending => return Poll::Pending,
421 };
422 Poll::Ready(serde_json::from_slice::<U>(&body).map_err(|e| {
423 Error::from(JsonPayloadError::from(e)).with_service(this.config.service())
424 }))
425 }
426}
427
428fn content_length(res: &ClientResponse) -> Result<Option<usize>, ClientPayloadError> {
430 res.headers()
431 .get(&CONTENT_LENGTH)
432 .map(|l| {
433 l.to_str()
434 .ok()
435 .and_then(|s| s.parse::<usize>().ok())
436 .ok_or(ClientPayloadError(PayloadError::UnknownLength))
437 })
438 .transpose()
439}
440
441#[derive(Debug)]
442struct ReadBody {
443 stream: Payload,
444 buf: BytesMut,
445 limit: usize,
446 timeout: Deadline,
447 config: Cfg<ClientConfig>,
448}
449
450impl ReadBody {
451 fn new(stream: Payload, limit: usize, timeout: Millis, config: Cfg<ClientConfig>) -> Self {
452 Self {
453 stream,
454 limit,
455 config,
456 buf: BytesMut::with_capacity(std::cmp::min(limit, 32768)),
457 timeout: Deadline::new(timeout),
458 }
459 }
460}
461
462impl Future for ReadBody {
463 type Output = Result<Bytes, Error<ClientPayloadError>>;
464
465 fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
466 let this = self.get_mut();
467
468 loop {
469 return match Pin::new(&mut this.stream).poll_next(cx) {
470 Poll::Ready(Some(Ok(chunk))) => {
471 if this.limit > 0 && (this.buf.len() + chunk.len()) > this.limit {
472 Poll::Ready(Err(Error::from(ClientPayloadError(PayloadError::Overflow))
473 .with_service(this.config.service())))
474 } else {
475 this.buf.extend_from_slice(&chunk);
476 continue;
477 }
478 }
479 Poll::Ready(None) => Poll::Ready(Ok(take_trimmed(&mut this.buf))),
480 Poll::Ready(Some(Err(err))) => Poll::Ready(Err(Error::from(ClientPayloadError(
481 err,
482 ))
483 .with_service(this.config.service()))),
484 Poll::Pending => {
485 if this.timeout.poll_elapsed(cx).is_ready() {
486 Poll::Ready(Err(Error::from(ClientPayloadError(
487 PayloadError::Incomplete(Some(std::io::Error::new(
488 std::io::ErrorKind::TimedOut,
489 "Operation timed out",
490 ))),
491 ))
492 .with_service(this.config.service())))
493 } else {
494 Poll::Pending
495 }
496 }
497 };
498 }
499 }
500}
501
502#[cfg(test)]
503mod tests {
504 use serde::{Deserialize, Serialize};
505
506 use super::*;
507 use crate::client::{error::ClientPayloadError, test::TestResponse};
508 use crate::http::header;
509
510 #[crate::rt_test]
511 async fn test_body() {
512 let req = TestResponse::with_header(header::CONTENT_LENGTH, "xxxx").build();
513 match &*req.body().await.err().unwrap().into_error() {
514 PayloadError::UnknownLength => (),
515 _ => unreachable!("error"),
516 }
517
518 let req = TestResponse::with_header(header::CONTENT_LENGTH, "1000000").build();
519 match &*req.body().await.err().unwrap().into_error() {
520 PayloadError::Overflow => (),
521 _ => unreachable!("error"),
522 }
523
524 let req = TestResponse::builder()
525 .set_payload(Bytes::from_static(b"test"))
526 .build();
527 assert_eq!(req.body().await.ok().unwrap(), Bytes::from_static(b"test"));
528
529 let req = TestResponse::builder()
530 .set_payload(Bytes::from_static(b"11111111111111"))
531 .build();
532 match &*req.body().limit(5).await.err().unwrap().into_error() {
533 PayloadError::Overflow => (),
534 _ => unreachable!("error"),
535 }
536
537 let req = TestResponse::builder()
539 .set_payload(Bytes::from(vec![b'x'; 1000]))
540 .build();
541 let mut body = req.body().await.ok().unwrap();
542 let ptr = body.as_ptr();
543 body.trimdown();
544 assert_eq!(body, [b'x'; 1000][..]);
545 assert_eq!(body.as_ptr(), ptr, "the body has no unused space");
546 }
547
548 #[derive(Serialize, Deserialize, PartialEq, Debug)]
549 struct MyObject {
550 name: String,
551 }
552
553 fn json_eq(err: &JsonPayloadError, other: &JsonPayloadError) -> bool {
554 match err {
555 JsonPayloadError::Payload(ClientPayloadError(PayloadError::Overflow)) => {
556 matches!(
557 other,
558 JsonPayloadError::Payload(ClientPayloadError(PayloadError::Overflow))
559 )
560 }
561 JsonPayloadError::Payload(ClientPayloadError(PayloadError::UnknownLength)) => {
562 matches!(
563 other,
564 JsonPayloadError::Payload(ClientPayloadError(PayloadError::UnknownLength))
565 )
566 }
567 JsonPayloadError::ContentType => matches!(other, JsonPayloadError::ContentType),
568 _ => false,
569 }
570 }
571
572 #[crate::rt_test]
573 async fn test_json_body() {
574 let req = TestResponse::builder().build();
575 let json = JsonBody::<MyObject>::new(&req).await;
576 assert!(json_eq(
577 &json.err().unwrap(),
578 &JsonPayloadError::ContentType
579 ));
580
581 let req = TestResponse::builder()
582 .header(
583 header::CONTENT_TYPE,
584 header::HeaderValue::from_static("application/text"),
585 )
586 .build();
587 let json = JsonBody::<MyObject>::new(&req).await;
588 assert!(json_eq(
589 &json.err().unwrap(),
590 &JsonPayloadError::ContentType
591 ));
592
593 let req = TestResponse::builder()
594 .header(
595 header::CONTENT_TYPE,
596 header::HeaderValue::from_static("application/json"),
597 )
598 .header(
599 header::CONTENT_LENGTH,
600 header::HeaderValue::from_static("10000"),
601 )
602 .build();
603
604 let json = JsonBody::<MyObject>::new(&req).limit(100).await;
605 assert!(json_eq(
606 &json.err().unwrap(),
607 &JsonPayloadError::Payload(ClientPayloadError(PayloadError::Overflow))
608 ));
609
610 let req = TestResponse::with_header(header::CONTENT_TYPE, "application/json")
611 .header(header::CONTENT_LENGTH, "xxxx")
612 .set_payload(Bytes::from_static(b"{\"name\": \"test\"}"))
613 .build();
614 let json = JsonBody::<MyObject>::new(&req).await;
615 assert!(json_eq(
616 &json.err().unwrap(),
617 &JsonPayloadError::Payload(ClientPayloadError(PayloadError::UnknownLength))
618 ));
619
620 let req = TestResponse::builder()
621 .header(
622 header::CONTENT_TYPE,
623 header::HeaderValue::from_static("application/json"),
624 )
625 .header(
626 header::CONTENT_LENGTH,
627 header::HeaderValue::from_static("16"),
628 )
629 .set_payload(Bytes::from_static(b"{\"name\": \"test\"}"))
630 .build();
631
632 let json = JsonBody::<MyObject>::new(&req).await;
633 assert_eq!(
634 json.ok().unwrap(),
635 MyObject {
636 name: "test".to_owned()
637 }
638 );
639 }
640
641 fn pending_payload() -> Payload {
642 Payload::Stream(Box::pin(futures_util::stream::pending()))
643 }
644
645 fn error_payload() -> Payload {
646 Payload::Stream(Box::pin(futures_util::stream::iter([
647 Ok(Bytes::from_static(b"{")),
648 Err(PayloadError::Incomplete(None)),
649 ])))
650 }
651
652 fn is_timeout(err: &PayloadError) -> bool {
653 matches!(err, PayloadError::Incomplete(Some(e)) if e.kind() == std::io::ErrorKind::TimedOut)
654 }
655
656 #[crate::rt_test]
657 async fn test_body_timeout_and_error() {
658 let res = TestResponse::builder().build();
659 res.set_payload(pending_payload());
660 let err = res.body().timeout(Millis(1)).await.unwrap_err();
661 assert!(is_timeout(&err.into_error().0));
662
663 let res = TestResponse::builder().build();
664 res.set_payload(error_payload());
665 let err = res.body().await.unwrap_err();
666 assert!(matches!(err.into_error().0, PayloadError::Incomplete(None)));
667
668 let res = TestResponse::builder()
670 .set_payload(b"data".as_ref())
671 .build();
672 let body = res.body();
673 assert_eq!(res.body().await.unwrap(), Bytes::new());
674 assert_eq!(body.await.unwrap(), Bytes::from_static(b"data"));
675 }
676
677 #[crate::rt_test]
678 async fn test_json_timeout_and_errors() {
679 let res = TestResponse::with_header(header::CONTENT_TYPE, "application/json").build();
680 res.set_payload(pending_payload());
681 let err = res.json::<MyObject>().timeout(Millis(1)).await.unwrap_err();
682 let JsonPayloadError::Payload(ClientPayloadError(err)) = &*err else {
683 panic!("{err:?}")
684 };
685 assert!(is_timeout(err));
686
687 let res = TestResponse::with_header(header::CONTENT_TYPE, "application/json").build();
688 res.set_payload(error_payload());
689 let err = res.json::<MyObject>().await.unwrap_err();
690 assert!(matches!(
691 &*err,
692 JsonPayloadError::Payload(ClientPayloadError(PayloadError::Incomplete(None)))
693 ));
694
695 let res = TestResponse::with_header(header::CONTENT_TYPE, "application/json")
696 .set_payload(b"{\"name\": 1}".as_ref())
697 .build();
698 let err = res.json::<MyObject>().await.unwrap_err();
699 assert!(matches!(&*err, JsonPayloadError::Deserialize(Some(_))));
700
701 let res = TestResponse::with_header(header::CONTENT_TYPE, "application/problem+json")
703 .set_payload(b"{\"name\": \"test\"}".as_ref())
704 .build();
705 assert_eq!(res.json::<MyObject>().await.unwrap().name, "test");
706 }
707
708 #[crate::rt_test]
709 async fn test_response_stream() {
710 use futures_util::StreamExt;
711
712 let mut res = TestResponse::builder()
713 .set_payload(b"data".as_ref())
714 .build();
715 assert_eq!(res.next().await.unwrap().unwrap(), "data");
716 assert!(res.next().await.is_none());
717
718 let mut res = TestResponse::builder().build();
719 res.set_payload(error_payload());
720 assert_eq!(res.next().await.unwrap().unwrap(), "{");
721 let err = res.next().await.unwrap().unwrap_err();
722 assert!(matches!(err.into_error().0, PayloadError::Incomplete(None)));
723
724 let mut res = TestResponse::builder()
726 .set_payload(b"data".as_ref())
727 .build();
728 assert!(matches!(
729 res.take_payload(),
730 Payload::Stream(_) | Payload::H1(_)
731 ));
732 assert!(matches!(res.take_payload(), Payload::None));
733 assert!(res.next().await.is_none());
734 }
735
736 #[test]
737 fn test_response_accessors() {
738 let mut res = TestResponse::with_header(header::CONTENT_TYPE, "text/plain")
739 .version(Version::HTTP_2)
740 .build();
741 res.headers_mut()
742 .insert(header::SERVER, HeaderValue::from_static("test"));
743 assert_eq!(res.header(header::SERVER).unwrap(), "test");
744 assert_eq!(res.version(), Version::HTTP_2);
745
746 res.extensions_mut().insert(10u32);
747 assert_eq!(res.extensions().get::<u32>(), Some(&10));
748
749 let s = format!("{res:?}");
750 assert!(s.contains("ClientResponse HTTP/2.0 200 OK"), "{s}");
751 assert!(s.contains("\"server\": \"test\""), "{s}");
752 }
753}