use std::{
marker::PhantomData,
mem,
sync::{Arc, atomic::AtomicUsize},
task::{Context, Poll},
};
use bytes::{Buf, Bytes, BytesMut};
use futures_util::future;
use http::request;
#[cfg(feature = "tracing")]
use tracing::{info, instrument, trace};
use super::stream::RequestStream;
use crate::{
connection::{self, ConnectionInner},
error::{
Code, ConnectionError, StreamError, connection_error_creators::CloseStream,
internal_error::InternalConnectionError,
},
frame::FrameStream,
proto::{frame::Frame, headers::Header, push::PushId},
qpack::{self, QpackDecoder, QpackEncoder},
quic::{self, SendStream, StreamId},
shared_state::{ConnectionState, SharedState},
stream::{self, BufRecvStream},
};
const MAX_RETAINED_QPACK_ENCODE_CAPACITY: usize = 4 * 1024;
fn clear_qpack_encode_buffer(buffer: &mut BytesMut) {
if buffer.capacity() > MAX_RETAINED_QPACK_ENCODE_CAPACITY {
*buffer = BytesMut::new();
} else {
buffer.clear();
}
}
fn take_qpack_encode_buffer(buffer: &mut BytesMut) -> Bytes {
if buffer.capacity() > MAX_RETAINED_QPACK_ENCODE_CAPACITY {
mem::take(buffer).freeze()
} else {
buffer.split().freeze()
}
}
pub struct SendRequest<T, B>
where
T: quic::OpenStreams<B>,
B: Buf,
{
pub(super) open: T,
pub(super) conn_state: Arc<SharedState>,
pub(super) decoder: Option<QpackDecoder>,
pub(super) encoder: Option<QpackEncoder>,
pub(super) max_field_section_size: u64, pub(super) max_qpack_decode_buffer_size: usize,
pub(super) sender_count: Arc<AtomicUsize>,
pub(super) _buf: PhantomData<fn(B)>,
pub(super) send_grease_frame: bool,
pub(super) qpack_encode_buffer: BytesMut,
}
impl<T, B> ConnectionState for SendRequest<T, B>
where
T: quic::OpenStreams<B>,
B: Buf,
{
fn shared_state(&self) -> &SharedState {
&self.conn_state
}
}
impl<T, B> CloseStream for SendRequest<T, B>
where
T: quic::OpenStreams<B>,
B: Buf,
{
}
impl<T, B> SendRequest<T, B>
where
T: quic::OpenStreams<B>,
B: Buf,
{
#[cfg_attr(feature = "tracing", instrument(skip_all, level = "trace"))]
pub async fn send_request(
&mut self,
req: http::Request<()>,
) -> Result<RequestStream<T::BidiStream, B>, StreamError> {
if let Some(error) = self.check_peer_connection_closing() {
return Err(error);
};
let (parts, _) = req.into_parts();
let request::Parts {
method,
uri,
headers,
extensions,
..
} = parts;
let headers = Header::request(method, uri, headers, extensions).map_err(|_e| {
self.handle_connection_error_on_stream(InternalConnectionError {
code: Code::H3_INTERNAL_ERROR,
message: "Failed to build request headers".to_string(),
})
})?;
let dynamic_encoder = match self.encoder.as_ref() {
Some(encoder) => match encoder.ready() {
Ok(true) => Some(encoder),
Ok(false) => None,
Err(error) => {
return Err(self.handle_connection_error_on_stream(
InternalConnectionError::new(
Code::H3_INTERNAL_ERROR,
format!("failed to access QPACK encoder: {error}"),
),
));
}
},
None => None,
};
let peer_max_field_section_size = self.settings().max_field_section_size;
let mut stream;
if let Some(encoder) = dynamic_encoder {
let mem_size = headers
.into_iter()
.try_fold(0_u64, |size, field| {
size.checked_add(field.mem_size() as u64)
})
.ok_or_else(|| {
self.handle_connection_error_on_stream(InternalConnectionError::new(
Code::H3_INTERNAL_ERROR,
"request field section size overflowed".to_string(),
))
})?;
stream = future::poll_fn(|cx| self.open.poll_open_bidi(cx))
.await
.map_err(|e| self.handle_quic_stream_error(e))?;
if mem_size > peer_max_field_section_size {
return Err(StreamError::HeaderTooBig {
actual_size: mem_size,
max_size: peer_max_field_section_size,
});
}
let encoder_instructions_queued =
match encoder.encode(stream.send_id(), &mut self.qpack_encode_buffer, &headers) {
Ok(encoded) => encoded,
Err(error) => {
clear_qpack_encode_buffer(&mut self.qpack_encode_buffer);
return Err(self.handle_connection_error_on_stream(
InternalConnectionError::new(
Code::H3_INTERNAL_ERROR,
format!("failed to encode request headers: {error}"),
),
));
}
};
drop(headers);
let block = take_qpack_encode_buffer(&mut self.qpack_encode_buffer);
if encoder_instructions_queued {
self.waker().wake();
}
stream::write(&mut stream, Frame::Headers(block))
.await
.map_err(|error| self.handle_quic_stream_error(error))?;
} else {
let mem_size = match qpack::encode_stateless(&mut self.qpack_encode_buffer, &headers) {
Ok(mem_size) => mem_size,
Err(_error) => {
clear_qpack_encode_buffer(&mut self.qpack_encode_buffer);
return Err(self.handle_connection_error_on_stream(
InternalConnectionError::new(
Code::H3_INTERNAL_ERROR,
"failed to encode request headers".to_string(),
),
));
}
};
drop(headers);
let block = take_qpack_encode_buffer(&mut self.qpack_encode_buffer);
stream = future::poll_fn(|cx| self.open.poll_open_bidi(cx))
.await
.map_err(|e| self.handle_quic_stream_error(e))?;
if mem_size > peer_max_field_section_size {
return Err(StreamError::HeaderTooBig {
actual_size: mem_size,
max_size: peer_max_field_section_size,
});
}
stream::write(&mut stream, Frame::Headers(block))
.await
.map_err(|e| self.handle_quic_stream_error(e))?;
}
let request_stream = RequestStream {
inner: connection::RequestStream::new(
FrameStream::new(BufRecvStream::new(stream)),
self.max_field_section_size,
self.max_qpack_decode_buffer_size,
self.send_grease_frame,
self.conn_state.clone(),
self.decoder.clone(),
),
};
self.send_grease_frame = false;
Ok(request_stream)
}
}
impl<T, B> Clone for SendRequest<T, B>
where
T: quic::OpenStreams<B> + Clone,
B: Buf,
{
fn clone(&self) -> Self {
self.sender_count
.fetch_add(1, std::sync::atomic::Ordering::Release);
Self {
conn_state: self.conn_state.clone(),
decoder: self.decoder.clone(),
encoder: self.encoder.clone(),
open: self.open.clone(),
max_field_section_size: self.max_field_section_size,
max_qpack_decode_buffer_size: self.max_qpack_decode_buffer_size,
sender_count: self.sender_count.clone(),
_buf: PhantomData,
send_grease_frame: self.send_grease_frame,
qpack_encode_buffer: BytesMut::new(),
}
}
}
impl<T, B> Drop for SendRequest<T, B>
where
T: quic::OpenStreams<B>,
B: Buf,
{
fn drop(&mut self) {
if self
.sender_count
.fetch_sub(1, std::sync::atomic::Ordering::AcqRel)
== 1
{
self.handle_connection_error_on_stream(InternalConnectionError::new(
Code::H3_NO_ERROR,
"Connection closed by client".to_string(),
));
}
}
}
pub struct Connection<C, B>
where
C: quic::Connection<B>,
B: Buf,
{
pub inner: ConnectionInner<C, B>,
pub(super) sent_closing: Option<PushId>,
pub(super) recv_closing: Option<StreamId>,
}
impl<C, B> ConnectionState for Connection<C, B>
where
C: quic::Connection<B>,
B: Buf,
{
fn shared_state(&self) -> &SharedState {
&self.inner.shared
}
}
impl<C, B> Connection<C, B>
where
C: quic::Connection<B>,
C::SendStream: quic::SendStreamUnframed<B>,
B: Buf,
{
#[cfg_attr(feature = "tracing", instrument(skip_all, level = "trace"))]
pub async fn shutdown(&mut self, _max_push: usize) -> Result<(), ConnectionError> {
self.inner.shutdown(&mut self.sent_closing, PushId(0)).await
}
#[cfg_attr(feature = "tracing", instrument(skip_all, level = "trace"))]
pub async fn wait_idle(&mut self) -> ConnectionError {
future::poll_fn(|cx| self.poll_close(cx)).await
}
#[cfg_attr(feature = "tracing", instrument(skip_all, level = "trace"))]
pub fn poll_close(&mut self, cx: &mut Context<'_>) -> Poll<ConnectionError> {
if let Err(err) = self.inner.poll_accept_recv(cx) {
return Poll::Ready(err);
}
if let Err(err) = self.inner.poll_qpack(cx) {
return Poll::Ready(err);
}
while let Poll::Ready(result) = self.inner.poll_accepted_control(cx) {
match result {
Ok(Frame::Settings(_)) => {
#[cfg(feature = "tracing")]
trace!("Got settings");
}
Ok(Frame::Goaway(id)) => {
if !StreamId::from(id).is_request() {
return Poll::Ready(self.inner.handle_connection_error(
InternalConnectionError::new(
Code::H3_ID_ERROR,
format!("non-request StreamId in a GoAway frame: {}", id),
),
));
}
if let Err(err) = self.inner.process_goaway(&mut self.recv_closing, id) {
return Poll::Ready(err);
}
#[cfg(feature = "tracing")]
info!("Server initiated graceful shutdown, last: StreamId({})", id);
}
Ok(frame) => {
return Poll::Ready(self.inner.handle_connection_error(
InternalConnectionError::new(
Code::H3_FRAME_UNEXPECTED,
format!("on client control stream: {:?}", frame),
),
));
}
Err(connection_error) => {
return Poll::Ready(connection_error);
}
}
}
if self.inner.poll_accept_bi(cx).is_ready() {
return Poll::Ready(
self.inner
.handle_connection_error(InternalConnectionError::new(
Code::H3_STREAM_CREATION_ERROR,
"client received a server-initiated bidirectional stream".to_string(),
)),
);
}
Poll::Pending
}
}
#[cfg(test)]
mod qpack_encode_buffer_tests {
use super::*;
#[test]
fn taken_blocks_remain_independent_and_large_storage_is_not_retained() {
let mut buffer = BytesMut::with_capacity(64);
buffer.extend_from_slice(b"first");
let first = take_qpack_encode_buffer(&mut buffer);
buffer.extend_from_slice(b"second");
let second = take_qpack_encode_buffer(&mut buffer);
assert_eq!(first, b"first"[..]);
assert_eq!(second, b"second"[..]);
let mut large = BytesMut::with_capacity(MAX_RETAINED_QPACK_ENCODE_CAPACITY + 1);
large.extend_from_slice(b"large");
assert_eq!(take_qpack_encode_buffer(&mut large), b"large"[..]);
assert_eq!(large.capacity(), 0);
}
#[test]
fn clearing_drops_only_oversized_storage() {
let mut small = BytesMut::with_capacity(64);
small.extend_from_slice(b"partial");
clear_qpack_encode_buffer(&mut small);
assert!(small.is_empty());
assert!(small.capacity() >= 64);
let mut large = BytesMut::with_capacity(MAX_RETAINED_QPACK_ENCODE_CAPACITY + 1);
large.extend_from_slice(b"partial");
clear_qpack_encode_buffer(&mut large);
assert!(large.is_empty());
assert_eq!(large.capacity(), 0);
}
}