use anyspawn::Spawner;
use fetch_hyper::HyperTransportBuilder;
use fetch_options::{SocketOptions, TransportOptions};
use fetch_tls::{TlsBackend, TlsBackendBuilder};
use http::uri::Scheme;
use http_extensions::{HttpError, Result};
use hyper_util::client::legacy::connect::HttpConnector;
use hyper_util::rt::TokioIo;
use seatbelt::RecoveryInfo;
use templated_uri::BaseUri;
use thread_aware::ThreadAware;
use tick::Clock;
use tower_service::Service as _;
use crate::custom::{CustomContext, CustomDeps, Isolation};
use crate::error_labels::LABEL_SCHEME_NOT_ALLOWED;
use crate::handlers::TransportHandler;
use crate::tls::TlsOptions;
use crate::{HttpClient, HttpClientBuilder};
#[derive(Debug, Clone, ThreadAware)]
#[fundle::deps]
pub struct TokioDeps {
pub clock: Clock,
pub global_pool: bytesbuf::mem::GlobalPool,
}
impl Default for TokioDeps {
fn default() -> Self {
Self::with_clock(&Clock::new_tokio())
}
}
impl TokioDeps {
#[must_use]
pub fn with_clock(clock: &Clock) -> Self {
Self {
global_pool: bytesbuf::mem::GlobalPool::new(),
clock: clock.clone(),
}
}
}
#[derive(Debug, Clone, Default)]
#[non_exhaustive]
pub struct TokioTransportOptions {
pub socket: SocketOptions,
}
impl TokioTransportOptions {
#[must_use]
pub fn socket(mut self, socket: SocketOptions) -> Self {
self.socket = socket;
self
}
}
impl HttpClient {
pub fn builder_tokio(deps: impl Into<TokioDeps>) -> HttpClientBuilder {
Self::builder_tokio_with_options(deps, TokioTransportOptions::default())
}
pub fn builder_tokio_with_options(deps: impl Into<TokioDeps>, options: TokioTransportOptions) -> HttpClientBuilder {
let deps = deps.into();
let clock = deps.clock.clone();
let global_pool = deps.global_pool.clone();
Self::builder_custom_internal(
crate::constants::TOKIO_RUNTIME_NAME,
crate::constants::HYPER_TRANSPORT_NAME,
move |cx| TransportHandler(build_tokio_handler(cx, &options).into()),
Isolation::Shared,
CustomDeps {
clock,
global_pool,
extras: deps,
},
)
}
#[must_use]
pub fn new_tokio() -> Self {
Self::builder_tokio(TokioDeps::default()).build()
}
}
#[derive(Clone, Debug)]
struct TokioConnector {
connector: HttpConnector,
}
impl TokioConnector {
fn new(options: SocketOptions) -> Self {
let mut connector = HttpConnector::new();
connector.enforce_http(false);
if let Some(no_delay) = options.no_delay {
connector.set_nodelay(no_delay);
}
connector.set_send_buffer_size(options.send_buffer_size.map(widen_socket_buffer_size));
connector.set_recv_buffer_size(options.receive_buffer_size.map(widen_socket_buffer_size));
Self { connector }
}
}
impl layered::Service<BaseUri> for TokioConnector {
type Out = Result<TokioIo<::tokio::net::TcpStream>>;
async fn execute(&self, input: BaseUri) -> Self::Out {
let scheme = input.origin().scheme();
if scheme != &Scheme::HTTP && scheme != &Scheme::HTTPS {
return Err(HttpError::other(
format!("the connector does not support the '{scheme}' scheme; only http and https are supported"),
RecoveryInfo::never(),
LABEL_SCHEME_NOT_ALLOWED,
));
}
let port = input.try_effective_port()?;
let uri = http::Uri::from(input.with_port(port));
let mut connector = self.connector.clone();
std::future::poll_fn(|cx| connector.poll_ready(cx))
.await
.map_err(map_connect_error)?;
connector.call(uri).await.map_err(map_connect_error)
}
}
#[inline]
fn widen_socket_buffer_size(size: u32) -> usize {
usize::try_from(size).expect("fetch supports only targets with pointers at least 32 bits wide")
}
fn map_connect_error<E: std::error::Error + Send + Sync + 'static>(error: E) -> HttpError {
let kind = io_error_kind(&error);
std::io::Error::new(kind, error).into()
}
fn io_error_kind(mut error: &(dyn std::error::Error + 'static)) -> std::io::ErrorKind {
loop {
if let Some(io_error) = error.downcast_ref::<std::io::Error>() {
return io_error.kind();
}
match error.source() {
Some(source) => error = source,
None => return std::io::ErrorKind::Other,
}
}
}
fn build_tokio_handler(cx: CustomContext<TokioDeps>, options: &TokioTransportOptions) -> fetch_hyper::HyperTransport {
let tls_backend = build_tls_backend(&cx.options, cx.tls);
let connector = TokioConnector::new(options.socket);
HyperTransportBuilder::new(connector, Spawner::new_tokio(), cx.clock, cx.options)
.body_builder(cx.body_builder)
.pool_index(cx.pool_index)
.meter(cx.meter)
.build(tls_backend)
}
fn build_tls_backend(options: &TransportOptions, tls: TlsOptions) -> TlsBackend {
let mut builder = TlsBackendBuilder::new();
if !options.supported_http_versions.is_empty() {
builder = builder.supported_http_versions(&options.supported_http_versions);
}
#[cfg(any(feature = "rustls", test))]
{
let provider = std::sync::Arc::new(::rustls::crypto::aws_lc_rs::default_provider());
let verifier = std::sync::Arc::new(
rustls_platform_verifier::Verifier::new(std::sync::Arc::clone(&provider))
.expect("the platform certificate verifier must initialize with the aws-lc-rs crypto provider"),
);
builder = builder.configure_rustls(provider, verifier);
}
#[cfg(all(feature = "native-tls", not(any(feature = "rustls", test))))]
{
builder = builder.defaults_to_native_tls();
}
builder
.build_backend(tls)
.expect("TLS backend construction must succeed for the configured TlsOptions")
}
#[cfg(test)]
#[cfg_attr(coverage_nightly, coverage(off))]
mod tests {
use http::StatusCode;
use http_extensions::FakeHandler;
use thread_aware::ThreadAware;
use thread_aware::affinity::pinned_affinities;
use tick::Clock;
use super::TokioDeps;
use crate::pipeline::Pipeline;
use crate::{HttpClient, HttpResponseBuilder};
#[cfg_attr(miri, ignore)]
#[tokio::test]
async fn test_builder_tokio() {
let clock = Clock::new_tokio();
let client = HttpClient::builder_tokio(TokioDeps::with_clock(&clock)).minimal_pipeline().build();
assert!(matches!(client.pipeline(), Pipeline::Minimal(_)));
if let Pipeline::Minimal(dispatch) = client.pipeline() {
assert!(matches!(dispatch.mode, crate::handlers::DispatchMode::Single(_)));
}
}
#[cfg_attr(miri, ignore)]
#[test]
fn tokio_transport_options_default() {
insta::assert_debug_snapshot!(super::TokioTransportOptions::default());
}
#[cfg_attr(miri, ignore)]
#[test]
fn configure_tokio_transport_options() {
let options = super::TokioTransportOptions::default().socket(fetch_options::SocketOptions::default().no_delay(true));
assert_eq!(options.socket.no_delay, Some(true));
}
#[cfg_attr(miri, ignore)]
#[test]
fn socket_buffer_sizes_are_widened_verbatim() {
assert_eq!(super::widen_socket_buffer_size(0), 0);
assert_eq!(super::widen_socket_buffer_size(1), 1);
assert_eq!(super::widen_socket_buffer_size(u32::MAX), u32::MAX as usize);
}
#[cfg_attr(miri, ignore)]
#[test]
fn assert_tokio_transport_options_type() {
static_assertions::assert_impl_all!(
super::TokioTransportOptions: Send,
Sync,
Clone,
std::fmt::Debug,
Default
);
}
#[cfg_attr(miri, ignore)]
#[tokio::test]
async fn test_new_tokio() {
let clock = Clock::new_tokio();
let client = HttpClient::builder_tokio(TokioDeps::with_clock(&clock)).build();
assert!(client.pipeline().is_standard());
}
#[cfg_attr(miri, ignore)]
#[tokio::test]
async fn new_tokio_uses_default_deps() {
let client = HttpClient::new_tokio();
assert!(client.pipeline().is_standard());
}
#[cfg_attr(miri, ignore)]
#[tokio::test]
async fn tokio_client_works_after_relocation() {
let affinities = pinned_affinities(&[2]);
let clock = Clock::new_tokio();
let mut client = HttpClient::builder_tokio(TokioDeps::with_clock(&clock))
.custom_pipeline(|_root, _ctx| FakeHandler::from_fn(|_request| HttpResponseBuilder::new_fake().status(StatusCode::OK).build()))
.build();
let response = client.get("https://example.com").fetch().await.unwrap();
assert_eq!(response.status(), StatusCode::OK);
client.relocate(None, affinities[0]);
let response = client.get("https://example.com/after-relocation").fetch().await.unwrap();
assert_eq!(response.status(), StatusCode::OK);
}
#[cfg_attr(miri, ignore)]
#[test]
fn build_tls_backend_skips_empty_supported_http_versions() {
use fetch_options::TransportOptions;
use crate::tls::TlsOptions;
let mut options = TransportOptions::default();
options.supported_http_versions = Vec::new();
let _backend = super::build_tls_backend(&options, TlsOptions::default());
}
async fn connect_with(options: fetch_options::SocketOptions, port: u16) -> http_extensions::Result<tokio::net::TcpStream> {
connect_to(options, &format!("http://127.0.0.1:{port}")).await
}
async fn connect_to(options: fetch_options::SocketOptions, uri: &str) -> http_extensions::Result<tokio::net::TcpStream> {
use layered::Service as _;
let base_uri = templated_uri::BaseUri::try_from(uri)?;
super::TokioConnector::new(options)
.execute(base_uri)
.await
.map(hyper_util::rt::TokioIo::into_inner)
}
#[cfg_attr(miri, ignore)]
#[tokio::test]
async fn connector_connects_with_default_options() -> http_extensions::Result<()> {
use fetch_options::SocketOptions;
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await?;
let port = listener.local_addr()?.port();
let stream = connect_with(SocketOptions::default(), port).await?;
assert_eq!(stream.peer_addr()?.port(), port);
Ok(())
}
#[cfg_attr(miri, ignore)]
#[tokio::test]
async fn connector_applies_no_delay() -> http_extensions::Result<()> {
use fetch_options::SocketOptions;
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await?;
let port = listener.local_addr()?.port();
let stream = connect_with(SocketOptions::default().no_delay(true), port).await?;
assert!(stream.nodelay()?);
let stream = connect_with(SocketOptions::default().no_delay(false), port).await?;
assert!(!stream.nodelay()?);
Ok(())
}
#[cfg_attr(miri, ignore)]
#[tokio::test]
async fn connector_connects_with_all_options_set() -> http_extensions::Result<()> {
use fetch_options::SocketOptions;
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await?;
let port = listener.local_addr()?.port();
let options = SocketOptions::default()
.receive_buffer_size(64 * 1024)
.send_buffer_size(64 * 1024)
.no_delay(true);
let stream = connect_with(options, port).await?;
assert_eq!(stream.peer_addr()?.port(), port);
assert!(stream.nodelay()?);
Ok(())
}
#[cfg_attr(miri, ignore)]
#[tokio::test]
async fn connector_accepts_https_targets() -> http_extensions::Result<()> {
use fetch_options::SocketOptions;
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await?;
let port = listener.local_addr()?.port();
let stream = connect_to(SocketOptions::default(), &format!("https://127.0.0.1:{port}")).await?;
assert_eq!(stream.peer_addr()?.port(), port);
Ok(())
}
#[cfg_attr(miri, ignore)]
#[tokio::test]
async fn connector_rejects_non_http_schemes() -> http_extensions::Result<()> {
use ohno::Labeled as _;
use seatbelt::{Recovery as _, RecoveryKind};
let error = connect_to(fetch_options::SocketOptions::default(), "ftp://127.0.0.1:21")
.await
.expect_err("the connector rejects every non-HTTP scheme before dialing");
assert_eq!(error.label().as_str(), "scheme_not_allowed");
assert!(
error.to_string().contains("'ftp'"),
"the diagnostic must identify the rejected scheme: {error}"
);
assert_eq!(error.recovery().kind(), RecoveryKind::Never);
Ok(())
}
#[cfg_attr(miri, ignore)]
#[tokio::test]
async fn connector_surfaces_connect_failure_as_recoverable_io_error() -> http_extensions::Result<()> {
use seatbelt::{Recovery as _, RecoveryInfo};
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await?;
let port = listener.local_addr()?.port();
drop(listener);
let error = connect_with(fetch_options::SocketOptions::default().send_buffer_size(64 * 1024), port)
.await
.expect_err("the listener was dropped before dialing, so the connection must be refused");
assert!(
std::error::Error::source(&error).is_some(),
"the underlying I/O cause must be preserved, got: {error}"
);
assert_eq!(
error.recovery().kind(),
RecoveryInfo::from(std::io::ErrorKind::ConnectionRefused).kind(),
"the refused connection must keep the recovery classification of a plain I/O error"
);
Ok(())
}
#[derive(Debug)]
struct WrappingError(std::io::Error);
impl std::fmt::Display for WrappingError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str("tcp connect error")
}
}
impl std::error::Error for WrappingError {
fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
Some(&self.0)
}
}
#[cfg_attr(miri, ignore)]
#[test]
fn io_error_kind_walks_the_source_chain() {
let wrapped = WrappingError(std::io::Error::new(std::io::ErrorKind::TimedOut, "timed out"));
assert_eq!(super::io_error_kind(&wrapped), std::io::ErrorKind::TimedOut);
let direct = std::io::Error::new(std::io::ErrorKind::ConnectionRefused, "refused");
assert_eq!(super::io_error_kind(&direct), std::io::ErrorKind::ConnectionRefused);
let opaque = http::Uri::try_from("::not a uri::").unwrap_err();
assert_eq!(super::io_error_kind(&opaque), std::io::ErrorKind::Other);
}
}