mod control;
mod date;
mod error;
mod frame;
mod options;
pub mod qpack;
#[cfg(feature = "h3-quinn")]
pub mod quinn;
mod settings;
mod stream;
pub mod transport;
mod upgrade;
pub use error::{H3Error, TransportError};
pub use frame::{Frame, FrameDecoder, FrameError, Settings};
pub use options::*;
use std::{
pin::Pin,
rc::Rc,
sync::{
atomic::{AtomicBool, Ordering},
Arc,
},
task::{Context, Poll},
};
use bytes::Bytes;
use futures_util::stream::FuturesUnordered;
use futures_util::{ready, Future, FutureExt, StreamExt};
use http::{Request, Response, StatusCode};
use http_body::{Body, Frame as BodyFrame};
use http_body_util::BodyExt;
use tokio_util::sync::CancellationToken;
use crate::{
h3::{
control::{ControlEvent, ControlStreams},
date::DateCache,
stream::{RequestStream, SharedCodecs},
},
EarlyHints, HttpProtocol, Incoming, Upgrade, Upgraded,
};
const H3_NO_ERROR: u64 = 0x0100;
const H3_REQUEST_REJECTED: u64 = 0x010b;
type SharedRequest = Arc<tokio::sync::Mutex<RequestStream>>;
static HTTP3_INVALID_HEADERS: [http::header::HeaderName; 5] = [
http::header::HeaderName::from_static("keep-alive"),
http::header::HeaderName::from_static("proxy-connection"),
http::header::CONNECTION,
http::header::TRANSFER_ENCODING,
http::header::UPGRADE,
];
struct H3BodyState {
stream: SharedRequest,
data_done: bool,
send_continue_body: Option<Arc<AtomicBool>>,
}
pub(crate) struct H3Body {
inner: tokio::sync::Mutex<H3BodyState>,
}
impl H3Body {
#[inline]
fn new(stream: SharedRequest, send_continue_body: Option<Arc<AtomicBool>>) -> Self {
Self {
inner: tokio::sync::Mutex::new(H3BodyState {
stream,
data_done: false,
send_continue_body,
}),
}
}
}
impl Body for H3Body {
type Data = Bytes;
type Error = std::io::Error;
#[inline]
fn poll_frame(
self: Pin<&mut Self>,
cx: &mut Context<'_>,
) -> Poll<Option<Result<BodyFrame<Self::Data>, Self::Error>>> {
let mut inner = match std::pin::pin!(self.inner.lock()).poll_unpin(cx) {
Poll::Ready(inner) => inner,
Poll::Pending => return Poll::Pending,
};
if !inner.data_done {
let done = {
let mut stream = match std::pin::pin!(inner.stream.lock()).poll_unpin(cx) {
Poll::Ready(stream) => stream,
Poll::Pending => return Poll::Pending,
};
match stream.poll_recv_data(cx) {
Poll::Ready(Ok(Some(data))) => {
return Poll::Ready(Some(Ok(BodyFrame::data(data))));
}
Poll::Ready(Ok(None)) => true,
Poll::Ready(Err(err)) => {
return Poll::Ready(Some(Err(h3_stream_error_to_io(err))));
}
Poll::Pending => {
if let Some(scb) = inner.send_continue_body.as_ref() {
scb.store(true, std::sync::atomic::Ordering::Relaxed);
}
return Poll::Pending;
}
}
};
if done {
inner.data_done = true;
}
}
let mut stream = match std::pin::pin!(inner.stream.lock()).poll_unpin(cx) {
Poll::Ready(stream) => stream,
Poll::Pending => {
if let Some(scb) = inner.send_continue_body.as_ref() {
scb.store(true, std::sync::atomic::Ordering::Relaxed);
}
return Poll::Pending;
}
};
match stream.poll_recv_trailers(cx) {
Poll::Ready(Ok(Some(trailers))) => Poll::Ready(Some(Ok(BodyFrame::trailers(trailers)))),
Poll::Ready(Ok(None)) => Poll::Ready(None),
Poll::Ready(Err(err)) => Poll::Ready(Some(Err(h3_stream_error_to_io(err)))),
Poll::Pending => {
if let Some(scb) = inner.send_continue_body.as_ref() {
scb.store(true, std::sync::atomic::Ordering::Relaxed);
}
Poll::Pending
}
}
}
}
#[inline]
fn h3_control_error_to_io(error: control::ControlError) -> std::io::Error {
std::io::Error::other(error)
}
#[inline]
fn h3_transport_error_to_io(error: TransportError) -> std::io::Error {
std::io::Error::other(error)
}
#[inline]
fn h3_stream_error_to_io(error: stream::StreamError) -> std::io::Error {
std::io::Error::other(error)
}
#[inline]
fn remove_invalid_http3_headers(headers: &mut http::HeaderMap) {
for header in &HTTP3_INVALID_HEADERS {
headers.remove(header);
}
if headers
.get(http::header::TE)
.is_some_and(|v| v != "trailers")
{
headers.remove(http::header::TE);
}
}
#[inline]
async fn wait_for_encoder(shared: &Arc<parking_lot::Mutex<SharedCodecs>>, stream_id: u64) {
std::future::poll_fn(|cx| {
let mut shared = shared.lock();
if shared.encoder.is_some() {
shared.waiters.remove(&stream_id);
return Poll::Ready(());
}
shared.waiters.insert(stream_id, cx.waker().clone());
Poll::Pending
})
.await
}
#[inline]
async fn send_interim_response(
stream: &SharedRequest,
status: StatusCode,
) -> Result<(), std::io::Error> {
let mut guard = stream.lock().await;
std::future::poll_fn(|cx| guard.poll_send_response(cx, status, &http::HeaderMap::new()))
.await
.map_err(h3_stream_error_to_io)
}
#[inline]
async fn send_response(
stream: &SharedRequest,
shared: &Arc<parking_lot::Mutex<SharedCodecs>>,
stream_id: u64,
status: StatusCode,
headers: &http::HeaderMap,
) -> Result<(), std::io::Error> {
wait_for_encoder(shared, stream_id).await;
let mut guard = stream.lock().await;
let res = std::future::poll_fn(|cx| guard.poll_send_response(cx, status, headers))
.await
.map_err(h3_stream_error_to_io);
res
}
#[inline]
async fn send_data(stream: &SharedRequest, data: Bytes) -> Result<(), std::io::Error> {
let mut guard = stream.lock().await;
std::future::poll_fn(|cx| guard.poll_send_data(cx, data.clone()))
.await
.map_err(h3_stream_error_to_io)
}
#[inline]
async fn send_trailers(
stream: &SharedRequest,
trailers: &http::HeaderMap,
) -> Result<(), std::io::Error> {
let mut guard = stream.lock().await;
std::future::poll_fn(|cx| guard.poll_send_trailers(cx, trailers))
.await
.map_err(h3_stream_error_to_io)
}
#[inline]
async fn send_finish(stream: &SharedRequest) -> Result<(), std::io::Error> {
let mut guard = stream.lock().await;
std::future::poll_fn(|cx| guard.poll_finish(cx))
.await
.map_err(h3_stream_error_to_io)
}
#[allow(clippy::type_complexity)]
#[allow(clippy::too_many_arguments)]
async fn handle_request<F, Fut, ResB, ResBE, ResE>(
stream: SharedRequest,
shared: Arc<parking_lot::Mutex<SharedCodecs>>,
stream_id: u64,
request_fn: Rc<F>,
date_cache: DateCache,
send_continue_response: bool,
send_date_header: bool,
conn_close: Arc<parking_lot::Mutex<Option<u64>>>,
) where
F: Fn(Request<Incoming>) -> Fut,
Fut: std::future::Future<Output = Result<Response<ResB>, ResE>>,
ResB: Body<Data = Bytes, Error = ResBE> + Unpin,
ResE: std::error::Error,
ResBE: std::error::Error,
{
let request_headers = {
let mut guard = stream.lock().await;
std::future::poll_fn(|cx| guard.poll_headers(cx)).await
};
let request = match request_headers {
Ok(Some(request)) => request,
Ok(None) => return,
Err(err) => {
if !err.is_stream_scoped() {
*conn_close.lock() = Some(err.h3_code());
}
return;
}
};
let is_100_continue = send_continue_response
&& request
.headers()
.get(http::header::EXPECT)
.and_then(|v| v.to_str().ok())
.is_some_and(|v| v.eq_ignore_ascii_case("100-continue"));
let send_continue_body = is_100_continue.then(|| Arc::new(AtomicBool::new(false)));
let (request_parts, _) = request.into_parts();
let (request_body, upgrade) = if request_parts.method == http::Method::CONNECT {
(Incoming::Empty, Some(stream.clone()))
} else {
(
Incoming::Boxed(Box::pin(H3Body::new(
stream.clone(),
send_continue_body.clone(),
))),
None,
)
};
let mut request = Request::from_parts(request_parts, request_body);
let (early_hints, mut early_hints_rx) = EarlyHints::new_lazy();
request.extensions_mut().insert(early_hints);
let upgrade = if let Some(recv_stream) = upgrade {
let (upgrade_tx, upgrade_rx) = oneshot::async_channel();
let upgrade = Upgrade::new(upgrade_rx);
let upgraded = upgrade.upgraded.clone();
request.extensions_mut().insert(upgrade);
Some((upgrade_tx, upgraded, recv_stream))
} else {
None
};
let mut response_fut = std::pin::pin!(request_fn(request));
let mut early_hints_open = true;
let mut continue_sent = false;
let response_result = loop {
if !early_hints_open {
break response_fut.as_mut().await;
}
let next = std::future::poll_fn(|cx| {
if let Poll::Ready(res) = response_fut.as_mut().poll(cx) {
return Poll::Ready(Some(futures_util::future::Either::Left(res)));
}
match early_hints_rx.poll_recv(cx) {
Poll::Ready(Some(msg)) => {
return Poll::Ready(Some(futures_util::future::Either::Right(Ok(msg))))
}
Poll::Ready(None) => {
return Poll::Ready(Some(futures_util::future::Either::Right(Err(()))))
}
Poll::Pending => {}
}
if !continue_sent
&& is_100_continue
&& send_continue_body
.as_ref()
.is_some_and(|b| b.load(Ordering::Relaxed))
{
continue_sent = true;
return Poll::Ready(None);
}
Poll::Pending
})
.await;
match next {
Some(futures_util::future::Either::Left(response_result)) => {
break response_result;
}
Some(futures_util::future::Either::Right(Ok((headers, sender)))) => {
sender
.into_inner()
.send(
send_response(
&stream,
&shared,
stream_id,
StatusCode::EARLY_HINTS,
&headers,
)
.await,
)
.ok();
}
Some(futures_util::future::Either::Right(Err(()))) => {
early_hints_open = false;
}
None => {
if send_interim_response(&stream, StatusCode::CONTINUE)
.await
.is_err()
{
return;
}
}
}
};
let Ok(mut response) = response_result else {
return;
};
{
let response_headers = response.headers_mut();
if send_date_header {
if let Some(http_date) = date_cache.get_date_header_value() {
response_headers
.entry(http::header::DATE)
.or_insert(http_date);
}
}
remove_invalid_http3_headers(response_headers);
}
let response_is_end_stream = response.body().is_end_stream();
if !response_is_end_stream {
if let Some(content_length) = response.body().size_hint().exact() {
if !response
.headers()
.contains_key(http::header::CONTENT_LENGTH)
{
response
.headers_mut()
.insert(http::header::CONTENT_LENGTH, content_length.into());
}
}
}
if is_100_continue
&& !continue_sent
&& !response.status().is_client_error()
&& !response.status().is_server_error()
&& send_interim_response(&stream, StatusCode::CONTINUE)
.await
.is_err()
{
return;
}
let (response_parts, mut response_body) = response.into_parts();
if send_response(
&stream,
&shared,
stream_id,
response_parts.status,
&response_parts.headers,
)
.await
.is_err()
{
return;
}
if let Some((upgrade_tx, upgraded, recv_stream)) = upgrade {
if upgraded.load(Ordering::Relaxed) {
let (upgraded, task) = self::upgrade::pair(recv_stream);
let _ = upgrade_tx.send(Upgraded::new(upgraded, None));
task.await;
return;
}
}
if !response_is_end_stream {
while let Some(chunk) = response_body.frame().await {
match chunk {
Ok(frame) => {
if frame.is_data() {
match frame.into_data() {
Ok(data) => {
if send_data(&stream, data).await.is_err() {
return;
}
}
Err(_) => {
return;
}
}
} else if frame.is_trailers() {
match frame.into_trailers() {
Ok(mut trailers) => {
remove_invalid_http3_headers(&mut trailers);
if send_trailers(&stream, &trailers).await.is_err() {
return;
}
break;
}
Err(_) => {
return;
}
}
}
}
Err(_) => {
return;
}
}
}
}
let _ = send_finish(&stream).await;
}
pub struct Http3<Io> {
io_to_handshake: Option<Io>,
date_header_value_cached: DateCache,
options: Http3Options,
cancel_token: Option<CancellationToken>,
}
impl<Io> Http3<Io>
where
Io: transport::Connection + Unpin + 'static,
{
#[inline]
pub fn new(io: Io, options: Http3Options) -> Self {
Self {
io_to_handshake: Some(io),
date_header_value_cached: DateCache::default(),
options,
cancel_token: None,
}
}
#[inline]
pub fn graceful_shutdown_token(mut self, token: CancellationToken) -> Self {
self.cancel_token = Some(token);
self
}
}
impl<Io> HttpProtocol for Http3<Io>
where
Io: transport::Connection + Unpin + 'static,
{
#[allow(clippy::manual_async_fn)]
#[inline]
fn handle<F, Fut, ResB, ResBE, ResE>(
self,
request_fn: F,
) -> impl std::future::Future<Output = Result<(), std::io::Error>>
where
F: Fn(Request<super::Incoming>) -> Fut + 'static,
Fut: std::future::Future<Output = Result<Response<ResB>, ResE>> + 'static,
ResB: http_body::Body<Data = bytes::Bytes, Error = ResBE> + Unpin + 'static,
ResE: std::error::Error + 'static,
ResBE: std::error::Error + 'static,
{
async move {
let request_fn = Rc::new(request_fn);
let Http3 {
mut io_to_handshake,
date_header_value_cached,
options,
cancel_token,
} = self;
let mut conn = io_to_handshake
.take()
.ok_or_else(|| std::io::Error::other("no io to handshake"))?;
let date_cache = date_header_value_cached;
let send_continue_response = options.send_continue_response;
let send_date_header = options.send_date_header;
if let Some(timeout) = options.handshake_timeout {
vibeio::time::timeout(timeout, async {
while !conn.is_handshake_complete() {
vibeio::time::sleep(std::time::Duration::from_millis(1)).await;
}
})
.await
.map_err(|_| {
std::io::Error::new(std::io::ErrorKind::TimedOut, "handshake timeout")
})?;
} else {
while !conn.is_handshake_complete() {
vibeio::time::sleep(std::time::Duration::from_millis(1)).await;
}
}
let mut controls = ControlStreams::new(options.local_settings.clone());
let shared = controls.shared().clone();
let conn_close: Arc<parking_lot::Mutex<Option<u64>>> =
Arc::new(parking_lot::Mutex::new(None));
let mut ongoing: FuturesUnordered<oneshot::AsyncReceiver<()>> = FuturesUnordered::new();
let mut cancel_fut: Option<Pin<Box<dyn std::future::Future<Output = ()> + Send>>> =
None;
if let Some(token) = cancel_token.as_ref() {
cancel_fut = Some(Box::pin(token.cancelled()));
}
let mut accept_sleep: Option<Pin<Box<vibeio::time::Sleep>>> = None;
let mut shutdown_sleep: Option<Pin<Box<vibeio::time::Sleep>>> = None;
let mut drain_grace: Option<Pin<Box<vibeio::time::Sleep>>> = None;
let mut shutdown = false;
let mut control_dead = false;
let mut closing_with: Option<u64> = None;
let mut outcome: Option<Result<(), std::io::Error>> = None;
let mut last_request_id = 0u64;
std::future::poll_fn(|cx| -> Poll<Result<(), std::io::Error>> {
ready!(controls
.poll_init(&mut conn, cx)
.map_err(h3_control_error_to_io))?;
ready!(controls.poll_flush(cx).map_err(h3_control_error_to_io))?;
Poll::Ready(Ok(()))
})
.await?;
std::future::poll_fn(|cx| loop {
if let Some(code) = conn_close.lock().take() {
if closing_with.is_none() {
closing_with = Some(code);
shutdown = true;
control_dead = true;
}
}
let mut timeout_fired = false;
if let Some(sleep) = accept_sleep.as_mut() {
if let Poll::Ready(()) = sleep.as_mut().poll(cx) {
accept_sleep = None;
timeout_fired = true;
}
} else if let Some(accept_timeout) = options.accept_timeout {
accept_sleep = Some(Box::pin(vibeio::time::sleep(accept_timeout)));
continue;
}
if let Some(sleep) = shutdown_sleep.as_mut() {
if let Poll::Ready(()) = sleep.as_mut().poll(cx) {
shutdown_sleep = None;
}
}
let mut cancel_fired = false;
if let Some(fut) = cancel_fut.as_mut() {
if let Poll::Ready(()) = fut.as_mut().poll(cx) {
cancel_fired = true;
}
}
if !shutdown {
if cancel_fired {
shutdown = true;
outcome = Some(Ok(()));
} else if timeout_fired {
shutdown = true;
outcome = Some(Err(std::io::Error::new(
std::io::ErrorKind::TimedOut,
"accept timeout",
)));
}
}
{
let mut shared = shared.lock();
controls.queue_encoder_streams(&mut shared.encoder_stream);
}
if !control_dead {
match controls.poll_flush(cx) {
Poll::Ready(Ok(())) => {}
Poll::Ready(Err(err)) => {
if shutdown {
control_dead = true;
} else {
closing_with = Some(err.h3_code());
shutdown = true;
control_dead = true;
}
}
Poll::Pending => {}
}
}
if !control_dead {
loop {
match controls.poll_read(&mut conn, cx) {
Poll::Ready(Ok(Some(ControlEvent::Goaway { .. }))) => {
if !shutdown {
shutdown = true;
outcome = Some(Ok(()));
}
}
Poll::Ready(Ok(Some(_))) => {}
Poll::Ready(Ok(None)) => {}
Poll::Ready(Err(err)) => {
if shutdown {
control_dead = true;
break;
}
closing_with = Some(err.h3_code());
shutdown = true;
control_dead = true;
break;
}
Poll::Pending => break,
}
}
}
if shutdown {
if let Some(code) = closing_with {
ready!(conn
.poll_shutdown(cx, code)
.map_err(h3_transport_error_to_io))?;
return Poll::Ready(outcome.take().unwrap_or(Ok(())));
}
if controls.goaway_sent().is_none() {
controls.send_goaway(last_request_id);
}
if ongoing.is_empty() {
if let Some(grace) = drain_grace.as_mut() {
if grace.as_mut().poll(cx).is_ready() {
ready!(conn
.poll_shutdown(cx, H3_NO_ERROR)
.map_err(h3_transport_error_to_io))?;
return Poll::Ready(outcome.take().unwrap_or(Ok(())));
}
} else {
drain_grace = Some(Box::pin(vibeio::time::sleep(
std::time::Duration::from_millis(50),
)));
}
} else if shutdown_sleep.is_none() {
shutdown_sleep = Some(Box::pin(vibeio::time::sleep(
std::time::Duration::from_millis(10),
)));
}
}
match conn.poll_accept(cx) {
Poll::Ready(Ok(Some(stream))) => {
let id = stream.id();
last_request_id = last_request_id.max(id);
accept_sleep = None;
if shutdown
&& (controls.goaway_sent().is_none()
|| id > controls.goaway_sent().unwrap_or(u64::MAX))
{
let mut rejected = RequestStream::new(stream, shared.clone());
let _ = rejected.poll_reset(cx, H3_REQUEST_REJECTED);
} else {
let (end_tx, end_rx) = oneshot::async_channel();
ongoing.push(end_rx);
let request_stream = Arc::new(tokio::sync::Mutex::new(
RequestStream::new(stream, shared.clone()),
));
let request_fn = request_fn.clone();
let date_cache = date_cache.clone();
let shared = shared.clone();
let conn_close_for_task = conn_close.clone();
vibeio::spawn(async move {
let _end = end_tx;
handle_request(
request_stream,
shared.clone(),
id,
request_fn,
date_cache,
send_continue_response,
send_date_header,
conn_close_for_task,
)
.await;
});
}
}
Poll::Ready(Ok(None)) => {
return Poll::Ready(Ok(()));
}
Poll::Ready(Err(err)) => {
return Poll::Ready(Err(h3_transport_error_to_io(err)));
}
Poll::Pending => {}
}
match ongoing.poll_next_unpin(cx) {
Poll::Ready(Some(Ok(()))) => continue,
Poll::Ready(Some(Err(_))) => continue,
Poll::Ready(None) => {}
Poll::Pending => {}
}
if ongoing.is_empty() && shutdown {
continue;
}
return Poll::Pending;
})
.await
}
}
}