use std::net::SocketAddr;
use std::sync::Arc;
use axum::body::Body;
use axum::extract::ConnectInfo;
use axum::http::{header::ALT_SVC, HeaderValue, Request, Response};
use axum::Router;
use bytes::Buf;
use futures::StreamExt;
use rustls::pki_types::{CertificateDer, PrivateKeyDer};
use tower::ServiceExt;
#[derive(Debug, thiserror::Error)]
pub enum Http3Error {
#[error("QUIC TLS config: {0}")]
QuicTls(#[from] quinn::crypto::rustls::NoInitialCipherSuite),
#[error("TLS config: {0}")]
Rustls(#[from] rustls::Error),
#[error("QUIC endpoint I/O: {0}")]
Io(#[from] std::io::Error),
#[error("QUIC connection: {0}")]
Connection(#[from] quinn::ConnectionError),
#[error("HTTP/3 protocol: {0}")]
H3(#[from] h3::error::ConnectionError),
}
pub fn advertise_http3(router: Router, port: u16) -> Router {
let value = format!("h3=\":{port}\"; ma=86400");
router.layer(axum::middleware::from_fn(
move |req: axum::extract::Request, next: axum::middleware::Next| {
let value = value.clone();
async move {
let mut resp = next.run(req).await;
if let Ok(v) = HeaderValue::from_str(&value) {
resp.headers_mut().insert(ALT_SVC, v);
}
resp
}
},
))
}
pub fn quinn_server_config(
rustls_config: rustls::ServerConfig,
) -> Result<quinn::ServerConfig, Http3Error> {
let quic = quinn::crypto::rustls::QuicServerConfig::try_from(rustls_config)?;
Ok(quinn::ServerConfig::with_crypto(Arc::new(quic)))
}
pub fn http3_endpoint(
addr: SocketAddr,
server_config: quinn::ServerConfig,
) -> Result<quinn::Endpoint, Http3Error> {
Ok(quinn::Endpoint::server(server_config, addr)?)
}
pub async fn serve_http3_endpoint(
endpoint: quinn::Endpoint,
router: Router,
) -> Result<(), Http3Error> {
if let Ok(addr) = endpoint.local_addr() {
tracing::info!(%addr, "serving HTTP/3 (QUIC)");
}
let shutdown = crate::shutdown_signal();
tokio::pin!(shutdown);
loop {
tokio::select! {
incoming = endpoint.accept() => {
let Some(incoming) = incoming else { break };
let router = router.clone();
tokio::spawn(async move {
if let Err(err) = handle_connection(incoming, router).await {
tracing::debug!(error = %err, "http/3 connection ended");
}
});
}
_ = &mut shutdown => {
tracing::info!("HTTP/3 listener draining");
break;
}
}
}
endpoint.close(0u32.into(), b"shutdown");
let _ = tokio::time::timeout(std::time::Duration::from_secs(10), endpoint.wait_idle()).await;
Ok(())
}
pub async fn serve_http3(
addr: SocketAddr,
cert_chain: Vec<CertificateDer<'static>>,
key: PrivateKeyDer<'static>,
router: Router,
) -> Result<(), Http3Error> {
let mut tls = rustls::ServerConfig::builder()
.with_no_client_auth()
.with_single_cert(cert_chain, key)?;
tls.alpn_protocols = vec![b"h3".to_vec()];
let endpoint = http3_endpoint(addr, quinn_server_config(tls)?)?;
serve_http3_endpoint(endpoint, router).await
}
async fn handle_connection(incoming: quinn::Incoming, router: Router) -> Result<(), Http3Error> {
let conn = incoming.await?;
let peer = conn.remote_address();
let mut h3_conn = h3::server::Connection::new(h3_quinn::Connection::new(conn)).await?;
loop {
match h3_conn.accept().await {
Ok(Some(resolver)) => {
let (req, mut stream) = match resolver.resolve_request().await {
Ok(parts) => parts,
Err(err) => {
tracing::debug!(error = %err, "http/3 request resolve failed");
continue;
}
};
let router = router.clone();
tokio::spawn(async move {
let mut body = Vec::new();
loop {
match stream.recv_data().await {
Ok(Some(mut chunk)) => {
let bytes = chunk.copy_to_bytes(chunk.remaining());
body.extend_from_slice(&bytes);
}
Ok(None) => break,
Err(err) => {
tracing::debug!(error = %err, "http/3 body read failed");
return;
}
}
}
let (parts, _) = req.into_parts();
let mut request = Request::from_parts(parts, Body::from(body));
request.extensions_mut().insert(ConnectInfo(peer));
let response = match router.clone().oneshot(request).await {
Ok(response) => response,
Err(err) => {
tracing::debug!(error = %err, "http/3 routing failed");
return;
}
};
let (parts, body) = response.into_parts();
if let Err(err) = stream.send_response(Response::from_parts(parts, ())).await {
tracing::debug!(error = %err, "http/3 send_response failed");
return;
}
let mut data = body.into_data_stream();
while let Some(chunk) = data.next().await {
match chunk {
Ok(bytes) => {
if let Err(err) = stream.send_data(bytes).await {
tracing::debug!(error = %err, "http/3 send_data failed");
return;
}
}
Err(err) => {
tracing::debug!(error = %err, "http/3 body stream error");
return;
}
}
}
let _ = stream.finish().await;
});
}
Ok(None) => return Ok(()),
Err(err) => return Err(err.into()),
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use axum::routing::get;
#[tokio::test]
async fn advertise_http3_sets_alt_svc() {
let app = advertise_http3(Router::new().route("/", get(|| async { "ok" })), 8443);
let resp = app
.oneshot(Request::builder().uri("/").body(Body::empty()).unwrap())
.await
.unwrap();
let alt = resp
.headers()
.get(ALT_SVC)
.expect("Alt-Svc present")
.to_str()
.unwrap();
assert!(
alt.contains("h3=\":8443\""),
"advertises h3 on the port: {alt}"
);
}
}