use std::{
collections::{HashMap, HashSet},
sync::{
atomic::{AtomicBool, AtomicU64, AtomicUsize, Ordering},
Arc, Mutex, OnceLock,
},
};
use crate::config::Config;
use crate::credit::Credit;
use crate::sched::PriorityQueue;
use crate::transport::{Reader, Transport, Writer};
use crate::{
proto::varint_size, ApplicationClose, ConnectionClose, Error, Frame, ResetStream, StopSending,
Stream, StreamDir, StreamId, TransportParams, Version, MAX_FRAME_PAYLOAD,
};
use bytes::{Buf, BufMut, Bytes};
use tokio::sync::{mpsc, watch};
use web_transport_proto::VarInt;
use web_transport_trait as generic;
const DATAGRAM_RECV_BUFFER: usize = 1024;
const DATAGRAM_SEND_BUFFER: usize = 64;
#[derive(Default)]
struct Streams {
send: HashMap<StreamId, SendState>,
recv: HashMap<StreamId, RecvState>,
peer_initial_max_stream_data_uni: u64,
peer_initial_max_stream_data_bidi_remote: u64,
}
struct SessionGuard {
closed: watch::Sender<Option<Error>>,
}
impl Drop for SessionGuard {
fn drop(&mut self) {
note_closed(&self.closed, Error::Closed);
}
}
#[derive(Clone)]
pub struct Session {
is_server: bool,
config: Config,
outbound: PriorityQueue,
outbound_priority: mpsc::UnboundedSender<Frame>,
accept_bi: Arc<tokio::sync::Mutex<mpsc::Receiver<(SendStream, RecvStream)>>>,
accept_uni: Arc<tokio::sync::Mutex<mpsc::Receiver<RecvStream>>>,
streams: Arc<Mutex<Streams>>,
closed: watch::Sender<Option<Error>>,
negotiated: Arc<OnceLock<Option<String>>>,
established: watch::Receiver<bool>,
open_bi_credit: Credit,
open_uni_credit: Credit,
conn_send_credit: Credit,
conn_recv_credit: Credit,
recv_datagram: Arc<tokio::sync::Mutex<mpsc::Receiver<Bytes>>>,
outbound_datagram: mpsc::Sender<Bytes>,
datagram_max_size: Arc<AtomicUsize>,
_guard: Arc<SessionGuard>,
}
#[derive(Default)]
struct RecvOpen {
created_max: Option<u64>,
holes: HashSet<u64>,
}
impl RecvOpen {
fn is_closed(&self, index: u64) -> bool {
matches!(self.created_max, Some(max) if index <= max) && !self.holes.contains(&index)
}
fn record(&mut self, index: u64) {
match self.created_max {
Some(max) if index <= max => {
self.holes.remove(&index);
}
prev => {
let start = prev.map_or(0, |max| max + 1);
self.holes.extend(start..index);
self.created_max = Some(index);
}
}
}
}
struct SessionState<R: Reader> {
reader: R,
config: Config,
is_server: bool,
outbound: PriorityQueue,
control: mpsc::UnboundedSender<Frame>,
accept_bi: mpsc::Sender<(SendStream, RecvStream)>,
accept_uni: mpsc::Sender<RecvStream>,
streams: Arc<Mutex<Streams>>,
closed: watch::Sender<Option<Error>>,
negotiated: Arc<OnceLock<Option<String>>>,
established: watch::Sender<bool>,
conn_send_credit: Credit,
conn_recv_credit: Credit,
our_params: TransportParams,
peer_params: TransportParams,
params_received: bool,
open_bi_credit: Credit,
open_uni_credit: Credit,
recv_bi_credit: Credit,
recv_uni_credit: Credit,
recv_open_bi: RecvOpen,
recv_open_uni: RecvOpen,
base: tokio::time::Instant,
last_recv_at: Arc<AtomicU64>,
reader_backpressured: Arc<AtomicBool>,
recv_datagram: mpsc::Sender<Bytes>,
datagram_max_size: Arc<AtomicUsize>,
record_limit: Arc<AtomicU64>,
idle_timeout_ms: Arc<AtomicU64>,
last_ping_recv: Option<u64>,
pings_sent: Arc<AtomicU64>,
}
async fn next_outbound(
control: &mut mpsc::UnboundedReceiver<Frame>,
datagram: &mut mpsc::Receiver<Bytes>,
stream: &PriorityQueue,
) -> Option<Frame> {
tokio::select! {
biased;
Some(frame) = control.recv() => Some(frame),
Some(payload) = datagram.recv() => Some(Frame::Datagram(payload.into())),
frame = stream.pop() => frame,
}
}
fn negotiated_idle_timeout_ms(ours: u64, peer: u64) -> u64 {
match (ours, peer) {
(0, 0) => 0,
(a, 0) | (0, a) => a,
(a, b) => a.min(b),
}
}
fn millis_since(base: tokio::time::Instant, now: tokio::time::Instant) -> u64 {
now.saturating_duration_since(base).as_millis() as u64
}
fn instant_at(base: tokio::time::Instant, ms: u64) -> tokio::time::Instant {
base + std::time::Duration::from_millis(ms)
}
fn note_closed(closed: &watch::Sender<Option<Error>>, err: Error) {
closed.send_if_modified(|slot| {
if slot.is_none() {
*slot = Some(err);
true
} else {
false
}
});
}
struct WriterState<W: Writer> {
writer: W,
version: Version,
control: mpsc::UnboundedReceiver<Frame>,
datagrams: mpsc::Receiver<Bytes>,
outbound: PriorityQueue,
streams: Arc<Mutex<Streams>>,
record_limit: Arc<AtomicU64>,
writer_backpressured: Arc<AtomicBool>,
closed: watch::Sender<Option<Error>>,
base: tokio::time::Instant,
last_send_at: Arc<AtomicU64>,
}
enum Transmitted {
Ok,
Failed(Error),
Interrupted,
}
impl<W: Writer> WriterState<W> {
fn note_closed(&self, err: Error) {
note_closed(&self.closed, err);
}
async fn run(&mut self) {
let mut closed_rx = self.closed.subscribe();
let mut interrupted = false;
loop {
tokio::select! {
biased;
frame = next_outbound(&mut self.control, &mut self.datagrams, &self.outbound) => {
match frame {
Some(frame) => match self.transmit_or_teardown(frame, &mut closed_rx).await {
Transmitted::Ok => {}
Transmitted::Failed(err) => {
self.note_closed(err);
break;
}
Transmitted::Interrupted => {
interrupted = true;
break;
}
},
None => break,
}
}
result = self.writer.maintain() => {
if let Err(err) = result {
self.note_closed(err);
break;
}
}
_ = async { closed_rx.wait_for(|slot| slot.is_some()).await.ok(); } => {
while let Ok(frame) = self.control.try_recv() {
if self.transmit(frame).await.is_err() {
break;
}
}
break;
}
}
}
if !interrupted {
let _ = self.writer.close().await;
}
}
async fn transmit_or_teardown(
&mut self,
frame: Frame,
closed_rx: &mut watch::Receiver<Option<Error>>,
) -> Transmitted {
tokio::select! {
biased;
result = self.transmit(frame) => match result {
Ok(()) => Transmitted::Ok,
Err(err) => Transmitted::Failed(err),
},
_ = async { closed_rx.wait_for(|slot| slot.is_some()).await.ok(); } => {
Transmitted::Interrupted
}
}
}
async fn transmit(&mut self, mut frame: Frame) -> Result<(), Error> {
let transmitted_stream = match &frame {
Frame::Stream(stream) if !stream.fin => Some((stream.id, stream.data.len() as u64)),
_ => None,
};
match &mut frame {
Frame::ResetStream(reset) => {
if let Some(send) = self.streams.lock().unwrap().send.remove(&reset.id) {
reset.final_size = send.sent_offset;
}
}
Frame::Stream(stream) if stream.fin => {
self.streams.lock().unwrap().send.remove(&stream.id);
}
Frame::StopSending(stop) => {
self.streams.lock().unwrap().recv.remove(&stop.id);
}
_ => {}
}
let bytes = frame.encode(self.version)?;
if self.version.uses_records() {
let limit = self.record_limit.load(Ordering::Acquire);
if bytes.len() as u64 > limit {
return Err(Error::FrameTooLarge);
}
}
self.writer_backpressured.store(true, Ordering::Release);
let result = self.writer.send(bytes).await;
self.writer_backpressured.store(false, Ordering::Release);
result?;
if let Some((id, len)) = transmitted_stream {
if let Some(send) = self.streams.lock().unwrap().send.get_mut(&id) {
send.sent_offset += len;
}
}
self.last_send_at.store(
millis_since(self.base, tokio::time::Instant::now()),
Ordering::Release,
);
Ok(())
}
}
#[cfg(test)]
mod writer_final_size_tests {
use super::*;
struct CaptureWriter(Arc<Mutex<Vec<Bytes>>>);
impl Writer for CaptureWriter {
async fn send(&mut self, data: Bytes) -> Result<(), Error> {
self.0.lock().unwrap().push(data);
Ok(())
}
async fn close(&mut self) -> Result<(), Error> {
Ok(())
}
}
#[tokio::test]
async fn reset_uses_bytes_transmitted_not_frontend_offset() {
let sent = Arc::new(Mutex::new(Vec::new()));
let streams = Arc::new(Mutex::new(Streams::default()));
let id = StreamId::new(0, StreamDir::Uni, false);
let (stopped, _stopped_rx) = mpsc::unbounded_channel();
streams.lock().unwrap().send.insert(
id,
SendState {
inbound_stopped: stopped,
sent_offset: 0,
stream_credit: None,
},
);
let (_control_tx, control) = mpsc::unbounded_channel();
let (_datagram_tx, datagrams) = mpsc::channel(1);
let mut writer = WriterState {
writer: CaptureWriter(sent.clone()),
version: Version::QMux01,
control,
datagrams,
outbound: PriorityQueue::new(1),
streams,
record_limit: Arc::new(AtomicU64::new(u64::MAX)),
writer_backpressured: Arc::new(AtomicBool::new(false)),
closed: watch::Sender::new(None),
base: tokio::time::Instant::now(),
last_send_at: Arc::new(AtomicU64::new(0)),
};
writer
.transmit(
Stream {
id,
offset: 0,
data: Bytes::from_static(b"abc"),
fin: false,
}
.into(),
)
.await
.unwrap();
writer
.transmit(
ResetStream {
id,
code: VarInt::from_u32(0),
final_size: 99,
reliable_size: None,
}
.into(),
)
.await
.unwrap();
let wire = sent.lock().unwrap()[1].clone();
let decoded = Frame::decode(wire, Version::QMux01).unwrap().unwrap();
let Frame::ResetStream(reset) = decoded else {
panic!("expected RESET_STREAM");
};
assert_eq!(reset.final_size, 3);
}
#[tokio::test]
async fn reset_returns_credit_reserved_for_dropped_frames() {
let id = StreamId::new(0, StreamDir::Uni, false);
let outbound = PriorityQueue::new(4);
let (priority, _priority_rx) = mpsc::unbounded_channel();
let (_stopped, stopped_rx) = mpsc::unbounded_channel();
let stream_credit = Credit::new(3);
let conn_credit = Credit::new(3);
let mut send = SendStream {
id,
outbound,
outbound_priority: priority,
inbound_stopped: stopped_rx,
offset: 0,
priority: 0,
closed: None,
fin: false,
stream_credit: Some(stream_credit.clone()),
conn_credit: Some(conn_credit.clone()),
};
assert_eq!(
generic::SendStream::write(&mut send, b"abc").await.unwrap(),
3
);
assert_eq!(stream_credit.try_claim(1), 0);
assert_eq!(conn_credit.try_claim(1), 0);
generic::SendStream::reset(&mut send, 0);
assert_eq!(stream_credit.try_claim(3), 3);
assert_eq!(conn_credit.try_claim(3), 3);
}
}
struct TimerState {
base: tokio::time::Instant,
last_recv_at: Arc<AtomicU64>,
last_send_at: Arc<AtomicU64>,
reader_backpressured: Arc<AtomicBool>,
writer_backpressured: Arc<AtomicBool>,
idle_timeout_ms: Arc<AtomicU64>,
control: mpsc::UnboundedSender<Frame>,
closed: watch::Sender<Option<Error>>,
established: watch::Receiver<bool>,
pings_sent: Arc<AtomicU64>,
}
impl TimerState {
async fn run(mut self) {
let mut closed_rx = self.closed.subscribe();
tokio::select! {
biased;
_ = closed_rx.wait_for(|s| s.is_some()) => return,
res = self.established.wait_for(|&e| e) => {
if res.is_err() {
return; }
}
}
let idle_ms = self.idle_timeout_ms.load(Ordering::Acquire);
if idle_ms == 0 {
return;
}
let idle = std::time::Duration::from_millis(idle_ms);
let ping_every = std::time::Duration::from_millis((idle_ms / 3).max(1));
let mut deferred_since: Option<tokio::time::Instant> = None;
let mut last_ping_ms = self.last_send_at.load(Ordering::Acquire);
let mut next_ping_seq: u64 = 0;
loop {
let last_activity = instant_at(
self.base,
self.last_recv_at
.load(Ordering::Acquire)
.max(self.last_send_at.load(Ordering::Acquire)),
);
let ping_ref = instant_at(
self.base,
self.last_send_at.load(Ordering::Acquire).max(last_ping_ms),
);
let idle_wake = match deferred_since {
Some(since) => since + idle,
None => last_activity + idle,
};
let wake = idle_wake.min(ping_ref + ping_every);
tokio::select! {
biased;
_ = closed_rx.wait_for(|s| s.is_some()) => return,
_ = tokio::time::sleep_until(wake) => {}
}
let now = tokio::time::Instant::now();
if now >= ping_ref + ping_every {
if !self.writer_backpressured.load(Ordering::Acquire) {
let ping = Frame::Ping(crate::Ping {
sequence: next_ping_seq,
response: false,
});
next_ping_seq = next_ping_seq.wrapping_add(1);
self.pings_sent.store(next_ping_seq, Ordering::Release);
if self.control.send(ping).is_err() {
return; }
}
last_ping_ms = millis_since(self.base, now);
}
let last_activity = instant_at(
self.base,
self.last_recv_at
.load(Ordering::Acquire)
.max(self.last_send_at.load(Ordering::Acquire)),
);
if now < last_activity + idle {
deferred_since = None; continue;
}
let backpressured = self.writer_backpressured.load(Ordering::Acquire)
|| self.reader_backpressured.load(Ordering::Acquire);
match deferred_since {
Some(since) if now.duration_since(since) < idle => continue,
Some(_) => {}
None if backpressured => {
deferred_since = Some(now);
continue;
}
None => {}
}
tracing::debug!("idle timeout fired");
note_closed(&self.closed, Error::IdleTimeout);
return;
}
}
}
impl<R: Reader> SessionState<R> {
async fn run(&mut self) -> Result<(), Error> {
let mut closed = self.closed.subscribe();
loop {
tokio::select! {
biased;
result = self.reader.recv() => {
let data = result?;
self.last_recv_at.store(
millis_since(self.base, tokio::time::Instant::now()),
Ordering::Release,
);
if self.config.version.uses_records() {
for frame in Frame::decode_record(data)? {
self.recv_frame(frame).await?;
}
} else if let Some(frame) = Frame::decode(data, self.config.version)? {
self.recv_frame(frame).await?;
}
}
_ = async { closed.wait_for(|err| err.is_some()).await.ok(); } => {
return Err(closed.borrow().clone().unwrap_or(Error::Closed))
}
}
}
}
fn recv_open(&self, dir: StreamDir) -> &RecvOpen {
match dir {
StreamDir::Bi => &self.recv_open_bi,
StreamDir::Uni => &self.recv_open_uni,
}
}
async fn recv_frame(&mut self, frame: Frame) -> Result<(), Error> {
if self.config.version == Version::QMux02 {
let is_params = matches!(frame, Frame::TransportParameters(_));
if is_params == self.params_received {
return Err(Error::ProtocolViolation);
}
}
match frame {
Frame::TransportParameters(params) => {
self.recv_transport_parameters(params)?;
}
Frame::Stream(stream) => {
if stream.data.len() > MAX_FRAME_PAYLOAD {
return Err(Error::FrameTooLarge);
}
if !stream.id.can_recv(self.is_server) {
return Err(Error::InvalidStreamId);
}
let live = self.streams.lock().unwrap().recv.contains_key(&stream.id);
if self.config.version.is_qmux()
&& stream.id.server_initiated() != self.is_server
&& !live
&& self.recv_open(stream.id.dir()).is_closed(stream.id.index())
{
return Ok(());
}
let data_len = stream.data.len() as u64;
if data_len > 0 && !self.conn_recv_credit.receive(data_len) {
return Err(Error::FlowControlError);
}
{
let mut streams = self.streams.lock().unwrap();
if let Some(recv) = streams.recv.get_mut(&stream.id) {
if data_len > 0 && !recv.recv_credit.receive(data_len) {
return Err(Error::FlowControlError);
}
recv.recv_offset += data_len;
if data_len == 0 && !stream.fin {
return Ok(());
}
let id = stream.id;
let fin = stream.fin;
recv.inbound_data.send(stream).ok();
if fin {
streams.recv.remove(&id);
}
return Ok(());
}
}
if self.is_server == stream.id.server_initiated() {
return Ok(());
}
if self.config.version.is_qmux() {
let credit = match stream.id.dir() {
StreamDir::Bi => &self.recv_bi_credit,
StreamDir::Uni => &self.recv_uni_credit,
};
if !credit.receive_up_to(stream.id.index() + 1) {
return Err(Error::StreamLimitExceeded);
}
match stream.id.dir() {
StreamDir::Bi => &mut self.recv_open_bi,
StreamDir::Uni => &mut self.recv_open_uni,
}
.record(stream.id.index());
}
let (tx, rx) = mpsc::unbounded_channel();
let (tx2, rx2) = mpsc::unbounded_channel();
let recv_window = if self.config.version.is_qmux() {
match stream.id.dir() {
StreamDir::Bi => self.our_params.initial_max_stream_data_bidi_remote,
StreamDir::Uni => self.our_params.initial_max_stream_data_uni,
}
} else {
u64::MAX
};
let recv_credit = Credit::new(recv_window);
if data_len > 0 && !recv_credit.receive(data_len) {
return Err(Error::FlowControlError);
}
let recv_backend = RecvState {
inbound_data: tx,
inbound_reset: tx2,
recv_credit: recv_credit.clone(),
recv_offset: data_len,
};
let recv_streams_credit = if self.config.version.is_qmux() {
Some(match stream.id.dir() {
StreamDir::Bi => self.recv_bi_credit.clone(),
StreamDir::Uni => self.recv_uni_credit.clone(),
})
} else {
None
};
let recv_frontend = RecvStream {
id: stream.id,
inbound_data: rx,
inbound_reset: rx2,
outbound_priority: self.control.clone(),
buffer: Bytes::new(),
closed: None,
fin: false,
recv_credit,
conn_recv_credit: self.conn_recv_credit.clone(),
version: self.config.version,
recv_streams_credit,
};
match stream.id.dir() {
StreamDir::Uni => {
self.reader_backpressured.store(true, Ordering::Release);
let result = self.accept_uni.send(recv_frontend).await;
self.reader_backpressured.store(false, Ordering::Release);
result.map_err(|_| Error::Closed)?;
}
StreamDir::Bi => {
let (tx, rx) = mpsc::unbounded_channel();
let send_backend = SendState {
inbound_stopped: tx,
sent_offset: 0,
stream_credit: if self.config.version.is_qmux() {
Some(Credit::new(
self.peer_params.initial_max_stream_data_bidi_local,
))
} else {
None
},
};
let send_frontend = SendStream {
id: stream.id,
outbound: self.outbound.clone(),
outbound_priority: self.control.clone(),
inbound_stopped: rx,
offset: 0,
priority: 0,
closed: None,
fin: false,
stream_credit: send_backend.stream_credit.clone(),
conn_credit: if self.config.version.is_qmux() {
Some(self.conn_send_credit.clone())
} else {
None
},
};
self.streams
.lock()
.unwrap()
.send
.insert(stream.id, send_backend);
self.reader_backpressured.store(true, Ordering::Release);
let result = self.accept_bi.send((send_frontend, recv_frontend)).await;
self.reader_backpressured.store(false, Ordering::Release);
result.map_err(|_| Error::Closed)?;
}
};
let id = stream.id;
let fin = stream.fin;
if data_len > 0 || fin {
recv_backend.inbound_data.send(stream).ok();
}
if !fin {
self.streams.lock().unwrap().recv.insert(id, recv_backend);
}
}
Frame::ResetStream(reset) => {
if reset.reliable_size.is_some() && !self.our_params.reset_stream_at {
return Err(Error::ProtocolViolation);
}
if !reset.id.can_recv(self.is_server) {
return Err(Error::InvalidStreamId);
}
let reset_id = reset.id;
let peer_initiated = reset_id.server_initiated() != self.is_server;
let live = self.streams.lock().unwrap().recv.contains_key(&reset_id);
if !live {
if !peer_initiated {
return Ok(());
}
if self.config.version.is_qmux()
&& self.recv_open(reset_id.dir()).is_closed(reset_id.index())
{
return Ok(());
}
if self.config.version.is_qmux() {
let credit = match reset_id.dir() {
StreamDir::Bi => &self.recv_bi_credit,
StreamDir::Uni => &self.recv_uni_credit,
};
if !credit.receive_up_to(reset_id.index() + 1) {
return Err(Error::StreamLimitExceeded);
}
}
}
if self.config.version.is_qmux() {
let received = self
.streams
.lock()
.unwrap()
.recv
.get(&reset_id)
.map_or(0, |recv| recv.recv_offset);
let final_size = reset.final_size.max(received);
let gap = final_size - received;
let stream_ok = if live {
let mut streams = self.streams.lock().unwrap();
let recv = streams.recv.get_mut(&reset_id).expect("live recv stream");
let ok = recv.recv_credit.receive(gap);
if ok {
recv.recv_offset = final_size;
}
ok
} else {
let recv_max = match reset_id.dir() {
StreamDir::Bi => self.our_params.initial_max_stream_data_bidi_remote,
StreamDir::Uni => self.our_params.initial_max_stream_data_uni,
};
final_size <= recv_max
};
if !stream_ok || !self.conn_recv_credit.receive(gap) {
return Err(Error::FlowControlError);
}
if let Some(new_max) = self.conn_recv_credit.consume(gap) {
self.control.send(Frame::MaxData(new_max)).ok();
}
}
let delivered = {
let mut streams = self.streams.lock().unwrap();
if let Some(recv) = streams.recv.remove(&reset_id) {
recv.inbound_reset.send(reset).ok();
true
} else {
false
}
};
if !delivered && self.config.version.is_qmux() && peer_initiated {
match reset_id.dir() {
StreamDir::Bi => &mut self.recv_open_bi,
StreamDir::Uni => &mut self.recv_open_uni,
}
.record(reset_id.index());
let credit = match reset_id.dir() {
StreamDir::Bi => &self.recv_bi_credit,
StreamDir::Uni => &self.recv_uni_credit,
};
if let Some(new_max) = credit.consume(1) {
let frame = match reset_id.dir() {
StreamDir::Bi => Frame::MaxStreamsBidi(new_max),
StreamDir::Uni => Frame::MaxStreamsUni(new_max),
};
self.control.send(frame).ok();
}
}
}
Frame::StopSending(stop) => {
if !stop.id.can_send(self.is_server) {
return Err(Error::InvalidStreamId);
}
if let Some(send) = self.streams.lock().unwrap().send.get(&stop.id) {
send.inbound_stopped.send(stop).ok();
}
}
Frame::ApplicationClose(close) => {
self.closed
.send(Some(Error::ConnectionClosed {
code: close.code,
reason: close.reason,
}))
.ok();
}
Frame::ConnectionClose(close) => {
self.closed
.send(Some(Error::ConnectionReset {
code: close.code,
reason: close.reason,
}))
.ok();
}
Frame::MaxData(max) => {
self.conn_send_credit.increase_max(max)?;
}
Frame::MaxStreamData { id, max } => {
if let Some(send) = self.streams.lock().unwrap().send.get(&id) {
if let Some(credit) = &send.stream_credit {
credit.increase_max(max)?;
}
}
}
Frame::MaxStreamsBidi(max) => {
self.open_bi_credit.increase_max(max)?;
}
Frame::MaxStreamsUni(max) => {
self.open_uni_credit.increase_max(max)?;
}
Frame::DataBlocked(_)
| Frame::StreamDataBlocked { .. }
| Frame::StreamsBlockedBidi(_)
| Frame::StreamsBlockedUni(_) => {}
Frame::Ping(ping) => {
if self.config.version == Version::QMux02 {
if ping.response {
if ping.sequence >= self.pings_sent.load(Ordering::Acquire) {
return Err(Error::ProtocolViolation);
}
} else {
if self
.last_ping_recv
.is_some_and(|prev| ping.sequence <= prev)
{
return Err(Error::ProtocolViolation);
}
self.last_ping_recv = Some(ping.sequence);
}
}
if !ping.response {
let response = Frame::Ping(crate::Ping {
sequence: ping.sequence,
response: true,
});
self.control.send(response).ok();
}
}
Frame::Datagram(datagram) => {
if self.our_params.max_datagram_frame_size == 0 {
return Err(Error::DatagramsUnsupported);
}
if datagram.frame_size() > self.our_params.max_datagram_frame_size {
return Err(Error::FrameTooLarge);
}
let _ = self.recv_datagram.try_send(datagram.data);
}
}
Ok(())
}
fn recv_transport_parameters(&mut self, params: TransportParams) -> Result<(), Error> {
if self.params_received {
return Err(Error::FlowControlError);
}
self.params_received = true;
if self.config.version.uses_records()
&& params.max_record_size < crate::proto::DEFAULT_MAX_RECORD_SIZE
{
return Err(Error::TransportParameter);
}
match &self.config.protocol {
crate::Protocol::Negotiate(ours) => {
let agreed = negotiate_protocol(self.is_server, ours, ¶ms.protocols);
self.negotiated.set(agreed).ok();
}
crate::Protocol::None | crate::Protocol::Negotiated(_) => {
if !params.protocols.is_empty() {
return Err(Error::UnexpectedProtocols);
}
}
}
self.conn_send_credit
.increase_max(params.initial_max_data)
.ok();
self.open_bi_credit
.increase_max(params.initial_max_streams_bidi)
.ok();
self.open_uni_credit
.increase_max(params.initial_max_streams_uni)
.ok();
{
let mut streams = self.streams.lock().unwrap();
streams.peer_initial_max_stream_data_uni = params.initial_max_stream_data_uni;
streams.peer_initial_max_stream_data_bidi_remote =
params.initial_max_stream_data_bidi_remote;
for (id, send) in &streams.send {
if let Some(credit) = &send.stream_credit {
let initial = match id.dir() {
StreamDir::Bi => {
if id.server_initiated() == self.is_server {
params.initial_max_stream_data_bidi_remote
} else {
params.initial_max_stream_data_bidi_local
}
}
StreamDir::Uni => params.initial_max_stream_data_uni,
};
credit.increase_max(initial).ok();
}
}
}
let idle_ms = if self.config.version.uses_records() {
self.record_limit
.store(params.max_record_size, Ordering::Release);
negotiated_idle_timeout_ms(self.our_params.max_idle_timeout, params.max_idle_timeout)
} else {
0
};
self.idle_timeout_ms.store(idle_ms, Ordering::Release);
let datagram_max =
if !self.config.version.uses_records() || params.max_datagram_frame_size == 0 {
0
} else {
let cap = params.max_datagram_frame_size.min(params.max_record_size);
let overhead = 1 + varint_size(cap);
usize::try_from(cap.saturating_sub(overhead)).unwrap_or(usize::MAX)
};
self.datagram_max_size
.store(datagram_max, Ordering::Release);
self.peer_params = params;
self.established.send_replace(true);
Ok(())
}
}
impl Session {
pub async fn connect<T: Transport>(transport: T, config: Config) -> Result<Session, Error> {
let session = Self::new(transport, false, config);
session.established().await?;
Ok(session)
}
pub async fn accept<T: Transport>(transport: T, config: Config) -> Result<Session, Error> {
let session = Self::new(transport, true, config);
session.established().await?;
Ok(session)
}
async fn established(&self) -> Result<(), Error> {
let mut established = self.established.clone();
if *established.borrow() {
return Ok(());
}
let wait = established.wait_for(|&done| done);
let timeout = self.config.handshake_timeout;
let outcome = if timeout.is_zero() {
Some(wait.await)
} else {
tokio::time::timeout(timeout, wait).await.ok()
};
match outcome {
Some(Ok(_)) => Ok(()),
Some(Err(_)) => Err(self.closed.borrow().clone().unwrap_or(Error::Closed)),
None => {
let _ = self.outbound_priority.send(
ConnectionClose {
code: VarInt::from(0u32),
reason: "handshake timeout".to_string(),
}
.into(),
);
self.closed.send_replace(Some(Error::HandshakeTimeout));
Err(Error::HandshakeTimeout)
}
}
}
pub(crate) fn new<T: Transport>(transport: T, is_server: bool, config: Config) -> Self {
let version = config.version;
let our_params = config.to_transport_params();
let (accept_bi_tx, accept_bi_rx) = mpsc::channel(1024);
let (accept_uni_tx, accept_uni_rx) = mpsc::channel(1024);
let outbound = PriorityQueue::new(8);
let (control_tx, control_rx) = mpsc::unbounded_channel();
let (recv_datagram_tx, recv_datagram_rx) = mpsc::channel(DATAGRAM_RECV_BUFFER);
let (outbound_datagram_tx, outbound_datagram_rx) = mpsc::channel(DATAGRAM_SEND_BUFFER);
let datagram_max_size = Arc::new(AtomicUsize::new(0));
let streams: Arc<Mutex<Streams>> = Arc::new(Mutex::new(Streams::default()));
let record_limit = Arc::new(AtomicU64::new(crate::proto::DEFAULT_MAX_RECORD_SIZE));
let idle_timeout_ms = Arc::new(AtomicU64::new(0));
let pings_sent = Arc::new(AtomicU64::new(0));
let base = tokio::time::Instant::now();
let last_recv_at = Arc::new(AtomicU64::new(0));
let last_send_at = Arc::new(AtomicU64::new(0));
let reader_backpressured = Arc::new(AtomicBool::new(false));
let writer_backpressured = Arc::new(AtomicBool::new(false));
let closed = watch::Sender::new(None);
if version.is_qmux() {
control_tx
.send(Frame::TransportParameters(our_params.clone()))
.ok();
}
let (writer_half, reader_half) = transport.split();
let mut writer = WriterState {
writer: writer_half,
version,
control: control_rx,
datagrams: outbound_datagram_rx,
outbound: outbound.clone(),
streams: streams.clone(),
record_limit: record_limit.clone(),
writer_backpressured: writer_backpressured.clone(),
closed: closed.clone(),
base,
last_send_at: last_send_at.clone(),
};
tokio::spawn(async move { writer.run().await });
let negotiated: Arc<OnceLock<Option<String>>> = Arc::new(OnceLock::new());
match &config.protocol {
crate::Protocol::Negotiate(_) => {} crate::Protocol::Negotiated(name) => {
negotiated.set(Some(name.clone())).ok();
}
crate::Protocol::None => {
negotiated.set(None).ok();
}
}
let (established_tx, established_rx) = watch::channel(!version.is_qmux());
let open_bi_credit = Credit::new(if version.is_qmux() { 0 } else { u64::MAX });
let open_uni_credit = Credit::new(if version.is_qmux() { 0 } else { u64::MAX });
let conn_send_credit = Credit::new(if version.is_qmux() { 0 } else { u64::MAX });
let conn_recv_credit = Credit::new(if version.is_qmux() {
our_params.initial_max_data
} else {
u64::MAX
});
let recv_bi_credit = Credit::new(if version.is_qmux() {
config.max_streams_bidi
} else {
u64::MAX
});
let recv_uni_credit = Credit::new(if version.is_qmux() {
config.max_streams_uni
} else {
u64::MAX
});
let mut backend = SessionState {
reader: reader_half,
config: config.clone(),
is_server,
outbound: outbound.clone(),
control: control_tx.clone(),
accept_bi: accept_bi_tx,
accept_uni: accept_uni_tx,
streams: streams.clone(),
closed: closed.clone(),
negotiated: negotiated.clone(),
established: established_tx,
conn_send_credit: conn_send_credit.clone(),
conn_recv_credit: conn_recv_credit.clone(),
our_params: our_params.clone(),
peer_params: TransportParams::default(),
params_received: false,
open_bi_credit: open_bi_credit.clone(),
open_uni_credit: open_uni_credit.clone(),
recv_bi_credit: recv_bi_credit.clone(),
recv_uni_credit: recv_uni_credit.clone(),
recv_open_bi: RecvOpen::default(),
recv_open_uni: RecvOpen::default(),
base,
last_recv_at: last_recv_at.clone(),
reader_backpressured: reader_backpressured.clone(),
recv_datagram: recv_datagram_tx,
datagram_max_size: datagram_max_size.clone(),
record_limit: record_limit.clone(),
idle_timeout_ms: idle_timeout_ms.clone(),
last_ping_recv: None,
pings_sent: pings_sent.clone(),
};
if version.uses_records() {
let timer = TimerState {
base,
last_recv_at: last_recv_at.clone(),
last_send_at: last_send_at.clone(),
reader_backpressured: reader_backpressured.clone(),
writer_backpressured: writer_backpressured.clone(),
idle_timeout_ms: idle_timeout_ms.clone(),
control: control_tx.clone(),
closed: closed.clone(),
established: established_rx.clone(),
pings_sent: pings_sent.clone(),
};
tokio::spawn(timer.run());
}
tokio::spawn(async move {
let err = backend.run().await.err().unwrap_or(Error::Closed);
if let Some(code) = err.transport_close() {
let _ = backend.control.send(
ConnectionClose {
code: VarInt::from(code),
reason: err.to_string(),
}
.into(),
);
}
backend.open_bi_credit.close();
backend.open_uni_credit.close();
backend.conn_send_credit.close();
backend.conn_recv_credit.close();
backend.outbound.close();
for send in backend.streams.lock().unwrap().send.values() {
if let Some(credit) = &send.stream_credit {
credit.close();
}
}
backend.closed.send_replace(Some(err));
});
let guard = Arc::new(SessionGuard {
closed: closed.clone(),
});
Session {
is_server,
config,
outbound,
outbound_priority: control_tx,
accept_bi: Arc::new(tokio::sync::Mutex::new(accept_bi_rx)),
accept_uni: Arc::new(tokio::sync::Mutex::new(accept_uni_rx)),
streams,
closed,
negotiated,
established: established_rx,
open_bi_credit,
open_uni_credit,
conn_send_credit,
conn_recv_credit,
recv_datagram: Arc::new(tokio::sync::Mutex::new(recv_datagram_rx)),
datagram_max_size,
outbound_datagram: outbound_datagram_tx,
_guard: guard,
}
}
}
impl generic::Session for Session {
type SendStream = SendStream;
type RecvStream = RecvStream;
type Error = Error;
async fn accept_uni(&self) -> Result<Self::RecvStream, Self::Error> {
self.accept_uni
.lock()
.await
.recv()
.await
.ok_or(Error::Closed)
}
async fn accept_bi(&self) -> Result<(Self::SendStream, Self::RecvStream), Self::Error> {
self.accept_bi
.lock()
.await
.recv()
.await
.ok_or(Error::Closed)
}
async fn open_uni(&self) -> Result<Self::SendStream, Self::Error> {
let index = self.open_uni_credit.claim_index().await?;
let id = StreamId::new(index, StreamDir::Uni, self.is_server);
let (tx, rx) = mpsc::unbounded_channel();
let stream_credit = if self.config.version.is_qmux() {
Some(Credit::new(0)) } else {
None
};
let send_backend = SendState {
inbound_stopped: tx,
sent_offset: 0,
stream_credit: stream_credit.clone(),
};
let send_frontend = SendStream {
id,
outbound: self.outbound.clone(),
outbound_priority: self.outbound_priority.clone(),
inbound_stopped: rx,
offset: 0,
priority: 0,
closed: None,
fin: false,
stream_credit,
conn_credit: if self.config.version.is_qmux() {
Some(self.conn_send_credit.clone())
} else {
None
},
};
{
let mut streams = self.streams.lock().unwrap();
if let Some(credit) = &send_backend.stream_credit {
credit
.increase_max(streams.peer_initial_max_stream_data_uni)
.ok();
}
streams.send.insert(id, send_backend);
}
Ok(send_frontend)
}
async fn open_bi(&self) -> Result<(Self::SendStream, Self::RecvStream), Self::Error> {
let index = self.open_bi_credit.claim_index().await?;
let id = StreamId::new(index, StreamDir::Bi, self.is_server);
let (tx, rx) = mpsc::unbounded_channel();
let (tx2, rx2) = mpsc::unbounded_channel();
let stream_credit = if self.config.version.is_qmux() {
Some(Credit::new(0)) } else {
None
};
let send_backend = SendState {
inbound_stopped: tx,
sent_offset: 0,
stream_credit: stream_credit.clone(),
};
let send_frontend = SendStream {
id,
outbound: self.outbound.clone(),
outbound_priority: self.outbound_priority.clone(),
inbound_stopped: rx,
offset: 0,
priority: 0,
closed: None,
fin: false,
stream_credit,
conn_credit: if self.config.version.is_qmux() {
Some(self.conn_send_credit.clone())
} else {
None
},
};
let (tx, rx) = mpsc::unbounded_channel();
let recv_window = if self.config.version.is_qmux() {
self.config.max_stream_data_bidi_local
} else {
u64::MAX
};
let recv_credit = Credit::new(recv_window);
let recv_backend = RecvState {
inbound_data: tx,
inbound_reset: tx2,
recv_credit: recv_credit.clone(),
recv_offset: 0,
};
let recv_frontend = RecvStream {
id,
inbound_data: rx,
inbound_reset: rx2,
outbound_priority: self.outbound_priority.clone(),
buffer: Bytes::new(),
closed: None,
fin: false,
recv_credit,
conn_recv_credit: self.conn_recv_credit.clone(),
version: self.config.version,
recv_streams_credit: None, };
{
let mut streams = self.streams.lock().unwrap();
if let Some(credit) = &send_backend.stream_credit {
credit
.increase_max(streams.peer_initial_max_stream_data_bidi_remote)
.ok();
}
streams.send.insert(id, send_backend);
streams.recv.insert(id, recv_backend);
}
Ok((send_frontend, recv_frontend))
}
fn close(&self, code: u32, reason: &str) {
let frame = ApplicationClose {
code: VarInt::from(code),
reason: reason.to_string(),
};
let _ = self.outbound_priority.send(frame.into());
self.closed
.send(Some(Error::ConnectionClosed {
code: VarInt::from(code),
reason: reason.to_string(),
}))
.ok();
}
async fn closed(&self) -> Self::Error {
let mut closed = self.closed.subscribe();
closed
.wait_for(|err| err.is_some())
.await
.map(|e| e.clone().unwrap_or(Error::Closed))
.unwrap_or(Error::Closed)
}
fn send_datagram(&self, payload: Bytes) -> Result<(), Self::Error> {
let max = self.datagram_max_size.load(Ordering::Acquire);
if max == 0 {
return Err(Error::DatagramsUnsupported);
}
if payload.len() > max {
return Err(Error::FrameTooLarge);
}
match self.outbound_datagram.try_send(payload) {
Ok(()) => Ok(()),
Err(mpsc::error::TrySendError::Full(_)) => Ok(()),
Err(mpsc::error::TrySendError::Closed(_)) => Err(Error::Closed),
}
}
fn max_datagram_size(&self) -> usize {
self.datagram_max_size.load(Ordering::Acquire)
}
async fn recv_datagram(&self) -> Result<Bytes, Self::Error> {
self.recv_datagram
.lock()
.await
.recv()
.await
.ok_or(Error::Closed)
}
fn protocol(&self) -> Option<&str> {
self.negotiated.get().and_then(|p| p.as_deref())
}
}
fn negotiate_protocol(is_server: bool, ours: &[String], peers: &[String]) -> Option<String> {
let (server, client) = if is_server {
(ours, peers)
} else {
(peers, ours)
};
server.iter().find(|p| client.contains(p)).cloned()
}
struct SendState {
inbound_stopped: mpsc::UnboundedSender<StopSending>,
sent_offset: u64,
stream_credit: Option<Credit>,
}
pub struct SendStream {
id: StreamId,
outbound: PriorityQueue, outbound_priority: mpsc::UnboundedSender<Frame>, inbound_stopped: mpsc::UnboundedReceiver<StopSending>,
offset: u64,
priority: u8,
closed: Option<Error>,
fin: bool,
stream_credit: Option<Credit>,
conn_credit: Option<Credit>,
}
impl SendStream {
fn recv_stop(&mut self, code: VarInt) -> Error {
if let Some(error) = &self.closed {
return error.clone();
}
let error = Error::StreamStop(code);
if !self.fin {
let frame = ResetStream {
id: self.id,
code,
final_size: self.offset,
reliable_size: None,
};
let dropped = self.outbound.remove(self.id);
self.release_credit(dropped);
self.outbound_priority.send(frame.into()).ok();
}
self.closed = Some(error.clone());
error
}
fn release_credit(&self, amount: u64) {
if let Some(s) = &self.stream_credit {
s.release(amount);
}
if let Some(c) = &self.conn_credit {
c.release(amount);
}
}
async fn claim_credit(&mut self, desired: u64) -> Result<u64, Error> {
let (stream_credit, conn_credit) = match (&self.stream_credit, &self.conn_credit) {
(Some(s), Some(c)) => (s, c),
_ => return Ok(desired), };
loop {
let stream_claimed = stream_credit.try_claim(desired);
if stream_claimed == 0 {
tokio::select! {
result = stream_credit.claim(desired) => {
let claimed = result?;
stream_credit.release(claimed);
}
Some(stop) = self.inbound_stopped.recv() => {
return Err(self.recv_stop(stop.code));
}
}
continue;
}
let conn_claimed = conn_credit.try_claim(stream_claimed);
if conn_claimed == 0 {
stream_credit.release(stream_claimed);
tokio::select! {
result = conn_credit.claim(1) => {
let claimed = result?;
conn_credit.release(claimed); }
Some(stop) = self.inbound_stopped.recv() => {
return Err(self.recv_stop(stop.code));
}
}
continue;
}
if conn_claimed < stream_claimed {
stream_credit.release(stream_claimed - conn_claimed);
}
return Ok(conn_claimed);
}
}
}
impl Drop for SendStream {
fn drop(&mut self) {
if !self.fin && self.closed.is_none() {
generic::SendStream::reset(self, 0);
}
}
}
impl generic::SendStream for SendStream {
type Error = Error;
async fn write(&mut self, mut buf: &[u8]) -> Result<usize, Self::Error> {
let size = buf.len();
let b = &mut buf;
self.write_buf(b).await?;
Ok(size - b.len())
}
async fn write_buf<B: Buf + Send>(&mut self, buf: &mut B) -> Result<usize, Self::Error> {
if let Some(error) = &self.closed {
return Err(error.clone());
}
if self.fin {
return Err(Error::StreamClosed);
}
let mut total = 0;
while buf.has_remaining() {
let chunk_len = buf.chunk().len().min(MAX_FRAME_PAYLOAD) as u64;
let permit = tokio::select! {
result = self.outbound.reserve() => result?,
Some(stop) = self.inbound_stopped.recv() => return Err(self.recv_stop(stop.code)),
};
let allowed = self.claim_credit(chunk_len).await?;
let to_send = allowed as usize;
let frame = Stream {
id: self.id,
offset: self.offset,
data: buf.copy_to_bytes(to_send),
fin: false,
};
if let Err(err) = permit.send(self.priority, self.id, frame.into()) {
self.release_credit(to_send as u64);
return Err(err);
}
self.offset += to_send as u64;
total += to_send;
}
Ok(total)
}
fn set_priority(&mut self, order: u8) {
self.priority = order;
self.outbound.set_priority(self.id, order);
}
fn reset(&mut self, code: u32) {
if self.fin || self.closed.is_some() {
return;
}
let code = VarInt::from(code);
let frame = ResetStream {
id: self.id,
code,
final_size: self.offset,
reliable_size: None,
};
let dropped = self.outbound.remove(self.id);
self.release_credit(dropped);
self.outbound_priority.send(frame.into()).ok();
self.closed = Some(Error::StreamReset(code));
}
fn finish(&mut self) -> Result<(), Self::Error> {
if let Some(error) = &self.closed {
return Err(error.clone());
}
let frame = Stream {
id: self.id,
offset: self.offset,
data: Bytes::new(),
fin: true,
};
self.outbound
.push_now(self.priority, self.id, frame.into())?;
self.fin = true;
Ok(())
}
async fn closed(&mut self) -> Result<(), Self::Error> {
if let Some(error) = &self.closed {
return Err(error.clone());
}
match self.inbound_stopped.recv().await {
Some(stop) => Err(self.recv_stop(stop.code)),
None => Err(Error::Closed),
}
}
}
pub(crate) struct RecvState {
inbound_data: mpsc::UnboundedSender<Stream>,
inbound_reset: mpsc::UnboundedSender<ResetStream>,
recv_credit: Credit,
recv_offset: u64,
}
pub struct RecvStream {
id: StreamId,
version: Version,
outbound_priority: mpsc::UnboundedSender<Frame>, inbound_data: mpsc::UnboundedReceiver<Stream>,
inbound_reset: mpsc::UnboundedReceiver<ResetStream>,
buffer: Bytes,
closed: Option<Error>,
fin: bool,
recv_credit: Credit,
conn_recv_credit: Credit,
recv_streams_credit: Option<Credit>,
}
impl RecvStream {
fn recv_reset(&mut self, code: VarInt) -> Error {
if let Some(error) = &self.closed {
return error.clone();
}
self.closed = Some(Error::StreamReset(code));
Error::StreamReset(code)
}
fn report_consumed(&self, len: u64) {
if !self.version.is_qmux() {
return;
}
if let Some(new_max) = self.recv_credit.consume(len) {
let frame = Frame::MaxStreamData {
id: self.id,
max: new_max,
};
self.outbound_priority.send(frame).ok();
}
if let Some(new_max) = self.conn_recv_credit.consume(len) {
let frame = Frame::MaxData(new_max);
self.outbound_priority.send(frame).ok();
}
}
}
impl Drop for RecvStream {
fn drop(&mut self) {
if !self.fin && self.closed.is_none() {
generic::RecvStream::stop(self, 0);
}
if let Some(credit) = &self.recv_streams_credit {
if let Some(new_max) = credit.consume(1) {
let frame = match self.id.dir() {
StreamDir::Bi => Frame::MaxStreamsBidi(new_max),
StreamDir::Uni => Frame::MaxStreamsUni(new_max),
};
self.outbound_priority.send(frame).ok();
}
}
}
}
impl generic::RecvStream for RecvStream {
type Error = Error;
async fn read_chunk(&mut self, max: usize) -> Result<Option<Bytes>, Self::Error> {
loop {
if !self.buffer.is_empty() {
let to_read = max.min(self.buffer.len());
let data = self.buffer.split_to(to_read);
self.report_consumed(to_read as u64);
return Ok(Some(data));
}
if self.fin {
return Ok(None);
}
if let Some(error) = &self.closed {
return Err(error.clone());
}
tokio::select! {
Some(stream) = self.inbound_data.recv() => {
assert_eq!(stream.id, self.id);
self.fin = stream.fin;
self.buffer = stream.data;
}
Some(reset) = self.inbound_reset.recv() => {
return Err(self.recv_reset(reset.code));
}
else => return Err(Error::Closed),
}
}
}
async fn read_buf<B: BufMut + Send>(
&mut self,
buf: &mut B,
) -> Result<Option<usize>, Self::Error> {
if !self.buffer.is_empty() {
let to_read = buf.remaining_mut().min(self.buffer.len());
buf.put(self.buffer.split_to(to_read));
self.report_consumed(to_read as u64);
return Ok(Some(to_read));
}
Ok(match self.read_chunk(buf.remaining_mut()).await? {
Some(data) if !data.is_empty() => {
let size = data.len();
buf.put(data);
Some(size)
}
_ => None,
})
}
async fn read(&mut self, mut buf: &mut [u8]) -> Result<Option<usize>, Self::Error> {
self.read_buf(&mut buf).await
}
fn stop(&mut self, code: u32) {
let code = VarInt::from(code);
let frame = StopSending { id: self.id, code };
self.outbound_priority.send(frame.into()).ok();
self.closed = Some(Error::StreamStop(code));
}
async fn closed(&mut self) -> Result<(), Self::Error> {
if let Some(error) = &self.closed {
return Err(error.clone());
}
loop {
if self.fin {
return Ok(());
}
tokio::select! {
Some(reset) = self.inbound_reset.recv() => {
return Err(self.recv_reset(reset.code));
}
Some(stream) = self.inbound_data.recv() => {
assert_eq!(stream.id, self.id);
self.buffer = stream.data;
self.fin = stream.fin;
}
else => {
return Err(Error::Closed);
}
}
}
}
}
#[cfg(test)]
mod timer_tests {
use std::sync::{
atomic::{AtomicBool, AtomicU64, Ordering},
Arc,
};
use std::time::Duration;
use tokio::sync::{mpsc, watch};
use super::TimerState;
use crate::Error;
struct Harness {
reader_backpressured: Arc<AtomicBool>,
last_recv_at: Arc<AtomicU64>,
last_send_at: Arc<AtomicU64>,
closed: watch::Sender<Option<Error>>,
_control_rx: mpsc::UnboundedReceiver<crate::Frame>,
}
fn spawn_timer(idle_ms: u64) -> Harness {
let base = tokio::time::Instant::now();
let last_recv_at = Arc::new(AtomicU64::new(0));
let last_send_at = Arc::new(AtomicU64::new(0));
let reader_backpressured = Arc::new(AtomicBool::new(false));
let writer_backpressured = Arc::new(AtomicBool::new(false));
let idle_timeout_ms = Arc::new(AtomicU64::new(idle_ms));
let (control, _control_rx) = mpsc::unbounded_channel();
let closed = watch::Sender::new(None);
let (_est_tx, established) = watch::channel(true);
let timer = TimerState {
base,
last_recv_at: last_recv_at.clone(),
last_send_at: last_send_at.clone(),
reader_backpressured: reader_backpressured.clone(),
writer_backpressured: writer_backpressured.clone(),
idle_timeout_ms,
control,
closed: closed.clone(),
established,
pings_sent: Arc::new(AtomicU64::new(0)),
};
tokio::spawn(timer.run());
Harness {
reader_backpressured,
last_recv_at,
last_send_at,
closed,
_control_rx,
}
}
async fn closed_reason(h: &Harness) -> Error {
let mut rx = h.closed.subscribe();
rx.wait_for(|s| s.is_some()).await.unwrap();
let reason = rx.borrow().clone().unwrap();
reason
}
#[tokio::test]
async fn idle_close_when_silent() {
let h = spawn_timer(100);
let reason = tokio::time::timeout(Duration::from_millis(400), closed_reason(&h))
.await
.expect("silent session must idle-close");
assert!(matches!(reason, Error::IdleTimeout), "got {reason:?}");
}
#[tokio::test]
async fn idle_close_deferred_while_reader_backpressured() {
let h = spawn_timer(100);
h.reader_backpressured.store(true, Ordering::Release);
tokio::time::sleep(Duration::from_millis(150)).await;
assert!(
h.closed.borrow().is_none(),
"idle-close must be deferred while the reader is backpressured"
);
let reason = tokio::time::timeout(Duration::from_millis(400), closed_reason(&h))
.await
.expect("bounded deferral must eventually idle-close");
assert!(matches!(reason, Error::IdleTimeout), "got {reason:?}");
}
#[tokio::test]
async fn receive_progress_averts_idle_close() {
let h = spawn_timer(100);
let base = tokio::time::Instant::now();
for _ in 0..6 {
tokio::time::sleep(Duration::from_millis(50)).await;
let elapsed = base.elapsed().as_millis() as u64;
h.last_recv_at.store(elapsed, Ordering::Release);
}
assert!(
h.closed.borrow().is_none(),
"a peer that keeps sending must not be idle-closed"
);
}
#[tokio::test]
async fn send_progress_averts_idle_close() {
let h = spawn_timer(100);
let base = tokio::time::Instant::now();
for _ in 0..6 {
tokio::time::sleep(Duration::from_millis(50)).await;
let elapsed = base.elapsed().as_millis() as u64;
h.last_send_at.store(elapsed, Ordering::Release);
}
assert!(
h.closed.borrow().is_none(),
"a session that keeps sending must not be idle-closed"
);
}
}
#[cfg(test)]
mod negotiate_tests {
use super::negotiate_protocol;
fn v(items: &[&str]) -> Vec<String> {
items.iter().map(|s| s.to_string()).collect()
}
#[test]
fn server_preference_wins() {
let server = v(&["b", "a"]);
let client = v(&["a", "b"]);
assert_eq!(
negotiate_protocol(true, &server, &client).as_deref(),
Some("b")
);
assert_eq!(
negotiate_protocol(false, &client, &server).as_deref(),
Some("b")
);
}
#[test]
fn no_overlap_is_none() {
assert_eq!(negotiate_protocol(true, &v(&["a"]), &v(&["b"])), None);
assert_eq!(negotiate_protocol(true, &v(&["a"]), &[]), None);
}
}
#[cfg(test)]
mod send_offset_tests {
use bytes::Bytes;
use tokio::sync::mpsc;
use web_transport_trait::SendStream as _;
use super::SendStream;
use crate::sched::PriorityQueue;
use crate::{Frame, StreamDir, StreamId};
#[tokio::test]
async fn sequential_writes_and_fin_carry_send_offsets() {
let id = StreamId::new(0, StreamDir::Uni, false);
let outbound = PriorityQueue::new(3);
let (control, _control_rx) = mpsc::unbounded_channel();
let (_stop_tx, stop_rx) = mpsc::unbounded_channel();
let mut send = SendStream {
id,
outbound: outbound.clone(),
outbound_priority: control,
inbound_stopped: stop_rx,
offset: 0,
priority: 0,
closed: None,
fin: false,
stream_credit: None,
conn_credit: None,
};
send.write(&[1, 2, 3]).await.unwrap();
send.write(&[4, 5]).await.unwrap();
send.finish().unwrap();
let expected: &[(u64, &[u8], bool)] =
&[(0, &[1, 2, 3], false), (3, &[4, 5], false), (5, &[], true)];
for &(offset, data, fin) in expected {
let frame = outbound.pop().await.expect("queued STREAM frame");
match frame {
Frame::Stream(stream) => {
assert_eq!(stream.offset, offset);
assert_eq!(stream.data, Bytes::copy_from_slice(data));
assert_eq!(stream.fin, fin);
}
other => panic!("expected STREAM, got {other:?}"),
}
}
}
}
#[cfg(test)]
mod write_cancel_tests {
use bytes::{Buf, Bytes};
use tokio::sync::mpsc;
use web_transport_trait::SendStream as _;
use super::SendStream;
use crate::credit::Credit;
use crate::sched::PriorityQueue;
use crate::{Frame, StreamDir, StreamId};
fn send_stream(
outbound: PriorityQueue,
stream_credit: Option<Credit>,
conn_credit: Option<Credit>,
) -> SendStream {
let (control, _control_rx) = mpsc::unbounded_channel();
let (_stop_tx, stop_rx) = mpsc::unbounded_channel();
SendStream {
id: StreamId::new(0, StreamDir::Uni, false),
outbound,
outbound_priority: control,
inbound_stopped: stop_rx,
offset: 0,
priority: 0,
closed: None,
fin: false,
stream_credit,
conn_credit,
}
}
async fn cancel_writes(send: &mut SendStream, buf: &mut Bytes) -> usize {
let start = buf.remaining();
for _ in 0..50 {
tokio::select! {
result = send.write_buf(buf) => { result.expect("write_buf"); }
_ = tokio::task::yield_now() => {}
}
}
start - buf.remaining()
}
async fn drain(outbound: &PriorityQueue) -> usize {
outbound.close();
let mut queued = 0;
while let Some(frame) = outbound.pop().await {
if let Frame::Stream(stream) = frame {
queued += stream.data.len();
}
}
queued
}
#[tokio::test]
async fn cancelled_write_blocked_on_queue_consumes_nothing() {
let outbound = PriorityQueue::new(1);
let mut filler = send_stream(outbound.clone(), None, None);
filler.write(&[0xFF]).await.unwrap();
let mut send = send_stream(outbound.clone(), None, None);
let mut buf = Bytes::from(vec![0xAB_u8; 1024 * 1024]);
let consumed = cancel_writes(&mut send, &mut buf).await;
assert_eq!(
consumed, 0,
"write_buf never got a queue slot but still took {consumed} bytes \
from the buffer",
);
}
#[tokio::test]
async fn cancelled_write_blocked_on_queue_spends_no_credit() {
const WINDOW: u64 = 64 * 1024;
let outbound = PriorityQueue::new(1);
let mut filler = send_stream(outbound.clone(), None, None);
filler.write(&[0xFF]).await.unwrap();
let stream_credit = Credit::new(WINDOW);
let conn_credit = Credit::new(WINDOW);
let mut send = send_stream(
outbound,
Some(stream_credit.clone()),
Some(conn_credit.clone()),
);
let mut buf = Bytes::from(vec![0xAB_u8; 1024 * 1024]);
cancel_writes(&mut send, &mut buf).await;
assert_eq!(
stream_credit.try_claim(WINDOW),
WINDOW,
"cancelled writes leaked stream credit",
);
assert_eq!(
conn_credit.try_claim(WINDOW),
WINDOW,
"cancelled writes leaked connection credit",
);
}
#[tokio::test]
async fn cancelled_write_blocked_on_credit_queues_what_it_consumes() {
let outbound = PriorityQueue::new(64);
let stream_credit = Credit::new(32 * 1024);
let conn_credit = Credit::new(64 * 1024 * 1024);
let mut send = send_stream(outbound.clone(), Some(stream_credit), Some(conn_credit));
let mut buf = Bytes::from(vec![0xAB_u8; 1024 * 1024]);
let consumed = cancel_writes(&mut send, &mut buf).await;
let queued = drain(&outbound).await;
assert_eq!(
consumed,
queued,
"write_buf took {consumed} bytes from the buffer but only {queued} \
reached the queue — {} bytes silently dropped by cancellation",
consumed.saturating_sub(queued),
);
assert!(
consumed > 0,
"expected the window's worth of data to get through"
);
}
}
#[cfg(test)]
mod recv_open_tests {
use std::time::Duration;
use bytes::Bytes;
use tokio::sync::mpsc;
use web_transport_trait::{RecvStream as _, Session as _};
use web_transport_proto::VarInt;
use super::{Reader, Session, Transport, Writer};
use crate::proto::{Frame, ResetStream, Stream};
use crate::{Config, Error, StreamDir, StreamId, Version};
struct ScriptedTransport {
incoming: mpsc::UnboundedReceiver<Bytes>,
}
struct ScriptedWriter;
struct ScriptedReader {
incoming: mpsc::UnboundedReceiver<Bytes>,
}
impl Transport for ScriptedTransport {
type Writer = ScriptedWriter;
type Reader = ScriptedReader;
fn split(self) -> (ScriptedWriter, ScriptedReader) {
(
ScriptedWriter,
ScriptedReader {
incoming: self.incoming,
},
)
}
}
impl Writer for ScriptedWriter {
async fn send(&mut self, _data: Bytes) -> Result<(), Error> {
Ok(())
}
async fn close(&mut self) -> Result<(), Error> {
Ok(())
}
}
impl Reader for ScriptedReader {
async fn recv(&mut self) -> Result<Bytes, Error> {
match self.incoming.recv().await {
Some(bytes) => Ok(bytes),
None => std::future::pending().await,
}
}
}
fn scripted_session_with_config(config: Config) -> (Session, mpsc::UnboundedSender<Bytes>) {
let (tx, rx) = mpsc::unbounded_channel();
let session = Session::new(ScriptedTransport { incoming: rx }, false, config);
(session, tx)
}
fn scripted_session() -> (Session, mpsc::UnboundedSender<Bytes>) {
scripted_session_with_config(Config::new(Version::QMux01))
}
fn uni_stream(index: u64, data: &'static [u8], fin: bool) -> Bytes {
uni_stream_at(index, 0, data, fin)
}
fn uni_stream_at(index: u64, offset: u64, data: &'static [u8], fin: bool) -> Bytes {
Frame::Stream(Stream {
id: StreamId::new(index, StreamDir::Uni, true),
offset,
data: Bytes::from_static(data),
fin,
})
.encode(Version::QMux01)
.unwrap()
}
fn uni_reset(index: u64, final_size: u64) -> Bytes {
Frame::ResetStream(ResetStream {
id: StreamId::new(index, StreamDir::Uni, true),
code: VarInt::from_u32(0),
final_size,
reliable_size: None,
})
.encode(Version::QMux01)
.unwrap()
}
#[tokio::test]
async fn retired_recv_stream_is_not_resurrected() {
let (session, tx) = scripted_session();
tx.send(uni_stream(0, b"hello", false)).unwrap();
tx.send(uni_stream(0, b"", true)).unwrap();
tx.send(uni_stream(0, b"late", false)).unwrap();
let mut recv = tokio::time::timeout(Duration::from_secs(1), session.accept_uni())
.await
.expect("accept_uni timed out")
.expect("accept_uni failed");
assert_eq!(recv.read_all().await.unwrap().as_ref(), b"hello");
let second = tokio::time::timeout(Duration::from_millis(200), session.accept_uni()).await;
assert!(
second.is_err(),
"a late frame on a retired stream resurrected a new accepted stream"
);
}
#[tokio::test]
async fn recv_stream_offsets_are_not_enforced_yet() {
let (session, tx) = scripted_session();
tx.send(uni_stream_at(0, 0, b"hello", false)).unwrap();
tx.send(uni_stream_at(0, 99, b"world", true)).unwrap();
let mut recv = tokio::time::timeout(Duration::from_secs(1), session.accept_uni())
.await
.expect("accept_uni timed out")
.expect("accept_uni failed");
assert_eq!(recv.read_all().await.unwrap().as_ref(), b"helloworld");
}
#[tokio::test]
async fn implicitly_opened_lower_stream_is_still_accepted() {
let (session, tx) = scripted_session();
tx.send(uni_stream(10, b"", true)).unwrap();
tx.send(uni_stream(6, b"hello", true)).unwrap();
let mut first = tokio::time::timeout(Duration::from_secs(1), session.accept_uni())
.await
.expect("accept_uni timed out")
.expect("accept_uni failed");
assert_eq!(first.read_all().await.unwrap().as_ref(), b"");
let mut second = tokio::time::timeout(Duration::from_secs(1), session.accept_uni())
.await
.expect("stream 6 was wrongly dropped as already-closed")
.expect("accept_uni failed");
assert_eq!(second.read_all().await.unwrap().as_ref(), b"hello");
}
#[tokio::test]
async fn reset_as_first_frame_prevents_resurrection() {
let (session, tx) = scripted_session();
tx.send(uni_reset(5, 0)).unwrap();
tx.send(uni_stream(5, b"late", false)).unwrap();
let accepted = tokio::time::timeout(Duration::from_millis(200), session.accept_uni()).await;
assert!(
accepted.is_err(),
"a STREAM after a RESET-first stream resurrected a new accepted stream"
);
}
#[tokio::test]
async fn empty_non_fin_stream_frames_are_not_queued() {
const FLOOD: usize = 10_000;
let (session, tx) = scripted_session();
tx.send(uni_stream(0, b"", false)).unwrap();
let recv = tokio::time::timeout(Duration::from_secs(1), session.accept_uni())
.await
.expect("accept_uni timed out")
.expect("accept_uni failed");
for _ in 0..FLOOD {
tx.send(uni_stream(0, b"", false)).unwrap();
}
tx.send(uni_stream(1, b"", true)).unwrap();
let _marker = tokio::time::timeout(Duration::from_secs(1), session.accept_uni())
.await
.expect("marker accept_uni timed out")
.expect("marker accept_uni failed");
assert_eq!(
recv.inbound_data.len(),
0,
"empty non-FIN STREAM frames accumulated behind a stalled reader"
);
}
#[tokio::test]
async fn reset_final_size_consumes_connection_credit() {
let mut config = Config::new(Version::QMux01);
config.max_data = 4;
config.max_stream_data_uni = 10;
let (session, tx) = scripted_session_with_config(config);
tx.send(uni_reset(0, 5)).unwrap();
let err = tokio::time::timeout(Duration::from_secs(1), session.closed())
.await
.expect("session did not close on RESET_STREAM flow-control violation");
assert!(matches!(err, Error::FlowControlError), "got {err:?}");
}
#[tokio::test]
async fn reset_final_size_is_cumulative_with_later_stream_data() {
let mut config = Config::new(Version::QMux01);
config.max_data = 10;
config.max_stream_data_uni = 10;
let (session, tx) = scripted_session_with_config(config);
tx.send(uni_reset(0, 3)).unwrap();
tx.send(uni_stream(1, b"12345678", false)).unwrap();
let err = tokio::time::timeout(Duration::from_secs(1), session.closed())
.await
.expect("session did not close after cumulative MAX_DATA exhaustion");
assert!(matches!(err, Error::FlowControlError), "got {err:?}");
}
#[tokio::test]
async fn reset_final_size_replenishes_connection_credit() {
let mut config = Config::new(Version::QMux01);
config.max_data = 10;
config.max_stream_data_uni = 10;
let (session, tx) = scripted_session_with_config(config);
tx.send(uni_reset(0, 6)).unwrap();
tx.send(uni_stream(1, b"12345", true)).unwrap();
let mut recv = tokio::time::timeout(Duration::from_secs(1), session.accept_uni())
.await
.expect("connection credit was not replenished after reset")
.expect("accept_uni failed");
assert_eq!(recv.read_all().await.unwrap().as_ref(), b"12345");
}
#[tokio::test]
async fn legacy_zero_final_size_after_data_is_tolerated() {
let (session, tx) = scripted_session();
tx.send(uni_stream(0, b"hello", false)).unwrap();
let mut recv = tokio::time::timeout(Duration::from_secs(1), session.accept_uni())
.await
.expect("accept_uni timed out")
.expect("accept_uni failed");
tx.send(uni_reset(0, 0)).unwrap();
let err = tokio::time::timeout(Duration::from_secs(1), recv.closed())
.await
.expect("reset was not delivered")
.expect_err("reset should close the receive stream");
assert!(matches!(err, Error::StreamReset(_)), "got {err:?}");
tx.send(uni_stream(1, b"ok", true)).unwrap();
let mut next = tokio::time::timeout(Duration::from_secs(1), session.accept_uni())
.await
.expect("connection closed after legacy reset")
.expect("accept_uni failed");
assert_eq!(next.read_all().await.unwrap().as_ref(), b"ok");
}
}
#[cfg(all(test, feature = "tcp"))]
mod datagram_recv_tests {
use super::*;
use crate::transport::Stream;
use tokio::io::{AsyncWriteExt, DuplexStream};
use web_transport_trait::Session as _;
fn record(frame: Bytes) -> Bytes {
let mut buf = bytes::BytesMut::new();
VarInt::try_from(frame.len()).unwrap().encode(&mut buf);
buf.extend_from_slice(&frame);
buf.freeze()
}
async fn raw_peer(server_cfg: Config) -> (Session, DuplexStream) {
let (server_io, mut raw) = tokio::io::duplex(1024 * 1024);
let transport = Stream::new(server_io, Version::QMux01, server_cfg.max_record_size);
let accept = tokio::spawn(Session::accept(transport, server_cfg));
let client_params = Config::new(Version::QMux01).to_transport_params();
let params = Frame::TransportParameters(client_params)
.encode(Version::QMux01)
.unwrap();
raw.write_all(&record(params)).await.unwrap();
raw.flush().await.unwrap();
let server = accept.await.unwrap().unwrap();
(server, raw)
}
#[tokio::test]
async fn oversized_frame_closes_session() {
let mut cfg = Config::new(Version::QMux01);
cfg.max_datagram_frame_size = 100;
let (server, mut raw) = raw_peer(cfg).await;
let datagram = Frame::Datagram(Bytes::from(vec![0u8; 98]).into())
.encode(Version::QMux01)
.unwrap();
raw.write_all(&record(datagram)).await.unwrap();
raw.flush().await.unwrap();
assert!(matches!(server.closed().await, Error::FrameTooLarge));
}
#[tokio::test]
async fn unnegotiated_datagram_closes_session() {
let mut cfg = Config::new(Version::QMux01);
cfg.max_datagram_frame_size = 0;
let (server, mut raw) = raw_peer(cfg).await;
let datagram = Frame::Datagram(Bytes::from_static(b"hi").into())
.encode(Version::QMux01)
.unwrap();
raw.write_all(&record(datagram)).await.unwrap();
raw.flush().await.unwrap();
assert!(matches!(server.closed().await, Error::DatagramsUnsupported));
}
fn find_connection_close(buf: &[u8]) -> Option<ConnectionClose> {
let mut data = Bytes::copy_from_slice(buf);
while !data.is_empty() {
let len = VarInt::decode(&mut data).ok()?.into_inner() as usize;
if data.len() < len {
return None;
}
let record = data.split_to(len);
for frame in Frame::decode_record(record).ok()? {
if let Frame::ConnectionClose(c) = frame {
return Some(c);
}
}
}
None
}
#[tokio::test]
async fn violation_emits_connection_close_to_peer() {
use tokio::io::AsyncReadExt;
let mut cfg = Config::new(Version::QMux01);
cfg.max_datagram_frame_size = 100;
let (server, mut raw) = raw_peer(cfg).await;
let datagram = Frame::Datagram(Bytes::from(vec![0u8; 98]).into())
.encode(Version::QMux01)
.unwrap();
raw.write_all(&record(datagram)).await.unwrap();
raw.flush().await.unwrap();
assert!(matches!(server.closed().await, Error::FrameTooLarge));
let mut buf = Vec::new();
tokio::time::timeout(std::time::Duration::from_secs(1), raw.read_to_end(&mut buf))
.await
.expect("reading the server's output timed out")
.unwrap();
let close = find_connection_close(&buf)
.expect("a violation must emit a transport CONNECTION_CLOSE (0x1c)");
assert_eq!(close.code.into_inner(), 1002, "protocol-violation code");
}
#[tokio::test]
async fn peer_application_close_is_graceful() {
let (server, mut raw) = raw_peer(Config::new(Version::QMux01)).await;
let close = Frame::ApplicationClose(ApplicationClose {
code: VarInt::from_u32(42),
reason: "bye".to_string(),
})
.encode(Version::QMux01)
.unwrap();
raw.write_all(&record(close)).await.unwrap();
raw.flush().await.unwrap();
match server.closed().await {
Error::ConnectionClosed { code, reason } => {
assert_eq!(code.into_inner(), 42);
assert_eq!(reason, "bye");
}
other => panic!("expected a graceful ConnectionClosed, got {other:?}"),
}
}
#[tokio::test]
async fn peer_connection_close_is_abnormal() {
use web_transport_trait::Error as _;
let (server, mut raw) = raw_peer(Config::new(Version::QMux01)).await;
let close = Frame::ConnectionClose(ConnectionClose {
code: VarInt::from_u32(1002),
reason: "protocol violation".to_string(),
})
.encode(Version::QMux01)
.unwrap();
raw.write_all(&record(close)).await.unwrap();
raw.flush().await.unwrap();
let err = server.closed().await;
assert!(
matches!(err, Error::ConnectionReset { .. }),
"a peer CONNECTION_CLOSE must be abnormal, got {err:?}"
);
assert!(err.session_error().is_none());
}
#[tokio::test]
async fn frame_at_limit_delivered() {
let mut cfg = Config::new(Version::QMux01);
cfg.max_datagram_frame_size = 100;
let (server, mut raw) = raw_peer(cfg).await;
let payload = vec![7u8; 97];
let datagram = Frame::Datagram(Bytes::from(payload.clone()).into())
.encode(Version::QMux01)
.unwrap();
raw.write_all(&record(datagram)).await.unwrap();
raw.flush().await.unwrap();
assert_eq!(server.recv_datagram().await.unwrap().as_ref(), &payload[..]);
}
#[tokio::test]
async fn no_length_datagram_uses_exact_frame_size() {
let mut cfg = Config::new(Version::QMux01);
cfg.max_datagram_frame_size = 100;
let (server, mut raw) = raw_peer(cfg).await;
let payload = vec![3u8; 99];
let mut frame = bytes::BytesMut::new();
frame.put_u8(0x30);
frame.extend_from_slice(&payload);
raw.write_all(&record(frame.freeze())).await.unwrap();
raw.flush().await.unwrap();
assert_eq!(server.recv_datagram().await.unwrap().as_ref(), &payload[..]);
}
}
#[cfg(all(test, feature = "tcp"))]
mod qmux02_recv_tests {
use super::*;
use crate::transport::Stream as ByteStream;
use tokio::io::{AsyncWriteExt, DuplexStream};
use web_transport_trait::Session as _;
fn record(frame: &[u8]) -> Bytes {
let mut buf = bytes::BytesMut::new();
VarInt::try_from(frame.len()).unwrap().encode(&mut buf);
buf.extend_from_slice(frame);
buf.freeze()
}
fn qmux02_params() -> Bytes {
Frame::TransportParameters(Config::new(Version::QMux02).to_transport_params())
.encode(Version::QMux02)
.unwrap()
}
fn raw_accept(
server_cfg: Config,
) -> (
tokio::task::JoinHandle<Result<Session, Error>>,
DuplexStream,
) {
let (server_io, raw) = tokio::io::duplex(1024 * 1024);
let transport = ByteStream::new(server_io, Version::QMux02, server_cfg.max_record_size);
(tokio::spawn(Session::accept(transport, server_cfg)), raw)
}
async fn established_peer() -> (Session, DuplexStream) {
let (accept, mut raw) = raw_accept(Config::new(Version::QMux02));
raw.write_all(&record(&qmux02_params())).await.unwrap();
raw.flush().await.unwrap();
(accept.await.unwrap().unwrap(), raw)
}
#[tokio::test]
async fn transport_parameters_must_be_first() {
let (accept, mut raw) = raw_accept(Config::new(Version::QMux02));
let stream = Frame::Stream(Stream {
id: StreamId::new(0, StreamDir::Uni, false),
offset: 0,
data: Bytes::from_static(b"hi"),
fin: false,
})
.encode(Version::QMux02)
.unwrap();
raw.write_all(&record(&stream)).await.unwrap();
raw.flush().await.unwrap();
assert!(matches!(
accept.await.unwrap(),
Err(Error::ProtocolViolation)
));
}
#[tokio::test]
async fn duplicate_transport_parameters_rejected() {
let (server, mut raw) = established_peer().await;
raw.write_all(&record(&qmux02_params())).await.unwrap();
raw.flush().await.unwrap();
assert!(matches!(server.closed().await, Error::ProtocolViolation));
}
#[tokio::test]
async fn max_record_size_below_default_rejected() {
for version in [Version::QMux01, Version::QMux02] {
let (accept, mut raw) = raw_accept(Config::new(version));
let mut params = Config::new(version).to_transport_params();
params.max_record_size = 100; let frame = Frame::TransportParameters(params).encode(version).unwrap();
raw.write_all(&record(&frame)).await.unwrap();
raw.flush().await.unwrap();
assert!(matches!(
accept.await.unwrap(),
Err(Error::TransportParameter)
));
}
}
#[tokio::test]
async fn ping_response_for_unsent_sequence_closes() {
let (server, mut raw) = established_peer().await;
let ping = Frame::Ping(crate::Ping {
sequence: 5,
response: true,
})
.encode(Version::QMux02)
.unwrap();
raw.write_all(&record(&ping)).await.unwrap();
raw.flush().await.unwrap();
assert!(matches!(server.closed().await, Error::ProtocolViolation));
}
#[tokio::test]
async fn ping_request_not_increasing_closes() {
let (server, mut raw) = established_peer().await;
for _ in 0..2 {
let ping = Frame::Ping(crate::Ping {
sequence: 5, response: false,
})
.encode(Version::QMux02)
.unwrap();
raw.write_all(&record(&ping)).await.unwrap();
}
raw.flush().await.unwrap();
assert!(matches!(server.closed().await, Error::ProtocolViolation));
}
#[tokio::test]
async fn reset_stream_at_resets_accepted_stream() {
use web_transport_trait::RecvStream as _;
let (server, mut raw) = established_peer().await;
let id = StreamId::new(0, StreamDir::Uni, false);
let stream = Frame::Stream(Stream {
id,
offset: 0,
data: Bytes::from_static(b"hi"),
fin: false,
})
.encode(Version::QMux02)
.unwrap();
raw.write_all(&record(&stream)).await.unwrap();
let mut reset_at = bytes::BytesMut::new();
reset_at.put_u8(0x24);
id.0.encode(&mut reset_at);
VarInt::from_u32(0).encode(&mut reset_at); VarInt::from_u32(2).encode(&mut reset_at); VarInt::from_u32(0).encode(&mut reset_at); raw.write_all(&record(&reset_at)).await.unwrap();
raw.flush().await.unwrap();
let mut recv = server.accept_uni().await.unwrap();
loop {
match recv.read_chunk(64).await {
Ok(Some(data)) => assert_eq!(data.as_ref(), b"hi"),
Ok(None) => panic!("stream finished cleanly instead of resetting"),
Err(err) => {
assert!(matches!(err, Error::StreamReset(_)), "got {err:?}");
break;
}
}
}
}
#[tokio::test]
async fn reset_stream_at_without_negotiation_rejected() {
let server_cfg = Config::new(Version::QMux01);
let (server_io, mut raw) = tokio::io::duplex(1024 * 1024);
let transport = ByteStream::new(server_io, Version::QMux01, server_cfg.max_record_size);
let accept = tokio::spawn(Session::accept(transport, server_cfg));
let params = Frame::TransportParameters(Config::new(Version::QMux01).to_transport_params())
.encode(Version::QMux01)
.unwrap();
raw.write_all(&record(¶ms)).await.unwrap();
raw.flush().await.unwrap();
let server = accept.await.unwrap().unwrap();
let id = StreamId::new(0, StreamDir::Uni, false);
let mut reset_at = bytes::BytesMut::new();
reset_at.put_u8(0x24);
id.0.encode(&mut reset_at);
VarInt::from_u32(0).encode(&mut reset_at); VarInt::from_u32(0).encode(&mut reset_at); VarInt::from_u32(0).encode(&mut reset_at); raw.write_all(&record(&reset_at)).await.unwrap();
raw.flush().await.unwrap();
assert!(matches!(server.closed().await, Error::ProtocolViolation));
}
}
#[cfg(test)]
mod teardown_tests {
use std::time::Duration;
use bytes::Bytes;
use tokio::sync::mpsc;
use super::{Reader, Session, Transport, Writer};
use crate::{Config, Error, Version};
struct WedgedTransport {
entered_send: mpsc::UnboundedSender<()>,
dropped: mpsc::UnboundedSender<()>,
}
struct WedgedWriter {
entered_send: mpsc::UnboundedSender<()>,
_dropped: DropSignal,
}
struct DropSignal(mpsc::UnboundedSender<()>);
impl Drop for DropSignal {
fn drop(&mut self) {
let _ = self.0.send(());
}
}
struct WedgedReader;
impl Transport for WedgedTransport {
type Writer = WedgedWriter;
type Reader = WedgedReader;
fn split(self) -> (WedgedWriter, WedgedReader) {
(
WedgedWriter {
entered_send: self.entered_send,
_dropped: DropSignal(self.dropped),
},
WedgedReader,
)
}
}
impl Writer for WedgedWriter {
async fn send(&mut self, _data: Bytes) -> Result<(), Error> {
let _ = self.entered_send.send(());
std::future::pending().await
}
async fn close(&mut self) -> Result<(), Error> {
Ok(())
}
}
impl Reader for WedgedReader {
async fn recv(&mut self) -> Result<Bytes, Error> {
std::future::pending().await
}
}
#[tokio::test]
async fn wedged_writer_tears_down_on_last_drop() {
let (entered_tx, mut entered_rx) = mpsc::unbounded_channel();
let (dropped_tx, mut dropped_rx) = mpsc::unbounded_channel();
let session = Session::new(
WedgedTransport {
entered_send: entered_tx,
dropped: dropped_tx,
},
false,
Config::new(Version::QMux01),
);
tokio::time::timeout(Duration::from_secs(1), entered_rx.recv())
.await
.expect("writer never entered send()")
.expect("entered_send channel closed unexpectedly");
drop(session);
tokio::time::timeout(Duration::from_secs(1), dropped_rx.recv())
.await
.expect("writer task did not tear down while wedged in send()")
.expect("dropped channel closed unexpectedly");
}
}