#![allow(dead_code)]
use std::collections::VecDeque;
use std::sync::Arc;
use std::task::{Context, Poll};
use bytes::{Bytes, BytesMut};
use futures_util::ready;
use parking_lot::Mutex;
use crate::h3::error::{H3Error, TransportError};
use crate::h3::frame::{self, Frame, FrameDecoder, FrameError};
use crate::h3::qpack::{Encoder, QpackError, UnblockedSection};
use crate::h3::settings::{LocalSettings, PeerSettings};
use crate::h3::stream::SharedCodecs;
use crate::h3::transport::{Connection, UniStream};
pub(crate) const STREAM_TYPE_CONTROL: u64 = 0x0;
pub(crate) const STREAM_TYPE_PUSH: u64 = 0x1;
pub(crate) const STREAM_TYPE_QPACK_ENCODER: u64 = 0x2;
pub(crate) const STREAM_TYPE_QPACK_DECODER: u64 = 0x3;
#[derive(Debug)]
pub(crate) enum ControlError {
Transport(TransportError),
ClosedCriticalStream,
MissingSettings,
StreamCreation,
FrameUnexpected,
Frame,
Settings,
Id,
Qpack(QpackError),
}
impl ControlError {
#[inline]
pub(crate) fn h3_code(&self) -> u64 {
match self {
ControlError::Transport(_) => H3Error::GeneralProtocol.code(),
ControlError::ClosedCriticalStream => H3Error::ClosedCriticalStream.code(),
ControlError::MissingSettings => H3Error::MissingSettings.code(),
ControlError::StreamCreation => H3Error::StreamCreation.code(),
ControlError::FrameUnexpected => H3Error::FrameUnexpected.code(),
ControlError::Frame => H3Error::FrameError.code(),
ControlError::Settings => H3Error::Settings.code(),
ControlError::Id => H3Error::Id.code(),
ControlError::Qpack(err) => u64::from(err.code()),
}
}
}
impl From<TransportError> for ControlError {
#[inline]
fn from(err: TransportError) -> Self {
ControlError::Transport(err)
}
}
impl std::fmt::Display for ControlError {
#[inline]
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
std::fmt::Debug::fmt(self, f)
}
}
impl std::error::Error for ControlError {}
#[inline]
fn map_frame_error(err: FrameError) -> ControlError {
match err {
FrameError::Frame => ControlError::Frame,
FrameError::Unexpected(_) => ControlError::FrameUnexpected,
FrameError::Settings => ControlError::Settings,
}
}
#[derive(Debug)]
pub(crate) enum ControlEvent {
Settings(PeerSettings),
Goaway { id: u64 },
MaxPushId { id: u64 },
CancelPush { push_id: u64 },
}
struct PeerControl {
stream: Box<dyn UniStream>,
decoder: FrameDecoder,
}
struct PendingUni {
stream: Box<dyn UniStream>,
buf: BytesMut,
}
pub(crate) struct ControlStreams {
local: LocalSettings,
peer: PeerSettings,
out_control: Option<Box<dyn UniStream>>,
control_buf: BytesMut,
out_encoder: Option<Box<dyn UniStream>>,
out_decoder: Option<Box<dyn UniStream>>,
encoder_pending: VecDeque<Bytes>,
decoder_pending: VecDeque<Bytes>,
in_control: Option<PeerControl>,
in_encoder: Option<Box<dyn UniStream>>,
in_decoder: Option<Box<dyn UniStream>>,
pending_uni: Option<PendingUni>,
in_discard: Option<Box<dyn UniStream>>,
shared: Arc<Mutex<SharedCodecs>>,
settings_received: bool,
max_push_id: Option<u64>,
goaway_sent: Option<u64>,
events: VecDeque<ControlEvent>,
}
impl ControlStreams {
#[inline]
pub(crate) fn new(local: LocalSettings) -> Self {
let shared = Arc::new(Mutex::new(SharedCodecs::new(&local)));
Self {
local,
peer: PeerSettings::default(),
out_control: None,
control_buf: BytesMut::new(),
out_encoder: None,
out_decoder: None,
encoder_pending: VecDeque::new(),
decoder_pending: VecDeque::new(),
in_control: None,
in_encoder: None,
in_decoder: None,
pending_uni: None,
in_discard: None,
shared,
settings_received: false,
max_push_id: None,
goaway_sent: None,
events: VecDeque::new(),
}
}
#[inline]
pub(crate) fn peer_settings(&self) -> &PeerSettings {
&self.peer
}
#[inline]
pub(crate) fn shared(&self) -> &Arc<Mutex<SharedCodecs>> {
&self.shared
}
#[inline]
pub(crate) fn take_unblocked(&mut self) -> Vec<UnblockedSection> {
std::mem::take(&mut self.shared.lock().unblocked)
}
#[inline]
pub(crate) fn settings_received(&self) -> bool {
self.settings_received
}
#[inline]
pub(crate) fn max_push_id(&self) -> Option<u64> {
self.max_push_id
}
#[inline]
pub(crate) fn goaway_sent(&self) -> Option<u64> {
self.goaway_sent
}
#[inline]
pub(crate) fn shutting_down(&self) -> bool {
self.goaway_sent.is_some()
}
#[inline]
pub(crate) fn poll_init(
&mut self,
conn: &mut dyn Connection,
cx: &mut Context<'_>,
) -> Poll<Result<(), ControlError>> {
if self.out_control.is_none() {
let stream = ready!(conn.poll_open_uni(cx).map_err(ControlError::from))?;
self.out_control = Some(stream);
let mut settings = BytesMut::new();
frame::write_varint(STREAM_TYPE_CONTROL, &mut settings);
Frame::Settings(self.local.to_frame()).encode(&mut settings);
self.control_buf.extend_from_slice(&settings);
}
if self.out_encoder.is_none() {
let stream = ready!(conn.poll_open_uni(cx).map_err(ControlError::from))?;
self.out_encoder = Some(stream);
self.encoder_pending
.push_back(Bytes::from_static(&[STREAM_TYPE_QPACK_ENCODER as u8]));
}
if self.out_decoder.is_none() {
let stream = ready!(conn.poll_open_uni(cx).map_err(ControlError::from))?;
self.out_decoder = Some(stream);
self.decoder_pending
.push_back(Bytes::from_static(&[STREAM_TYPE_QPACK_DECODER as u8]));
}
Poll::Ready(Ok(()))
}
#[inline]
pub(crate) fn poll_flush(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), ControlError>> {
if let Some(stream) = self.out_control.as_mut() {
if !self.control_buf.is_empty() {
ready!(stream
.poll_send(cx, &self.control_buf)
.map_err(ControlError::from)?);
self.control_buf.clear();
}
}
if let Some(stream) = self.out_encoder.as_mut() {
while let Some(bytes) = self.encoder_pending.front() {
match stream.poll_send(cx, bytes).map_err(ControlError::from)? {
Poll::Ready(()) => {
self.encoder_pending.pop_front();
}
Poll::Pending => return Poll::Pending,
}
}
}
if let Some(stream) = self.out_decoder.as_mut() {
while let Some(bytes) = self.decoder_pending.front() {
match stream.poll_send(cx, bytes).map_err(ControlError::from)? {
Poll::Ready(()) => {
self.decoder_pending.pop_front();
}
Poll::Pending => return Poll::Pending,
}
}
}
Poll::Ready(Ok(()))
}
#[inline]
pub(crate) fn send_goaway(&mut self, id: u64) {
if self.goaway_sent.is_some() {
return;
}
let mut buf = BytesMut::new();
Frame::Goaway(id).encode(&mut buf);
self.control_buf.extend_from_slice(&buf);
self.goaway_sent = Some(id);
}
#[inline]
pub(crate) fn send_max_push_id(&mut self, id: u64) {
let mut buf = BytesMut::new();
Frame::MaxPushId(id).encode(&mut buf);
self.control_buf.extend_from_slice(&buf);
}
#[inline]
pub(crate) fn send_cancel_push(&mut self, push_id: u64) {
let mut buf = BytesMut::new();
Frame::CancelPush(push_id).encode(&mut buf);
self.control_buf.extend_from_slice(&buf);
}
#[inline]
pub(crate) fn queue_encoder_stream(&mut self, bytes: Bytes) {
if !bytes.is_empty() {
self.encoder_pending.push_back(bytes);
}
}
#[inline]
pub(crate) fn queue_encoder_streams(&mut self, bytes: &mut VecDeque<Bytes>) {
self.encoder_pending.append(bytes);
}
#[inline]
pub(crate) fn poll_read(
&mut self,
conn: &mut dyn Connection,
cx: &mut Context<'_>,
) -> Poll<Result<Option<ControlEvent>, ControlError>> {
loop {
if let Some(event) = self.events.pop_front() {
return Poll::Ready(Ok(Some(event)));
}
let mut progressed = false;
if self.pending_uni.is_none() {
match conn.poll_accept_uni(cx).map_err(ControlError::from)? {
Poll::Ready(Some(stream)) => {
self.pending_uni = Some(PendingUni {
stream,
buf: BytesMut::new(),
});
progressed = true;
}
Poll::Ready(None) | Poll::Pending => {}
}
}
if self.pending_uni.is_some() {
if let Poll::Ready(()) = self.classify_uni(cx)? {
progressed = true;
}
}
if let Some(control) = self.in_control.as_mut() {
match control.stream.poll_recv(cx).map_err(ControlError::from)? {
Poll::Ready(Some(chunk)) => {
control.decoder.extend(chunk);
progressed = true;
}
Poll::Ready(None) => {
return Poll::Ready(Err(ControlError::ClosedCriticalStream));
}
Poll::Pending => {}
}
}
loop {
let frame = {
let control = match self.in_control.as_mut() {
Some(control) => control,
None => break,
};
match control.decoder.next_frame() {
Ok(Some(frame)) => Some(frame),
Ok(None) => None,
Err(err) => return Poll::Ready(Err(map_frame_error(err))),
}
};
match frame {
Some(frame) => {
progressed = true;
self.handle_control_frame(frame)?;
}
None => break,
}
}
if let Some(stream) = self.in_encoder.as_mut() {
match stream.poll_recv(cx).map_err(ControlError::from)? {
Poll::Ready(Some(chunk)) => {
let mut shared = self.shared.lock();
match shared.decoder.feed_encoder_stream(&chunk) {
Ok(mut unblocked) => shared.unblocked.append(&mut unblocked),
Err(err) => return Poll::Ready(Err(ControlError::Qpack(err))),
}
let acks = shared.decoder.take_decoder_stream();
let waiters = shared.take_waiters();
drop(shared);
for waker in waiters {
waker.wake();
}
if !acks.is_empty() {
self.decoder_pending.push_back(acks);
}
progressed = true;
}
Poll::Ready(None) => {
return Poll::Ready(Err(ControlError::ClosedCriticalStream));
}
Poll::Pending => {}
}
}
if let Some(stream) = self.in_decoder.as_mut() {
match stream.poll_recv(cx).map_err(ControlError::from)? {
Poll::Ready(Some(chunk)) => {
let mut shared = self.shared.lock();
if let Some(encoder) = shared.encoder.as_mut() {
if let Err(err) = encoder.feed_decoder_stream(&chunk) {
drop(shared);
return Poll::Ready(Err(ControlError::Qpack(err)));
}
}
if let Err(err) = shared.decoder.feed_decoder_stream(&chunk) {
drop(shared);
return Poll::Ready(Err(ControlError::Qpack(err)));
}
progressed = true;
}
Poll::Ready(None) => {
return Poll::Ready(Err(ControlError::ClosedCriticalStream));
}
Poll::Pending => {}
}
}
if let Some(stream) = self.in_discard.as_mut() {
match stream.poll_recv(cx).map_err(ControlError::from)? {
Poll::Ready(Some(_)) => progressed = true,
Poll::Ready(None) => {
self.in_discard = None;
progressed = true;
}
Poll::Pending => {}
}
}
if !progressed {
return Poll::Pending;
}
}
}
#[inline]
fn classify_uni(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), ControlError>> {
loop {
let chunk = {
let pending = self.pending_uni.as_mut().expect("pending uni stream");
match pending.stream.poll_recv(cx).map_err(ControlError::from)? {
Poll::Ready(Some(chunk)) => chunk,
Poll::Ready(None) => {
self.pending_uni = None;
return Poll::Ready(Err(ControlError::StreamCreation));
}
Poll::Pending => return Poll::Pending,
}
};
self.pending_uni
.as_mut()
.expect("pending uni stream")
.buf
.extend_from_slice(&chunk);
let (ty, used) = {
let pending = self.pending_uni.as_ref().expect("pending uni stream");
match frame::parse_varint(&pending.buf).map_err(map_frame_error)? {
Some((ty, used)) => (ty, used),
None => continue,
}
};
let mut pending = self.pending_uni.take().expect("pending uni stream");
let leftover = pending.buf.split_off(used);
let stream = pending.stream;
match ty {
STREAM_TYPE_CONTROL => {
if self.in_control.is_some() {
return Poll::Ready(Err(ControlError::StreamCreation));
}
let mut decoder = FrameDecoder::new();
decoder.extend(leftover.freeze());
self.in_control = Some(PeerControl { stream, decoder });
}
STREAM_TYPE_PUSH => return Poll::Ready(Err(ControlError::StreamCreation)),
STREAM_TYPE_QPACK_ENCODER => {
if self.in_encoder.is_some() {
return Poll::Ready(Err(ControlError::StreamCreation));
}
if !leftover.is_empty() {
let mut shared = self.shared.lock();
match shared.decoder.feed_encoder_stream(&leftover) {
Ok(mut unblocked) => shared.unblocked.append(&mut unblocked),
Err(err) => return Poll::Ready(Err(ControlError::Qpack(err))),
}
let acks = shared.decoder.take_decoder_stream();
let waiters = shared.take_waiters();
drop(shared);
for waker in waiters {
waker.wake();
}
if !acks.is_empty() {
self.decoder_pending.push_back(acks);
}
}
self.in_encoder = Some(stream);
}
STREAM_TYPE_QPACK_DECODER => {
if self.in_decoder.is_some() {
return Poll::Ready(Err(ControlError::StreamCreation));
}
if !leftover.is_empty() {
let mut shared = self.shared.lock();
match shared.decoder.feed_decoder_stream(&leftover) {
Ok(()) => {}
Err(err) => return Poll::Ready(Err(ControlError::Qpack(err))),
}
}
self.in_decoder = Some(stream);
}
_ => {
self.in_discard = Some(stream);
}
}
return Poll::Ready(Ok(()));
}
}
#[inline]
fn handle_control_frame(&mut self, frame: Frame) -> Result<(), ControlError> {
if !self.settings_received && !matches!(&frame, Frame::Settings(_)) {
return Err(ControlError::MissingSettings);
}
match frame {
Frame::Settings(settings) => {
if self.settings_received {
return Err(ControlError::FrameUnexpected);
}
self.settings_received = true;
self.peer.apply(&settings);
let waiters = {
let mut shared = self.shared.lock();
shared.encoder = Some(Encoder::new(self.peer.qpack_max_table_capacity(), true));
shared.peer_max_field_section_size = self.peer.max_field_section_size();
shared.take_waiters()
};
for waker in waiters {
waker.wake();
}
self.events
.push_back(ControlEvent::Settings(self.peer.clone()));
}
Frame::Goaway(id) => {
self.events.push_back(ControlEvent::Goaway { id });
}
Frame::MaxPushId(id) => {
if let Some(prev) = self.max_push_id {
if id < prev {
return Err(ControlError::Id);
}
}
self.max_push_id = Some(id);
self.events.push_back(ControlEvent::MaxPushId { id });
}
Frame::CancelPush(_push_id) => {
return Err(ControlError::Id);
}
Frame::Data(_) | Frame::Headers(_) | Frame::PushPromise { .. } => {
return Err(ControlError::FrameUnexpected);
}
}
Ok(())
}
#[inline]
fn queue_decoder_acks(&mut self) {
let acks = self.shared.lock().decoder.take_decoder_stream();
if !acks.is_empty() {
self.decoder_pending.push_back(acks);
}
}
}
#[cfg(test)]
mod tests;