use std::{error::Error as StdError, fmt, net, rc::Rc};
use base64::{Engine, engine::general_purpose::STANDARD as base64};
#[cfg(feature = "cookie")]
use coo_kie::{Cookie, CookieJar};
use serde::Serialize;
use urly::Url;
use crate::error::Error;
use crate::http::error::HttpError;
use crate::http::header::{self, HeaderMap, HeaderName, HeaderValue};
use crate::http::{ConnectionType, Method, Version, body::Body};
use crate::{Cfg, PipelineBinding, time::Millis, util::Bytes, util::Stream};
use super::error::{ClientError, InvalidUrl};
use super::{ClientConfig, ClientResponse, ServiceRequest, ServiceResponse};
pub struct ClientRequest {
request: ServiceRequest,
svc: PipelineBinding<ServiceRequest, ServiceResponse, Error<ClientError>>,
err: Option<ClientError>,
cfg: Cfg<ClientConfig>,
#[cfg(feature = "cookie")]
cookies: Option<CookieJar>,
}
impl ClientRequest {
pub(super) fn new<U>(
method: Method,
uri: U,
cfg: Cfg<ClientConfig>,
svc: PipelineBinding<ServiceRequest, ServiceResponse, Error<ClientError>>,
) -> Self
where
Url: TryFrom<U>,
<Url as TryFrom<U>>::Error: Into<InvalidUrl>,
{
ClientRequest {
svc,
cfg,
request: ServiceRequest::new(),
err: None,
#[cfg(feature = "cookie")]
cookies: None,
}
.method(method)
.uri(uri)
}
#[inline]
#[must_use]
pub fn uri<U>(mut self, uri: U) -> Self
where
Url: TryFrom<U>,
<Url as TryFrom<U>>::Error: Into<InvalidUrl>,
{
match Url::try_from(uri) {
Ok(uri) => self.request.head.uri = uri,
Err(e) => self.err = Some(e.into().into()),
}
self
}
pub fn get_uri(&self) -> &Url {
&self.request.head.uri
}
#[must_use]
pub fn address(mut self, addr: net::SocketAddr) -> Self {
self.request.addr = Some(addr);
self
}
#[inline]
#[must_use]
pub fn method(mut self, method: Method) -> Self {
self.request.head.method = method;
self
}
#[inline]
#[must_use]
pub fn get_method(&self) -> &Method {
&self.request.head.method
}
#[inline]
#[must_use]
pub fn version(mut self, version: Version) -> Self {
self.request.head.version = version;
self
}
#[inline]
pub fn get_version(&self) -> &Version {
&self.request.head.version
}
#[inline]
pub fn headers(&self) -> &HeaderMap {
&self.request.head.headers
}
#[inline]
pub fn headers_mut(&mut self) -> &mut HeaderMap {
&mut self.request.head.headers
}
#[must_use]
pub fn header<K, V>(mut self, key: K, value: V) -> Self
where
HeaderName: TryFrom<K>,
HeaderValue: TryFrom<V>,
<HeaderName as TryFrom<K>>::Error: Into<HttpError>,
<HeaderValue as TryFrom<V>>::Error: Into<HttpError>,
{
match HeaderName::try_from(key) {
Ok(key) => match HeaderValue::try_from(value) {
Ok(value) => self.request.head.headers.append(key, value),
Err(e) => self.err = Some(ClientError::Http(e.into())),
},
Err(e) => self.err = Some(ClientError::Http(e.into())),
}
self
}
#[must_use]
pub fn set_header<K, V>(mut self, key: K, value: V) -> Self
where
HeaderName: TryFrom<K>,
HeaderValue: TryFrom<V>,
<HeaderName as TryFrom<K>>::Error: Into<HttpError>,
<HeaderValue as TryFrom<V>>::Error: Into<HttpError>,
{
match HeaderName::try_from(key) {
Ok(key) => match HeaderValue::try_from(value) {
Ok(value) => self.request.head.headers.insert(key, value),
Err(e) => self.err = Some(ClientError::Http(e.into())),
},
Err(e) => self.err = Some(ClientError::Http(e.into())),
}
self
}
#[must_use]
pub fn set_header_if_none<K, V>(mut self, key: K, value: V) -> Self
where
HeaderName: TryFrom<K>,
HeaderValue: TryFrom<V>,
<HeaderName as TryFrom<K>>::Error: Into<HttpError>,
<HeaderValue as TryFrom<V>>::Error: Into<HttpError>,
{
match HeaderName::try_from(key) {
Ok(key) => {
if !self.request.head.headers.contains_key(&key) {
match HeaderValue::try_from(value) {
Ok(value) => self.request.head.headers.insert(key, value),
Err(e) => self.err = Some(ClientError::Http(e.into())),
}
}
}
Err(e) => self.err = Some(ClientError::Http(e.into())),
}
self
}
#[inline]
#[must_use]
pub fn set_connection_type(mut self, ctype: ConnectionType) -> Self {
self.request.head.set_connection_type(ctype);
self
}
#[inline]
#[must_use]
pub fn force_close(mut self) -> Self {
self.request.head.set_connection_type(ConnectionType::Close);
self
}
#[inline]
#[must_use]
pub fn content_type<V>(mut self, value: V) -> Self
where
HeaderValue: TryFrom<V>,
<HeaderValue as TryFrom<V>>::Error: Into<HttpError>,
{
match HeaderValue::try_from(value) {
Ok(value) => self
.request
.head
.headers
.insert(header::CONTENT_TYPE, value),
Err(e) => self.err = Some(ClientError::Http(e.into())),
}
self
}
#[inline]
#[must_use]
pub fn content_length(self, len: u64) -> Self {
self.set_header(header::CONTENT_LENGTH, len)
}
#[must_use]
pub fn basic_auth<U>(self, username: U, password: Option<&str>) -> Self
where
U: fmt::Display,
{
let auth = match password {
Some(password) => format!("{username}:{password}"),
None => format!("{username}:"),
};
self.set_header(
header::AUTHORIZATION,
format!("Basic {}", base64.encode(auth)),
)
}
#[must_use]
pub fn bearer_auth<T>(self, token: T) -> Self
where
T: fmt::Display,
{
self.set_header(header::AUTHORIZATION, format!("Bearer {token}"))
}
#[must_use]
#[cfg(feature = "cookie")]
pub fn cookie<C>(mut self, cookie: C) -> Self
where
C: Into<Cookie<'static>>,
{
if let Some(cookies) = &mut self.cookies {
cookies.add(cookie.into());
} else {
let mut jar = CookieJar::new();
jar.add(cookie.into());
self.cookies = Some(jar);
}
self
}
#[must_use]
pub fn no_decompress(mut self) -> Self {
self.request.response_decompress = false;
self
}
#[must_use]
pub fn timeout<T: Into<Millis>>(mut self, timeout: T) -> Self {
self.request.timeout = Some(timeout.into());
self
}
#[must_use]
pub fn if_true<F>(self, value: bool, f: F) -> Self
where
F: FnOnce(ClientRequest) -> ClientRequest,
{
if value { f(self) } else { self }
}
#[must_use]
pub fn if_some<T, F>(self, value: Option<T>, f: F) -> Self
where
F: FnOnce(T, ClientRequest) -> ClientRequest,
{
if let Some(val) = value { f(val, self) } else { self }
}
#[must_use]
pub fn query<T: Serialize>(mut self, query: &T) -> Self {
let query = match serde_urlencoded::to_string(query) {
Ok(query) => query,
Err(err) => {
self.err = Some(ClientError::Error(Rc::new(err)));
return self;
}
};
self.request.head.uri.set_query(Some(&query));
self
}
}
impl ClientRequest {
pub async fn send_body<B>(mut self, body: B) -> Result<ClientResponse, Error<ClientError>>
where
B: Into<Body>,
{
self.prep_for_sending()?;
*self.request.body() = body.into();
self.svc.call(self.request).await.map(Into::into)
}
pub async fn send_json<T: Serialize>(
mut self,
value: &T,
) -> Result<ClientResponse, Error<ClientError>> {
self.prep_for_sending()?;
self.request.set_json(value)?;
self.svc.call(self.request).await.map(Into::into)
}
pub async fn send_form<T: Serialize>(
mut self,
value: &T,
) -> Result<ClientResponse, Error<ClientError>> {
self.prep_for_sending()?;
self.request.set_form(value)?;
self.svc.call(self.request).await.map(Into::into)
}
pub async fn send_stream<T, E>(
mut self,
stream: T,
) -> Result<ClientResponse, Error<ClientError>>
where
T: Stream<Item = Result<Bytes, E>> + Unpin + 'static,
E: StdError + 'static,
{
self.prep_for_sending()?;
self.request.set_stream(stream);
self.svc.call(self.request).await.map(Into::into)
}
pub async fn send(mut self) -> Result<ClientResponse, Error<ClientError>> {
self.prep_for_sending()?;
self.svc.call(self.request).await.map(Into::into)
}
fn prep_for_sending(&mut self) -> Result<(), Error<ClientError>> {
self.prep_for_sending_inner()
.map_err(|e| e.with_service(self.cfg.service()))
}
fn prep_for_sending_inner(&mut self) -> Result<(), Error<ClientError>> {
if let Some(e) = self.err.take() {
return Err(e.into());
}
let uri = &self.request.head.uri;
if uri.host().is_none() {
return Err(ClientError::from(InvalidUrl::MissingHost).into());
}
match uri.scheme_str() {
Some("http" | "ws" | "https" | "wss") => (),
Some(_) => return Err(ClientError::from(InvalidUrl::UnknownScheme).into()),
None => return Err(ClientError::from(InvalidUrl::MissingScheme).into()),
}
#[cfg(feature = "cookie")]
{
if let Some(ref jar) = self.cookies {
let headers = &mut self.request.head.headers;
let mut cookie = headers
.get(header::COOKIE)
.map(|v| v.as_bytes().to_vec())
.unwrap_or_default();
for c in jar.iter() {
crate::http::helpers::push_cookie(&mut cookie, c.name(), c.value());
}
if let Ok(val) = HeaderValue::from_bytes(&cookie) {
headers.insert(header::COOKIE, val);
}
}
}
#[cfg(feature = "compress")]
if self.request.response_decompress
&& !self
.request
.head
.headers
.contains_key(&header::ACCEPT_ENCODING)
{
const COMPRESSION: HeaderValue = HeaderValue::from_static("gzip, deflate");
self.request
.head
.headers
.insert(header::ACCEPT_ENCODING, COMPRESSION);
}
Ok(())
}
}
impl fmt::Debug for ClientRequest {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
writeln!(
f,
"\nClientRequest {:?} {}:{}",
self.request.head.version, self.request.head.method, self.request.head.uri
)?;
writeln!(f, " headers:")?;
for (key, val) in &self.request.head.headers {
if key == header::AUTHORIZATION {
writeln!(f, " {key:?}: <REDACTED>")?;
} else {
writeln!(f, " {key:?}: {val:?}")?;
}
}
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::{SharedCfg, client::Client};
struct InvalidQuery;
impl Serialize for InvalidQuery {
fn serialize<S>(&self, _: S) -> Result<S::Ok, S::Error>
where
S: serde::Serializer,
{
Err(serde::ser::Error::custom("invalid query"))
}
}
#[crate::rt_test]
async fn test_debug() {
let request = Client::new().get("/").header("x-test", "111");
let repr = format!("{request:?}");
assert!(repr.contains("ClientRequest"));
assert!(repr.contains("x-test"));
}
#[crate::rt_test]
async fn test_basics() {
let mut req = Client::new()
.put("/")
.version(Version::HTTP_2)
.header(header::DATE, "data")
.content_type("plain/text")
.if_true(true, |req| req.header(header::SERVER, "awc"))
.if_true(false, |req| req.header(header::EXPECT, "awc"))
.if_some(Some("server"), |val, req| {
req.header(header::USER_AGENT, val)
})
.if_some(Option::<&str>::None, |_, req| {
req.header(header::ALLOW, "1")
})
.content_length(100);
assert!(req.headers().contains_key(header::CONTENT_TYPE));
assert!(req.headers().contains_key(header::DATE));
assert!(req.headers().contains_key(header::SERVER));
assert!(req.headers().contains_key(header::USER_AGENT));
assert!(!req.headers().contains_key(header::ALLOW));
assert!(!req.headers().contains_key(header::EXPECT));
assert_eq!(req.request.head.version, Version::HTTP_2);
assert_eq!(req.get_version(), &Version::HTTP_2);
assert_eq!(req.get_method(), Method::PUT);
let _ = req.headers_mut();
let _ = req.send_body("").await;
}
#[cfg(feature = "cookie")]
#[crate::rt_test]
async fn cookies_extend_cookie_header() {
use coo_kie::Cookie;
fn cookie_header(req: &mut ClientRequest) -> String {
req.prep_for_sending_inner().unwrap();
let val = req.request.head.headers.get(header::COOKIE).unwrap();
val.to_str().unwrap().to_string()
}
let mut req = Client::new()
.get("http://localhost/")
.cookie(Cookie::build(("c1", "v1")))
.cookie(Cookie::build(("c2", "v2")));
let cookie = cookie_header(&mut req);
let mut cookies: Vec<_> = cookie.split("; ").collect();
cookies.sort_unstable();
assert_eq!(cookies, ["c1=v1", "c2=v2"]);
let mut req = Client::new()
.get("http://localhost/")
.header(header::COOKIE, "c0=v0")
.cookie(Cookie::build(("c1", "v1")));
assert_eq!(cookie_header(&mut req), "c0=v0; c1=v1");
}
#[crate::rt_test]
async fn test_client_header() {
let req = Client::builder()
.build(
SharedCfg::new("H").add(
ClientConfig::new()
.set_header(header::CONTENT_TYPE, "111")
.unwrap(),
),
)
.get("/");
assert_eq!(
req.request
.head
.headers
.get(header::CONTENT_TYPE)
.unwrap()
.to_str()
.unwrap(),
"111"
);
}
#[crate::rt_test]
async fn test_client_header_override() {
let req = Client::builder()
.build(
SharedCfg::new("H").add(
ClientConfig::new()
.set_header(header::CONTENT_TYPE, "111")
.unwrap(),
),
)
.get("/")
.set_header(header::CONTENT_TYPE, "222");
assert_eq!(
req.request
.head
.headers
.get(header::CONTENT_TYPE)
.unwrap()
.to_str()
.unwrap(),
"222"
);
}
#[crate::rt_test]
async fn client_basic_auth() {
let req = Client::new()
.get("/")
.basic_auth("username", Some("password"));
assert_eq!(
req.request
.head
.headers
.get(header::AUTHORIZATION)
.unwrap()
.to_str()
.unwrap(),
"Basic dXNlcm5hbWU6cGFzc3dvcmQ="
);
let req = Client::new().get("/").basic_auth("username", None);
assert_eq!(
req.request
.head
.headers
.get(header::AUTHORIZATION)
.unwrap()
.to_str()
.unwrap(),
"Basic dXNlcm5hbWU6"
);
}
#[crate::rt_test]
async fn client_bearer_auth() {
let req = Client::new().get("/").bearer_auth("someS3cr3tAutht0k3n");
assert_eq!(
req.request
.head
.headers
.get(header::AUTHORIZATION)
.unwrap()
.to_str()
.unwrap(),
"Bearer someS3cr3tAutht0k3n"
);
}
#[crate::rt_test]
async fn client_auth_replaces_header() {
let client = Client::builder().build(
SharedCfg::new("TEST").add(ClientConfig::new().set_bearer_auth("token").unwrap()),
);
let req = client
.get("/")
.basic_auth("username", Some("password"))
.content_length(1)
.content_length(2);
let headers = &req.request.head.headers;
let auth: Vec<_> = headers
.get_all(header::AUTHORIZATION)
.map(|v| v.to_str().unwrap())
.collect();
assert_eq!(auth, ["Basic dXNlcm5hbWU6cGFzc3dvcmQ="]);
let len: Vec<_> = headers
.get_all(header::CONTENT_LENGTH)
.map(|v| v.to_str().unwrap())
.collect();
assert_eq!(len, ["2"]);
let req = Client::new()
.get("/")
.basic_auth("a", None)
.bearer_auth("b");
assert_eq!(
req.request
.head
.headers
.get_all(header::AUTHORIZATION)
.count(),
1
);
let cfg = ClientConfig::new()
.set_basic_auth("a", None)
.unwrap()
.set_bearer_auth("b")
.unwrap();
assert_eq!(cfg.headers().get_all(header::AUTHORIZATION).count(), 1);
}
#[crate::rt_test]
async fn client_invalid_url() {
let req = Client::new().get("http://local host/");
assert!(matches!(
req.err,
Some(ClientError::Url(InvalidUrl::Parse(_)))
));
let err = req.send().await.unwrap_err();
assert!(matches!(
err.into_error(),
ClientError::Url(InvalidUrl::Parse(_))
));
let req = Client::new().get("/").header("bad header", "1");
assert!(matches!(req.err, Some(ClientError::Http(_))));
}
#[crate::rt_test]
async fn client_query() {
let req = Client::new()
.get("/")
.query(&[("key1", "val1"), ("key2", "val2")]);
assert_eq!(req.get_uri().query().unwrap(), "key1=val1&key2=val2");
let req = Client::new().get("/").query(&InvalidQuery);
assert!(matches!(req.err, Some(ClientError::Error(_))));
let req = Client::new().get("http://localhost").query(&[("k", "v")]);
assert!(req.err.is_none());
assert_eq!(req.get_uri(), "http://localhost/?k=v");
let req = Client::new()
.get("http://localhost/p?a=1")
.query(&[("k", "v")]);
assert_eq!(req.get_uri(), "http://localhost/p?k=v");
}
#[crate::rt_test]
async fn client_invalid_headers() {
let bad_value = "a\nb";
for req in [
Client::new().get("/").header("x-test", bad_value),
Client::new().get("/").set_header("bad header", "1"),
Client::new().get("/").set_header("x-test", bad_value),
Client::new().get("/").set_header_if_none("bad header", "1"),
Client::new()
.get("/")
.set_header_if_none("x-test", bad_value),
Client::new().get("/").content_type(bad_value),
] {
assert!(matches!(req.err, Some(ClientError::Http(_))), "{req:?}");
let err = req.send().await.unwrap_err();
assert!(matches!(err.into_error(), ClientError::Http(_)));
}
let req = Client::new()
.get("/")
.header("x-test", "1")
.set_header_if_none("x-test", bad_value);
assert!(req.err.is_none());
assert_eq!(req.headers().get("x-test").unwrap(), "1");
}
#[crate::rt_test]
async fn client_url_validation() {
for (url, expected) in [
("/path", "missing-host"),
("//localhost:8080/", "missing-scheme"),
("localhost:8080", "missing-host"),
("ftp://localhost/", "unknown-scheme"),
] {
let err = Client::new().get(url).send().await.unwrap_err();
let kind = match err.into_error() {
ClientError::Url(InvalidUrl::MissingHost) => "missing-host",
ClientError::Url(InvalidUrl::MissingScheme) => "missing-scheme",
ClientError::Url(InvalidUrl::UnknownScheme) => "unknown-scheme",
err => panic!("{url}: {err:?}"),
};
assert_eq!(kind, expected, "{url}");
}
}
#[crate::rt_test]
async fn test_debug_redacts_authorization() {
let req = Client::new()
.get("http://localhost/")
.basic_auth("user", Some("secret"))
.address("127.0.0.1:1".parse().unwrap());
assert_eq!(req.request.addr, Some("127.0.0.1:1".parse().unwrap()));
let repr = format!("{req:?}");
assert!(repr.contains("\"authorization\": <REDACTED>"), "{repr}");
assert!(!repr.contains("Basic"), "{repr}");
}
}