use std::{
error::Error as StdError,
fmt, io,
pin::Pin,
task::{Context, Poll},
time::Duration,
};
use bytes::Bytes;
use futures_core::Stream;
use hyper::{
Request, Uri,
body::{Body, Frame, Incoming, SizeHint},
header::{HeaderName, HeaderValue},
};
use hyper_util::{
client::legacy::{Client, connect::HttpConnector},
rt::TokioExecutor,
};
use super::{Backend, HttpRequest, HttpResponse, Method};
use crate::{Error, ErrorKind, body::BodyStream};
#[cfg(any(feature = "tls", feature = "tls-aws-lc"))]
type Connector = hyper_rustls::HttpsConnector<HttpConnector>;
#[cfg(not(any(feature = "tls", feature = "tls-aws-lc")))]
type Connector = HttpConnector;
#[cfg(feature = "tls-aws-lc")]
fn provider() -> rustls::crypto::CryptoProvider {
rustls::crypto::aws_lc_rs::default_provider()
}
#[cfg(all(feature = "tls", not(feature = "tls-aws-lc")))]
fn provider() -> rustls::crypto::CryptoProvider {
rustls::crypto::ring::default_provider()
}
pub struct Hyper {
client: Client<Connector, SizedBody>,
}
impl Hyper {
pub(crate) fn new(connect_timeout: Duration) -> Result<Self, Error> {
let mut http = HttpConnector::new();
http.set_connect_timeout(Some(connect_timeout));
http.set_nodelay(true);
http.enforce_http(false);
#[cfg(any(feature = "tls", feature = "tls-aws-lc"))]
let connector = hyper_rustls::HttpsConnectorBuilder::new()
.with_provider_and_webpki_roots(provider())
.map_err(|source| {
Error::new(ErrorKind::Config)
.with_detail("TLS could not be set up")
.with_source(source)
})?
.https_or_http()
.enable_all_versions()
.wrap_connector(http);
#[cfg(not(any(feature = "tls", feature = "tls-aws-lc")))]
let connector = http;
Ok(Self {
client: Client::builder(TokioExecutor::new()).build(connector),
})
}
}
impl fmt::Debug for Hyper {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("Hyper").finish_non_exhaustive()
}
}
impl Backend for Hyper {
type Body = HyperBody;
async fn send(&self, request: HttpRequest) -> Result<HttpResponse<HyperBody>, Error> {
let invalid = |what: &'static str| Error::new(ErrorKind::Config).with_detail(what);
let uri: Uri = request
.url
.parse()
.map_err(|_| invalid("the URL is not valid"))?;
let mut builder = Request::builder().uri(uri).method(match request.method {
Method::Get => hyper::Method::GET,
Method::Post => hyper::Method::POST,
});
for (name, value) in &request.headers {
let header = HeaderName::from_bytes(name.as_bytes())
.ok()
.zip(HeaderValue::from_str(value).ok());
let Some((name, mut value)) = header else {
return Err(invalid("a header is not valid"));
};
value.set_sensitive(!super::is_plain(name.as_str()));
builder = builder.header(name, value);
}
let body = match request.body {
Some(body) => SizedBody {
stream: Some(body.stream),
remaining: body.length,
},
None => SizedBody {
stream: None,
remaining: 0,
},
};
let request = builder
.body(body)
.map_err(|_| invalid("the request is not valid"))?;
let response = self
.client
.request(request)
.await
.map_err(|error| failed(&error))?;
let (parts, incoming) = response.into_parts();
let headers = parts
.headers
.iter()
.map(|(name, value)| {
(
name.as_str().to_owned(),
String::from_utf8_lossy(value.as_bytes()).into_owned(),
)
})
.collect();
let body = HyperBody {
incoming: Some(incoming),
};
Ok(HttpResponse::new(parts.status.as_u16(), headers, body))
}
}
pub struct HyperBody {
incoming: Option<Incoming>,
}
impl Stream for HyperBody {
type Item = Result<Bytes, Error>;
fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
loop {
let Some(incoming) = self.incoming.as_mut() else {
return Poll::Ready(None);
};
match Pin::new(incoming).poll_frame(cx) {
Poll::Pending => return Poll::Pending,
Poll::Ready(Some(Ok(frame))) => {
if let Ok(data) = frame.into_data() {
return Poll::Ready(Some(Ok(data)));
}
}
Poll::Ready(Some(Err(error))) => {
self.incoming = None;
return Poll::Ready(Some(Err(failed(&error))));
}
Poll::Ready(None) => {
self.incoming = None;
return Poll::Ready(None);
}
}
}
}
}
impl fmt::Debug for HyperBody {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("HyperBody").finish_non_exhaustive()
}
}
struct SizedBody {
stream: Option<BodyStream>,
remaining: u64,
}
impl Body for SizedBody {
type Data = Bytes;
type Error = Error;
fn poll_frame(
mut self: Pin<&mut Self>,
cx: &mut Context<'_>,
) -> Poll<Option<Result<Frame<Bytes>, Error>>> {
let Some(stream) = self.stream.as_mut() else {
return Poll::Ready(None);
};
match Pin::new(stream).poll_next(cx) {
Poll::Ready(Some(Ok(bytes))) => {
self.remaining = self.remaining.saturating_sub(bytes.len() as u64);
Poll::Ready(Some(Ok(Frame::data(bytes))))
}
Poll::Ready(Some(Err(error))) => {
self.stream = None;
Poll::Ready(Some(Err(error)))
}
Poll::Ready(None) => {
self.stream = None;
Poll::Ready(None)
}
Poll::Pending => Poll::Pending,
}
}
fn is_end_stream(&self) -> bool {
self.stream.is_none()
}
fn size_hint(&self) -> SizeHint {
SizeHint::with_exact(self.remaining)
}
}
fn failed(error: &(dyn StdError + 'static)) -> Error {
let unsent = error
.downcast_ref::<hyper_util::client::legacy::Error>()
.is_some_and(hyper_util::client::legacy::Error::is_connect);
let mut timed_out = false;
let mut causes = String::new();
let mut cause: Option<&(dyn StdError + 'static)> = Some(error);
while let Some(current) = cause {
if let Some(own) = current.downcast_ref::<Error>() {
let copy = Error::new(own.kind());
return match own.detail() {
Some(detail) => copy.with_detail(detail.to_owned()),
None => copy,
};
}
if let Some(io) = current.downcast_ref::<io::Error>() {
timed_out |= io.kind() == io::ErrorKind::TimedOut;
}
let message = current.to_string();
if !causes.ends_with(&message) {
if !causes.is_empty() {
causes.push_str(": ");
}
causes.push_str(&message);
}
cause = current.source();
}
let mapped = if timed_out {
Error::new(ErrorKind::Timeout).with_detail("connecting to the model server timed out")
} else {
Error::new(ErrorKind::Transport).with_detail(causes)
};
if unsent { mapped.with_unsent() } else { mapped }
}