use crate::{Error, RequestOptions};
use bytes::Bytes;
use core::future::poll_fn;
use core::pin::{Pin, pin};
use core::task::{Poll, ready};
use http::uri::Scheme;
use http::{Request, Response};
use http_body::Body;
use std::time::Duration;
use tokio::io::{AsyncRead, AsyncWrite};
use tokio::net::TcpStream;
trait TokioStream: AsyncRead + AsyncWrite + Send + Sync + Unpin + 'static {
fn boxed(self) -> Box<dyn TokioStream>
where
Self: Sized,
{
Box::new(self)
}
}
impl<T> TokioStream for T where T: AsyncRead + AsyncWrite + Send + Sync + Unpin + 'static {}
pub async fn default_send_request(
mut req: Request<impl Body<Data = Bytes, Error = Error> + Send + 'static>,
options: Option<RequestOptions>,
) -> Result<
(
Response<impl Body<Data = Bytes, Error = Error>>,
impl Future<Output = Result<(), Error>> + Send,
),
Error,
> {
let uri = req.uri();
let authority = uri.authority().ok_or(Error::HttpRequestUriInvalid)?;
let use_tls = uri.scheme() == Some(&Scheme::HTTPS);
let authority = if authority.port().is_some() {
authority.to_string()
} else {
let port = if use_tls { 443 } else { 80 };
format!("{authority}:{port}")
};
let connect_timeout = options
.and_then(|r| r.connect_timeout)
.unwrap_or(Duration::from_secs(600));
let first_byte_timeout = options
.and_then(|r| r.first_byte_timeout)
.unwrap_or(Duration::from_secs(600));
let between_bytes_timeout = options
.and_then(|r| r.between_bytes_timeout)
.unwrap_or(Duration::from_secs(600));
let stream = match tokio::time::timeout(connect_timeout, TcpStream::connect(&authority)).await {
Ok(stream) => stream.map_err(Error::Connect)?,
Err(..) => return Err(Error::ConnectionTimeout),
};
let stream = if use_tls {
let root_cert_store = rustls::RootCertStore {
roots: webpki_roots::TLS_SERVER_ROOTS.into(),
};
let config = rustls::ClientConfig::builder()
.with_root_certificates(root_cert_store)
.with_no_client_auth();
let connector = tokio_rustls::TlsConnector::from(std::sync::Arc::new(config));
let domain = tls_server_name(&authority)?;
let stream = connector
.connect(domain, stream)
.await
.map_err(Error::Tls)?;
stream.boxed()
} else {
stream.boxed()
};
let (mut sender, conn) = tokio::time::timeout(
connect_timeout,
hyper::client::conn::http1::Builder::new().handshake(crate::io::TokioIo::new(stream)),
)
.await
.map_err(|_| Error::ConnectionTimeout)??;
*req.uri_mut() = http::Uri::builder()
.path_and_query(
req.uri()
.path_and_query()
.map(|p| p.as_str())
.unwrap_or("/"),
)
.build()
.expect("comes from valid request");
let send = async move {
use core::task::Context;
struct IncomingResponseBody {
incoming: hyper::body::Incoming,
timeout: tokio::time::Interval,
}
impl http_body::Body for IncomingResponseBody {
type Data = <hyper::body::Incoming as http_body::Body>::Data;
type Error = Error;
fn poll_frame(
mut self: Pin<&mut Self>,
cx: &mut Context<'_>,
) -> Poll<Option<Result<http_body::Frame<Self::Data>, Self::Error>>> {
match Pin::new(&mut self.as_mut().incoming).poll_frame(cx) {
Poll::Ready(None) => Poll::Ready(None),
Poll::Ready(Some(Err(err))) => {
let err = if err.is_timeout() {
Error::HttpResponseTimeout
} else {
Error::from(err)
};
Poll::Ready(Some(Err(err)))
}
Poll::Ready(Some(Ok(frame))) => {
self.timeout.reset();
Poll::Ready(Some(Ok(frame)))
}
Poll::Pending => {
ready!(self.timeout.poll_tick(cx));
Poll::Ready(Some(Err(Error::ConnectionReadTimeout)))
}
}
}
fn is_end_stream(&self) -> bool {
self.incoming.is_end_stream()
}
fn size_hint(&self) -> http_body::SizeHint {
self.incoming.size_hint()
}
}
let res = tokio::time::timeout(first_byte_timeout, sender.send_request(req))
.await
.map_err(|_| Error::ConnectionReadTimeout)?
.map_err(Error::from)?;
let mut timeout = tokio::time::interval(between_bytes_timeout);
timeout.reset();
Ok(res.map(|incoming| IncomingResponseBody { incoming, timeout }))
};
let mut send = pin!(send);
let mut conn = Some(conn);
let res = poll_fn(|cx| match send.as_mut().poll(cx) {
Poll::Ready(Ok(res)) => Poll::Ready(Ok(res)),
Poll::Ready(Err(err)) => Poll::Ready(Err(err)),
Poll::Pending => {
let Some(fut) = conn.as_mut() else {
return Poll::Pending;
};
let res = ready!(Pin::new(fut).poll(cx));
conn = None;
match res {
Ok(()) => send.as_mut().poll(cx),
Err(err) => Poll::Ready(Err(Error::from(err))),
}
}
})
.await?;
Ok((res, async move {
let Some(conn) = conn.take() else {
return Ok(());
};
if let Err(err) = conn.await {
if err.is_timeout() {
return Err(Error::HttpResponseTimeout);
}
return Err(err.into());
}
Ok(())
}))
}
pub(crate) fn tls_server_name(
authority: &str,
) -> Result<rustls::pki_types::ServerName<'static>, rustls::pki_types::InvalidDnsNameError> {
use rustls::pki_types::ServerName;
if let Ok(addr) = authority.parse::<std::net::SocketAddr>() {
return Ok(ServerName::from(addr.ip()));
}
let host = match authority.split_once(':') {
Some((host, _port)) => host,
None => authority,
};
Ok(ServerName::try_from(host)?.to_owned())
}
#[cfg(test)]
mod tls_server_name_tests {
use super::tls_server_name;
use rustls::pki_types::ServerName;
#[test]
fn resolves_server_name_from_authority() {
assert_eq!(
tls_server_name("example.com:443").unwrap(),
ServerName::try_from("example.com").unwrap()
);
assert_eq!(
tls_server_name("example.com").unwrap(),
ServerName::try_from("example.com").unwrap()
);
assert_eq!(
tls_server_name("127.0.0.1:80").unwrap(),
ServerName::from(std::net::Ipv4Addr::LOCALHOST)
);
assert_eq!(
tls_server_name("[::1]:443").unwrap(),
ServerName::from(std::net::Ipv6Addr::LOCALHOST)
);
assert_eq!(
tls_server_name("[2001:db8::1]:8443").unwrap(),
ServerName::from("2001:db8::1".parse::<std::net::Ipv6Addr>().unwrap())
);
}
}