use std::{
fmt,
future::Future,
net::{IpAddr, SocketAddr},
pin::Pin,
str::FromStr,
sync::{
Arc,
atomic::{AtomicBool, Ordering},
},
task::{Context, Poll},
time::{Duration, Instant},
};
use bytes::Bytes;
use tower::Service;
use super::{Connected, Connection};
pub use crate::error::H3Error;
use crate::{
client::{core::http3::Http3Options, http::ConnectRequest},
dns::{self, Name, resolve},
};
type H3Future = Pin<Box<dyn Future<Output = Result<H3Connection, H3Error>> + Send>>;
pub struct H3Connection {
pub send_request: hpx_h3::client::SendRequest<hpx_h3_quinn::OpenStreams, Bytes>,
pub close_rx: tokio::sync::mpsc::Receiver<hpx_h3::error::ConnectionError>,
pub idle_at: Instant,
pub is_broken: Arc<AtomicBool>,
pub used_0rtt: Arc<AtomicBool>,
}
impl fmt::Debug for H3Connection {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("H3Connection")
.field("send_request", &"<hpx_h3::client::SendRequest>")
.field("close_rx", &self.close_rx)
.field("idle_at", &self.idle_at)
.field("is_broken", &self.is_broken.load(Ordering::Relaxed))
.field("used_0rtt", &self.used_0rtt.load(Ordering::Relaxed))
.finish()
}
}
impl H3Connection {
pub fn is_valid(&self, idle_timeout: Duration) -> bool {
if self.is_broken.load(Ordering::Acquire) {
return false;
}
if idle_timeout == Duration::MAX {
return true;
}
Instant::now().saturating_duration_since(self.idle_at) <= idle_timeout
}
#[must_use]
pub fn used_0rtt(&self) -> bool {
self.used_0rtt.load(Ordering::Acquire)
}
}
impl Connection for H3Connection {
fn connected(&self) -> Connected {
Connected::new()
}
}
#[derive(Clone)]
pub struct QuicConnector {
endpoint: quinn::Endpoint,
transport_config: Arc<quinn::TransportConfig>,
tls_config: Arc<rustls::ClientConfig>,
h3_options: Http3Options,
}
impl fmt::Debug for QuicConnector {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("QuicConnector")
.field("endpoint", &self.endpoint)
.field("transport_config", &self.transport_config)
.field("tls_config", &self.tls_config)
.field("h3_options", &self.h3_options)
.finish_non_exhaustive()
}
}
impl QuicConnector {
#[must_use]
pub fn new(
endpoint: quinn::Endpoint,
transport_config: Arc<quinn::TransportConfig>,
tls_config: Arc<rustls::ClientConfig>,
h3_options: Http3Options,
) -> Self {
Self {
endpoint,
transport_config,
tls_config,
h3_options,
}
}
async fn drive_connection(
mut conn: hpx_h3::client::Connection<hpx_h3_quinn::Connection, Bytes>,
close_tx: tokio::sync::mpsc::Sender<hpx_h3::error::ConnectionError>,
is_broken: Arc<AtomicBool>,
max_idle_timeout: Option<Duration>,
) {
let err = match max_idle_timeout {
Some(idle_dur) => {
tokio::select! {
err = futures_util::future::poll_fn(|cx| conn.poll_close(cx)) => {
err
}
() = tokio::time::sleep(idle_dur) => {
trace!("h3 idle timeout expired, sending GOAWAY");
let _ = conn.shutdown(0).await;
futures_util::future::poll_fn(|cx| conn.poll_close(cx)).await
}
}
}
None => {
futures_util::future::poll_fn(|cx| conn.poll_close(cx)).await
}
};
is_broken.store(true, Ordering::Release);
let _ = close_tx.send(err).await;
}
}
impl Service<ConnectRequest> for QuicConnector {
type Response = H3Connection;
type Error = H3Error;
type Future = H3Future;
fn poll_ready(&mut self, _cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
Poll::Ready(Ok(()))
}
fn call(&mut self, req: ConnectRequest) -> Self::Future {
let endpoint = self.endpoint.clone();
let transport_config = self.transport_config.clone();
let tls_config = self.tls_config.clone();
let h3_options = self.h3_options.clone();
Box::pin(async move {
let uri = req.uri().clone();
let host = uri
.host()
.ok_or_else(|| -> H3Error { H3Error::Other("URI missing host".into()) })?;
let host = host.trim_start_matches('[').trim_end_matches(']');
let port = uri.port_u16().unwrap_or(443);
let addrs: Vec<SocketAddr> = if let Ok(ip) = IpAddr::from_str(host) {
vec![SocketAddr::new(ip, port)]
} else {
let mut resolver = dns::GaiResolver::new();
let name = Name::new(host.into());
let resolved = resolve(&mut resolver, name)
.await
.map_err(|e| H3Error::Other(Box::new(e)))?;
resolved
.map(|mut addr| {
addr.set_port(port);
addr
})
.collect()
};
if addrs.is_empty() {
return Err(H3Error::Other("no addresses resolved".into()));
}
let quic_client_config = quinn::crypto::rustls::QuicClientConfig::try_from(tls_config)
.map_err(|e| H3Error::Other(Box::new(e)))?;
let mut client_config = quinn::ClientConfig::new(Arc::new(quic_client_config));
client_config.transport_config(transport_config);
let mut last_err: Option<H3Error> = None;
for addr in &addrs {
let connecting = match endpoint.connect_with(client_config.clone(), *addr, host) {
Ok(c) => c,
Err(source) => {
last_err = Some(H3Error::Other(Box::new(source)));
continue;
}
};
let (quinn_conn, used_0rtt) = match connecting.into_0rtt() {
Ok((conn, accepted)) => {
let used = Arc::new(AtomicBool::new(false));
let used_clone = used.clone();
drop(tokio::spawn(async move {
used_clone.store(accepted.await, Ordering::Release);
}));
(conn, used)
}
Err(connecting) => {
match connecting.await {
Ok(quinn_conn) => (quinn_conn, Arc::new(AtomicBool::new(false))),
Err(source) => {
last_err = Some(H3Error::Handshake { source });
continue;
}
}
}
};
let h3_quinn_conn = hpx_h3_quinn::Connection::new(quinn_conn);
let mut builder = hpx_h3::client::builder();
if let Some(max_field_section_size) = h3_options.max_field_section_size {
builder.max_field_section_size(max_field_section_size);
}
builder.send_grease(h3_options.send_grease);
builder.enable_extended_connect(h3_options.enable_connect_protocol);
let (h3_conn, send_request) = builder
.build(h3_quinn_conn)
.await
.map_err(|e| H3Error::Framing { source: e })?;
let (close_tx, close_rx) =
tokio::sync::mpsc::channel::<hpx_h3::error::ConnectionError>(1);
let is_broken = Arc::new(AtomicBool::new(false));
drop(tokio::spawn(Self::drive_connection(
h3_conn,
close_tx,
is_broken.clone(),
h3_options.max_idle_timeout,
)));
return Ok(H3Connection {
send_request,
close_rx,
idle_at: Instant::now(),
is_broken,
used_0rtt,
});
}
match last_err {
Some(err) => Err(err),
None => Err(H3Error::Other("no addresses available".into())),
}
})
}
}
#[cfg(test)]
mod tests {
use std::{
net::SocketAddr,
sync::Arc,
task::{Context, Poll},
time::Duration,
};
use http::Uri;
use quinn::{Endpoint, TransportConfig};
use rustls::ClientConfig;
use tower::Service;
use super::*;
use crate::client::{core::http3::Http3Options, http::ConnectRequest};
type TestResult<T> = Result<T, Box<dyn std::error::Error + Send + Sync>>;
fn make_test_tls_config() -> TestResult<Arc<ClientConfig>> {
let provider = Arc::new(rustls::crypto::ring::default_provider());
let config = ClientConfig::builder_with_provider(provider)
.with_safe_default_protocol_versions()
.map_err(|e| -> Box<dyn std::error::Error + Send + Sync> {
format!("bad protocol versions: {e}").into()
})?
.with_root_certificates(rustls::RootCertStore::empty())
.with_no_client_auth();
Ok(Arc::new(config))
}
fn make_test_endpoint() -> TestResult<Endpoint> {
let addr: SocketAddr =
"127.0.0.1:0"
.parse()
.map_err(|e| -> Box<dyn std::error::Error + Send + Sync> {
format!("bad addr: {e}").into()
})?;
Ok(Endpoint::client(addr)?)
}
#[tokio::test]
async fn quic_connector_new_constructs() -> TestResult<()> {
let endpoint = make_test_endpoint()?;
let transport_config = Arc::new(TransportConfig::default());
let tls_config = make_test_tls_config()?;
let h3_options = Http3Options::default();
let _connector = QuicConnector::new(endpoint, transport_config, tls_config, h3_options);
Ok(())
}
#[tokio::test]
async fn quic_connector_poll_ready_returns_ok() -> TestResult<()> {
let endpoint = make_test_endpoint()?;
let transport_config = Arc::new(TransportConfig::default());
let tls_config = make_test_tls_config()?;
let h3_options = Http3Options::default();
let mut connector = QuicConnector::new(endpoint, transport_config, tls_config, h3_options);
let waker = futures_util::task::noop_waker();
let mut cx = Context::from_waker(&waker);
let poll = <QuicConnector as Service<ConnectRequest>>::poll_ready(&mut connector, &mut cx);
match poll {
Poll::Ready(Ok(())) => Ok(()),
other => Err(format!("expected Poll::Ready(Ok(())), got {other:?}").into()),
}
}
#[test]
fn h3_error_handshake_display() -> TestResult<()> {
let err = H3Error::Handshake {
source: quinn::ConnectionError::TimedOut,
};
let s = format!("{err}");
assert!(
s.contains("QUIC handshake failed"),
"expected Display to mention 'QUIC handshake failed', got: {s}"
);
assert!(
s.contains("TimedOut"),
"expected Display to mention the underlying 'TimedOut', got: {s}"
);
Ok(())
}
#[tokio::test]
async fn quic_connector_call_with_closed_endpoint_returns_error() -> TestResult<()> {
let endpoint = make_test_endpoint()?;
endpoint.close(quinn::VarInt::from(0u32), &[]);
let transport_config = Arc::new(TransportConfig::default());
let tls_config = make_test_tls_config()?;
let h3_options = Http3Options::default();
let mut connector = QuicConnector::new(endpoint, transport_config, tls_config, h3_options);
let uri: Uri = "https://127.0.0.1:443".parse()?;
let req = ConnectRequest::new(uri, None);
let waker = futures_util::task::noop_waker();
let mut cx = Context::from_waker(&waker);
let poll = <QuicConnector as Service<ConnectRequest>>::poll_ready(&mut connector, &mut cx);
match poll {
Poll::Ready(Ok(())) => {}
other => return Err(format!("poll_ready should be Ok, got {other:?}").into()),
}
let fut = connector.call(req);
let result = tokio::time::timeout(Duration::from_secs(2), fut).await;
let conn_result = match result {
Ok(r) => r,
Err(_) => return Err("call future should resolve within 2s".into()),
};
assert!(
conn_result.is_err(),
"expected an error from call() with a closed endpoint"
);
Ok(())
}
}