use std::collections::VecDeque;
use bytes::{Bytes, BytesMut};
#[derive(Debug, Clone, Copy, PartialEq)]
pub struct H3Limits {
pub max_message_size: u64,
pub max_message_body_size: u64,
pub max_decompressed_body_size: u64,
pub max_headers_size: u64,
pub max_header_count: u16,
pub max_concurrent_streams: u32,
pub max_connection_buffer_size: u64,
pub max_premature_resets: u32,
pub max_requests_per_connection: u64,
pub max_encoder_table_size: u64,
pub qpack_block_timeout: f64,
pub max_peer_uni_streams: u32,
pub max_outstanding_sections: u32,
pub max_blocked_streams: u32,
pub tunnel_backlog: u32,
pub command_backlog: u32,
pub idle_capacity: u64,
pub receive_timeout: f64,
pub send_timeout: f64,
}
impl Default for H3Limits {
fn default() -> Self {
Limits::default().into()
}
}
impl From<Limits> for H3Limits {
fn from(limits: Limits) -> Self {
Self {
max_message_size: limits.max_message_size,
max_message_body_size: limits.max_message_body_size,
max_decompressed_body_size: limits.max_decompressed_body_size,
max_headers_size: limits.max_headers_size,
max_header_count: limits.max_header_count,
max_concurrent_streams: limits.max_concurrent_streams,
max_connection_buffer_size: limits.max_connection_buffer_size,
max_premature_resets: limits.max_premature_resets,
max_requests_per_connection: limits.max_requests_per_connection,
max_encoder_table_size: limits.max_encoder_table_size,
qpack_block_timeout: limits.qpack_block_timeout,
max_peer_uni_streams: limits.max_peer_uni_streams,
max_outstanding_sections: limits.max_outstanding_sections,
max_blocked_streams: limits.max_blocked_streams,
tunnel_backlog: limits.tunnel_backlog,
command_backlog: limits.command_backlog,
idle_capacity: limits.idle_capacity,
receive_timeout: limits.receive_timeout,
send_timeout: limits.send_timeout,
}
}
}
pub mod frames;
pub use frames::{Code, Frame, FrameType, Settings, StreamKind};
use crate::helpers::compression::Compression;
use crate::helpers::fields::HeaderField;
use crate::helpers::qpack::{self, Decoder, Encoder, EncoderInstruction};
use crate::models::{Body, ConnectionID, Limits, Message, Method, Role, StreamID, Version};
use crate::tls::Security;
use crate::protocol::base::{Connection, Stream};
use crate::protocol::common::{self, Error};
use crate::protocol::quic::{Handshake, QUICApplication, QUICConnection, QUICError, QUICGuard, QUICHandshake, QUICOutcome, QUICTransport, QUICStreamID, StreamRead, StreamWrite, Varint};
use crate::helpers::sync::{Lock, Timeout};
#[derive(Default)]
pub struct StreamState {
pub buffer: BytesMut,
pub body: BytesMut,
pub message: Option<Message>,
pub method: Option<Method>,
pub accepted: Option<Compression>,
pub pending: Option<Bytes>,
pub eof: bool,
pub delivered: bool,
pub finished: bool,
pub raw: bool,
pub responded: bool,
}
impl StreamState {
pub fn spent(&self) -> bool {
self.finished && self.eof && self.delivered && !self.raw
}
}
pub struct H3Session {
pub role: Role,
pub id: ConnectionID,
pub client: Option<std::net::SocketAddr>,
pub limits: H3Limits,
pub settings_local: Settings,
pub settings_remote: Option<Settings>,
pub encoder: Encoder,
pub decoder: Decoder,
pub streams: common::StreamMap<StreamID, StreamState>,
pub blocked_since: common::StreamMap<StreamID, std::time::Instant>,
pub ready: VecDeque<Message>,
pub buffered_bound: u64,
pub control_recv: BytesMut,
pub next_stream_id: u64,
pub highest_peer_stream_id: u64,
pub total_streams: u64,
pub goaway: Option<u64>,
pub goaway_sent: Option<u64>,
pub fields: Vec<HeaderField>,
pub block: Vec<u8>,
pub security: std::sync::Arc<std::sync::Mutex<Security>>,
}
impl H3Session {
pub fn new(role: Role, id: ConnectionID, limits: impl Into<H3Limits>) -> Self {
let limits: H3Limits = limits.into();
let settings_local = Settings { qpack_blocked_streams: limits.max_blocked_streams as u64, ..Settings::default() };
let mut decoder = Decoder::new();
decoder.set_max_capacity(settings_local.qpack_max_table_capacity as usize);
decoder.set_max_decoded_size(limits.max_headers_size as usize);
decoder.set_max_instruction_size(limits.max_headers_size as usize);
decoder.set_max_blocked_streams(limits.max_blocked_streams as usize);
decoder.set_idle_capacity(limits.idle_capacity as usize);
let mut encoder = Encoder::new();
encoder.set_max_outstanding_sections(limits.max_outstanding_sections as usize);
encoder.set_max_instruction_size(limits.max_headers_size as usize);
encoder.set_idle_capacity(limits.idle_capacity as usize);
if let Some(instruction) = encoder.set_capacity_limit(limits.max_encoder_table_size as usize) {
encoder.queue(&[instruction]);
}
let next_stream_id = QUICStreamID::first_bidi(role);
Self {
role,
id,
client: None,
limits,
settings_local,
settings_remote: None,
encoder,
decoder,
streams: common::StreamMap::default(),
blocked_since: common::StreamMap::default(),
ready: VecDeque::new(),
buffered_bound: 0,
control_recv: BytesMut::new(),
next_stream_id,
highest_peer_stream_id: 0,
total_streams: 0,
goaway: None,
goaway_sent: None,
fields: Vec::new(),
block: Vec::new(),
security: std::sync::Arc::new(std::sync::Mutex::new(Security::quic(None))),
}
}
pub fn with_client(mut self, client: Option<std::net::SocketAddr>) -> Self {
self.client = client;
self
}
pub const BLOCK_FLOOR: usize = 256;
pub const FRAMES_PER_MESSAGE: usize = 3;
pub fn control_frame(&self) -> Bytes {
let mut out = BytesMut::new();
Frame::Settings(self.settings_local.parameters()).encode_into(&mut out);
out.freeze()
}
pub fn open(&mut self) -> StreamID {
let stream_id = StreamID(self.next_stream_id);
self.next_stream_id += QUICStreamID::STEP;
self.streams.entry(stream_id).or_default();
stream_id
}
pub fn stream_ceiling(&self) -> usize {
(self.limits.max_concurrent_streams as usize).saturating_mul(2).max(2)
}
pub fn forget(&mut self, stream_id: StreamID) -> Option<StreamState> {
self.blocked_since.remove(&stream_id);
self.encoder.cancel(stream_id.0);
self.decoder.cancel(stream_id.0);
self.streams.remove(&stream_id)
}
pub fn retire(&mut self, stream_id: StreamID) {
if self.streams.get(&stream_id).is_some_and(StreamState::spent) {
self.forget(stream_id);
}
}
pub fn encode_message(&mut self, stream_id: StreamID, message: &mut Message) -> Result<(Bytes, bool), Error> {
let mut out = BytesMut::with_capacity(self.message_room(message));
let fin = self.encode_message_into(stream_id, message, &mut out)?;
Ok((out.freeze(), fin))
}
pub fn message_room(&self, message: &Message) -> usize {
let body = message.body.as_ref().and_then(Body::len).unwrap_or(0);
let frames = Self::FRAMES_PER_MESSAGE * 2 * Varint::len(Varint::MAXIMUM);
self.block.capacity().max(Self::BLOCK_FLOOR) + body + frames
}
pub fn encode_message_into(&mut self, stream_id: StreamID, message: &mut Message, out: &mut BytesMut) -> Result<bool, Error> {
let start = out.len();
match self.frame_message(stream_id, message, out) {
Ok(fin) => Ok(fin),
Err(error) => {
out.truncate(start);
Err(error)
}
}
}
pub fn write_block(
&mut self,
stream_id: StreamID,
kind: FrameType,
out: &mut BytesMut,
gather: impl FnOnce(&mut Vec<HeaderField>) -> Result<(), Error>,
) -> Result<(), Error> {
self.fields.clear();
gather(&mut self.fields)?;
self.block.clear();
self.encoder.encode_into(&mut self.block, stream_id.0, &self.fields);
Frame::write(kind, &self.block, out);
let idle = self.limits.idle_capacity as usize;
common::Buffer::reclaim_octets(&mut self.block, idle);
Ok(())
}
pub fn frame_message(&mut self, stream_id: StreamID, message: &mut Message, out: &mut BytesMut) -> Result<bool, Error> {
let accepted = message.is_response().then(|| self.streams.get(&stream_id)?.accepted).flatten();
message.compress(accepted)?;
let framed = !message.bodyless(self.streams.get(&stream_id).and_then(|state| state.method));
let message = &*message;
self.write_block(stream_id, FrameType::Headers, out, |fields| common::Fields::write(message, fields))?;
if let Some(body) = message.body.as_ref().filter(|_| framed) {
let body = body
.inline()
.ok_or_else(|| Error::Protocol("a file body must be materialised before HTTP/3 encoding".into()))?;
if !body.is_empty() {
Frame::Data(body).encode_into(out);
}
}
if let Some(trailers) = message.trailers.as_ref().filter(|trailers| framed && !trailers.is_empty()) {
self.write_block(stream_id, FrameType::Headers, out, |fields| {
fields.extend_from_slice(trailers.fields());
Ok(())
})?;
}
let state = self.streams.entry(stream_id).or_default();
if message.is_request() {
state.method = message.method;
}
let tunneling = message.tunneling(state.method);
state.finished |= !tunneling;
Ok(!tunneling)
}
pub fn on_encoder_bytes(&mut self, bytes: &[u8]) -> Result<(), Error> {
self.decoder.on_encoder_stream(bytes).map_err(|err| match err {
qpack::Error::InstructionTooLarge => {
Error::Limit(format!("an encoder instruction exceeds {} octets", self.limits.max_headers_size))
}
err => err.into(),
})?;
if self.blocked_since.is_empty() {
return Ok(());
}
for stream_id in self.decoder.unblocked() {
self.advance(StreamID(stream_id))?;
}
Ok(())
}
pub fn on_decoder_bytes(&mut self, bytes: &[u8]) -> Result<(), Error> {
self.encoder.on_decoder_stream(bytes).map_err(|err| match err {
qpack::Error::InstructionTooLarge => {
Error::Limit(format!("a decoder instruction exceeds {} octets", self.limits.max_headers_size))
}
err => err.into(),
})
}
pub fn on_control_bytes(&mut self, bytes: &[u8]) -> Result<(), Error> {
self.control_recv.extend_from_slice(bytes);
let limit = self.limits.max_headers_size;
if self.control_recv.len() as u64 > limit {
return Err(Error::Limit(format!("a control frame exceeds {limit} octets")));
}
while let Some(frame) = Frame::parse(&mut self.control_recv)? {
match frame {
Frame::Settings(parameters) => {
if self.settings_remote.is_some() {
return Err(Error::Protocol("a second SETTINGS frame arrived on the control stream".into()));
}
let mut settings = Settings::peer();
for (id, value) in parameters {
settings.apply(id, value)?;
}
self.apply_peer_settings(settings);
}
Frame::GoAway { id } => {
self.goaway = Some(self.goaway.map_or(id, |earlier| earlier.min(id)));
}
Frame::MaxPushID { .. } | Frame::CancelPush { .. } => {}
_ => return Err(Error::Protocol("an unexpected frame arrived on the control stream".into())),
}
}
Ok(())
}
pub fn apply_peer_settings(&mut self, settings: Settings) {
let permitted = usize::try_from(settings.qpack_max_table_capacity).unwrap_or(usize::MAX);
if let Some(instruction) = self.encoder.set_max_capacity(permitted) {
self.encoder.queue(&[instruction]);
}
self.settings_remote = Some(settings);
}
pub fn on_stream_bytes(&mut self, stream_id: StreamID, bytes: &[u8], fin: bool) -> Result<(), Error> {
let created = !self.streams.contains_key(&stream_id);
if created && self.streams.len() >= self.stream_ceiling() {
let reason = format!("more than {} streams are held open at once", self.stream_ceiling());
return Err(Error::stream(stream_id, Code::EXCESSIVE_LOAD, reason));
}
if created {
self.total_streams += 1;
self.highest_peer_stream_id = self.highest_peer_stream_id.max(stream_id.0);
}
let state = self.streams.entry(stream_id).or_default();
state.buffer.extend_from_slice(bytes);
if fin {
state.eof = true;
}
let unparsed = state.buffer.len() as u64;
self.buffered_bound = self.buffered_bound.saturating_add(bytes.len() as u64);
let limit = self.limits.max_message_size;
if unparsed > limit {
let reason = format!("unparsed stream data exceeds {limit} octets");
return Err(Error::stream(stream_id, Code::EXCESSIVE_LOAD, reason));
}
if self.overbuffered() {
let limit = self.limits.max_connection_buffer_size;
return Err(Error::Limit(format!("buffered messages exceed {limit} octets")));
}
self.advance(stream_id)
}
pub fn buffered(&self) -> u64 {
self.streams.values().map(|state| (state.buffer.len() + state.body.len()) as u64).sum()
}
pub fn overbuffered(&mut self) -> bool {
let limit = self.limits.max_connection_buffer_size;
if self.buffered_bound <= limit {
return false;
}
self.buffered_bound = self.buffered();
self.buffered_bound > limit
}
pub fn advance(&mut self, stream_id: StreamID) -> Result<(), Error> {
loop {
if self.streams.get(&stream_id).is_some_and(|state| state.raw) {
return Ok(());
}
if let Some(block) = self.streams.get(&stream_id).and_then(|state| state.pending.clone()) {
match self.decoder.decode(stream_id.0, &block) {
Ok((fields, acknowledgment)) => {
if let Some(acknowledgment) = acknowledgment {
self.decoder.queue(&[acknowledgment]);
}
let state = self.streams.get_mut(&stream_id).ok_or(Error::Closed)?;
state.pending = None;
self.blocked_since.remove(&stream_id);
self.absorb_headers(stream_id, fields)?;
continue;
}
Err(qpack::Error::Blocked) => return Ok(()),
Err(err) => return Err(err.into()),
}
}
let Some(state) = self.streams.get_mut(&stream_id) else {
return Ok(());
};
let Some(frame) = Frame::parse(&mut state.buffer)? else {
if state.eof {
break;
}
return Ok(());
};
match frame {
Frame::Headers(block) => {
if block.len() as u64 > self.limits.max_headers_size {
let reason = format!("field section exceeds {} octets", self.limits.max_headers_size);
return Err(Error::stream(stream_id, Code::EXCESSIVE_LOAD, reason));
}
match self.decoder.decode(stream_id.0, &block) {
Ok((fields, acknowledgment)) => {
if let Some(acknowledgment) = acknowledgment {
self.decoder.queue(&[acknowledgment]);
}
self.absorb_headers(stream_id, fields)?;
}
Err(qpack::Error::Blocked) => {
let state = self.streams.get_mut(&stream_id).ok_or(Error::Closed)?;
state.pending = Some(block);
self.blocked_since.entry(stream_id).or_insert_with(std::time::Instant::now);
return Ok(());
}
Err(err) => return Err(err.into()),
}
}
Frame::Data(data) => {
let state = self.streams.get_mut(&stream_id).ok_or(Error::Closed)?;
if state.message.is_none() {
return Err(Error::Protocol("DATA arrived before HEADERS".into()));
}
state.body.extend_from_slice(&data);
let limit = self.limits.max_message_body_size;
if state.body.len() as u64 > limit {
return Err(Error::stream(stream_id, Code::EXCESSIVE_LOAD, format!("body exceeds {limit} octets")));
}
}
Frame::Settings(_) | Frame::GoAway { .. } | Frame::MaxPushID { .. } | Frame::CancelPush { .. } => {
return Err(Error::Protocol("a control frame arrived on a request stream".into()));
}
Frame::PushPromise { .. } => {
return Err(Error::Protocol("PUSH_PROMISE arrived with push disabled".into()));
}
}
}
let state = self.streams.get_mut(&stream_id).ok_or(Error::Closed)?;
if state.delivered {
self.retire(stream_id);
return Ok(());
}
let Some(mut message) = state.message.take() else {
return Err(Error::stream(stream_id, Code::REQUEST_INCOMPLETE, "request stream carried no field section"));
};
if !state.body.is_empty() {
message.body = Some(Body::Data(std::mem::take(&mut state.body).freeze()));
}
state.delivered = true;
self.ready.push_back(message);
self.retire(stream_id);
Ok(())
}
pub fn absorb_headers(&mut self, stream_id: StreamID, fields: Vec<HeaderField>) -> Result<(), Error> {
if fields.len() > self.limits.max_header_count as usize {
let reason = format!("more than {} header fields", self.limits.max_header_count);
return Err(Error::stream(stream_id, Code::EXCESSIVE_LOAD, reason));
}
let id = self.id.clone();
let client = self.client;
let state = self.streams.get_mut(&stream_id).ok_or(Error::Closed)?;
if state.message.is_some() {
let trailers = common::Fields::into_trailers(fields).map_err(|err| err.on_stream(stream_id, Code::MESSAGE_ERROR))?;
let state = self.streams.get_mut(&stream_id).ok_or(Error::Closed)?;
if let Some(message) = state.message.as_mut() {
message.trailers = Some(trailers);
}
return Ok(());
}
let mut message = common::Fields::into_message(fields, Version::V3_0).map_err(|err| err.on_stream(stream_id, Code::MESSAGE_ERROR))?;
message.stream_id = Some(stream_id);
message.connection_id = Some(id);
message.client = client;
Lock::on(&self.security).apply(&mut message);
if message.is_request() {
state.method = message.method;
state.accepted = message.accepted();
}
if message.tunneling(state.method) {
state.delivered = true;
state.raw = true;
self.ready.push_back(message);
return Ok(());
}
state.message = Some(message);
Ok(())
}
pub fn take_ready(&mut self) -> Option<Message> {
self.ready.pop_front()
}
pub fn take_encoder_out(&mut self) -> Bytes {
Bytes::from(self.encoder.take_encoder_stream())
}
pub fn take_decoder_out(&mut self) -> Bytes {
Bytes::from(self.decoder.take_decoder_stream())
}
}
#[allow(clippy::large_enum_variant)]
pub enum H3Command {
Send(Message),
Open(tokio::sync::oneshot::Sender<StreamID>),
OpenUni(StreamKind, tokio::sync::oneshot::Sender<StreamID>),
Tunnel(StreamID, tokio::sync::mpsc::Sender<(Bytes, bool)>),
WriteEncoder(Bytes),
Reset(StreamID, u64),
Close,
}
impl H3Command {
pub fn buffers(&self) -> bool {
matches!(self, Self::Send(_) | Self::WriteEncoder(_))
}
}
#[allow(clippy::large_enum_variant)]
pub enum H3Event {
Message(Message),
Failed(Error),
}
pub struct H3Connection {
pub commands: tokio::sync::mpsc::Sender<H3Command>,
pub events: tokio::sync::mpsc::Receiver<H3Event>,
pub raw: tokio::sync::mpsc::Sender<(StreamID, Bytes, bool)>,
pub id: ConnectionID,
pub role: Role,
pub client: Option<std::net::SocketAddr>,
pub limits: H3Limits,
pub guard: Option<std::sync::Arc<QUICGuard>>,
pub settings_local: Settings,
pub request_finalizer: crate::finalizer::RequestFinalizer,
pub response_finalizer: crate::finalizer::ResponseFinalizer,
pub security: std::sync::Arc<std::sync::Mutex<Security>>,
}
impl H3Connection {
pub fn pair(session: H3Session) -> (Self, H3Worker) {
let backlog = (session.limits.command_backlog as usize).max(1);
let (commands, commands_receiver) = tokio::sync::mpsc::channel(backlog);
let (events_sender, events) = tokio::sync::mpsc::channel(backlog);
let (raw, raw_receiver) = tokio::sync::mpsc::channel(session.limits.tunnel_backlog as usize);
let connection = Self {
commands,
events,
raw,
id: session.id.clone(),
role: session.role,
client: session.client,
limits: session.limits,
guard: None,
settings_local: session.settings_local,
request_finalizer: crate::finalizer::RequestFinalizer::default(),
response_finalizer: crate::finalizer::ResponseFinalizer::new(None),
security: std::sync::Arc::clone(&session.security),
};
(connection, H3Worker::new(session, commands_receiver, events_sender, raw_receiver))
}
pub fn limits(&self) -> H3Limits {
self.limits
}
pub fn settings_local(&self) -> &Settings {
&self.settings_local
}
pub fn with_request_finalizer(mut self, finalizer: crate::finalizer::RequestFinalizer) -> Self {
self.request_finalizer = finalizer;
self
}
pub fn with_response_finalizer(mut self, finalizer: crate::finalizer::ResponseFinalizer) -> Self {
self.response_finalizer = finalizer;
self
}
pub fn with_security(self, security: Security) -> Self {
*Lock::on(&self.security) = security;
self
}
pub fn with_guard(mut self, guard: std::sync::Arc<QUICGuard>) -> Self {
self.guard = Some(guard);
self
}
pub async fn send_message(&mut self, message: Message) -> Result<(), Error> {
let mut message = message;
self.request_finalizer.finalize(self.role, &mut message);
self.response_finalizer.finalize(self.role, Lock::on(&self.security).secure, &mut message);
message.materialize().await?;
self.commands.send(H3Command::Send(message)).await.map_err(|_| Error::Closed)
}
pub async fn receive_message(&mut self) -> Result<Message, Error> {
match self.events.recv().await {
Some(H3Event::Message(mut message)) => {
message.decompress(self.limits.max_decompressed_body_size)?;
Ok(message)
}
Some(H3Event::Failed(error)) => Err(error),
None => Err(Error::Closed),
}
}
pub async fn start(&mut self) -> Result<(), Error> {
Ok(())
}
pub async fn open(&mut self) -> Result<StreamID, Error> {
let (reply, opened) = tokio::sync::oneshot::channel();
self.commands.send(H3Command::Open(reply)).await.map_err(|_| Error::Closed)?;
opened.await.map_err(|_| Error::Closed)
}
pub async fn open_uni(&mut self, kind: StreamKind) -> Result<H3Stream, Error> {
if kind.code().is_none() {
return Err(Error::Protocol(format!("{kind:?} is not a unidirectional stream type")));
}
let (reply, opened) = tokio::sync::oneshot::channel();
self.commands.send(H3Command::OpenUni(kind, reply)).await.map_err(|_| Error::Closed)?;
let stream_id = opened.await.map_err(|_| Error::Closed)?;
let (_, silent) = tokio::sync::mpsc::channel(1);
Ok(H3Stream::new(stream_id, self.raw.clone(), silent, self.guard.clone()))
}
pub fn tunnel(&mut self, stream_id: StreamID) -> Result<H3Stream, Error> {
let (sink, reads) = tokio::sync::mpsc::channel(self.limits.tunnel_backlog as usize);
self.commands.try_send(H3Command::Tunnel(stream_id, sink)).map_err(|_| Error::Closed)?;
Ok(H3Stream::new(stream_id, self.raw.clone(), reads, self.guard.clone()).with_commands(self.commands.clone()))
}
pub async fn write_encoder(&mut self, instructions: &[EncoderInstruction]) -> Result<(), Error> {
let mut bytes = BytesMut::new();
for instruction in instructions {
bytes.extend_from_slice(&instruction.encode());
}
self.commands.send(H3Command::WriteEncoder(bytes.freeze())).await.map_err(|_| Error::Closed)
}
}
pub type RawWrite = (StreamID, Bytes, bool);
pub type RawPermit = tokio::sync::mpsc::OwnedPermit<RawWrite>;
pub type Reserving = std::pin::Pin<Box<dyn std::future::Future<Output = Result<RawPermit, tokio::sync::mpsc::error::SendError<()>>> + Send>>;
pub struct H3Stream {
pub id: StreamID,
pub writes: tokio::sync::mpsc::Sender<RawWrite>,
pub reads: tokio::sync::mpsc::Receiver<(Bytes, bool)>,
pub reserving: Option<Reserving>,
pub buffer: BytesMut,
pub eof: bool,
pub guard: Option<std::sync::Arc<QUICGuard>>,
pub commands: Option<tokio::sync::mpsc::Sender<H3Command>>,
}
impl H3Stream {
pub fn new(id: StreamID, writes: tokio::sync::mpsc::Sender<RawWrite>, reads: tokio::sync::mpsc::Receiver<(Bytes, bool)>, guard: Option<std::sync::Arc<QUICGuard>>) -> Self {
Self { id, writes, reads, reserving: None, buffer: BytesMut::new(), eof: false, guard, commands: None }
}
pub fn with_commands(mut self, commands: tokio::sync::mpsc::Sender<H3Command>) -> Self {
self.commands = Some(commands);
self
}
pub fn reserve(&mut self, context: &mut std::task::Context<'_>) -> std::task::Poll<Option<RawPermit>> {
use std::task::Poll;
loop {
if let Some(reserving) = &mut self.reserving {
let reserved = std::future::Future::poll(reserving.as_mut(), context);
return match reserved {
Poll::Ready(permit) => {
self.reserving = None;
Poll::Ready(permit.ok())
}
Poll::Pending => Poll::Pending,
};
}
match self.writes.clone().try_reserve_owned() {
Ok(permit) => return Poll::Ready(Some(permit)),
Err(tokio::sync::mpsc::error::TrySendError::Full(sender)) => {
self.reserving = Some(Box::pin(sender.reserve_owned()));
}
Err(tokio::sync::mpsc::error::TrySendError::Closed(_)) => return Poll::Ready(None),
}
}
}
pub fn guard(&self) -> Option<&std::sync::Arc<QUICGuard>> {
self.guard.as_ref()
}
}
impl Stream for H3Stream {
fn id(&self) -> StreamID {
self.id
}
async fn reset(&mut self, code: u64) {
self.eof = true;
if let Some(commands) = &self.commands {
let _ = commands.send(H3Command::Reset(self.id, code)).await;
}
}
}
impl tokio::io::AsyncRead for H3Stream {
fn poll_read(mut self: std::pin::Pin<&mut Self>, context: &mut std::task::Context<'_>, out: &mut tokio::io::ReadBuf<'_>) -> std::task::Poll<std::io::Result<()>> {
loop {
if !self.buffer.is_empty() {
let count = self.buffer.len().min(out.remaining());
out.put_slice(&self.buffer.split_to(count));
return std::task::Poll::Ready(Ok(()));
}
if self.eof {
return std::task::Poll::Ready(Ok(()));
}
match self.reads.poll_recv(context) {
std::task::Poll::Ready(Some((bytes, fin))) => {
self.buffer.extend_from_slice(&bytes);
self.eof |= fin;
}
std::task::Poll::Ready(None) => {
self.eof = true;
}
std::task::Poll::Pending => return std::task::Poll::Pending,
}
}
}
}
impl tokio::io::AsyncWrite for H3Stream {
fn poll_write(mut self: std::pin::Pin<&mut Self>, context: &mut std::task::Context<'_>, data: &[u8]) -> std::task::Poll<std::io::Result<usize>> {
let id = self.id;
match std::task::ready!(self.reserve(context)) {
Some(permit) => {
permit.send((id, Bytes::copy_from_slice(data), false));
std::task::Poll::Ready(Ok(data.len()))
}
None => std::task::Poll::Ready(Err(std::io::Error::from(std::io::ErrorKind::BrokenPipe))),
}
}
fn poll_flush(self: std::pin::Pin<&mut Self>, _context: &mut std::task::Context<'_>) -> std::task::Poll<std::io::Result<()>> {
std::task::Poll::Ready(Ok(()))
}
fn poll_shutdown(mut self: std::pin::Pin<&mut Self>, context: &mut std::task::Context<'_>) -> std::task::Poll<std::io::Result<()>> {
let id = self.id;
if let Some(permit) = std::task::ready!(self.reserve(context)) {
permit.send((id, Bytes::new(), true));
}
std::task::Poll::Ready(Ok(()))
}
}
impl Connection for H3Connection {
fn version(&self) -> Version {
Version::V3_0
}
fn role(&self) -> Role {
self.role
}
fn id(&self) -> ConnectionID {
self.id.clone()
}
fn security(&self) -> Security {
*Lock::on(&self.security)
}
fn client(&self) -> Option<std::net::SocketAddr> {
self.client
}
async fn send(&mut self, message: Message) -> Result<(), Error> {
let timeout = self.limits.send_timeout;
let sending = std::pin::pin!(self.send_message(message));
Timeout::within(timeout, sending).await?
}
async fn receive(&mut self) -> Result<Message, Error> {
let timeout = self.limits.receive_timeout;
let receiving = std::pin::pin!(self.receive_message());
Timeout::within(timeout, receiving).await?
}
async fn close(&mut self) {
let _ = self.commands.send(H3Command::Close).await;
}
}
#[derive(Default)]
pub struct PeerUni {
pub kind: Option<StreamKind>,
pub prefix: Vec<u8>,
pub abandoned: bool,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum Side {
Encoder,
Decoder,
}
pub struct H3Worker {
pub session: H3Session,
pub commands: tokio::sync::mpsc::Receiver<H3Command>,
pub events: tokio::sync::mpsc::Sender<H3Event>,
pub raw_writes: tokio::sync::mpsc::Receiver<(StreamID, Bytes, bool)>,
pub pending: Option<H3Command>,
pub pending_raw: Option<(StreamID, Bytes, bool)>,
pub orphaned: bool,
pub established: bool,
pub next_uni: u64,
pub local_control: u64,
pub local_encoder: u64,
pub local_decoder: u64,
pub peer_uni: common::StreamMap<u64, PeerUni>,
pub tunnels: common::StreamMap<u64, tokio::sync::mpsc::Sender<(Bytes, bool)>>,
pub outbound: common::StreamMap<u64, (BytesMut, bool)>,
pub outbound_bytes: u64,
pub premature_resets: u32,
pub rejected: u32,
pub rejected_floor: u64,
pub drained: bool,
pub scratch: Vec<u8>,
pub read: Vec<u8>,
pub readable: Vec<u64>,
pub flushing: Vec<u64>,
}
impl H3Worker {
pub fn boxed(error: Error) -> QUICError {
Box::new(error)
}
pub fn new(session: H3Session, commands: tokio::sync::mpsc::Receiver<H3Command>, events: tokio::sync::mpsc::Sender<H3Event>, raw_writes: tokio::sync::mpsc::Receiver<(StreamID, Bytes, bool)>) -> Self {
let next_uni = QUICStreamID::first_uni(session.role);
Self {
session,
commands,
events,
raw_writes,
pending: None,
pending_raw: None,
orphaned: false,
established: false,
next_uni,
local_control: 0,
local_encoder: 0,
local_decoder: 0,
peer_uni: common::StreamMap::default(),
tunnels: common::StreamMap::default(),
outbound: common::StreamMap::default(),
outbound_bytes: 0,
premature_resets: 0,
rejected: 0,
rejected_floor: 0,
drained: false,
scratch: vec![0u8; 64 * 1024],
read: vec![0u8; 64 * 1024],
readable: Vec::new(),
flushing: Vec::new(),
}
}
pub fn forget_stream(&mut self, stream_id: u64) {
self.session.forget(StreamID(stream_id));
self.peer_uni.remove(&stream_id);
self.tunnels.remove(&stream_id);
if let Some((buffer, _)) = self.outbound.remove(&stream_id) {
self.outbound_bytes = self.outbound_bytes.saturating_sub(buffer.len() as u64);
}
}
pub fn outbound_limit(&self) -> u64 {
self.session.limits.max_connection_buffer_size
}
pub fn accepting_writes(&self) -> bool {
self.outbound_bytes < self.outbound_limit()
}
pub fn fail(&mut self, error: Error) -> QUICError {
let description = error.to_string();
let _ = self.events.try_send(H3Event::Failed(error));
Box::new(std::io::Error::other(description))
}
pub fn block_deadline(&self) -> Option<std::time::Instant> {
if self.session.blocked_since.is_empty() {
return None;
}
let wait = Timeout::duration(self.session.limits.qpack_block_timeout)?;
let earliest = self.session.blocked_since.values().min()?;
Some(*earliest + wait)
}
pub fn expired_block(&self) -> Option<Error> {
let deadline = self.block_deadline()?;
(std::time::Instant::now() >= deadline).then(|| {
Error::Timeout(format!(
"a QPACK block stayed blocked beyond {}s",
self.session.limits.qpack_block_timeout
))
})
}
pub fn alloc_uni(&mut self) -> u64 {
let id = self.next_uni;
self.next_uni += QUICStreamID::STEP;
id
}
pub fn open_uni(&mut self, transport: &mut impl QUICTransport, stream_id: u64, kind: StreamKind) -> Result<(), Error> {
let code = kind
.code()
.ok_or_else(|| Error::Protocol(format!("{kind:?} is not a unidirectional stream type")))?;
let mut prefix = BytesMut::new();
Varint::encode(&mut prefix, code);
self.write(transport, stream_id, &prefix, false)
}
pub fn write(&mut self, transport: &mut impl QUICTransport, stream_id: u64, data: &[u8], fin: bool) -> Result<(), Error> {
let entry = self.outbound.entry(stream_id).or_default();
entry.0.extend_from_slice(data);
entry.1 |= fin;
self.outbound_bytes = self.outbound_bytes.saturating_add(data.len() as u64);
self.flush_stream(transport, stream_id)
}
pub fn flush_stream(&mut self, transport: &mut impl QUICTransport, stream_id: u64) -> Result<(), Error> {
let Some((buffer, fin)) = self.outbound.get_mut(&stream_id) else {
return Ok(());
};
if buffer.is_empty() && !*fin {
return Ok(());
}
match transport.send(stream_id, buffer, *fin)? {
StreamWrite::Sent(sent) => {
let _ = buffer.split_to(sent);
self.outbound_bytes = self.outbound_bytes.saturating_sub(sent as u64);
if buffer.is_empty() {
self.outbound.remove(&stream_id);
}
Ok(())
}
StreamWrite::Blocked => Ok(()),
StreamWrite::Stopped(code) => {
if let Some((buffer, _)) = self.outbound.remove(&stream_id) {
self.outbound_bytes = self.outbound_bytes.saturating_sub(buffer.len() as u64);
}
let error = Error::stream(StreamID(stream_id), code, "the peer stopped the stream");
let _ = self.events.try_send(H3Event::Failed(error));
Ok(())
}
}
}
pub fn execute(&mut self, transport: &mut impl QUICTransport, command: H3Command) -> Result<(), Error> {
match command {
H3Command::Send(message) => {
let mut message = message;
let stream_id = message.stream_id.unwrap_or_else(|| self.session.open());
self.session.streams.entry(stream_id).or_default().responded = true;
let entry = self.outbound.entry(stream_id.0).or_default();
let before = entry.0.len();
let fin = self.session.encode_message_into(stream_id, &mut message, &mut entry.0)?;
entry.1 |= fin;
let framed = entry.0.len().saturating_sub(before) as u64;
self.outbound_bytes = self.outbound_bytes.saturating_add(framed);
self.flush_stream(transport, stream_id.0)?;
self.session.retire(stream_id);
Ok(())
}
H3Command::Open(reply) => {
let _ = reply.send(self.session.open());
Ok(())
}
H3Command::OpenUni(kind, reply) => {
let stream_id = self.alloc_uni();
self.open_uni(transport, stream_id, kind)?;
let _ = reply.send(StreamID(stream_id));
Ok(())
}
H3Command::Tunnel(stream_id, sink) => {
if let Some(state) = self.session.streams.get_mut(&stream_id) {
let buffered = std::mem::take(&mut state.buffer);
if !buffered.is_empty() || state.eof {
let _ = sink.try_send((buffered.freeze(), state.eof));
}
}
self.tunnels.insert(stream_id.0, sink);
Ok(())
}
H3Command::WriteEncoder(bytes) => self.write(transport, self.local_encoder, &bytes, false),
H3Command::Reset(stream_id, code) => {
self.outbound.remove(&stream_id.0);
self.tunnels.remove(&stream_id.0);
self.session.retire(stream_id);
let _ = transport.shutdown_write(stream_id.0, code);
let _ = transport.shutdown_read(stream_id.0, code);
Ok(())
}
H3Command::Close => {
let _ = transport.close(Code::NO_ERROR, b"");
Ok(())
}
}
}
pub fn dispatch(&mut self, transport: &mut impl QUICTransport, stream_id: u64, data: &[u8], fin: bool) -> Result<(), Error> {
if let Some(sink) = self.tunnels.get(&stream_id) {
if sink.try_send((Bytes::copy_from_slice(data), fin)).is_err() {
return Err(Error::stream(StreamID(stream_id), Code::EXCESSIVE_LOAD, "the tunnel could not take the octets"));
}
if fin {
self.tunnels.remove(&stream_id);
self.session.forget(StreamID(stream_id));
}
return Ok(());
}
if QUICStreamID::is_bidi(stream_id) {
if self.session.goaway_sent.is_some_and(|id| stream_id >= id) && !self.session.streams.contains_key(&StreamID(stream_id)) {
return self.reject(transport, stream_id);
}
self.session.on_stream_bytes(StreamID(stream_id), data, fin)?;
return self.goaway(transport);
}
self.feed_uni(transport, stream_id, data, fin)
}
pub fn goaway(&mut self, transport: &mut impl QUICTransport) -> Result<(), Error> {
let limit = self.session.limits.max_requests_per_connection;
if self.session.role.is_client() || self.session.goaway_sent.is_some() || limit == 0 || self.session.total_streams < limit {
return Ok(());
}
let id = self.session.highest_peer_stream_id + QUICStreamID::STEP;
let mut frame = BytesMut::new();
Frame::GoAway { id }.encode_into(&mut frame);
self.write(transport, self.local_control, &frame, false)?;
self.session.goaway_sent = Some(id);
Ok(())
}
pub fn reject(&mut self, transport: &mut impl QUICTransport, stream_id: u64) -> Result<(), Error> {
if stream_id >= self.rejected_floor {
self.rejected_floor = stream_id + QUICStreamID::STEP;
self.rejected = self.rejected.saturating_add(1);
}
if self.rejected as usize > self.session.stream_ceiling() {
let reason = format!("more than {} streams arrived after GOAWAY", self.session.stream_ceiling());
return Err(self.overloaded(transport, reason));
}
let _ = transport.shutdown_write(stream_id, Code::REQUEST_REJECTED);
let _ = transport.shutdown_read(stream_id, Code::REQUEST_REJECTED);
Ok(())
}
pub fn wind_down(&mut self, transport: &mut impl QUICTransport) {
if self.drained || self.session.goaway.is_none() {
return;
}
if !self.session.streams.is_empty() || !self.session.ready.is_empty() || !self.outbound.is_empty() || !self.tunnels.is_empty() {
return;
}
self.drained = true;
let _ = transport.close(Code::NO_ERROR, b"");
let _ = self.events.try_send(H3Event::Failed(Error::Closed));
}
pub fn feed_uni(&mut self, transport: &mut impl QUICTransport, stream_id: u64, data: &[u8], fin: bool) -> Result<(), Error> {
let outcome = self.feed_uni_bytes(transport, stream_id, data);
if fin {
self.peer_uni.remove(&stream_id);
}
outcome
}
pub fn feed_uni_bytes(&mut self, transport: &mut impl QUICTransport, stream_id: u64, data: &[u8]) -> Result<(), Error> {
if self.peer_uni.get(&stream_id).is_some_and(|uni| uni.abandoned) {
return Ok(());
}
if let Some(kind) = self.peer_uni.get(&stream_id).and_then(|uni| uni.kind) {
return self.feed_uni_kind(kind, data);
}
let ceiling = self.session.limits.max_peer_uni_streams as usize;
if !self.peer_uni.contains_key(&stream_id) && self.peer_uni.len() >= ceiling {
let reason = format!("more than {ceiling} unidirectional streams are open at once");
return Err(Error::Limit(reason));
}
let uni = self.peer_uni.entry(stream_id).or_default();
uni.prefix.extend_from_slice(data);
let (consumed, code) = Varint::decode(&uni.prefix);
if consumed == 0 {
if uni.prefix.len() > Varint::MAX_SIZE {
return Err(Error::Protocol("a unidirectional stream carries no readable type".into()));
}
return Ok(());
}
let Some(kind) = StreamKind::from_code(code) else {
uni.prefix = Vec::new();
uni.abandoned = true;
let _ = transport.shutdown_read(stream_id, Code::STREAM_CREATION_ERROR);
return Ok(());
};
uni.kind = Some(kind);
let payload = uni.prefix.split_off(consumed);
uni.prefix = Vec::new();
self.feed_uni_kind(kind, &payload)
}
pub fn feed_uni_kind(&mut self, kind: StreamKind, payload: &[u8]) -> Result<(), Error> {
match kind {
StreamKind::Control => self.session.on_control_bytes(payload),
StreamKind::QPACKEncoder => self.session.on_encoder_bytes(payload),
StreamKind::QPACKDecoder => self.session.on_decoder_bytes(payload),
_ => Ok(()),
}
}
pub fn overloaded(&mut self, transport: &mut impl QUICTransport, reason: impl Into<String>) -> Error {
let _ = transport.close(Code::EXCESSIVE_LOAD, b"");
Error::Limit(reason.into())
}
pub fn reset_stream(&mut self, transport: &mut impl QUICTransport, stream_id: u64, error: &Error) {
self.forget_stream(stream_id);
let code = match error {
Error::Stream { code, .. } => *code,
_ => Code::MESSAGE_ERROR,
};
let _ = transport.shutdown_write(stream_id, code);
let _ = transport.shutdown_read(stream_id, code);
}
pub fn drain_side_channels(&mut self, transport: &mut impl QUICTransport) -> Result<(), Error> {
self.drain_side_channel(transport, Side::Encoder)?;
self.drain_side_channel(transport, Side::Decoder)
}
pub fn drain_side_channel(&mut self, transport: &mut impl QUICTransport, side: Side) -> Result<(), Error> {
let (pending, stream_id) = match side {
Side::Encoder => (!self.session.encoder.encoder_stream().is_empty(), self.local_encoder),
Side::Decoder => (!self.session.decoder.decoder_stream().is_empty(), self.local_decoder),
};
if !pending {
return Ok(());
}
let queued = match side {
Side::Encoder => self.session.encoder.take_encoder_stream(),
Side::Decoder => self.session.decoder.take_decoder_stream(),
};
let outcome = self.write(transport, stream_id, &queued, false);
match side {
Side::Encoder => self.session.encoder.reclaim_encoder_stream(queued),
Side::Decoder => self.session.decoder.reclaim_decoder_stream(queued),
}
outcome
}
}
impl QUICApplication for H3Worker {
fn on_conn_established(&mut self, qconn: &mut QUICConnection, _handshake: &QUICHandshake) -> QUICOutcome<()> {
self.established = true;
*Lock::on(&self.session.security) = Handshake::of(qconn).security();
self.local_control = self.alloc_uni();
self.local_encoder = self.alloc_uni();
self.local_decoder = self.alloc_uni();
let opened = (|| {
self.open_uni(qconn, self.local_control, StreamKind::Control)?;
let settings = self.session.control_frame();
self.write(qconn, self.local_control, &settings, false)?;
self.open_uni(qconn, self.local_encoder, StreamKind::QPACKEncoder)?;
self.open_uni(qconn, self.local_decoder, StreamKind::QPACKDecoder)
})();
opened.map_err(|error| self.fail(error))
}
fn should_act(&self) -> bool {
true
}
fn buffer(&mut self) -> &mut [u8] {
&mut self.scratch
}
async fn wait_for_data(&mut self, _qconn: &mut QUICConnection) -> QUICOutcome<()> {
let deadline = self.block_deadline();
tokio::select! {
command = self.commands.recv(), if !self.orphaned && self.pending.is_none() => match command {
Some(command) => {
self.pending = Some(command);
Ok(())
}
None => {
self.orphaned = true;
Ok(())
}
},
raw = self.raw_writes.recv(), if self.pending_raw.is_none() && self.accepting_writes() => match raw {
Some(raw) => {
self.pending_raw = Some(raw);
Ok(())
}
None => Err(Self::boxed(Error::Closed)),
},
_ = tokio::time::sleep_until(tokio::time::Instant::from_std(deadline.unwrap_or_else(std::time::Instant::now))), if deadline.is_some() => Ok(()),
else => {
std::future::pending::<()>().await;
Ok(())
}
}
}
fn process_reads(&mut self, qconn: &mut QUICConnection) -> QUICOutcome<()> {
let mut read = std::mem::take(&mut self.read);
let mut readable = std::mem::take(&mut self.readable);
readable.clear();
readable.extend(QUICTransport::readable(qconn));
let outcome = self.drain_reads(qconn, &readable, &mut read);
self.read = read;
self.readable = readable;
outcome?;
self.deliver_ready();
self.wind_down(qconn);
Ok(())
}
fn process_writes(&mut self, qconn: &mut QUICConnection) -> QUICOutcome<()> {
if let Some(error) = self.expired_block() {
let _ = QUICTransport::close(qconn, Code::QPACK_DECOMPRESSION_FAILED, b"");
return Err(self.fail(error));
}
if let Some(command) = self.pending.take() {
match command.buffers() && !self.accepting_writes() {
true => self.pending = Some(command),
false => {
if let Err(error) = self.execute(qconn, command) {
return Err(self.fail(error));
}
}
}
}
while self.pending.is_none() {
let Ok(command) = self.commands.try_recv() else {
break;
};
if command.buffers() && !self.accepting_writes() {
self.pending = Some(command);
break;
}
if let Err(error) = self.execute(qconn, command) {
return Err(self.fail(error));
}
}
while self.accepting_writes() {
let Some((stream_id, bytes, fin)) = self.pending_raw.take().or_else(|| self.raw_writes.try_recv().ok()) else {
break;
};
if let Err(error) = self.write(qconn, stream_id.0, &bytes, fin) {
return Err(self.fail(error));
}
}
if let Err(error) = self.drain_side_channels(qconn) {
return Err(self.fail(error));
}
let mut flushing = std::mem::take(&mut self.flushing);
flushing.clear();
flushing.extend(self.outbound.keys().copied());
let mut outcome = Ok(());
for stream_id in &flushing {
if let Err(error) = self.flush_stream(qconn, *stream_id) {
outcome = Err(error);
break;
}
}
self.flushing = flushing;
if let Err(error) = outcome {
return Err(self.fail(error));
}
self.deliver_ready();
self.wind_down(qconn);
Ok(())
}
}
impl H3Worker {
pub fn deliver_ready(&mut self) {
while let Some(message) = self.session.take_ready() {
let Err(tokio::sync::mpsc::error::TrySendError::Full(event)) = self.events.try_send(H3Event::Message(message)) else {
continue;
};
if let H3Event::Message(message) = event {
self.session.ready.push_front(message);
}
return;
}
}
pub fn drain_reads(&mut self, transport: &mut impl QUICTransport, readable: &[u64], read: &mut [u8]) -> Result<(), QUICError> {
for stream_id in readable.iter().copied() {
loop {
if self.tunnels.get(&stream_id).is_some_and(|sink| sink.capacity() == 0) {
break;
}
let outcome = match transport.receive(stream_id, read) {
Ok(outcome) => outcome,
Err(error) => return Err(self.fail(error)),
};
match outcome {
StreamRead::Data { len, fin } => {
match self.dispatch(transport, stream_id, &read[..len], fin) {
Ok(()) => {}
Err(error) if matches!(error, Error::Stream { .. }) => {
self.reset_stream(transport, stream_id, &error);
let _ = self.events.try_send(H3Event::Failed(error));
break;
}
Err(error) => return Err(self.fail(error)),
}
if fin || len == 0 {
break;
}
}
StreamRead::Done => break,
StreamRead::Reset(code) => {
let premature = self.session.forget(StreamID(stream_id)).is_some_and(|state| !state.responded);
self.forget_stream(stream_id);
if premature {
self.premature_resets = self.premature_resets.saturating_add(1);
if self.premature_resets > self.session.limits.max_premature_resets {
let reason = format!(
"more than {} streams were reset before a response was sent",
self.session.limits.max_premature_resets
);
let error = self.overloaded(transport, reason);
return Err(self.fail(error));
}
}
let error = Error::stream(StreamID(stream_id), code, "the peer reset the stream");
let _ = self.events.try_send(H3Event::Failed(error));
break;
}
}
}
}
Ok(())
}
}