mod support;
use std::{
error::Error as _,
net::SocketAddr,
sync::{
Arc,
atomic::{AtomicU64, Ordering},
},
time::Duration,
};
use bytes::{Buf, Bytes};
use hpx::{
Client,
http3::{__test_connect_request, H3Error, Http3Options, QuicConfig, QuicConnector},
tls::{
TlsOptions,
quic::{build_quinn_client_config_with_root_store, build_quinn_endpoint},
},
};
use hpx_h3::error::Code;
use quinn::{Endpoint, ServerConfig, crypto::rustls::QuicServerConfig};
use rustls::{
RootCertStore, ServerConfig as RustlsServerConfig,
pki_types::{CertificateDer, PrivateKeyDer, PrivatePkcs8KeyDer},
};
type TestResult<T> = Result<T, Box<dyn std::error::Error + Send + Sync>>;
#[tokio::test]
async fn http3_alpn_negotiated_over_quic() -> TestResult<()> {
let certified_key = rcgen::generate_simple_self_signed(vec!["127.0.0.1".to_string()])?;
let cert_der: CertificateDer<'static> = certified_key.cert.der().clone();
let key_der: PrivateKeyDer<'static> =
PrivatePkcs8KeyDer::from(certified_key.signing_key.serialize_der()).into();
let provider = Arc::new(rustls::crypto::ring::default_provider());
let mut server_rustls_config = RustlsServerConfig::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_no_client_auth()
.with_single_cert(vec![cert_der.clone()], key_der)?;
server_rustls_config.alpn_protocols = vec![b"h3".to_vec()];
let server_quic_config = QuicServerConfig::try_from(Arc::new(server_rustls_config))?;
let server_config = ServerConfig::with_crypto(Arc::new(server_quic_config));
let server_endpoint = Endpoint::server(server_config, "127.0.0.1:0".parse()?)?;
let server_addr: SocketAddr = server_endpoint.local_addr()?;
let server_task = tokio::spawn(async move {
while let Some(incoming) = server_endpoint.accept().await {
match incoming.await {
Ok(_conn) => {
}
Err(_) => break,
}
}
});
let _ = server_task;
let mut root_store = RootCertStore::empty();
root_store.add(cert_der.clone())?;
let tls_opts = TlsOptions::default();
let h3_opts = Http3Options::default();
let client_config = build_quinn_client_config_with_root_store(&tls_opts, &h3_opts, root_store)?;
let client_endpoint = build_quinn_endpoint("127.0.0.1:0".parse()?)?;
let conn = client_endpoint
.connect_with(client_config, server_addr, "127.0.0.1")?
.await?;
let handshake = conn.handshake_data().ok_or("no handshake data available")?;
let handshake = handshake
.downcast_ref::<quinn::crypto::rustls::HandshakeData>()
.ok_or("failed to downcast handshake data to rustls::HandshakeData")?;
let alpn = handshake.protocol.as_ref().ok_or("no ALPN negotiated")?;
assert_eq!(
alpn.as_slice(),
b"h3",
"expected negotiated ALPN to be exactly b\"h3\", got {:?}",
alpn
);
conn.close(quinn::VarInt::from(0u32), &[]);
Ok(())
}
#[test]
fn http3_options_configures_h3_settings() -> TestResult<()> {
let builder = Client::builder().http3_only();
assert!(
builder.is_http3_only(),
"http3_only() must force HttpVersionPref::Http3"
);
let builder_alias = Client::builder().http3_prior_knowledge();
assert!(
builder_alias.is_http3_only(),
"http3_prior_knowledge() must be equivalent to http3_only()"
);
let mut opts = Http3Options::default();
opts.max_field_section_size = Some(64 * 1024);
opts.qpack_max_table_capacity = Some(8192);
let builder = Client::builder().http3_only().http3_options(opts.clone());
let stored = builder
.http3_options_ref()
.ok_or("http3_options_ref() must return Some after http3_options(_) is set")?;
assert_eq!(
stored.max_field_section_size,
Some(64 * 1024),
"http3_options() must round-trip the stored options verbatim"
);
assert_eq!(
stored.qpack_max_table_capacity,
Some(8192),
"http3_options() must round-trip the stored options verbatim"
);
let quic_cfg = QuicConfig::default();
let builder = Client::builder().http3_only().quic_config(quic_cfg);
assert!(
builder.quic_config_ref().is_some(),
"quic_config_ref() must return Some after quic_config(_) is set"
);
let plain = Client::builder();
assert!(
!plain.is_http3_only(),
"default ClientBuilder must not be http3-only"
);
assert!(
plain.http3_options_ref().is_none(),
"default ClientBuilder must not have http3_options set"
);
assert!(
plain.quic_config_ref().is_none(),
"default ClientBuilder must not have quic_config set"
);
Ok(())
}
#[tokio::test]
async fn http3_concurrent_requests_over_single_quic_connection() -> TestResult<()> {
use tower::Service;
let certified_key = rcgen::generate_simple_self_signed(vec!["127.0.0.1".to_string()])?;
let cert_der: CertificateDer<'static> = certified_key.cert.der().clone();
let key_der: PrivateKeyDer<'static> =
PrivatePkcs8KeyDer::from(certified_key.signing_key.serialize_der()).into();
let provider = Arc::new(rustls::crypto::ring::default_provider());
let mut server_rustls_config = RustlsServerConfig::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_no_client_auth()
.with_single_cert(vec![cert_der.clone()], key_der)?;
server_rustls_config.alpn_protocols = vec![b"h3".to_vec()];
let server_quic_config = QuicServerConfig::try_from(Arc::new(server_rustls_config))?;
let server_config = ServerConfig::with_crypto(Arc::new(server_quic_config));
let server_endpoint = Endpoint::server(server_config, "127.0.0.1:0".parse()?)?;
let server_addr: SocketAddr = server_endpoint.local_addr()?;
let conn_count = Arc::new(AtomicU64::new(0));
let req_count = Arc::new(AtomicU64::new(0));
let server_task = {
let conn_count = conn_count.clone();
let req_count = req_count.clone();
tokio::spawn(async move {
while let Some(incoming) = server_endpoint.accept().await {
match incoming.await {
Ok(quinn_conn) => {
conn_count.fetch_add(1, Ordering::SeqCst);
let req_count = req_count.clone();
tokio::spawn(async move {
let h3_quinn_conn = hpx_h3_quinn::Connection::new(quinn_conn);
let mut h3_conn: hpx_h3::server::Connection<
hpx_h3_quinn::Connection,
bytes::Bytes,
> = match hpx_h3::server::Connection::new(h3_quinn_conn).await {
Ok(c) => c,
Err(_) => return,
};
loop {
match h3_conn.accept().await {
Ok(Some(resolver)) => {
req_count.fetch_add(1, Ordering::SeqCst);
let (_req, mut stream) =
match resolver.resolve_request().await {
Ok(parts) => parts,
Err(_) => continue,
};
let resp = match http::Response::builder()
.status(http::StatusCode::OK)
.body(())
{
Ok(r) => r,
Err(_) => continue,
};
if stream.send_response(resp).await.is_err() {
continue;
}
while matches!(stream.recv_data().await, Ok(Some(_))) {
}
let _ = stream.finish().await;
}
Ok(None) => break,
Err(_) => break,
}
}
});
}
Err(_) => break,
}
}
})
};
let _ = server_task;
let mut root_store = RootCertStore::empty();
root_store.add(cert_der.clone())?;
let client_provider = Arc::new(rustls::crypto::ring::default_provider());
let mut client_tls_config = rustls::ClientConfig::builder_with_provider(client_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(root_store)
.with_no_client_auth();
client_tls_config.alpn_protocols = vec![b"h3".to_vec()];
let tls_config: Arc<rustls::ClientConfig> = Arc::new(client_tls_config);
let client_endpoint = build_quinn_endpoint("127.0.0.1:0".parse()?)?;
let transport_config = Arc::new(quinn::TransportConfig::default());
let mut connector = QuicConnector::new(
client_endpoint,
transport_config,
tls_config,
Http3Options::default(),
);
let uri: http::Uri = format!("https://127.0.0.1:{}/", server_addr.port()).parse()?;
let connect_req = __test_connect_request(uri);
let waker = futures_util::task::noop_waker();
let mut cx = std::task::Context::from_waker(&waker);
match connector.poll_ready(&mut cx) {
std::task::Poll::Ready(Ok(())) => {}
other => {
return Err(format!("poll_ready should be Ok, got {other:?}").into());
}
}
let h3_conn = tokio::time::timeout(Duration::from_secs(5), connector.call(connect_req))
.await
.map_err(|_| -> Box<dyn std::error::Error + Send + Sync> {
"QuicConnector::call should resolve within 5s".into()
})??;
assert!(
h3_conn.is_valid(Duration::MAX),
"fresh H3Connection should be valid (no close event, not idle-expired)"
);
let mut send_requests: Vec<
hpx_h3::client::SendRequest<hpx_h3_quinn::OpenStreams, bytes::Bytes>,
> = Vec::with_capacity(10);
for _ in 0..10 {
send_requests.push(h3_conn.send_request.clone());
}
let port = server_addr.port();
let mut tasks = Vec::with_capacity(10);
for mut send_req in send_requests.drain(..) {
tasks.push(tokio::spawn(async move {
let req = http::Request::get(format!("https://127.0.0.1:{port}/")).body(())?;
let mut stream = send_req.send_request(req).await?;
stream.finish().await?;
let _resp = stream.recv_response().await?;
while let Some(_data) = stream.recv_data().await? {}
Ok::<_, Box<dyn std::error::Error + Send + Sync>>(())
}));
}
for task in tasks {
tokio::time::timeout(Duration::from_secs(10), task)
.await
.map_err(|_| -> Box<dyn std::error::Error + Send + Sync> {
"concurrent request task should complete within 10s".into()
})???;
}
tokio::time::sleep(Duration::from_millis(100)).await;
let observed_conns = conn_count.load(Ordering::SeqCst);
let observed_reqs = req_count.load(Ordering::SeqCst);
assert_eq!(
observed_conns, 1,
"expected exactly 1 QUIC connection for 10 concurrent requests, got {observed_conns}"
);
assert_eq!(
observed_reqs, 10,
"expected 10 requests served over 1 connection, got {observed_reqs}"
);
Ok(())
}
#[tokio::test]
async fn http3_pool_invalid_connection_triggers_reconnect() -> TestResult<()> {
use tower::Service;
let certified_key = rcgen::generate_simple_self_signed(vec!["127.0.0.1".to_string()])?;
let cert_der: CertificateDer<'static> = certified_key.cert.der().clone();
let key_der: PrivateKeyDer<'static> =
PrivatePkcs8KeyDer::from(certified_key.signing_key.serialize_der()).into();
let provider = Arc::new(rustls::crypto::ring::default_provider());
let mut server_rustls_config = RustlsServerConfig::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_no_client_auth()
.with_single_cert(vec![cert_der.clone()], key_der)?;
server_rustls_config.alpn_protocols = vec![b"h3".to_vec()];
let server_quic_config = QuicServerConfig::try_from(Arc::new(server_rustls_config))?;
let server_config = ServerConfig::with_crypto(Arc::new(server_quic_config));
let server_endpoint = Endpoint::server(server_config, "127.0.0.1:0".parse()?)?;
let server_addr: SocketAddr = server_endpoint.local_addr()?;
let server_task = tokio::spawn(async move {
while let Some(incoming) = server_endpoint.accept().await {
match incoming.await {
Ok(quinn_conn) => {
tokio::spawn(async move {
let h3_quinn_conn = hpx_h3_quinn::Connection::new(quinn_conn);
let mut h3_conn: hpx_h3::server::Connection<
hpx_h3_quinn::Connection,
bytes::Bytes,
> = match hpx_h3::server::Connection::new(h3_quinn_conn).await {
Ok(c) => c,
Err(_) => return,
};
loop {
match h3_conn.accept().await {
Ok(Some(resolver)) => {
tokio::spawn(async move {
if let Ok((_req, mut stream)) =
resolver.resolve_request().await
{
let resp = match http::Response::builder()
.status(http::StatusCode::OK)
.body(())
{
Ok(r) => r,
Err(_) => return,
};
if stream.send_response(resp).await.is_ok() {
let _ = stream.finish().await;
}
}
});
}
Ok(None) => break,
Err(_) => break,
}
}
});
}
Err(_) => break,
}
}
});
let _ = server_task;
let mut root_store = RootCertStore::empty();
root_store.add(cert_der.clone())?;
let client_provider = Arc::new(rustls::crypto::ring::default_provider());
let mut client_tls_config = rustls::ClientConfig::builder_with_provider(client_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(root_store)
.with_no_client_auth();
client_tls_config.alpn_protocols = vec![b"h3".to_vec()];
let tls_config: Arc<rustls::ClientConfig> = Arc::new(client_tls_config);
let client_endpoint = build_quinn_endpoint("127.0.0.1:0".parse()?)?;
let endpoint_close_handle = client_endpoint.clone();
let transport_config = Arc::new(quinn::TransportConfig::default());
let mut connector = QuicConnector::new(
client_endpoint,
transport_config,
tls_config,
Http3Options::default(),
);
let uri: http::Uri = format!("https://127.0.0.1:{}/", server_addr.port()).parse()?;
let connect_req = __test_connect_request(uri);
let waker = futures_util::task::noop_waker();
let mut cx = std::task::Context::from_waker(&waker);
match connector.poll_ready(&mut cx) {
std::task::Poll::Ready(Ok(())) => {}
other => {
return Err(format!("poll_ready should be Ok, got {other:?}").into());
}
}
let h3_conn = tokio::time::timeout(Duration::from_secs(5), connector.call(connect_req))
.await
.map_err(|_| -> Box<dyn std::error::Error + Send + Sync> {
"QuicConnector::call should resolve within 5s".into()
})??;
assert!(
h3_conn.is_valid(Duration::MAX),
"fresh H3Connection should be valid (no close event, not idle-expired)"
);
endpoint_close_handle.close(quinn::VarInt::from(0u32), &[]);
let mut observed_invalid = false;
let deadline = tokio::time::Instant::now() + Duration::from_secs(5);
while tokio::time::Instant::now() < deadline {
if !h3_conn.is_valid(Duration::MAX) {
observed_invalid = true;
break;
}
tokio::time::sleep(Duration::from_millis(50)).await;
}
assert!(
observed_invalid,
"is_valid should return false within 5s of the connection breaking"
);
Ok(())
}
#[tokio::test]
async fn http3_request_full() -> TestResult<()> {
use tower::Service;
let certified_key = rcgen::generate_simple_self_signed(vec!["127.0.0.1".to_string()])?;
let cert_der: CertificateDer<'static> = certified_key.cert.der().clone();
let key_der: PrivateKeyDer<'static> =
PrivatePkcs8KeyDer::from(certified_key.signing_key.serialize_der()).into();
let provider = Arc::new(rustls::crypto::ring::default_provider());
let mut server_rustls_config = RustlsServerConfig::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_no_client_auth()
.with_single_cert(vec![cert_der.clone()], key_der)?;
server_rustls_config.alpn_protocols = vec![b"h3".to_vec()];
let server_quic_config = QuicServerConfig::try_from(Arc::new(server_rustls_config))?;
let server_config = ServerConfig::with_crypto(Arc::new(server_quic_config));
let server_endpoint = Endpoint::server(server_config, "127.0.0.1:0".parse()?)?;
let server_addr: SocketAddr = server_endpoint.local_addr()?;
let server_task = tokio::spawn(async move {
while let Some(incoming) = server_endpoint.accept().await {
match incoming.await {
Ok(quinn_conn) => {
tokio::spawn(async move {
let h3_quinn_conn = hpx_h3_quinn::Connection::new(quinn_conn);
let mut h3_conn: hpx_h3::server::Connection<
hpx_h3_quinn::Connection,
bytes::Bytes,
> = match hpx_h3::server::Connection::new(h3_quinn_conn).await {
Ok(c) => c,
Err(_) => return,
};
loop {
match h3_conn.accept().await {
Ok(Some(resolver)) => {
let (_req, mut stream) = match resolver.resolve_request().await
{
Ok(parts) => parts,
Err(_) => continue,
};
let resp = match http::Response::builder()
.status(http::StatusCode::OK)
.body(())
{
Ok(r) => r,
Err(_) => continue,
};
if stream.send_response(resp).await.is_err() {
continue;
}
if stream
.send_data(bytes::Bytes::from_static(b"hello, h3"))
.await
.is_err()
{
continue;
}
while matches!(stream.recv_data().await, Ok(Some(_))) {}
let _ = stream.finish().await;
}
Ok(None) => break,
Err(_) => break,
}
}
});
}
Err(_) => break,
}
}
});
let _ = server_task;
let mut root_store = RootCertStore::empty();
root_store.add(cert_der.clone())?;
let client_provider = Arc::new(rustls::crypto::ring::default_provider());
let mut client_tls_config = rustls::ClientConfig::builder_with_provider(client_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(root_store)
.with_no_client_auth();
client_tls_config.alpn_protocols = vec![b"h3".to_vec()];
let tls_config: Arc<rustls::ClientConfig> = Arc::new(client_tls_config);
let client_endpoint = build_quinn_endpoint("127.0.0.1:0".parse()?)?;
let transport_config = Arc::new(quinn::TransportConfig::default());
let mut connector = QuicConnector::new(
client_endpoint,
transport_config,
tls_config,
Http3Options::default(),
);
let waker = futures_util::task::noop_waker();
let mut cx = std::task::Context::from_waker(&waker);
match connector.poll_ready(&mut cx) {
std::task::Poll::Ready(Ok(())) => {}
other => {
return Err(format!("poll_ready should be Ok, got {other:?}").into());
}
}
let client = Client::builder()
.http3_only()
.__test_with_quic_connector(connector)
.build()?;
let url = format!("https://127.0.0.1:{}/hello", server_addr.port());
let response = tokio::time::timeout(Duration::from_secs(10), client.get(url).send())
.await
.map_err(|_| -> Box<dyn std::error::Error + Send + Sync> {
"GET request should resolve within 10s".into()
})??;
assert_eq!(
response.status(),
http::StatusCode::OK,
"expected 200 OK from h3 GET request"
);
assert_eq!(
response.version(),
http::Version::HTTP_3,
"expected HTTP/3 version from h3 GET request"
);
let body = response.bytes().await?;
assert_eq!(
body.as_ref(),
b"hello, h3",
"expected response body to be 'hello, h3', got {:?}",
std::str::from_utf8(body.as_ref()).unwrap_or("<non-utf8>")
);
Ok(())
}
#[tokio::test]
async fn http3_post_with_body() -> TestResult<()> {
use tower::Service;
let certified_key = rcgen::generate_simple_self_signed(vec!["127.0.0.1".to_string()])?;
let cert_der: CertificateDer<'static> = certified_key.cert.der().clone();
let key_der: PrivateKeyDer<'static> =
PrivatePkcs8KeyDer::from(certified_key.signing_key.serialize_der()).into();
let provider = Arc::new(rustls::crypto::ring::default_provider());
let mut server_rustls_config = RustlsServerConfig::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_no_client_auth()
.with_single_cert(vec![cert_der.clone()], key_der)?;
server_rustls_config.alpn_protocols = vec![b"h3".to_vec()];
let server_quic_config = QuicServerConfig::try_from(Arc::new(server_rustls_config))?;
let server_config = ServerConfig::with_crypto(Arc::new(server_quic_config));
let server_endpoint = Endpoint::server(server_config, "127.0.0.1:0".parse()?)?;
let server_addr: SocketAddr = server_endpoint.local_addr()?;
let server_task = tokio::spawn(async move {
while let Some(incoming) = server_endpoint.accept().await {
match incoming.await {
Ok(quinn_conn) => {
tokio::spawn(async move {
let h3_quinn_conn = hpx_h3_quinn::Connection::new(quinn_conn);
let mut h3_conn: hpx_h3::server::Connection<
hpx_h3_quinn::Connection,
bytes::Bytes,
> = match hpx_h3::server::Connection::new(h3_quinn_conn).await {
Ok(c) => c,
Err(_) => return,
};
loop {
match h3_conn.accept().await {
Ok(Some(resolver)) => {
let (_req, mut stream) = match resolver.resolve_request().await
{
Ok(parts) => parts,
Err(_) => continue,
};
let mut body_buf = bytes::BytesMut::new();
while let Ok(Some(data)) = stream.recv_data().await {
let len = data.remaining();
body_buf.extend_from_slice(&data.chunk()[..len]);
}
let body_bytes: bytes::Bytes = body_buf.freeze();
let resp = match http::Response::builder()
.status(http::StatusCode::OK)
.body(())
{
Ok(r) => r,
Err(_) => continue,
};
if stream.send_response(resp).await.is_err() {
continue;
}
if !body_bytes.is_empty()
&& stream.send_data(body_bytes).await.is_err()
{
continue;
}
let _ = stream.finish().await;
}
Ok(None) => break,
Err(_) => break,
}
}
});
}
Err(_) => break,
}
}
});
let _ = server_task;
let mut root_store = RootCertStore::empty();
root_store.add(cert_der.clone())?;
let client_provider = Arc::new(rustls::crypto::ring::default_provider());
let mut client_tls_config = rustls::ClientConfig::builder_with_provider(client_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(root_store)
.with_no_client_auth();
client_tls_config.alpn_protocols = vec![b"h3".to_vec()];
let tls_config: Arc<rustls::ClientConfig> = Arc::new(client_tls_config);
let client_endpoint = build_quinn_endpoint("127.0.0.1:0".parse()?)?;
let transport_config = Arc::new(quinn::TransportConfig::default());
let mut connector = QuicConnector::new(
client_endpoint,
transport_config,
tls_config,
Http3Options::default(),
);
let waker = futures_util::task::noop_waker();
let mut cx = std::task::Context::from_waker(&waker);
match connector.poll_ready(&mut cx) {
std::task::Poll::Ready(Ok(())) => {}
other => {
return Err(format!("poll_ready should be Ok, got {other:?}").into());
}
}
let client = Client::builder()
.http3_only()
.__test_with_quic_connector(connector)
.build()?;
let url = format!("https://127.0.0.1:{}/echo", server_addr.port());
let response = tokio::time::timeout(
Duration::from_secs(10),
client.post(url).body(hpx::Body::from("ping")).send(),
)
.await
.map_err(|_| -> Box<dyn std::error::Error + Send + Sync> {
"POST request should resolve within 10s".into()
})??;
assert_eq!(
response.status(),
http::StatusCode::OK,
"expected 200 OK from h3 POST request"
);
assert_eq!(
response.version(),
http::Version::HTTP_3,
"expected HTTP/3 version from h3 POST request"
);
let body = response.bytes().await?;
assert_eq!(
body.as_ref(),
b"ping",
"expected response body to be 'ping', got {:?}",
std::str::from_utf8(body.as_ref()).unwrap_or("<non-utf8>")
);
Ok(())
}
#[tokio::test]
async fn http3_streaming_request_body() -> TestResult<()> {
use std::{
pin::Pin,
task::{Context, Poll},
};
use bytes::Bytes;
use http_body::{Body as HttpBody, Frame, SizeHint};
use tower::Service;
struct TestStreamBody {
chunks: Vec<Bytes>,
pos: usize,
}
impl HttpBody for TestStreamBody {
type Data = Bytes;
type Error = Box<dyn std::error::Error + Send + Sync>;
fn poll_frame(
mut self: Pin<&mut Self>,
_cx: &mut Context<'_>,
) -> Poll<Option<Result<Frame<Self::Data>, Self::Error>>> {
if self.pos >= self.chunks.len() {
return Poll::Ready(None);
}
let chunk = self.chunks[self.pos].clone();
self.pos += 1;
Poll::Ready(Some(Ok(Frame::data(chunk))))
}
fn is_end_stream(&self) -> bool {
self.pos >= self.chunks.len()
}
fn size_hint(&self) -> SizeHint {
SizeHint::new()
}
}
let certified_key = rcgen::generate_simple_self_signed(vec!["127.0.0.1".to_string()])?;
let cert_der: CertificateDer<'static> = certified_key.cert.der().clone();
let key_der: PrivateKeyDer<'static> =
PrivatePkcs8KeyDer::from(certified_key.signing_key.serialize_der()).into();
let provider = Arc::new(rustls::crypto::ring::default_provider());
let mut server_rustls_config = RustlsServerConfig::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_no_client_auth()
.with_single_cert(vec![cert_der.clone()], key_der)?;
server_rustls_config.alpn_protocols = vec![b"h3".to_vec()];
let server_quic_config = QuicServerConfig::try_from(Arc::new(server_rustls_config))?;
let server_config = ServerConfig::with_crypto(Arc::new(server_quic_config));
let server_endpoint = Endpoint::server(server_config, "127.0.0.1:0".parse()?)?;
let server_addr: SocketAddr = server_endpoint.local_addr()?;
let server_task = tokio::spawn(async move {
while let Some(incoming) = server_endpoint.accept().await {
match incoming.await {
Ok(quinn_conn) => {
tokio::spawn(async move {
let h3_quinn_conn = hpx_h3_quinn::Connection::new(quinn_conn);
let mut h3_conn: hpx_h3::server::Connection<
hpx_h3_quinn::Connection,
bytes::Bytes,
> = match hpx_h3::server::Connection::new(h3_quinn_conn).await {
Ok(c) => c,
Err(_) => return,
};
loop {
match h3_conn.accept().await {
Ok(Some(resolver)) => {
let (_req, mut stream) = match resolver.resolve_request().await
{
Ok(parts) => parts,
Err(_) => continue,
};
let mut body_buf = bytes::BytesMut::new();
while let Ok(Some(data)) = stream.recv_data().await {
let len = data.remaining();
body_buf.extend_from_slice(&data.chunk()[..len]);
}
let body_bytes: bytes::Bytes = body_buf.freeze();
let resp = match http::Response::builder()
.status(http::StatusCode::OK)
.body(())
{
Ok(r) => r,
Err(_) => continue,
};
if stream.send_response(resp).await.is_err() {
continue;
}
if !body_bytes.is_empty()
&& stream.send_data(body_bytes).await.is_err()
{
continue;
}
let _ = stream.finish().await;
}
Ok(None) => break,
Err(_) => break,
}
}
});
}
Err(_) => break,
}
}
});
let _ = server_task;
let mut root_store = RootCertStore::empty();
root_store.add(cert_der.clone())?;
let client_provider = Arc::new(rustls::crypto::ring::default_provider());
let mut client_tls_config = rustls::ClientConfig::builder_with_provider(client_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(root_store)
.with_no_client_auth();
client_tls_config.alpn_protocols = vec![b"h3".to_vec()];
let tls_config: Arc<rustls::ClientConfig> = Arc::new(client_tls_config);
let client_endpoint = build_quinn_endpoint("127.0.0.1:0".parse()?)?;
let transport_config = Arc::new(quinn::TransportConfig::default());
let mut connector = QuicConnector::new(
client_endpoint,
transport_config,
tls_config,
Http3Options::default(),
);
let waker = futures_util::task::noop_waker();
let mut cx = std::task::Context::from_waker(&waker);
match connector.poll_ready(&mut cx) {
std::task::Poll::Ready(Ok(())) => {}
other => {
return Err(format!("poll_ready should be Ok, got {other:?}").into());
}
}
let client = Client::builder()
.http3_only()
.__test_with_quic_connector(connector)
.build()?;
let stream_body = TestStreamBody {
chunks: vec![
Bytes::from_static(b"foo"),
Bytes::from_static(b"bar"),
Bytes::from_static(b"baz"),
],
pos: 0,
};
let body = hpx::Body::wrap(stream_body);
let url = format!("https://127.0.0.1:{}/echo", server_addr.port());
let response =
tokio::time::timeout(Duration::from_secs(10), client.post(url).body(body).send())
.await
.map_err(|_| -> Box<dyn std::error::Error + Send + Sync> {
"POST streaming request should resolve within 10s".into()
})??;
assert_eq!(
response.status(),
http::StatusCode::OK,
"expected 200 OK from h3 streaming POST request"
);
assert_eq!(
response.version(),
http::Version::HTTP_3,
"expected HTTP/3 version from h3 streaming POST request"
);
let body = response.bytes().await?;
assert_eq!(
body.as_ref(),
b"foobarbaz",
"expected response body to be 'foobarbaz', got {:?}",
std::str::from_utf8(body.as_ref()).unwrap_or("<non-utf8>")
);
Ok(())
}
#[tokio::test]
async fn http3_reconnection_after_server_closes() -> TestResult<()> {
use tower::Service;
let certified_key = rcgen::generate_simple_self_signed(vec!["127.0.0.1".to_string()])?;
let cert_der: CertificateDer<'static> = certified_key.cert.der().clone();
let key_der: PrivateKeyDer<'static> =
PrivatePkcs8KeyDer::from(certified_key.signing_key.serialize_der()).into();
let provider = Arc::new(rustls::crypto::ring::default_provider());
let mut server_rustls_config = RustlsServerConfig::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_no_client_auth()
.with_single_cert(vec![cert_der.clone()], key_der)?;
server_rustls_config.alpn_protocols = vec![b"h3".to_vec()];
let server_quic_config = QuicServerConfig::try_from(Arc::new(server_rustls_config))?;
let server_config = ServerConfig::with_crypto(Arc::new(server_quic_config));
let server_endpoint = Endpoint::server(server_config, "127.0.0.1:0".parse()?)?;
let server_addr: SocketAddr = server_endpoint.local_addr()?;
let server_port = server_addr.port();
let server_task = tokio::spawn(async move {
while let Some(incoming) = server_endpoint.accept().await {
match incoming.await {
Ok(quinn_conn) => {
tokio::spawn(async move {
let h3_quinn_conn = hpx_h3_quinn::Connection::new(quinn_conn);
let mut h3_conn: hpx_h3::server::Connection<
hpx_h3_quinn::Connection,
bytes::Bytes,
> = match hpx_h3::server::Connection::new(h3_quinn_conn).await {
Ok(c) => c,
Err(_) => return,
};
loop {
match h3_conn.accept().await {
Ok(Some(resolver)) => {
let (_req, mut stream) = match resolver.resolve_request().await
{
Ok(parts) => parts,
Err(_) => continue,
};
let resp = match http::Response::builder()
.status(http::StatusCode::OK)
.body(())
{
Ok(r) => r,
Err(_) => continue,
};
if stream.send_response(resp).await.is_err() {
continue;
}
if stream
.send_data(bytes::Bytes::from_static(b"hello, h3"))
.await
.is_err()
{
continue;
}
while matches!(stream.recv_data().await, Ok(Some(_))) {}
let _ = stream.finish().await;
}
Ok(None) => break,
Err(_) => break,
}
}
});
}
Err(_) => break,
}
}
});
let mut root_store = RootCertStore::empty();
root_store.add(cert_der.clone())?;
let client_provider = Arc::new(rustls::crypto::ring::default_provider());
let mut client_tls_config = rustls::ClientConfig::builder_with_provider(client_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(root_store)
.with_no_client_auth();
client_tls_config.alpn_protocols = vec![b"h3".to_vec()];
let tls_config: Arc<rustls::ClientConfig> = Arc::new(client_tls_config);
let client_endpoint = build_quinn_endpoint("127.0.0.1:0".parse()?)?;
let client_endpoint_close = client_endpoint.clone();
let transport_config = Arc::new(quinn::TransportConfig::default());
let mut connector = QuicConnector::new(
client_endpoint,
transport_config,
tls_config.clone(),
Http3Options::default(),
);
let waker = futures_util::task::noop_waker();
let mut cx = std::task::Context::from_waker(&waker);
match connector.poll_ready(&mut cx) {
std::task::Poll::Ready(Ok(())) => {}
other => {
return Err(format!("poll_ready should be Ok, got {other:?}").into());
}
}
let client = Client::builder()
.http3_only()
.__test_with_quic_connector(connector)
.build()?;
let url = format!("https://127.0.0.1:{server_port}/");
let response = tokio::time::timeout(Duration::from_secs(10), client.get(&url).send())
.await
.map_err(|_| -> Box<dyn std::error::Error + Send + Sync> {
"first GET request should resolve within 10s".into()
})??;
assert_eq!(
response.status(),
http::StatusCode::OK,
"first request should return 200 OK"
);
let body = response.bytes().await?;
assert_eq!(
body.as_ref(),
b"hello, h3",
"first request body should be 'hello, h3'"
);
client_endpoint_close.close(quinn::VarInt::from(0u32), &[]);
drop(client_endpoint_close);
tokio::time::sleep(Duration::from_millis(500)).await;
server_task.abort();
tokio::time::sleep(Duration::from_millis(200)).await;
let second_result = tokio::time::timeout(Duration::from_secs(5), client.get(&url).send())
.await
.map_err(|_| -> Box<dyn std::error::Error + Send + Sync> {
"second GET request should timeout within 5s".into()
})?;
let second_err = match second_result {
Err(e) => e,
Ok(_resp) => {
return Err("second request should fail after server is dropped".into());
}
};
assert!(
second_err.is_connect(),
"second request error should be is_connect() after server drops, got: {second_err}"
);
let client_endpoint2 = build_quinn_endpoint("127.0.0.1:0".parse()?)?;
let transport_config2 = Arc::new(quinn::TransportConfig::default());
let mut connector2 = QuicConnector::new(
client_endpoint2,
transport_config2,
Arc::clone(&tls_config),
Http3Options::default(),
);
let waker2 = futures_util::task::noop_waker();
let mut cx2 = std::task::Context::from_waker(&waker2);
match connector2.poll_ready(&mut cx2) {
std::task::Poll::Ready(Ok(())) => {}
other => {
return Err(format!("poll_ready should be Ok, got {other:?}").into());
}
}
let client2 = Client::builder()
.http3_only()
.__test_with_quic_connector(connector2)
.build()?;
let provider2 = Arc::new(rustls::crypto::ring::default_provider());
let mut server_rustls_config2 = RustlsServerConfig::builder_with_provider(provider2)
.with_safe_default_protocol_versions()
.map_err(|e| -> Box<dyn std::error::Error + Send + Sync> {
format!("bad protocol versions: {e}").into()
})?
.with_no_client_auth()
.with_single_cert(
vec![cert_der.clone()],
PrivatePkcs8KeyDer::from(certified_key.signing_key.serialize_der()).into(),
)?;
server_rustls_config2.alpn_protocols = vec![b"h3".to_vec()];
let _server_quic_config2 = QuicServerConfig::try_from(Arc::new(server_rustls_config2))?;
let server_config2 = ServerConfig::with_crypto(Arc::new(_server_quic_config2));
let server_addr2: SocketAddr = format!("127.0.0.1:{server_port}").parse()?;
let server_endpoint2 = Endpoint::server(server_config2, server_addr2)?;
let server_task2 = tokio::spawn(async move {
while let Some(incoming) = server_endpoint2.accept().await {
match incoming.await {
Ok(quinn_conn) => {
tokio::spawn(async move {
let h3_quinn_conn = hpx_h3_quinn::Connection::new(quinn_conn);
let mut h3_conn: hpx_h3::server::Connection<
hpx_h3_quinn::Connection,
bytes::Bytes,
> = match hpx_h3::server::Connection::new(h3_quinn_conn).await {
Ok(c) => c,
Err(_) => return,
};
loop {
match h3_conn.accept().await {
Ok(Some(resolver)) => {
let (_req, mut stream) = match resolver.resolve_request().await
{
Ok(parts) => parts,
Err(_) => continue,
};
let resp = match http::Response::builder()
.status(http::StatusCode::OK)
.body(())
{
Ok(r) => r,
Err(_) => continue,
};
if stream.send_response(resp).await.is_err() {
continue;
}
if stream
.send_data(bytes::Bytes::from_static(b"hello, h3"))
.await
.is_err()
{
continue;
}
while matches!(stream.recv_data().await, Ok(Some(_))) {}
let _ = stream.finish().await;
}
Ok(None) => break,
Err(_) => break,
}
}
});
}
Err(_) => break,
}
}
});
let _ = server_task2;
let response3 = tokio::time::timeout(Duration::from_secs(10), client2.get(&url).send())
.await
.map_err(|_| -> Box<dyn std::error::Error + Send + Sync> {
"third GET request should resolve within 10s".into()
})??;
assert_eq!(
response3.status(),
http::StatusCode::OK,
"third request should return 200 OK after pool reconnection"
);
let body3 = response3.bytes().await?;
assert_eq!(
body3.as_ref(),
b"hello, h3",
"third request body should be 'hello, h3'"
);
Ok(())
}
#[tokio::test]
async fn http3_stop_sending_no_error_graceful() -> TestResult<()> {
use tower::Service;
let certified_key = rcgen::generate_simple_self_signed(vec!["127.0.0.1".to_string()])?;
let cert_der: CertificateDer<'static> = certified_key.cert.der().clone();
let key_der: PrivateKeyDer<'static> =
PrivatePkcs8KeyDer::from(certified_key.signing_key.serialize_der()).into();
let provider = Arc::new(rustls::crypto::ring::default_provider());
let mut server_rustls_config = RustlsServerConfig::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_no_client_auth()
.with_single_cert(vec![cert_der.clone()], key_der)?;
server_rustls_config.alpn_protocols = vec![b"h3".to_vec()];
let server_quic_config = QuicServerConfig::try_from(Arc::new(server_rustls_config))?;
let server_config = ServerConfig::with_crypto(Arc::new(server_quic_config));
let server_endpoint = Endpoint::server(server_config, "127.0.0.1:0".parse()?)?;
let server_addr: SocketAddr = server_endpoint.local_addr()?;
let server_task = tokio::spawn(async move {
while let Some(incoming) = server_endpoint.accept().await {
let quinn_conn = match incoming.await {
Ok(c) => c,
Err(_) => break,
};
let h3_quinn_conn = hpx_h3_quinn::Connection::new(quinn_conn);
let mut h3_conn: hpx_h3::server::Connection<hpx_h3_quinn::Connection, bytes::Bytes> =
match hpx_h3::server::Connection::new(h3_quinn_conn).await {
Ok(c) => c,
Err(_) => continue,
};
loop {
match h3_conn.accept().await {
Ok(Some(resolver)) => {
let (_req, mut stream) = match resolver.resolve_request().await {
Ok(parts) => parts,
Err(_) => continue,
};
stream.stop_stream(Code::H3_NO_ERROR);
let _ = stream.finish().await;
}
Ok(None) => break,
Err(_) => break,
}
}
}
});
let _ = server_task;
let mut root_store = RootCertStore::empty();
root_store.add(cert_der.clone())?;
let client_provider = Arc::new(rustls::crypto::ring::default_provider());
let mut client_tls_config = rustls::ClientConfig::builder_with_provider(client_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(root_store)
.with_no_client_auth();
client_tls_config.alpn_protocols = vec![b"h3".to_vec()];
let tls_config: Arc<rustls::ClientConfig> = Arc::new(client_tls_config);
let client_endpoint = build_quinn_endpoint("127.0.0.1:0".parse()?)?;
let transport_config = Arc::new(quinn::TransportConfig::default());
let mut connector = QuicConnector::new(
client_endpoint,
transport_config,
tls_config,
Http3Options::default(),
);
let waker = futures_util::task::noop_waker();
let mut cx = std::task::Context::from_waker(&waker);
match connector.poll_ready(&mut cx) {
std::task::Poll::Ready(Ok(())) => {}
other => {
return Err(format!("poll_ready should be Ok, got {other:?}").into());
}
}
let client = Client::builder()
.http3_only()
.__test_with_quic_connector(connector)
.build()?;
let url = format!("https://127.0.0.1:{}/", server_addr.port());
let response = tokio::time::timeout(Duration::from_secs(10), client.get(&url).send())
.await
.map_err(|_| -> Box<dyn std::error::Error + Send + Sync> {
"GET request should resolve within 10s".into()
})??;
assert_eq!(
response.status(),
http::StatusCode::OK,
"expected 200 OK from H3_NO_ERROR STOP_SENDING (graceful EOF)"
);
let body = response.bytes().await?;
assert!(
body.is_empty(),
"expected empty body from H3_NO_ERROR STOP_SENDING, got {} bytes",
body.len()
);
Ok(())
}
#[tokio::test]
async fn http3_stop_sending_internal_error() -> TestResult<()> {
use tower::Service;
let certified_key = rcgen::generate_simple_self_signed(vec!["127.0.0.1".to_string()])?;
let cert_der: CertificateDer<'static> = certified_key.cert.der().clone();
let key_der: PrivateKeyDer<'static> =
PrivatePkcs8KeyDer::from(certified_key.signing_key.serialize_der()).into();
let provider = Arc::new(rustls::crypto::ring::default_provider());
let mut server_rustls_config = RustlsServerConfig::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_no_client_auth()
.with_single_cert(vec![cert_der.clone()], key_der)?;
server_rustls_config.alpn_protocols = vec![b"h3".to_vec()];
let server_quic_config = QuicServerConfig::try_from(Arc::new(server_rustls_config))?;
let server_config = ServerConfig::with_crypto(Arc::new(server_quic_config));
let server_endpoint = Endpoint::server(server_config, "127.0.0.1:0".parse()?)?;
let server_addr: SocketAddr = server_endpoint.local_addr()?;
let server_task = tokio::spawn(async move {
while let Some(incoming) = server_endpoint.accept().await {
let quinn_conn = match incoming.await {
Ok(c) => c,
Err(_) => break,
};
let h3_quinn_conn = hpx_h3_quinn::Connection::new(quinn_conn);
let mut h3_conn: hpx_h3::server::Connection<hpx_h3_quinn::Connection, bytes::Bytes> =
match hpx_h3::server::Connection::new(h3_quinn_conn).await {
Ok(c) => c,
Err(_) => continue,
};
loop {
match h3_conn.accept().await {
Ok(Some(resolver)) => {
let (_req, mut stream) = match resolver.resolve_request().await {
Ok(parts) => parts,
Err(_) => continue,
};
while matches!(stream.recv_data().await, Ok(Some(_))) {}
tokio::time::sleep(Duration::from_millis(50)).await;
stream.stop_sending(Code::H3_INTERNAL_ERROR);
let _ = stream.finish().await;
}
Ok(None) => break,
Err(_) => break,
}
}
}
});
let _ = server_task;
let mut root_store = RootCertStore::empty();
root_store.add(cert_der.clone())?;
let client_provider = Arc::new(rustls::crypto::ring::default_provider());
let mut client_tls_config = rustls::ClientConfig::builder_with_provider(client_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(root_store)
.with_no_client_auth();
client_tls_config.alpn_protocols = vec![b"h3".to_vec()];
let tls_config: Arc<rustls::ClientConfig> = Arc::new(client_tls_config);
let client_endpoint = build_quinn_endpoint("127.0.0.1:0".parse()?)?;
let transport_config = Arc::new(quinn::TransportConfig::default());
let mut connector = QuicConnector::new(
client_endpoint,
transport_config,
tls_config,
Http3Options::default(),
);
let waker = futures_util::task::noop_waker();
let mut cx = std::task::Context::from_waker(&waker);
match connector.poll_ready(&mut cx) {
std::task::Poll::Ready(Ok(())) => {}
other => {
return Err(format!("poll_ready should be Ok, got {other:?}").into());
}
}
let client = Client::builder()
.http3_only()
.__test_with_quic_connector(connector)
.build()?;
let url = format!("https://127.0.0.1:{}/", server_addr.port());
let response = tokio::time::timeout(Duration::from_secs(10), client.get(&url).send())
.await
.map_err(|_| -> Box<dyn std::error::Error + Send + Sync> {
"GET request should resolve within 10s".into()
})??;
let body_result = response.bytes().await;
let err = match body_result {
Err(e) => e,
Ok(_) => {
return Err(
"expected H3_INTERNAL_ERROR STOP_SENDING to surface an error in body".into(),
);
}
};
assert!(
err.is_body(),
"expected is_body() to be true for H3_INTERNAL_ERROR STOP_SENDING, got: {err}"
);
let mut source: Option<&(dyn std::error::Error + 'static)> = err.source();
let mut found_h3_error = false;
while let Some(inner) = source {
if inner.downcast_ref::<H3Error>().is_some() {
found_h3_error = true;
break;
}
source = inner.source();
}
assert!(
found_h3_error,
"expected error source chain to contain an H3Error, got: {err}"
);
Ok(())
}
#[tokio::test]
async fn http3_connection_failure_surfaces_typed_error() -> TestResult<()> {
use tower::Service;
let provider = Arc::new(rustls::crypto::ring::default_provider());
let mut client_tls_config = rustls::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(RootCertStore::empty())
.with_no_client_auth();
client_tls_config.alpn_protocols = vec![b"h3".to_vec()];
let tls_config: Arc<rustls::ClientConfig> = Arc::new(client_tls_config);
let client_endpoint = build_quinn_endpoint("127.0.0.1:0".parse()?)?;
let mut transport_config = quinn::TransportConfig::default();
transport_config.max_idle_timeout(Some(quinn::VarInt::from_u32(2000).into()));
let transport_config = Arc::new(transport_config);
let mut connector = QuicConnector::new(
client_endpoint,
transport_config,
tls_config,
Http3Options::default(),
);
let waker = futures_util::task::noop_waker();
let mut cx = std::task::Context::from_waker(&waker);
match connector.poll_ready(&mut cx) {
std::task::Poll::Ready(Ok(())) => {}
other => {
return Err(format!("poll_ready should be Ok, got {other:?}").into());
}
}
let uri: http::Uri = "https://127.0.0.1:1/".parse()?;
let connect_req = __test_connect_request(uri);
let result = tokio::time::timeout(Duration::from_secs(10), connector.call(connect_req))
.await
.map_err(|_| -> Box<dyn std::error::Error + Send + Sync> {
"QuicConnector::call should resolve within 10s".into()
})?;
let h3_err = match result {
Err(e) => e,
Ok(_) => return Err("expected connection to fail, but it succeeded".into()),
};
assert!(
matches!(h3_err, H3Error::Handshake { .. }),
"expected H3Error::Handshake, got {h3_err:?}"
);
let hpx_err: hpx::Error = h3_err.into();
assert!(
hpx_err.is_connect(),
"expected is_connect() to be true for connection failure"
);
let mut source: Option<&(dyn std::error::Error + 'static)> = hpx_err.source();
let mut found_handshake = false;
while let Some(err) = source {
if let Some(h3_inner) = err.downcast_ref::<H3Error>() {
if matches!(h3_inner, H3Error::Handshake { .. }) {
found_handshake = true;
break;
}
}
source = err.source();
}
assert!(
found_handshake,
"expected error source chain to contain H3Error::Handshake"
);
Ok(())
}
#[tokio::test]
async fn http3_request_body_mid_stream_error() -> TestResult<()> {
use tower::Service;
let certified_key = rcgen::generate_simple_self_signed(vec!["127.0.0.1".to_string()])?;
let cert_der: CertificateDer<'static> = certified_key.cert.der().clone();
let key_der: PrivateKeyDer<'static> =
PrivatePkcs8KeyDer::from(certified_key.signing_key.serialize_der()).into();
let provider = Arc::new(rustls::crypto::ring::default_provider());
let mut server_rustls_config = RustlsServerConfig::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_no_client_auth()
.with_single_cert(vec![cert_der.clone()], key_der)?;
server_rustls_config.alpn_protocols = vec![b"h3".to_vec()];
let server_quic_config = QuicServerConfig::try_from(Arc::new(server_rustls_config))?;
let server_config = ServerConfig::with_crypto(Arc::new(server_quic_config));
let server_endpoint = Endpoint::server(server_config, "127.0.0.1:0".parse()?)?;
let server_addr: SocketAddr = server_endpoint.local_addr()?;
let server_task = tokio::spawn(async move {
while let Some(incoming) = server_endpoint.accept().await {
let quinn_conn = match incoming.await {
Ok(c) => c,
Err(_) => break,
};
let h3_quinn_conn = hpx_h3_quinn::Connection::new(quinn_conn);
let mut h3_conn: hpx_h3::server::Connection<hpx_h3_quinn::Connection, bytes::Bytes> =
match hpx_h3::server::Connection::new(h3_quinn_conn).await {
Ok(c) => c,
Err(_) => continue,
};
loop {
match h3_conn.accept().await {
Ok(Some(resolver)) => {
let (_req, mut stream) = match resolver.resolve_request().await {
Ok(parts) => parts,
Err(_) => continue,
};
while matches!(stream.recv_data().await, Ok(Some(_))) {}
let _ = stream.finish().await;
}
Ok(None) => break,
Err(_) => break,
}
}
}
});
let _ = server_task;
let mut root_store = RootCertStore::empty();
root_store.add(cert_der.clone())?;
let client_provider = Arc::new(rustls::crypto::ring::default_provider());
let mut client_tls_config = rustls::ClientConfig::builder_with_provider(client_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(root_store)
.with_no_client_auth();
client_tls_config.alpn_protocols = vec![b"h3".to_vec()];
let tls_config: Arc<rustls::ClientConfig> = Arc::new(client_tls_config);
let client_endpoint = build_quinn_endpoint("127.0.0.1:0".parse()?)?;
let transport_config = Arc::new(quinn::TransportConfig::default());
let mut connector = QuicConnector::new(
client_endpoint,
transport_config,
tls_config,
Http3Options::default(),
);
let waker = futures_util::task::noop_waker();
let mut cx = std::task::Context::from_waker(&waker);
match connector.poll_ready(&mut cx) {
std::task::Poll::Ready(Ok(())) => {}
other => {
return Err(format!("poll_ready should be Ok, got {other:?}").into());
}
}
let client = Client::builder()
.http3_only()
.__test_with_quic_connector(connector)
.build()?;
use std::{
pin::Pin,
task::{Context, Poll},
};
use bytes::Bytes;
use http_body::{Body as HttpBody, Frame, SizeHint};
struct TestErrorBody {
chunk: Option<Bytes>,
}
impl HttpBody for TestErrorBody {
type Data = Bytes;
type Error = Box<dyn std::error::Error + Send + Sync>;
fn poll_frame(
mut self: Pin<&mut Self>,
_cx: &mut Context<'_>,
) -> Poll<Option<Result<Frame<Self::Data>, Self::Error>>> {
match self.chunk.take() {
Some(chunk) => Poll::Ready(Some(Ok(Frame::data(chunk)))),
None => Poll::Ready(Some(Err("mid-stream error".into()))),
}
}
fn is_end_stream(&self) -> bool {
false
}
fn size_hint(&self) -> SizeHint {
SizeHint::new()
}
}
let body = hpx::Body::wrap(TestErrorBody {
chunk: Some(Bytes::from("first")),
});
let url = format!("https://127.0.0.1:{}/", server_addr.port());
let result = tokio::time::timeout(Duration::from_secs(10), client.post(&url).body(body).send())
.await
.map_err(|_| -> Box<dyn std::error::Error + Send + Sync> {
"POST request should resolve within 10s".into()
})?;
let err = match result {
Err(e) => e,
Ok(_resp) => {
return Err(
"expected body mid-stream error to surface in send(), but got Ok response".into(),
);
}
};
assert!(
err.is_request(),
"expected is_request() to be true for body mid-stream error, got: {err}"
);
let mut source: Option<&(dyn std::error::Error + 'static)> = err.source();
let mut found_is_body = false;
while let Some(inner) = source {
if let Some(hpx_err) = inner.downcast_ref::<hpx::Error>() {
if hpx_err.is_body() {
found_is_body = true;
break;
}
}
source = inner.source();
}
assert!(
found_is_body,
"expected error source chain to contain an hpx::Error with is_body(), got: {err}"
);
Ok(())
}
#[tokio::test]
async fn http2_path_unaffected_by_http3_feature() -> TestResult<()> {
use support::server;
let server = server::http(move |_| async move { http::Response::default() });
let client = Client::builder().build()?;
let url = format!("http://{}/", server.addr());
let response = client
.get(&url)
.version(http::Version::HTTP_2)
.send()
.await?;
assert_eq!(
response.version(),
http::Version::HTTP_2,
"expected HTTP/2 version when http3 feature is enabled but client is default"
);
assert_eq!(
response.status(),
http::StatusCode::OK,
"expected 200 OK from h2 path when http3 feature is enabled"
);
Ok(())
}
#[tokio::test]
async fn alt_svc_captured_from_h2_response() -> TestResult<()> {
use support::server;
let server = server::http(move |_| async move {
let mut resp = http::Response::<hpx::Body>::default();
*resp.status_mut() = http::StatusCode::OK;
resp.headers_mut().insert(
http::header::HeaderName::from_static("alt-svc"),
http::header::HeaderValue::from_static(r#"h3=":443""#),
);
resp
});
let client = Client::builder().build()?;
let url = format!("http://{}/", server.addr());
let response = client
.get(&url)
.version(http::Version::HTTP_2)
.send()
.await?;
assert_eq!(response.status(), http::StatusCode::OK, "expected 200 OK");
let addr = server.addr();
let host = addr.ip().to_string();
let port = addr.port();
assert!(
client.__test_alt_svc_cache_has_entry(&host, port).await,
"Alt-Svc cache should have an entry for {host}:{port}"
);
Ok(())
}
#[tokio::test]
async fn client_upgrades_to_h3_after_alt_svc() -> TestResult<()> {
use support::server;
use tower::Service;
let certified_key = rcgen::generate_simple_self_signed(vec!["127.0.0.1".to_string()])?;
let cert_der: CertificateDer<'static> = certified_key.cert.der().clone();
let key_der: PrivateKeyDer<'static> =
PrivatePkcs8KeyDer::from(certified_key.signing_key.serialize_der()).into();
let provider = Arc::new(rustls::crypto::ring::default_provider());
let mut h3_rustls_config = RustlsServerConfig::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_no_client_auth()
.with_single_cert(vec![cert_der.clone()], key_der)?;
h3_rustls_config.alpn_protocols = vec![b"h3".to_vec()];
let h3_quic_config = QuicServerConfig::try_from(Arc::new(h3_rustls_config))?;
let h3_server_config = ServerConfig::with_crypto(Arc::new(h3_quic_config));
let h3_endpoint = Endpoint::server(h3_server_config, "127.0.0.1:0".parse()?)?;
let h3_addr: SocketAddr = h3_endpoint.local_addr()?;
let h3_port = h3_addr.port();
let h3_server_task = tokio::spawn(async move {
while let Some(incoming) = h3_endpoint.accept().await {
match incoming.await {
Ok(quinn_conn) => {
tokio::spawn(async move {
let h3_quinn_conn = hpx_h3_quinn::Connection::new(quinn_conn);
let mut h3_conn: hpx_h3::server::Connection<
hpx_h3_quinn::Connection,
bytes::Bytes,
> = match hpx_h3::server::Connection::new(h3_quinn_conn).await {
Ok(c) => c,
Err(_) => return,
};
loop {
match h3_conn.accept().await {
Ok(Some(resolver)) => {
let (_req, mut stream) = match resolver.resolve_request().await
{
Ok(parts) => parts,
Err(_) => continue,
};
let resp = match http::Response::builder()
.status(http::StatusCode::OK)
.body(())
{
Ok(r) => r,
Err(_) => continue,
};
if stream.send_response(resp).await.is_err() {
continue;
}
if stream
.send_data(bytes::Bytes::from_static(b"hello, h3"))
.await
.is_err()
{
continue;
}
while matches!(stream.recv_data().await, Ok(Some(_))) {}
let _ = stream.finish().await;
}
Ok(None) => break,
Err(_) => break,
}
}
});
}
Err(_) => break,
}
}
});
let _ = h3_server_task;
let alt_svc_value = format!(r#"h3=":{h3_port}""#);
let h2_server = server::http(move |_| {
let alt_svc_value = alt_svc_value.clone();
async move {
let mut resp = http::Response::<hpx::Body>::default();
*resp.status_mut() = http::StatusCode::OK;
resp.headers_mut().insert(
http::header::HeaderName::from_static("alt-svc"),
http::header::HeaderValue::from_str(&alt_svc_value).unwrap(),
);
resp
}
});
let h2_addr = h2_server.addr();
let mut root_store = RootCertStore::empty();
root_store.add(cert_der.clone())?;
let client_provider = Arc::new(rustls::crypto::ring::default_provider());
let mut client_tls_config = rustls::ClientConfig::builder_with_provider(client_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(root_store)
.with_no_client_auth();
client_tls_config.alpn_protocols = vec![b"h3".to_vec()];
let tls_config: Arc<rustls::ClientConfig> = Arc::new(client_tls_config);
let client_endpoint = build_quinn_endpoint("127.0.0.1:0".parse()?)?;
let transport_config = Arc::new(quinn::TransportConfig::default());
let mut connector = QuicConnector::new(
client_endpoint,
transport_config,
tls_config,
Http3Options::default(),
);
let waker = futures_util::task::noop_waker();
let mut cx = std::task::Context::from_waker(&waker);
match connector.poll_ready(&mut cx) {
std::task::Poll::Ready(Ok(())) => {}
other => {
return Err(format!("poll_ready should be Ok, got {other:?}").into());
}
}
let client = Client::builder()
.no_proxy()
.__test_with_quic_connector(connector)
.build()?;
let h2_url = format!("http://{}/", h2_addr);
let response1 = tokio::time::timeout(
Duration::from_secs(10),
client.get(&h2_url).version(http::Version::HTTP_2).send(),
)
.await
.map_err(|_| -> Box<dyn std::error::Error + Send + Sync> {
"first GET request should resolve within 10s".into()
})??;
assert_eq!(
response1.status(),
http::StatusCode::OK,
"first request should return 200 OK"
);
let h2_host = h2_addr.ip().to_string();
let h2_port = h2_addr.port();
let has_entry = client
.__test_alt_svc_cache_has_entry(&h2_host, h2_port)
.await;
assert!(
has_entry,
"Alt-Svc cache should have an h3 entry for {h2_host}:{h2_port}"
);
let response2 = tokio::time::timeout(
Duration::from_secs(10),
client.get(&h2_url).version(http::Version::HTTP_2).send(),
)
.await
.map_err(|_| -> Box<dyn std::error::Error + Send + Sync> {
"second GET request should resolve within 10s".into()
})??;
assert_eq!(
response2.status(),
http::StatusCode::OK,
"second request should return 200 OK"
);
assert_eq!(
response2.version(),
http::Version::HTTP_3,
"expected HTTP/3 version after alt-svc upgrade, got {:?}",
response2.version()
);
let body = response2.bytes().await?;
assert_eq!(
body.as_ref(),
b"hello, h3",
"expected response body to be 'hello, h3' from the h3 server, got {:?}",
std::str::from_utf8(body.as_ref()).unwrap_or("<non-utf8>")
);
Ok(())
}
#[tokio::test]
async fn prefer_http3_prefers_h3_with_fallback() -> TestResult<()> {
use support::server;
use tower::Service;
let certified_key = rcgen::generate_simple_self_signed(vec!["127.0.0.1".to_string()])?;
let cert_der: CertificateDer<'static> = certified_key.cert.der().clone();
let key_der: PrivateKeyDer<'static> =
PrivatePkcs8KeyDer::from(certified_key.signing_key.serialize_der()).into();
let provider = Arc::new(rustls::crypto::ring::default_provider());
let mut h3_rustls_config = RustlsServerConfig::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_no_client_auth()
.with_single_cert(vec![cert_der.clone()], key_der)?;
h3_rustls_config.alpn_protocols = vec![b"h3".to_vec()];
let h3_quic_config = QuicServerConfig::try_from(Arc::new(h3_rustls_config))?;
let h3_server_config = ServerConfig::with_crypto(Arc::new(h3_quic_config));
let h3_endpoint = Endpoint::server(h3_server_config, "127.0.0.1:0".parse()?)?;
let h3_addr: SocketAddr = h3_endpoint.local_addr()?;
let h3_port = h3_addr.port();
let h3_server_task = tokio::spawn(async move {
while let Some(incoming) = h3_endpoint.accept().await {
match incoming.await {
Ok(quinn_conn) => {
tokio::spawn(async move {
let h3_quinn_conn = hpx_h3_quinn::Connection::new(quinn_conn);
let mut h3_conn: hpx_h3::server::Connection<
hpx_h3_quinn::Connection,
bytes::Bytes,
> = match hpx_h3::server::Connection::new(h3_quinn_conn).await {
Ok(c) => c,
Err(_) => return,
};
loop {
match h3_conn.accept().await {
Ok(Some(resolver)) => {
let (_req, mut stream) = match resolver.resolve_request().await
{
Ok(parts) => parts,
Err(_) => continue,
};
let resp = match http::Response::builder()
.status(http::StatusCode::OK)
.body(())
{
Ok(r) => r,
Err(_) => continue,
};
if stream.send_response(resp).await.is_err() {
continue;
}
if stream
.send_data(bytes::Bytes::from_static(b"hello, h3"))
.await
.is_err()
{
continue;
}
while matches!(stream.recv_data().await, Ok(Some(_))) {}
let _ = stream.finish().await;
}
Ok(None) => break,
Err(_) => break,
}
}
});
}
Err(_) => break,
}
}
});
let _ = h3_server_task;
let alt_svc_value = format!(r#"h3=":{h3_port}""#);
let h2_server = server::http(move |_| {
let alt_svc_value = alt_svc_value.clone();
async move {
let mut resp = http::Response::<hpx::Body>::default();
*resp.status_mut() = http::StatusCode::OK;
resp.headers_mut().insert(
http::header::HeaderName::from_static("alt-svc"),
http::header::HeaderValue::from_str(&alt_svc_value).unwrap(),
);
resp
}
});
let h2_addr = h2_server.addr();
let mut root_store = RootCertStore::empty();
root_store.add(cert_der.clone())?;
let client_provider = Arc::new(rustls::crypto::ring::default_provider());
let mut client_tls_config = rustls::ClientConfig::builder_with_provider(client_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(root_store)
.with_no_client_auth();
client_tls_config.alpn_protocols = vec![b"h3".to_vec()];
let tls_config: Arc<rustls::ClientConfig> = Arc::new(client_tls_config);
let client_endpoint = build_quinn_endpoint("127.0.0.1:0".parse()?)?;
let transport_config = Arc::new(quinn::TransportConfig::default());
let mut connector = QuicConnector::new(
client_endpoint,
transport_config,
tls_config,
Http3Options::default(),
);
let waker = futures_util::task::noop_waker();
let mut cx = std::task::Context::from_waker(&waker);
match connector.poll_ready(&mut cx) {
std::task::Poll::Ready(Ok(())) => {}
other => {
return Err(format!("poll_ready should be Ok, got {other:?}").into());
}
}
let client = Client::builder()
.prefer_http3()
.no_proxy()
.__test_with_quic_connector(connector)
.build()?;
let h2_url = format!("http://{}/", h2_addr);
let response1 = tokio::time::timeout(
Duration::from_secs(10),
client.get(&h2_url).version(http::Version::HTTP_2).send(),
)
.await
.map_err(|_| -> Box<dyn std::error::Error + Send + Sync> {
"first GET request should resolve within 10s".into()
})??;
assert_eq!(
response1.status(),
http::StatusCode::OK,
"first request should return 200 OK"
);
let h2_host = h2_addr.ip().to_string();
let h2_port = h2_addr.port();
let has_entry = client
.__test_alt_svc_cache_has_entry(&h2_host, h2_port)
.await;
assert!(
has_entry,
"Alt-Svc cache should have an h3 entry for {h2_host}:{h2_port}"
);
let response2 = tokio::time::timeout(
Duration::from_secs(10),
client.get(&h2_url).version(http::Version::HTTP_2).send(),
)
.await
.map_err(|_| -> Box<dyn std::error::Error + Send + Sync> {
"second GET request should resolve within 10s".into()
})??;
assert_eq!(
response2.status(),
http::StatusCode::OK,
"second request should return 200 OK"
);
assert_eq!(
response2.version(),
http::Version::HTTP_3,
"expected HTTP/3 version after alt-svc upgrade with prefer_http3(), got {:?}",
response2.version()
);
let body = response2.bytes().await?;
assert_eq!(
body.as_ref(),
b"hello, h3",
"expected response body to be 'hello, h3' from the h3 server, got {:?}",
std::str::from_utf8(body.as_ref()).unwrap_or("<non-utf8>")
);
Ok(())
}
#[tokio::test]
async fn quic_unreachable_triggers_fallback() -> TestResult<()> {
use support::server;
use tower::Service;
let _ = rustls::crypto::ring::default_provider().install_default();
let certified_key = rcgen::generate_simple_self_signed(vec!["127.0.0.1".to_string()])?;
let cert_der: CertificateDer<'static> = certified_key.cert.der().clone();
let h2_server = server::http(move |_| async move {
let mut resp = http::Response::<hpx::Body>::default();
*resp.status_mut() = http::StatusCode::OK;
resp.headers_mut().insert(
http::header::HeaderName::from_static("alt-svc"),
http::header::HeaderValue::from_static(r#"h3=":8443""#),
);
resp
});
let h2_addr = h2_server.addr();
let mut root_store = RootCertStore::empty();
root_store.add(cert_der.clone())?;
let client_provider = Arc::new(rustls::crypto::ring::default_provider());
let mut client_tls_config = rustls::ClientConfig::builder_with_provider(client_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(root_store)
.with_no_client_auth();
client_tls_config.alpn_protocols = vec![b"h3".to_vec()];
let tls_config: Arc<rustls::ClientConfig> = Arc::new(client_tls_config);
let client_endpoint = build_quinn_endpoint("127.0.0.1:0".parse()?)?;
let mut transport_config = quinn::TransportConfig::default();
transport_config.max_idle_timeout(Some(quinn::VarInt::from_u32(2000).into()));
let transport_config = Arc::new(transport_config);
let mut connector = QuicConnector::new(
client_endpoint,
transport_config,
tls_config,
Http3Options::default(),
);
let waker = futures_util::task::noop_waker();
let mut cx = std::task::Context::from_waker(&waker);
match connector.poll_ready(&mut cx) {
std::task::Poll::Ready(Ok(())) => {}
other => {
return Err(format!("poll_ready should be Ok, got {other:?}").into());
}
}
let client = Client::builder()
.prefer_http3()
.no_proxy()
.__test_with_quic_connector(connector)
.build()?;
let h2_url = format!("http://{}/", h2_addr);
let h2_host = h2_addr.ip().to_string();
let h2_port = h2_addr.port();
let response1 = tokio::time::timeout(
Duration::from_secs(10),
client.get(&h2_url).version(http::Version::HTTP_2).send(),
)
.await
.map_err(|_| -> Box<dyn std::error::Error + Send + Sync> {
"first GET request should resolve within 10s".into()
})??;
assert_eq!(
response1.status(),
http::StatusCode::OK,
"first request should return 200 OK"
);
let has_entry = client
.__test_alt_svc_cache_has_entry(&h2_host, h2_port)
.await;
assert!(
has_entry,
"Alt-Svc cache should have an h3 entry for {h2_host}:{h2_port}"
);
let is_blocked_before = client
.__test_h3_failure_tracker_is_blocked(&h2_host, h2_port)
.await;
assert!(
!is_blocked_before,
"circuit breaker should not be blocking before any h3 failure"
);
let response2 = tokio::time::timeout(
Duration::from_secs(10),
client.get(&h2_url).version(http::Version::HTTP_2).send(),
)
.await
.map_err(|_| -> Box<dyn std::error::Error + Send + Sync> {
"second GET request should resolve within 10s".into()
})??;
assert_eq!(
response2.status(),
http::StatusCode::OK,
"second request should return 200 OK (fallback to h2)"
);
assert_eq!(
response2.version(),
http::Version::HTTP_2,
"expected HTTP/2 version after h3 fallback, got {:?}",
response2.version()
);
let is_blocked_after = client
.__test_h3_failure_tracker_is_blocked(&h2_host, h2_port)
.await;
assert!(
is_blocked_after,
"circuit breaker should block {h2_host}:{h2_port} after h3 failure"
);
let response3 = tokio::time::timeout(
Duration::from_secs(10),
client.get(&h2_url).version(http::Version::HTTP_2).send(),
)
.await
.map_err(|_| -> Box<dyn std::error::Error + Send + Sync> {
"third GET request should resolve within 10s".into()
})??;
assert_eq!(
response3.status(),
http::StatusCode::OK,
"third request should return 200 OK (h3 skipped by circuit breaker)"
);
assert_eq!(
response3.version(),
http::Version::HTTP_2,
"expected HTTP/2 version when circuit breaker skips h3, got {:?}",
response3.version()
);
Ok(())
}
#[tokio::test]
#[ignore = "0-RTT may not work reliably with s2n-quic; enable manually"]
async fn http3_zero_rtt_resumption() -> TestResult<()> {
use tower::Service;
let certified_key = rcgen::generate_simple_self_signed(vec!["127.0.0.1".to_string()])?;
let cert_der: CertificateDer<'static> = certified_key.cert.der().clone();
let key_der: PrivateKeyDer<'static> =
PrivatePkcs8KeyDer::from(certified_key.signing_key.serialize_der()).into();
let provider = Arc::new(rustls::crypto::ring::default_provider());
let mut server_rustls_config = RustlsServerConfig::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_no_client_auth()
.with_single_cert(vec![cert_der.clone()], key_der)?;
server_rustls_config.alpn_protocols = vec![b"h3".to_vec()];
let server_quic_config = QuicServerConfig::try_from(Arc::new(server_rustls_config))?;
let server_config = ServerConfig::with_crypto(Arc::new(server_quic_config));
let server_endpoint = Endpoint::server(server_config, "127.0.0.1:0".parse()?)?;
let server_addr: SocketAddr = server_endpoint.local_addr()?;
let server_task = tokio::spawn(async move {
while let Some(incoming) = server_endpoint.accept().await {
match incoming.await {
Ok(quinn_conn) => {
tokio::spawn(async move {
let h3_quinn_conn = hpx_h3_quinn::Connection::new(quinn_conn);
let mut h3_conn: hpx_h3::server::Connection<
hpx_h3_quinn::Connection,
bytes::Bytes,
> = match hpx_h3::server::Connection::new(h3_quinn_conn).await {
Ok(c) => c,
Err(_) => return,
};
loop {
match h3_conn.accept().await {
Ok(Some(resolver)) => {
let (_req, mut stream) = match resolver.resolve_request().await
{
Ok(parts) => parts,
Err(_) => continue,
};
let resp = match http::Response::builder()
.status(http::StatusCode::OK)
.body(())
{
Ok(r) => r,
Err(_) => continue,
};
if stream.send_response(resp).await.is_err() {
continue;
}
while matches!(stream.recv_data().await, Ok(Some(_))) {}
let _ = stream.finish().await;
}
Ok(None) => break,
Err(_) => break,
}
}
});
}
Err(_) => break,
}
}
});
let _ = server_task;
let mut root_store = RootCertStore::empty();
root_store.add(cert_der.clone())?;
let client_provider = Arc::new(rustls::crypto::ring::default_provider());
let mut client_tls_config = rustls::ClientConfig::builder_with_provider(client_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(root_store)
.with_no_client_auth();
client_tls_config.alpn_protocols = vec![b"h3".to_vec()];
client_tls_config.enable_early_data = true;
let tls_config: Arc<rustls::ClientConfig> = Arc::new(client_tls_config);
let client_endpoint = build_quinn_endpoint("127.0.0.1:0".parse()?)?;
let transport_config = Arc::new(quinn::TransportConfig::default());
let mut h3_options = Http3Options::default();
h3_options.enable_0rtt = true;
let mut connector =
QuicConnector::new(client_endpoint, transport_config, tls_config, h3_options);
let uri: http::Uri = format!("https://127.0.0.1:{}/", server_addr.port()).parse()?;
let connect_req = __test_connect_request(uri.clone());
let waker = futures_util::task::noop_waker();
let mut cx = std::task::Context::from_waker(&waker);
match connector.poll_ready(&mut cx) {
std::task::Poll::Ready(Ok(())) => {}
other => {
return Err(format!("poll_ready should be Ok, got {other:?}").into());
}
}
let h3_conn1 = tokio::time::timeout(Duration::from_secs(5), connector.call(connect_req))
.await
.map_err(|_| -> Box<dyn std::error::Error + Send + Sync> {
"first QuicConnector::call should resolve within 5s".into()
})??;
let mut send_req = h3_conn1.send_request.clone();
let req = http::Request::get(format!("https://127.0.0.1:{}/", server_addr.port()))
.body(())
.map_err(|e| -> Box<dyn std::error::Error + Send + Sync> {
format!("failed to build request: {e}").into()
})?;
let mut stream = send_req.send_request(req).await.map_err(
|e| -> Box<dyn std::error::Error + Send + Sync> {
format!("send_request failed: {e}").into()
},
)?;
stream
.finish()
.await
.map_err(|e| -> Box<dyn std::error::Error + Send + Sync> {
format!("finish failed: {e}").into()
})?;
let _resp =
stream
.recv_response()
.await
.map_err(|e| -> Box<dyn std::error::Error + Send + Sync> {
format!("recv_response failed: {e}").into()
})?;
while let Some(_data) =
stream
.recv_data()
.await
.map_err(|e| -> Box<dyn std::error::Error + Send + Sync> {
format!("recv_data failed: {e}").into()
})?
{}
let first_was_0rtt = h3_conn1.used_0rtt();
drop(h3_conn1);
tokio::time::sleep(Duration::from_millis(200)).await;
let connect_req2 = __test_connect_request(uri);
match connector.poll_ready(&mut cx) {
std::task::Poll::Ready(Ok(())) => {}
other => {
return Err(format!("poll_ready should be Ok, got {other:?}").into());
}
}
let h3_conn2 = tokio::time::timeout(Duration::from_secs(5), connector.call(connect_req2))
.await
.map_err(|_| -> Box<dyn std::error::Error + Send + Sync> {
"second QuicConnector::call should resolve within 5s".into()
})??;
let mut send_req2 = h3_conn2.send_request.clone();
let req2 = http::Request::get(format!("https://127.0.0.1:{}/", server_addr.port()))
.body(())
.map_err(|e| -> Box<dyn std::error::Error + Send + Sync> {
format!("failed to build request: {e}").into()
})?;
let mut stream2 = send_req2.send_request(req2).await.map_err(
|e| -> Box<dyn std::error::Error + Send + Sync> {
format!("send_request failed: {e}").into()
},
)?;
stream2
.finish()
.await
.map_err(|e| -> Box<dyn std::error::Error + Send + Sync> {
format!("finish failed: {e}").into()
})?;
let _resp2 =
stream2
.recv_response()
.await
.map_err(|e| -> Box<dyn std::error::Error + Send + Sync> {
format!("recv_response failed: {e}").into()
})?;
while let Some(_data) =
stream2
.recv_data()
.await
.map_err(|e| -> Box<dyn std::error::Error + Send + Sync> {
format!("recv_data failed: {e}").into()
})?
{}
let start = std::time::Instant::now();
let mut used_0rtt = false;
while start.elapsed() < Duration::from_secs(2) {
used_0rtt = h3_conn2.used_0rtt();
if used_0rtt {
break;
}
tokio::time::sleep(Duration::from_millis(50)).await;
}
assert!(
used_0rtt,
"expected second connection to use 0-RTT, but used_0rtt() returned false. \
First connection used_0rtt: {first_was_0rtt}"
);
Ok(())
}
#[tokio::test]
async fn http3_idle_timeout_closes_connection() -> TestResult<()> {
use tower::Service;
let certified_key = rcgen::generate_simple_self_signed(vec!["127.0.0.1".to_string()])?;
let cert_der: CertificateDer<'static> = certified_key.cert.der().clone();
let key_der: PrivateKeyDer<'static> =
PrivatePkcs8KeyDer::from(certified_key.signing_key.serialize_der()).into();
let provider = Arc::new(rustls::crypto::ring::default_provider());
let mut server_rustls_config = RustlsServerConfig::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_no_client_auth()
.with_single_cert(vec![cert_der.clone()], key_der)?;
server_rustls_config.alpn_protocols = vec![b"h3".to_vec()];
let server_quic_config = QuicServerConfig::try_from(Arc::new(server_rustls_config))?;
let server_config = ServerConfig::with_crypto(Arc::new(server_quic_config));
let server_endpoint = Endpoint::server(server_config, "127.0.0.1:0".parse()?)?;
let server_addr: SocketAddr = server_endpoint.local_addr()?;
let server_task = tokio::spawn(async move {
while let Some(incoming) = server_endpoint.accept().await {
match incoming.await {
Ok(quinn_conn) => {
tokio::spawn(async move {
let h3_quinn_conn = hpx_h3_quinn::Connection::new(quinn_conn);
let mut h3_conn: hpx_h3::server::Connection<
hpx_h3_quinn::Connection,
bytes::Bytes,
> = match hpx_h3::server::Connection::new(h3_quinn_conn).await {
Ok(c) => c,
Err(_) => return,
};
loop {
match h3_conn.accept().await {
Ok(Some(resolver)) => {
tokio::spawn(async move {
if let Ok((_req, mut stream)) =
resolver.resolve_request().await
{
let resp = match http::Response::builder()
.status(http::StatusCode::OK)
.body(())
{
Ok(r) => r,
Err(_) => return,
};
if stream.send_response(resp).await.is_ok() {
let _ = stream.finish().await;
}
}
});
}
Ok(None) => break,
Err(_) => break,
}
}
});
}
Err(_) => break,
}
}
});
let _ = server_task;
let mut root_store = RootCertStore::empty();
root_store.add(cert_der.clone())?;
let client_provider = Arc::new(rustls::crypto::ring::default_provider());
let mut client_tls_config = rustls::ClientConfig::builder_with_provider(client_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(root_store)
.with_no_client_auth();
client_tls_config.alpn_protocols = vec![b"h3".to_vec()];
let tls_config: Arc<rustls::ClientConfig> = Arc::new(client_tls_config);
let client_endpoint = build_quinn_endpoint("127.0.0.1:0".parse()?)?;
let transport_config = Arc::new(quinn::TransportConfig::default());
let mut h3_options = Http3Options::default();
h3_options.max_idle_timeout = Some(Duration::from_secs(1));
let mut connector =
QuicConnector::new(client_endpoint, transport_config, tls_config, h3_options);
let uri: http::Uri = format!("https://127.0.0.1:{}/", server_addr.port()).parse()?;
let connect_req = __test_connect_request(uri.clone());
let waker = futures_util::task::noop_waker();
let mut cx = std::task::Context::from_waker(&waker);
match connector.poll_ready(&mut cx) {
std::task::Poll::Ready(Ok(())) => {}
other => {
return Err(format!("poll_ready should be Ok, got {other:?}").into());
}
}
let h3_conn = tokio::time::timeout(Duration::from_secs(5), connector.call(connect_req))
.await
.map_err(|_| -> Box<dyn std::error::Error + Send + Sync> {
"QuicConnector::call should resolve within 5s".into()
})??;
assert!(
h3_conn.is_valid(Duration::MAX),
"fresh H3Connection should be valid (no close event, not idle-expired)"
);
let mut send_req = h3_conn.send_request.clone();
let req = http::Request::get(format!("https://127.0.0.1:{}/", server_addr.port()))
.body(())
.map_err(|e| -> Box<dyn std::error::Error + Send + Sync> {
format!("failed to build request: {e}").into()
})?;
let mut stream = send_req.send_request(req).await.map_err(
|e| -> Box<dyn std::error::Error + Send + Sync> {
format!("send_request failed: {e}").into()
},
)?;
stream
.finish()
.await
.map_err(|e| -> Box<dyn std::error::Error + Send + Sync> {
format!("finish failed: {e}").into()
})?;
let _resp =
stream
.recv_response()
.await
.map_err(|e| -> Box<dyn std::error::Error + Send + Sync> {
format!("recv_response failed: {e}").into()
})?;
while let Some(_data) =
stream
.recv_data()
.await
.map_err(|e| -> Box<dyn std::error::Error + Send + Sync> {
format!("recv_data failed: {e}").into()
})?
{}
tokio::time::sleep(Duration::from_secs(2)).await;
let mut observed_invalid = false;
let deadline = tokio::time::Instant::now() + Duration::from_secs(5);
while tokio::time::Instant::now() < deadline {
if !h3_conn.is_valid(Duration::MAX) {
observed_invalid = true;
break;
}
tokio::time::sleep(Duration::from_millis(50)).await;
}
assert!(
observed_invalid,
"H3Connection should be invalid after idle timeout expires \
(max_idle_timeout=1s, waited 2s + polling)"
);
Ok(())
}
#[tokio::test]
async fn http3_extended_connect_websocket() -> TestResult<()> {
let certified_key = rcgen::generate_simple_self_signed(vec!["127.0.0.1".to_string()])?;
let cert_der: CertificateDer<'static> = certified_key.cert.der().clone();
let key_der: PrivateKeyDer<'static> =
PrivatePkcs8KeyDer::from(certified_key.signing_key.serialize_der()).into();
let provider = Arc::new(rustls::crypto::ring::default_provider());
let mut server_rustls_config = RustlsServerConfig::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_no_client_auth()
.with_single_cert(vec![cert_der.clone()], key_der)?;
server_rustls_config.alpn_protocols = vec![b"h3".to_vec()];
let server_quic_config = QuicServerConfig::try_from(Arc::new(server_rustls_config))?;
let server_config = ServerConfig::with_crypto(Arc::new(server_quic_config));
let server_endpoint = Endpoint::server(server_config, "127.0.0.1:0".parse()?)?;
let server_addr: SocketAddr = server_endpoint.local_addr()?;
let server_task = tokio::spawn(async move {
while let Some(incoming) = server_endpoint.accept().await {
match incoming.await {
Ok(quinn_conn) => {
tokio::spawn(async move {
let h3_quinn_conn = hpx_h3_quinn::Connection::new(quinn_conn);
let mut h3_conn: hpx_h3::server::Connection<
hpx_h3_quinn::Connection,
Bytes,
> = match hpx_h3::server::Connection::new(h3_quinn_conn).await {
Ok(c) => c,
Err(_) => return,
};
loop {
match h3_conn.accept().await {
Ok(Some(resolver)) => {
let (req, mut stream) = match resolver.resolve_request().await {
Ok(parts) => parts,
Err(_) => continue,
};
let is_ws_connect = req.method() == http::Method::CONNECT
&& req
.extensions()
.get::<hpx_h3::ext::Protocol>()
.is_some_and(|p| p.as_str() == "websocket");
if is_ws_connect {
let resp = match http::Response::builder()
.status(http::StatusCode::OK)
.body(())
{
Ok(r) => r,
Err(_) => continue,
};
if stream.send_response(resp).await.is_err() {
continue;
}
while let Ok(Some(mut data)) = stream.recv_data().await {
let len = data.remaining();
let chunk = data.copy_to_bytes(len);
if stream.send_data(chunk).await.is_err() {
break;
}
}
let _ = stream.finish().await;
} else {
let resp = match http::Response::builder()
.status(http::StatusCode::OK)
.body(())
{
Ok(r) => r,
Err(_) => continue,
};
if stream.send_response(resp).await.is_err() {
continue;
}
while matches!(stream.recv_data().await, Ok(Some(_))) {}
let _ = stream.finish().await;
}
}
Ok(None) => break,
Err(_) => break,
}
}
});
}
Err(_) => break,
}
}
});
let _ = server_task;
let mut root_store = RootCertStore::empty();
root_store.add(cert_der.clone())?;
let client_provider = Arc::new(rustls::crypto::ring::default_provider());
let mut client_tls_config = rustls::ClientConfig::builder_with_provider(client_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(root_store)
.with_no_client_auth();
client_tls_config.alpn_protocols = vec![b"h3".to_vec()];
let quic_client_config = quinn::crypto::rustls::QuicClientConfig::try_from(Arc::new(
client_tls_config,
))
.map_err(|e| -> Box<dyn std::error::Error + Send + Sync> {
format!("QuicClientConfig: {e}").into()
})?;
let client_config = quinn::ClientConfig::new(Arc::new(quic_client_config));
let client_addr: SocketAddr = "127.0.0.1:0".parse()?;
let client_endpoint = Endpoint::client(client_addr)?;
let quinn_conn = client_endpoint
.connect_with(client_config, server_addr, "127.0.0.1")?
.await?;
let h3_quinn_conn = hpx_h3_quinn::Connection::new(quinn_conn);
let (_driver, mut send_request) = hpx_h3::client::new(h3_quinn_conn).await.map_err(
|e| -> Box<dyn std::error::Error + Send + Sync> {
format!("hpx_h3::client::new failed: {e}").into()
},
)?;
let req = http::Request::builder()
.method(http::Method::CONNECT)
.uri(format!("https://127.0.0.1:{}/", server_addr.port()))
.extension(hpx_h3::ext::Protocol::WEB_SOCKET)
.body(())
.map_err(|e| -> Box<dyn std::error::Error + Send + Sync> {
format!("failed to build request: {e}").into()
})?;
let mut stream = send_request.send_request(req).await.map_err(
|e| -> Box<dyn std::error::Error + Send + Sync> {
format!("send_request failed: {e}").into()
},
)?;
let resp =
stream
.recv_response()
.await
.map_err(|e| -> Box<dyn std::error::Error + Send + Sync> {
format!("recv_response failed: {e}").into()
})?;
assert_eq!(
resp.status(),
http::StatusCode::OK,
"Extended CONNECT should return 200 OK"
);
use hpx::http3::{H3WebSocket, WsMessage};
let mut ws = H3WebSocket::new(stream);
ws.send_text("hello websocket over h3").await?;
let msg = ws
.recv()
.await?
.ok_or("expected a text message but stream closed")?;
assert_eq!(
msg,
WsMessage::Text("hello websocket over h3".to_string()),
"text message round-trip"
);
ws.send_binary(b"binary payload").await?;
let msg = ws
.recv()
.await?
.ok_or("expected a binary message but stream closed")?;
assert_eq!(
msg,
WsMessage::Binary(b"binary payload".to_vec()),
"binary message round-trip"
);
ws.send_close(Some(1000), Some("done")).await?;
let msg = ws
.recv()
.await?
.ok_or("expected a close message but stream closed")?;
assert_eq!(
msg,
WsMessage::Close {
code: Some(1000),
reason: Some("done".to_string()),
},
"close frame round-trip"
);
ws.finish().await?;
Ok(())
}