use std::collections::HashMap;
use std::time::Duration;
use crate::quiche;
use quiche::h3::frame::Frame;
use quiche::h3::Header;
use quiche::ConnectionError;
use serde::Deserialize;
use serde::Serialize;
use serde_with::serde_as;
use crate::encode_header_block;
use crate::encode_header_block_literal;
#[derive(Clone, Debug, Default, PartialEq, Eq)]
pub enum ExpectedStreamSendResult {
#[default]
Ok,
OkExact(usize),
Error(quiche::Error),
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum Action {
SendFrame {
stream_id: u64,
fin_stream: bool,
frame: Frame,
expected_result: ExpectedStreamSendResult,
},
SendHeadersFrame {
stream_id: u64,
fin_stream: bool,
literal_headers: bool,
headers: Vec<Header>,
frame: Frame,
expected_result: ExpectedStreamSendResult,
},
StreamBytes {
stream_id: u64,
fin_stream: bool,
bytes: Vec<u8>,
expected_result: ExpectedStreamSendResult,
},
SendDatagram {
payload: Vec<u8>,
},
OpenUniStream {
stream_id: u64,
fin_stream: bool,
stream_type: u64,
expected_result: ExpectedStreamSendResult,
},
ResetStream {
stream_id: u64,
error_code: u64,
},
StopSending {
stream_id: u64,
error_code: u64,
},
ConnectionClose {
error: ConnectionError,
},
FlushPackets,
Wait {
wait_type: WaitType,
},
}
#[serde_as]
#[derive(Clone, Debug, PartialEq, Eq, Hash, Serialize, Deserialize)]
pub enum WaitType {
#[serde(rename = "duration")]
WaitDuration(
#[serde_as(as = "serde_with::DurationMilliSecondsWithFrac<f64>")]
Duration,
),
StreamEvent(StreamEvent),
CanOpenNumStreams(RequiredStreamsQuota),
}
impl From<WaitType> for Action {
fn from(value: WaitType) -> Self {
Self::Wait { wait_type: value }
}
}
#[derive(Copy, Clone, Debug, PartialEq, Eq, Hash, Serialize, Deserialize)]
#[serde(rename = "snake_case")]
pub struct StreamEvent {
pub stream_id: u64,
#[serde(rename = "type")]
pub event_type: StreamEventType,
}
#[derive(Copy, Clone, Debug, PartialEq, Eq, Hash, Serialize, Deserialize)]
#[serde(rename = "snake_case")]
pub struct RequiredStreamsQuota {
pub num: u64,
pub bidi: bool,
}
#[derive(Copy, Clone, Debug, PartialEq, Eq, Hash, Serialize, Deserialize)]
#[serde(rename_all = "lowercase")]
pub enum StreamEventType {
Headers,
Data,
Finished,
}
#[derive(Debug, Default)]
pub(crate) struct WaitingFor {
stream_events: HashMap<u64, Vec<StreamEvent>>,
required_stream_quota: Option<RequiredStreamsQuota>,
}
impl WaitingFor {
pub(crate) fn is_empty(&self) -> bool {
self.stream_events.values().all(|v| v.is_empty()) &&
self.required_stream_quota.is_none()
}
pub(crate) fn add_wait(&mut self, stream_event: &StreamEvent) {
self.stream_events
.entry(stream_event.stream_id)
.or_default()
.push(*stream_event);
}
pub(crate) fn set_required_stream_quota(
&mut self, required: RequiredStreamsQuota,
) {
self.required_stream_quota = Some(required);
}
pub(crate) fn check_can_open_num_streams<F: quiche::BufFactory>(
&mut self, conn: &quiche::Connection<F>,
) {
if let Some(streams_required) = self.required_stream_quota {
let available_streams = if streams_required.bidi {
conn.peer_streams_left_bidi()
} else {
conn.peer_streams_left_uni()
};
if available_streams >= streams_required.num {
log::info!(
"required_stream_quota condition met \
(needed={}, available={}, bidi={})",
streams_required.num,
available_streams,
streams_required.bidi,
);
self.required_stream_quota = None;
}
}
}
pub(crate) fn remove_wait(&mut self, stream_event: StreamEvent) {
if let Some(waits) = self.stream_events.get_mut(&stream_event.stream_id) {
let old_len = waits.len();
waits.retain(|wait| wait != &stream_event);
let new_len = waits.len();
if old_len != new_len {
log::info!("No longer waiting for {stream_event:?}");
}
}
}
pub(crate) fn clear_waits_on_stream(&mut self, stream_id: u64) {
if let Some(waits) = self.stream_events.get_mut(&stream_id) {
if !waits.is_empty() {
log::info!("Clearing all waits for stream {stream_id}");
waits.clear();
}
}
}
}
pub fn send_headers_frame(
stream_id: u64, fin_stream: bool, headers: Vec<Header>,
) -> Action {
let header_block = encode_header_block(&headers).unwrap();
Action::SendHeadersFrame {
stream_id,
fin_stream,
headers,
literal_headers: false,
frame: Frame::Headers { header_block },
expected_result: ExpectedStreamSendResult::Ok,
}
}
pub fn send_headers_frame_with_expected_result(
stream_id: u64, fin_stream: bool, headers: Vec<Header>,
expected_result: ExpectedStreamSendResult,
) -> Action {
let header_block = encode_header_block(&headers).unwrap();
Action::SendHeadersFrame {
stream_id,
fin_stream,
headers,
literal_headers: false,
frame: Frame::Headers { header_block },
expected_result,
}
}
pub fn send_headers_frame_literal(
stream_id: u64, fin_stream: bool, headers: Vec<Header>,
) -> Action {
let header_block = encode_header_block_literal(&headers).unwrap();
Action::SendHeadersFrame {
stream_id,
fin_stream,
headers,
literal_headers: true,
frame: Frame::Headers { header_block },
expected_result: ExpectedStreamSendResult::Ok,
}
}
pub fn send_headers_frame_literal_with_expected_result(
stream_id: u64, fin_stream: bool, headers: Vec<Header>,
expected_result: ExpectedStreamSendResult,
) -> Action {
let header_block = encode_header_block_literal(&headers).unwrap();
Action::SendHeadersFrame {
stream_id,
fin_stream,
headers,
literal_headers: true,
frame: Frame::Headers { header_block },
expected_result,
}
}