#![allow(dead_code)]
use std::collections::VecDeque;
use std::sync::Arc;
use std::task::{Context, Poll, Waker};
use bytes::{Bytes, BytesMut};
use futures_util::ready;
use http::header::{HeaderMap, HeaderName, HeaderValue};
use http::{Method, Request, StatusCode, Uri, Version};
use parking_lot::Mutex;
use rustc_hash::FxHashMap;
use crate::h3::error::{H3Error, TransportError};
use crate::h3::frame::{
write_varint, Frame, FrameDecoder, FrameError, FRAME_CANCEL_PUSH, FRAME_DATA, FRAME_GOAWAY,
FRAME_HEADERS, FRAME_MAX_PUSH_ID, FRAME_PUSH_PROMISE, FRAME_SETTINGS,
};
use crate::h3::qpack::{Decoder, Encoder, QpackError, UnblockedSection};
use crate::h3::settings::LocalSettings;
use crate::h3::transport::BidiStream;
#[derive(Debug, Clone, PartialEq, Eq)]
pub(crate) enum StreamError {
Transport(TransportError),
Frame,
Message,
FrameUnexpected,
Qpack(QpackError),
HeadersTooBig { size: u64, limit: u64 },
}
impl StreamError {
#[inline]
pub(crate) fn h3_code(&self) -> u64 {
match self {
StreamError::Transport(TransportError::Closed { code }) => *code,
StreamError::Transport(_) => H3Error::Internal.code(),
StreamError::Frame => H3Error::FrameError.code(),
StreamError::Message => H3Error::Message.code(),
StreamError::FrameUnexpected => H3Error::FrameUnexpected.code(),
StreamError::Qpack(err) => u64::from(err.code()),
StreamError::HeadersTooBig { .. } => H3Error::Message.code(),
}
}
#[inline]
pub(crate) fn is_stream_scoped(&self) -> bool {
matches!(
self,
StreamError::Transport(TransportError::Reset { .. } | TransportError::Stopped { .. })
)
}
}
impl From<TransportError> for StreamError {
#[inline]
fn from(err: TransportError) -> Self {
StreamError::Transport(err)
}
}
impl std::fmt::Display for StreamError {
#[inline]
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
std::fmt::Debug::fmt(self, f)
}
}
impl std::error::Error for StreamError {}
#[inline]
fn map_frame_error(err: FrameError) -> StreamError {
match err {
FrameError::Frame => StreamError::Frame,
FrameError::Unexpected(_) | FrameError::Settings => StreamError::FrameUnexpected,
}
}
#[derive(Debug)]
pub(crate) struct SharedCodecs {
pub(crate) decoder: Decoder,
pub(crate) encoder: Option<Encoder>,
pub(crate) encoder_stream: VecDeque<Bytes>,
pub(crate) unblocked: Vec<UnblockedSection>,
pub(crate) peer_max_field_section_size: Option<u64>,
pub(crate) waiters: FxHashMap<u64, Waker>,
}
impl SharedCodecs {
#[inline]
pub(crate) fn new(local: &LocalSettings) -> Self {
let mut decoder = Decoder::new(
local.qpack_max_table_capacity,
local.qpack_blocked_streams as usize,
);
decoder.set_max_field_section_size(
local
.max_field_section_size
.map(|v| v as usize)
.unwrap_or(usize::MAX),
);
Self {
decoder,
encoder: None,
encoder_stream: VecDeque::new(),
unblocked: Vec::new(),
peer_max_field_section_size: None,
waiters: FxHashMap::default(),
}
}
#[inline]
pub(crate) fn take_waiters(&mut self) -> Vec<Waker> {
self.waiters.drain().map(|(_, waker)| waker).collect()
}
}
enum PendingSend {
None,
ResponseHeaders(StatusCode),
Data,
Trailers,
}
const WRITE_INLINE_LIMIT: usize = 1024;
pub(crate) struct RequestStream {
stream: Box<dyn BidiStream>,
stream_id: u64,
frame_decoder: FrameDecoder,
shared: Arc<Mutex<SharedCodecs>>,
awaiting_headers: bool,
headers_blocked: bool,
recv_data: VecDeque<Bytes>,
trailers_lines: Option<Vec<(Bytes, Bytes)>>,
trailers_blocked: bool,
trailers_done: bool,
recv_finished: bool,
sent_headers: bool,
sent_trailers: bool,
sent_fin: bool,
pending_send: PendingSend,
send_buf: VecDeque<Bytes>,
}
impl RequestStream {
#[inline]
pub(crate) fn new(stream: Box<dyn BidiStream>, shared: Arc<Mutex<SharedCodecs>>) -> Self {
let stream_id = stream.id();
Self {
stream,
stream_id,
frame_decoder: FrameDecoder::new(),
shared,
awaiting_headers: true,
headers_blocked: false,
recv_data: VecDeque::new(),
trailers_lines: None,
trailers_blocked: false,
trailers_done: false,
recv_finished: false,
sent_headers: false,
sent_trailers: false,
sent_fin: false,
pending_send: PendingSend::None,
send_buf: VecDeque::new(),
}
}
#[allow(dead_code)] #[inline]
pub(crate) fn id(&self) -> u64 {
self.stream_id
}
#[inline]
pub(crate) fn poll_headers(
&mut self,
cx: &mut Context<'_>,
) -> Poll<Result<Option<Request<()>>, StreamError>> {
if self.headers_blocked {
let section = take_unblocked_for(&self.shared, self.stream_id, cx);
if let Some(section) = section {
self.headers_blocked = false;
self.awaiting_headers = false;
return Poll::Ready(build_request(section.headers).map(Some));
}
return Poll::Pending;
}
if !self.awaiting_headers {
return Poll::Ready(Ok(None));
}
match ready!(self.poll_frame(cx)?) {
None => {
self.awaiting_headers = false;
self.finish_recv();
Poll::Ready(Ok(None))
}
Some(Frame::Headers(block)) => match self.decode_block(&block) {
Ok(Some(headers)) => match build_request(headers) {
Ok(request) => {
self.awaiting_headers = false;
Poll::Ready(Ok(Some(request)))
}
Err(err) => Poll::Ready(Err(err)),
},
Ok(None) => {
self.headers_blocked = true;
Poll::Pending
}
Err(err) => Poll::Ready(Err(err)),
},
Some(_) => {
Poll::Ready(Err(StreamError::FrameUnexpected))
}
}
}
#[inline]
pub(crate) fn poll_recv_data(
&mut self,
cx: &mut Context<'_>,
) -> Poll<Result<Option<Bytes>, StreamError>> {
if let Some(data) = self.recv_data.pop_front() {
return Poll::Ready(Ok(Some(data)));
}
if self.recv_finished {
return Poll::Ready(Ok(None));
}
if self.awaiting_headers || self.headers_blocked || self.trailers_blocked {
return Poll::Pending;
}
if self.trailers_done {
match ready!(self.poll_after_trailers(cx)?) {
Some(()) => {
self.finish_recv();
return Poll::Ready(Ok(None));
}
None => return Poll::Pending,
}
}
match ready!(self.poll_frame(cx)?) {
None => {
self.finish_recv();
Poll::Ready(Ok(None))
}
Some(Frame::Data(data)) => Poll::Ready(Ok(Some(data))),
Some(Frame::Headers(block)) => {
match self.decode_block(&block) {
Ok(Some(headers)) => {
self.trailers_done = true;
self.trailers_lines = Some(headers);
Poll::Ready(Ok(None))
}
Ok(None) => {
self.trailers_blocked = true;
Poll::Ready(Ok(None))
}
Err(err) => Poll::Ready(Err(err)),
}
}
Some(_) => Poll::Ready(Err(StreamError::FrameUnexpected)),
}
}
#[inline]
pub(crate) fn poll_recv_trailers(
&mut self,
cx: &mut Context<'_>,
) -> Poll<Result<Option<HeaderMap>, StreamError>> {
if let Some(lines) = self.trailers_lines.take() {
return Poll::Ready(header_map(lines).map(Some));
}
if self.trailers_blocked {
if let Some(section) = take_unblocked_for(&self.shared, self.stream_id, cx) {
self.trailers_blocked = false;
self.trailers_done = true;
return Poll::Ready(header_map(section.headers).map(Some));
}
return Poll::Pending;
}
if self.recv_finished {
return Poll::Ready(Ok(None));
}
if self.awaiting_headers || self.headers_blocked {
return Poll::Pending;
}
if self.trailers_done {
match ready!(self.poll_after_trailers(cx)?) {
Some(()) => {
self.finish_recv();
return Poll::Ready(Ok(None));
}
None => return Poll::Pending,
}
}
match ready!(self.poll_frame(cx)?) {
None => {
self.finish_recv();
Poll::Ready(Ok(None))
}
Some(Frame::Headers(block)) => match self.decode_block(&block) {
Ok(Some(headers)) => {
self.trailers_done = true;
Poll::Ready(header_map(headers).map(Some))
}
Ok(None) => {
self.trailers_blocked = true;
Poll::Pending
}
Err(err) => Poll::Ready(Err(err)),
},
Some(Frame::Data(_)) => {
Poll::Ready(Err(StreamError::FrameUnexpected))
}
Some(_) => Poll::Ready(Err(StreamError::FrameUnexpected)),
}
}
#[inline]
pub(crate) fn poll_send_response(
&mut self,
cx: &mut Context<'_>,
status: StatusCode,
headers: &HeaderMap,
) -> Poll<Result<(), StreamError>> {
if self.sent_headers && !status.is_informational() {
return Poll::Ready(Err(StreamError::Message));
}
match &self.pending_send {
PendingSend::ResponseHeaders(_) => {
ready!(self.poll_write(cx, Bytes::new()))?;
let pending = std::mem::replace(&mut self.pending_send, PendingSend::None);
if let PendingSend::ResponseHeaders(s) = pending {
if !s.is_informational() {
self.sent_headers = true;
}
}
Poll::Ready(Ok(()))
}
PendingSend::None => {
let mut lines = Vec::with_capacity(headers.len() + 1);
lines.push((
Bytes::from_static(b":status"),
Bytes::copy_from_slice(status.as_str().as_bytes()),
));
for (name, value) in headers {
if name.as_str().starts_with(':') {
return Poll::Ready(Err(StreamError::Message));
}
lines.push((
Bytes::copy_from_slice(name.as_str().as_bytes()),
Bytes::copy_from_slice(value.as_bytes()),
));
}
let (frame_header, field_section) = self.encode_headers(&lines)?;
self.pending_send = PendingSend::ResponseHeaders(status);
ready!(self.poll_write_parts(cx, frame_header, field_section))?;
let pending = std::mem::replace(&mut self.pending_send, PendingSend::None);
if let PendingSend::ResponseHeaders(s) = pending {
if !s.is_informational() {
self.sent_headers = true;
}
}
Poll::Ready(Ok(()))
}
_ => Poll::Ready(Err(StreamError::Message)),
}
}
#[inline]
pub(crate) fn poll_send_data(
&mut self,
cx: &mut Context<'_>,
data: Bytes,
) -> Poll<Result<(), StreamError>> {
match &self.pending_send {
PendingSend::Data => {
ready!(self.poll_write(cx, Bytes::new()))?;
self.pending_send = PendingSend::None;
Poll::Ready(Ok(()))
}
PendingSend::None => {
let mut frame_header =
BytesMut::with_capacity(1 + crate::h3::frame::varint_size(data.len() as u64));
write_varint(FRAME_DATA, &mut frame_header);
write_varint(data.len() as u64, &mut frame_header);
self.pending_send = PendingSend::Data;
ready!(self.poll_write_parts(cx, frame_header.freeze(), data))?;
self.pending_send = PendingSend::None;
Poll::Ready(Ok(()))
}
_ => Poll::Ready(Err(StreamError::Message)),
}
}
#[inline]
pub(crate) fn poll_send_trailers(
&mut self,
cx: &mut Context<'_>,
trailers: &HeaderMap,
) -> Poll<Result<(), StreamError>> {
if !self.sent_headers {
return Poll::Ready(Err(StreamError::Message));
}
match &self.pending_send {
PendingSend::Trailers => {
ready!(self.poll_write(cx, Bytes::new()))?;
self.pending_send = PendingSend::None;
Poll::Ready(Ok(()))
}
PendingSend::None => {
let mut lines = Vec::with_capacity(trailers.len());
for (name, value) in trailers {
if name.as_str().starts_with(':') {
return Poll::Ready(Err(StreamError::Message));
}
lines.push((
Bytes::copy_from_slice(name.as_str().as_bytes()),
Bytes::copy_from_slice(value.as_bytes()),
));
}
let (frame_header, field_section) = self.encode_headers(&lines)?;
self.pending_send = PendingSend::Trailers;
ready!(self.poll_write_parts(cx, frame_header, field_section))?;
self.pending_send = PendingSend::None;
self.sent_trailers = true;
Poll::Ready(Ok(()))
}
_ => Poll::Ready(Err(StreamError::Message)),
}
}
#[inline]
pub(crate) fn poll_finish(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), StreamError>> {
ready!(self.poll_write(cx, Bytes::new())?);
self.stream.poll_finish(cx).map_err(StreamError::Transport)
}
#[inline]
pub(crate) fn poll_reset(
&mut self,
cx: &mut Context<'_>,
code: u64,
) -> Poll<Result<(), StreamError>> {
self.stream
.poll_reset(cx, code)
.map_err(StreamError::Transport)
}
#[inline]
pub(crate) fn poll_stop_sending(
&mut self,
cx: &mut Context<'_>,
code: u64,
) -> Poll<Result<(), StreamError>> {
self.stream
.poll_stop_sending(cx, code)
.map_err(StreamError::Transport)
}
#[inline]
fn encode_headers(&mut self, lines: &[(Bytes, Bytes)]) -> Result<(Bytes, Bytes), StreamError> {
let mut shared = self.shared.lock();
if shared.encoder.is_none() {
return Err(StreamError::Message);
}
let encoder = shared.encoder.as_mut().expect("encoder ensured");
let section = encoder.encode_section_with_ack_base(self.stream_id, lines);
let size = section.block.len() as u64;
if let Some(limit) = shared.peer_max_field_section_size {
if size > limit {
return Err(StreamError::HeadersTooBig { size, limit });
}
}
if !section.encoder_stream.is_empty() {
shared.encoder_stream.push_back(section.encoder_stream);
}
let mut frame_header =
BytesMut::with_capacity(1 + crate::h3::frame::varint_size(section.block.len() as u64));
write_varint(FRAME_HEADERS, &mut frame_header);
write_varint(section.block.len() as u64, &mut frame_header);
Ok((frame_header.freeze(), section.block))
}
#[inline]
fn decode_block(&self, block: &[u8]) -> Result<Option<Vec<(Bytes, Bytes)>>, StreamError> {
self.shared
.lock()
.decoder
.decode_block(block, self.stream_id, now())
.map_err(StreamError::Qpack)
}
#[inline]
fn finish_recv(&mut self) {
self.recv_finished = true;
self.shared.lock().decoder.stream_finished(self.stream_id);
}
#[inline]
fn poll_after_trailers(
&mut self,
cx: &mut Context<'_>,
) -> Poll<Result<Option<()>, StreamError>> {
loop {
match ready!(self.poll_frame(cx)?) {
None => return Poll::Ready(Ok(Some(()))),
Some(frame) => {
if frame.is_known() {
return Poll::Ready(Err(StreamError::FrameUnexpected));
}
}
}
}
}
#[inline]
fn poll_frame(&mut self, cx: &mut Context<'_>) -> Poll<Result<Option<Frame>, StreamError>> {
loop {
if let Some(ty) = self.frame_decoder.peek_frame_type() {
if matches!(
ty,
FRAME_CANCEL_PUSH
| FRAME_SETTINGS
| FRAME_PUSH_PROMISE
| FRAME_GOAWAY
| FRAME_MAX_PUSH_ID
) {
return Poll::Ready(Err(StreamError::FrameUnexpected));
}
}
match self.frame_decoder.next_frame() {
Ok(Some(frame)) => return Poll::Ready(Ok(Some(frame))),
Ok(None) => {}
Err(err) => return Poll::Ready(Err(map_frame_error(err))),
}
match ready!(self.stream.poll_recv(cx))? {
Some(chunk) => self.frame_decoder.extend(chunk),
None => {
if self.frame_decoder.buffered() != 0 {
return Poll::Ready(Err(StreamError::Frame));
}
self.recv_finished = true;
return Poll::Ready(Ok(None));
}
}
}
}
#[inline]
fn poll_write(&mut self, cx: &mut Context<'_>, bytes: Bytes) -> Poll<Result<(), StreamError>> {
if !bytes.is_empty() {
self.send_buf.push_back(bytes);
}
while let Some(bytes) = self.send_buf.front() {
match self.stream.poll_send(cx, bytes) {
Poll::Ready(Ok(())) => {
self.send_buf.pop_front();
}
Poll::Ready(Err(err)) => return Poll::Ready(Err(err.into())),
Poll::Pending => return Poll::Pending,
}
}
Poll::Ready(Ok(()))
}
#[inline]
fn poll_write_parts(
&mut self,
cx: &mut Context<'_>,
first: Bytes,
second: Bytes,
) -> Poll<Result<(), StreamError>> {
if first.is_empty() {
if second.is_empty() {
return self.poll_write(cx, Bytes::new());
}
self.send_buf.push_back(second);
return self.poll_write(cx, Bytes::new());
}
if second.is_empty() {
self.send_buf.push_back(first);
return self.poll_write(cx, Bytes::new());
}
if first.len() + second.len() <= WRITE_INLINE_LIMIT {
let mut combined = BytesMut::with_capacity(first.len() + second.len());
combined.extend_from_slice(&first);
combined.extend_from_slice(&second);
self.send_buf.push_back(combined.freeze());
} else {
self.send_buf.push_back(first);
self.send_buf.push_back(second);
}
self.poll_write(cx, Bytes::new())
}
}
#[inline]
fn take_unblocked_for(
shared: &Arc<Mutex<SharedCodecs>>,
stream_id: u64,
cx: &mut Context<'_>,
) -> Option<UnblockedSection> {
let mut shared = shared.lock();
match shared
.unblocked
.iter()
.position(|section| section.stream_id == stream_id)
{
Some(index) => Some(shared.unblocked.remove(index)),
None => {
shared.waiters.insert(stream_id, cx.waker().clone());
None
}
}
}
#[inline]
fn now() -> u64 {
std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap_or_default()
.as_nanos() as u64
}
#[inline]
fn build_request(headers: Vec<(Bytes, Bytes)>) -> Result<Request<()>, StreamError> {
let mut method = None;
let mut scheme = None;
let mut authority = None;
let mut path = None;
let mut protocol = None;
let mut regular = Vec::new();
let mut pseudo_done = false;
for (name, value) in headers {
if name.first() == Some(&b':') {
if pseudo_done {
return Err(StreamError::Message);
}
match name.as_ref() {
b":method" => take_pseudo(&mut method, value)?,
b":scheme" => take_pseudo(&mut scheme, value)?,
b":authority" => take_pseudo(&mut authority, value)?,
b":path" => take_pseudo(&mut path, value)?,
b":protocol" => take_pseudo(&mut protocol, value)?,
_ => return Err(StreamError::Message),
}
} else {
pseudo_done = true;
regular.push((
HeaderName::from_bytes(&name).map_err(|_| StreamError::Message)?,
HeaderValue::from_bytes(&value).map_err(|_| StreamError::Message)?,
));
}
}
let method = Method::from_bytes(&method.ok_or(StreamError::Message)?)
.map_err(|_| StreamError::Message)?;
let connect = method == Method::CONNECT;
let extended = protocol.is_some();
if extended && !connect {
return Err(StreamError::Message);
}
if connect {
if extended {
if scheme.is_none() || authority.is_none() || path.is_some() {
return Err(StreamError::Message);
}
} else {
if authority.is_none() || scheme.is_some() || path.is_some() {
return Err(StreamError::Message);
}
}
} else if scheme.is_none() || path.is_none() || authority.is_none() {
return Err(StreamError::Message);
}
let uri = build_uri(connect, &scheme, &authority, &path)?;
let mut request = Request::builder()
.method(method)
.uri(uri)
.version(Version::HTTP_3)
.body(())
.expect("validated parts build a request");
for (name, value) in regular {
request.headers_mut().append(name, value);
}
Ok(request)
}
#[inline]
fn take_pseudo(slot: &mut Option<Bytes>, value: Bytes) -> Result<(), StreamError> {
if slot.is_some() {
return Err(StreamError::Message);
}
*slot = Some(value);
Ok(())
}
#[inline]
fn build_uri(
connect: bool,
scheme: &Option<Bytes>,
authority: &Option<Bytes>,
path: &Option<Bytes>,
) -> Result<Uri, StreamError> {
if connect {
let authority = authority.as_ref().expect("CONNECT authority validated");
let scheme = scheme
.as_ref()
.map(|s| String::from_utf8_lossy(s))
.unwrap_or_else(|| std::borrow::Cow::from("http"));
return Uri::builder()
.scheme(scheme.as_ref())
.authority(String::from_utf8_lossy(authority).as_ref())
.path_and_query("")
.build()
.map_err(|_| StreamError::Message);
}
let scheme = scheme.as_ref().expect("scheme validated");
let path = path.as_ref().expect("path validated");
match authority {
Some(authority) => Uri::builder()
.scheme(String::from_utf8_lossy(scheme).as_ref())
.authority(String::from_utf8_lossy(authority).as_ref())
.path_and_query(String::from_utf8_lossy(path).as_ref())
.build()
.map_err(|_| StreamError::Message),
None => Uri::builder()
.path_and_query(String::from_utf8_lossy(path).as_ref())
.build()
.map_err(|_| StreamError::Message),
}
}
#[inline]
fn header_map(headers: Vec<(Bytes, Bytes)>) -> Result<HeaderMap, StreamError> {
let mut map = HeaderMap::with_capacity(headers.len());
for (name, value) in headers {
if name.first() == Some(&b':') {
return Err(StreamError::Message);
}
map.append(
HeaderName::from_bytes(&name).map_err(|_| StreamError::Message)?,
HeaderValue::from_bytes(&value).map_err(|_| StreamError::Message)?,
);
}
Ok(map)
}
#[cfg(test)]
mod tests;