1use std::{error::Error as StdError, fmt, net, rc::Rc};
2
3use base64::{Engine, engine::general_purpose::STANDARD as base64};
4#[cfg(feature = "cookie")]
5use coo_kie::{Cookie, CookieJar};
6use serde::Serialize;
7use urly::Url;
8
9use crate::error::Error;
10use crate::http::error::HttpError;
11use crate::http::header::{self, HeaderMap, HeaderName, HeaderValue};
12use crate::http::{ConnectionType, Method, Version, body::Body};
13use crate::{Cfg, PipelineBinding, time::Millis, util::Bytes, util::Stream};
14
15use super::error::{ClientError, InvalidUrl};
16use super::{ClientConfig, ClientResponse, ServiceRequest, ServiceResponse};
17
18pub struct ClientRequest {
41 request: ServiceRequest,
42 svc: PipelineBinding<ServiceRequest, ServiceResponse, Error<ClientError>>,
43 err: Option<ClientError>,
44 cfg: Cfg<ClientConfig>,
45 #[cfg(feature = "cookie")]
46 cookies: Option<CookieJar>,
47}
48
49impl ClientRequest {
50 pub(super) fn new<U>(
52 method: Method,
53 uri: U,
54 cfg: Cfg<ClientConfig>,
55 svc: PipelineBinding<ServiceRequest, ServiceResponse, Error<ClientError>>,
56 ) -> Self
57 where
58 Url: TryFrom<U>,
59 <Url as TryFrom<U>>::Error: Into<InvalidUrl>,
60 {
61 ClientRequest {
62 svc,
63 cfg,
64 request: ServiceRequest::new(),
65 err: None,
66 #[cfg(feature = "cookie")]
67 cookies: None,
68 }
69 .method(method)
70 .uri(uri)
71 }
72
73 #[inline]
77 #[must_use]
78 pub fn uri<U>(mut self, uri: U) -> Self
79 where
80 Url: TryFrom<U>,
81 <Url as TryFrom<U>>::Error: Into<InvalidUrl>,
82 {
83 match Url::try_from(uri) {
84 Ok(uri) => self.request.head.uri = uri,
85 Err(e) => self.err = Some(e.into().into()),
86 }
87 self
88 }
89
90 pub fn get_uri(&self) -> &Url {
92 &self.request.head.uri
93 }
94
95 #[must_use]
96 pub fn address(mut self, addr: net::SocketAddr) -> Self {
100 self.request.addr = Some(addr);
101 self
102 }
103
104 #[inline]
106 #[must_use]
107 pub fn method(mut self, method: Method) -> Self {
108 self.request.head.method = method;
109 self
110 }
111
112 #[inline]
113 #[must_use]
114 pub fn get_method(&self) -> &Method {
116 &self.request.head.method
117 }
118
119 #[inline]
124 #[must_use]
125 pub fn version(mut self, version: Version) -> Self {
126 self.request.head.version = version;
127 self
128 }
129
130 #[inline]
131 pub fn get_version(&self) -> &Version {
133 &self.request.head.version
134 }
135
136 #[inline]
137 pub fn headers(&self) -> &HeaderMap {
139 &self.request.head.headers
140 }
141
142 #[inline]
143 pub fn headers_mut(&mut self) -> &mut HeaderMap {
145 &mut self.request.head.headers
146 }
147
148 #[must_use]
149 pub fn header<K, V>(mut self, key: K, value: V) -> Self
169 where
170 HeaderName: TryFrom<K>,
171 HeaderValue: TryFrom<V>,
172 <HeaderName as TryFrom<K>>::Error: Into<HttpError>,
173 <HeaderValue as TryFrom<V>>::Error: Into<HttpError>,
174 {
175 match HeaderName::try_from(key) {
176 Ok(key) => match HeaderValue::try_from(value) {
177 Ok(value) => self.request.head.headers.append(key, value),
178 Err(e) => self.err = Some(ClientError::Http(e.into())),
179 },
180 Err(e) => self.err = Some(ClientError::Http(e.into())),
181 }
182 self
183 }
184
185 #[must_use]
186 pub fn set_header<K, V>(mut self, key: K, value: V) -> Self
191 where
192 HeaderName: TryFrom<K>,
193 HeaderValue: TryFrom<V>,
194 <HeaderName as TryFrom<K>>::Error: Into<HttpError>,
195 <HeaderValue as TryFrom<V>>::Error: Into<HttpError>,
196 {
197 match HeaderName::try_from(key) {
198 Ok(key) => match HeaderValue::try_from(value) {
199 Ok(value) => self.request.head.headers.insert(key, value),
200 Err(e) => self.err = Some(ClientError::Http(e.into())),
201 },
202 Err(e) => self.err = Some(ClientError::Http(e.into())),
203 }
204 self
205 }
206
207 #[must_use]
208 pub fn set_header_if_none<K, V>(mut self, key: K, value: V) -> Self
214 where
215 HeaderName: TryFrom<K>,
216 HeaderValue: TryFrom<V>,
217 <HeaderName as TryFrom<K>>::Error: Into<HttpError>,
218 <HeaderValue as TryFrom<V>>::Error: Into<HttpError>,
219 {
220 match HeaderName::try_from(key) {
221 Ok(key) => {
222 if !self.request.head.headers.contains_key(&key) {
223 match HeaderValue::try_from(value) {
224 Ok(value) => self.request.head.headers.insert(key, value),
225 Err(e) => self.err = Some(ClientError::Http(e.into())),
226 }
227 }
228 }
229 Err(e) => self.err = Some(ClientError::Http(e.into())),
230 }
231 self
232 }
233
234 #[inline]
235 #[must_use]
236 pub fn set_connection_type(mut self, ctype: ConnectionType) -> Self {
241 self.request.head.set_connection_type(ctype);
242 self
243 }
244
245 #[inline]
249 #[must_use]
250 pub fn force_close(mut self) -> Self {
251 self.request.head.set_connection_type(ConnectionType::Close);
252 self
253 }
254
255 #[inline]
260 #[must_use]
261 pub fn content_type<V>(mut self, value: V) -> Self
262 where
263 HeaderValue: TryFrom<V>,
264 <HeaderValue as TryFrom<V>>::Error: Into<HttpError>,
265 {
266 match HeaderValue::try_from(value) {
267 Ok(value) => self
268 .request
269 .head
270 .headers
271 .insert(header::CONTENT_TYPE, value),
272 Err(e) => self.err = Some(ClientError::Http(e.into())),
273 }
274 self
275 }
276
277 #[inline]
279 #[must_use]
280 pub fn content_length(self, len: u64) -> Self {
281 self.set_header(header::CONTENT_LENGTH, len)
282 }
283
284 #[must_use]
285 pub fn basic_auth<U>(self, username: U, password: Option<&str>) -> Self
289 where
290 U: fmt::Display,
291 {
292 let auth = match password {
293 Some(password) => format!("{username}:{password}"),
294 None => format!("{username}:"),
295 };
296 self.set_header(
297 header::AUTHORIZATION,
298 format!("Basic {}", base64.encode(auth)),
299 )
300 }
301
302 #[must_use]
303 pub fn bearer_auth<T>(self, token: T) -> Self
307 where
308 T: fmt::Display,
309 {
310 self.set_header(header::AUTHORIZATION, format!("Bearer {token}"))
311 }
312
313 #[must_use]
314 #[cfg(feature = "cookie")]
315 pub fn cookie<C>(mut self, cookie: C) -> Self
338 where
339 C: Into<Cookie<'static>>,
340 {
341 if let Some(cookies) = &mut self.cookies {
342 cookies.add(cookie.into());
343 } else {
344 let mut jar = CookieJar::new();
345 jar.add(cookie.into());
346 self.cookies = Some(jar);
347 }
348 self
349 }
350
351 #[must_use]
352 pub fn no_decompress(mut self) -> Self {
354 self.request.response_decompress = false;
355 self
356 }
357
358 #[must_use]
359 pub fn timeout<T: Into<Millis>>(mut self, timeout: T) -> Self {
367 self.request.timeout = Some(timeout.into());
368 self
369 }
370
371 #[must_use]
372 pub fn if_true<F>(self, value: bool, f: F) -> Self
374 where
375 F: FnOnce(ClientRequest) -> ClientRequest,
376 {
377 if value { f(self) } else { self }
378 }
379
380 #[must_use]
381 pub fn if_some<T, F>(self, value: Option<T>, f: F) -> Self
384 where
385 F: FnOnce(T, ClientRequest) -> ClientRequest,
386 {
387 if let Some(val) = value { f(val, self) } else { self }
388 }
389
390 #[must_use]
394 pub fn query<T: Serialize>(mut self, query: &T) -> Self {
395 let query = match serde_urlencoded::to_string(query) {
396 Ok(query) => query,
397 Err(err) => {
398 self.err = Some(ClientError::Error(Rc::new(err)));
399 return self;
400 }
401 };
402
403 self.request.head.uri.set_query(Some(&query));
404 self
405 }
406}
407
408impl ClientRequest {
409 pub async fn send_body<B>(mut self, body: B) -> Result<ClientResponse, Error<ClientError>>
411 where
412 B: Into<Body>,
413 {
414 self.prep_for_sending()?;
415 *self.request.body() = body.into();
416 self.svc.call(self.request).await.map(Into::into)
417 }
418
419 pub async fn send_json<T: Serialize>(
421 mut self,
422 value: &T,
423 ) -> Result<ClientResponse, Error<ClientError>> {
424 self.prep_for_sending()?;
425 self.request.set_json(value)?;
426 self.svc.call(self.request).await.map(Into::into)
427 }
428
429 pub async fn send_form<T: Serialize>(
431 mut self,
432 value: &T,
433 ) -> Result<ClientResponse, Error<ClientError>> {
434 self.prep_for_sending()?;
435 self.request.set_form(value)?;
436 self.svc.call(self.request).await.map(Into::into)
437 }
438
439 pub async fn send_stream<T, E>(
441 mut self,
442 stream: T,
443 ) -> Result<ClientResponse, Error<ClientError>>
444 where
445 T: Stream<Item = Result<Bytes, E>> + Unpin + 'static,
446 E: StdError + 'static,
447 {
448 self.prep_for_sending()?;
449 self.request.set_stream(stream);
450 self.svc.call(self.request).await.map(Into::into)
451 }
452
453 pub async fn send(mut self) -> Result<ClientResponse, Error<ClientError>> {
455 self.prep_for_sending()?;
456 self.svc.call(self.request).await.map(Into::into)
457 }
458
459 fn prep_for_sending(&mut self) -> Result<(), Error<ClientError>> {
460 self.prep_for_sending_inner()
461 .map_err(|e| e.with_service(self.cfg.service()))
462 }
463
464 fn prep_for_sending_inner(&mut self) -> Result<(), Error<ClientError>> {
465 if let Some(e) = self.err.take() {
466 return Err(e.into());
467 }
468
469 let uri = &self.request.head.uri;
471 if uri.host().is_none() {
472 return Err(ClientError::from(InvalidUrl::MissingHost).into());
473 }
474 match uri.scheme_str() {
475 Some("http" | "ws" | "https" | "wss") => (),
476 Some(_) => return Err(ClientError::from(InvalidUrl::UnknownScheme).into()),
477 None => return Err(ClientError::from(InvalidUrl::MissingScheme).into()),
478 }
479
480 #[cfg(feature = "cookie")]
482 {
483 if let Some(ref jar) = self.cookies {
484 let headers = &mut self.request.head.headers;
485 let mut cookie = headers
486 .get(header::COOKIE)
487 .map(|v| v.as_bytes().to_vec())
488 .unwrap_or_default();
489 for c in jar.iter() {
490 crate::http::helpers::push_cookie(&mut cookie, c.name(), c.value());
491 }
492 if let Ok(val) = HeaderValue::from_bytes(&cookie) {
493 headers.insert(header::COOKIE, val);
494 }
495 }
496 }
497
498 #[cfg(feature = "compress")]
499 if self.request.response_decompress
500 && !self
501 .request
502 .head
503 .headers
504 .contains_key(&header::ACCEPT_ENCODING)
505 {
506 const COMPRESSION: HeaderValue = HeaderValue::from_static("gzip, deflate, zstd");
507 self.request
508 .head
509 .headers
510 .insert(header::ACCEPT_ENCODING, COMPRESSION);
511 }
512
513 Ok(())
514 }
515}
516
517impl fmt::Debug for ClientRequest {
518 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
519 writeln!(
520 f,
521 "\nClientRequest {:?} {}:{}",
522 self.request.head.version, self.request.head.method, self.request.head.uri
523 )?;
524 writeln!(f, " headers:")?;
525 for (key, val) in &self.request.head.headers {
526 if key == header::AUTHORIZATION {
527 writeln!(f, " {key:?}: <REDACTED>")?;
528 } else {
529 writeln!(f, " {key:?}: {val:?}")?;
530 }
531 }
532 Ok(())
533 }
534}
535
536#[cfg(test)]
537mod tests {
538 use super::*;
539 use crate::{SharedCfg, client::Client};
540
541 struct InvalidQuery;
542
543 impl Serialize for InvalidQuery {
544 fn serialize<S>(&self, _: S) -> Result<S::Ok, S::Error>
545 where
546 S: serde::Serializer,
547 {
548 Err(serde::ser::Error::custom("invalid query"))
549 }
550 }
551
552 #[crate::rt_test]
553 async fn test_debug() {
554 let request = Client::new().get("/").header("x-test", "111");
555 let repr = format!("{request:?}");
556 assert!(repr.contains("ClientRequest"));
557 assert!(repr.contains("x-test"));
558 }
559
560 #[crate::rt_test]
561 async fn test_basics() {
562 let mut req = Client::new()
563 .put("/")
564 .version(Version::HTTP_2)
565 .header(header::DATE, "data")
566 .content_type("plain/text")
567 .if_true(true, |req| req.header(header::SERVER, "awc"))
568 .if_true(false, |req| req.header(header::EXPECT, "awc"))
569 .if_some(Some("server"), |val, req| {
570 req.header(header::USER_AGENT, val)
571 })
572 .if_some(Option::<&str>::None, |_, req| {
573 req.header(header::ALLOW, "1")
574 })
575 .content_length(100);
576 assert!(req.headers().contains_key(header::CONTENT_TYPE));
577 assert!(req.headers().contains_key(header::DATE));
578 assert!(req.headers().contains_key(header::SERVER));
579 assert!(req.headers().contains_key(header::USER_AGENT));
580 assert!(!req.headers().contains_key(header::ALLOW));
581 assert!(!req.headers().contains_key(header::EXPECT));
582 assert_eq!(req.request.head.version, Version::HTTP_2);
583 assert_eq!(req.get_version(), &Version::HTTP_2);
584 assert_eq!(req.get_method(), Method::PUT);
585 let _ = req.headers_mut();
586 let _ = req.send_body("").await;
587 }
588
589 #[cfg(feature = "cookie")]
590 #[crate::rt_test]
591 async fn cookies_extend_cookie_header() {
592 use coo_kie::Cookie;
593
594 fn cookie_header(req: &mut ClientRequest) -> String {
595 req.prep_for_sending_inner().unwrap();
596 let val = req.request.head.headers.get(header::COOKIE).unwrap();
597 val.to_str().unwrap().to_string()
598 }
599
600 let mut req = Client::new()
601 .get("http://localhost/")
602 .cookie(Cookie::build(("c1", "v1")))
603 .cookie(Cookie::build(("c2", "v2")));
604 let cookie = cookie_header(&mut req);
605 let mut cookies: Vec<_> = cookie.split("; ").collect();
606 cookies.sort_unstable();
607 assert_eq!(cookies, ["c1=v1", "c2=v2"]);
608
609 let mut req = Client::new()
610 .get("http://localhost/")
611 .header(header::COOKIE, "c0=v0")
612 .cookie(Cookie::build(("c1", "v1")));
613 assert_eq!(cookie_header(&mut req), "c0=v0; c1=v1");
614 }
615
616 #[crate::rt_test]
617 async fn test_client_header() {
618 let req = Client::builder()
619 .build(
620 SharedCfg::new("H").add(
621 ClientConfig::new()
622 .set_header(header::CONTENT_TYPE, "111")
623 .unwrap(),
624 ),
625 )
626 .get("/");
627
628 assert_eq!(
629 req.request
630 .head
631 .headers
632 .get(header::CONTENT_TYPE)
633 .unwrap()
634 .to_str()
635 .unwrap(),
636 "111"
637 );
638 }
639
640 #[crate::rt_test]
641 async fn test_client_header_override() {
642 let req = Client::builder()
643 .build(
644 SharedCfg::new("H").add(
645 ClientConfig::new()
646 .set_header(header::CONTENT_TYPE, "111")
647 .unwrap(),
648 ),
649 )
650 .get("/")
651 .set_header(header::CONTENT_TYPE, "222");
652
653 assert_eq!(
654 req.request
655 .head
656 .headers
657 .get(header::CONTENT_TYPE)
658 .unwrap()
659 .to_str()
660 .unwrap(),
661 "222"
662 );
663 }
664
665 #[crate::rt_test]
666 async fn client_basic_auth() {
667 let req = Client::new()
668 .get("/")
669 .basic_auth("username", Some("password"));
670 assert_eq!(
671 req.request
672 .head
673 .headers
674 .get(header::AUTHORIZATION)
675 .unwrap()
676 .to_str()
677 .unwrap(),
678 "Basic dXNlcm5hbWU6cGFzc3dvcmQ="
679 );
680
681 let req = Client::new().get("/").basic_auth("username", None);
682 assert_eq!(
683 req.request
684 .head
685 .headers
686 .get(header::AUTHORIZATION)
687 .unwrap()
688 .to_str()
689 .unwrap(),
690 "Basic dXNlcm5hbWU6"
691 );
692 }
693
694 #[crate::rt_test]
695 async fn client_bearer_auth() {
696 let req = Client::new().get("/").bearer_auth("someS3cr3tAutht0k3n");
697 assert_eq!(
698 req.request
699 .head
700 .headers
701 .get(header::AUTHORIZATION)
702 .unwrap()
703 .to_str()
704 .unwrap(),
705 "Bearer someS3cr3tAutht0k3n"
706 );
707 }
708
709 #[crate::rt_test]
710 async fn client_auth_replaces_header() {
711 let client = Client::builder().build(
712 SharedCfg::new("TEST").add(ClientConfig::new().set_bearer_auth("token").unwrap()),
713 );
714 let req = client
715 .get("/")
716 .basic_auth("username", Some("password"))
717 .content_length(1)
718 .content_length(2);
719 let headers = &req.request.head.headers;
720 let auth: Vec<_> = headers
721 .get_all(header::AUTHORIZATION)
722 .map(|v| v.to_str().unwrap())
723 .collect();
724 assert_eq!(auth, ["Basic dXNlcm5hbWU6cGFzc3dvcmQ="]);
725 let len: Vec<_> = headers
726 .get_all(header::CONTENT_LENGTH)
727 .map(|v| v.to_str().unwrap())
728 .collect();
729 assert_eq!(len, ["2"]);
730
731 let req = Client::new()
732 .get("/")
733 .basic_auth("a", None)
734 .bearer_auth("b");
735 assert_eq!(
736 req.request
737 .head
738 .headers
739 .get_all(header::AUTHORIZATION)
740 .count(),
741 1
742 );
743
744 let cfg = ClientConfig::new()
745 .set_basic_auth("a", None)
746 .unwrap()
747 .set_bearer_auth("b")
748 .unwrap();
749 assert_eq!(cfg.headers().get_all(header::AUTHORIZATION).count(), 1);
750 }
751
752 #[crate::rt_test]
753 async fn client_invalid_url() {
754 let req = Client::new().get("http://local host/");
755 assert!(matches!(
756 req.err,
757 Some(ClientError::Url(InvalidUrl::Parse(_)))
758 ));
759 let err = req.send().await.unwrap_err();
760 assert!(matches!(
761 err.into_error(),
762 ClientError::Url(InvalidUrl::Parse(_))
763 ));
764
765 let req = Client::new().get("/").header("bad header", "1");
766 assert!(matches!(req.err, Some(ClientError::Http(_))));
767 }
768
769 #[crate::rt_test]
770 async fn client_query() {
771 let req = Client::new()
772 .get("/")
773 .query(&[("key1", "val1"), ("key2", "val2")]);
774 assert_eq!(req.get_uri().query().unwrap(), "key1=val1&key2=val2");
775
776 let req = Client::new().get("/").query(&InvalidQuery);
777 assert!(matches!(req.err, Some(ClientError::Error(_))));
778
779 let req = Client::new().get("http://localhost").query(&[("k", "v")]);
781 assert!(req.err.is_none());
782 assert_eq!(req.get_uri(), "http://localhost/?k=v");
783
784 let req = Client::new()
786 .get("http://localhost/p?a=1")
787 .query(&[("k", "v")]);
788 assert_eq!(req.get_uri(), "http://localhost/p?k=v");
789 }
790
791 #[crate::rt_test]
792 async fn client_invalid_headers() {
793 let bad_value = "a\nb";
794 for req in [
795 Client::new().get("/").header("x-test", bad_value),
796 Client::new().get("/").set_header("bad header", "1"),
797 Client::new().get("/").set_header("x-test", bad_value),
798 Client::new().get("/").set_header_if_none("bad header", "1"),
799 Client::new()
800 .get("/")
801 .set_header_if_none("x-test", bad_value),
802 Client::new().get("/").content_type(bad_value),
803 ] {
804 assert!(matches!(req.err, Some(ClientError::Http(_))), "{req:?}");
805 let err = req.send().await.unwrap_err();
806 assert!(matches!(err.into_error(), ClientError::Http(_)));
807 }
808
809 let req = Client::new()
811 .get("/")
812 .header("x-test", "1")
813 .set_header_if_none("x-test", bad_value);
814 assert!(req.err.is_none());
815 assert_eq!(req.headers().get("x-test").unwrap(), "1");
816 }
817
818 #[crate::rt_test]
819 async fn client_url_validation() {
820 for (url, expected) in [
821 ("/path", "missing-host"),
822 ("//localhost:8080/", "missing-scheme"),
823 ("localhost:8080", "missing-host"),
824 ("ftp://localhost/", "unknown-scheme"),
825 ] {
826 let err = Client::new().get(url).send().await.unwrap_err();
827 let kind = match err.into_error() {
828 ClientError::Url(InvalidUrl::MissingHost) => "missing-host",
829 ClientError::Url(InvalidUrl::MissingScheme) => "missing-scheme",
830 ClientError::Url(InvalidUrl::UnknownScheme) => "unknown-scheme",
831 err => panic!("{url}: {err:?}"),
832 };
833 assert_eq!(kind, expected, "{url}");
834 }
835 }
836
837 #[crate::rt_test]
838 async fn test_debug_redacts_authorization() {
839 let req = Client::new()
840 .get("http://localhost/")
841 .basic_auth("user", Some("secret"))
842 .address("127.0.0.1:1".parse().unwrap());
843 assert_eq!(req.request.addr, Some("127.0.0.1:1".parse().unwrap()));
844 let repr = format!("{req:?}");
845 assert!(repr.contains("\"authorization\": <REDACTED>"), "{repr}");
846 assert!(!repr.contains("Basic"), "{repr}");
847 }
848}