use std::cell::RefCell;
use std::collections::VecDeque;
use std::rc::Rc;
use std::task::{Context, Poll, ready};
use bytes::{Buf, Bytes, BytesMut};
use web_transport_proto as proto;
use super::{Connection, Error};
const FRAME_WEBTRANSPORT: u64 = 0x41;
const STREAM_WEBTRANSPORT: u64 = 0x54;
const STREAM_CONTROL: u64 = 0x00;
const STREAM_QPACK_ENCODER: u64 = 0x02;
const STREAM_QPACK_DECODER: u64 = 0x03;
const FRAME_DATA: u64 = 0x00;
const CLOSE_GRACE: std::time::Duration = std::time::Duration::from_secs(1);
const MESSAGE_LIMIT: usize = 64 * 1024;
const HANDSHAKE_STREAMS: usize = 64;
const HELD_STREAMS: usize = 8;
const H3_GENERAL_PROTOCOL_ERROR: u64 = 0x0101;
const H3_NO_ERROR: u64 = 0x0100;
pub struct Request {
conn: Connection,
guard: Guard,
request: proto::ConnectRequest,
send: super::SendStream,
recv: super::RecvStream,
held: Vec<super::RecvStream>,
control: super::SendStream,
early: Vec<(u64, super::RecvStream)>,
pending: Vec<PendingUni>,
}
struct Guard {
conn: Option<Connection>,
}
impl Guard {
fn new(conn: Connection) -> Self {
Self { conn: Some(conn) }
}
fn disarm(&mut self) {
self.conn = None;
}
fn close(&mut self, reason: &str) {
if let Some(conn) = self.conn.take() {
conn.close_code(H3_GENERAL_PROTOCOL_ERROR, reason);
}
}
}
impl Drop for Guard {
fn drop(&mut self) {
self.close("webtransport handshake abandoned");
}
}
impl Request {
pub async fn accept(conn: Connection) -> Result<Self, Error> {
let mut guard = Guard::new(conn.clone());
match Self::handshake(conn).await {
Ok(request) => {
guard.disarm();
Ok(request)
}
Err(err) => {
guard.close(&err.to_string());
Err(err)
}
}
}
async fn handshake(mut conn: Connection) -> Result<Self, Error> {
let mut control = open_uni(&mut conn).await?;
let mut settings = proto::Settings::default();
settings.enable_webtransport(1);
let mut buf = Vec::new();
settings.encode(&mut buf);
write_all(&mut control, &buf).await?;
let mut pending: Vec<PendingUni> = Vec::new();
let mut held = Vec::new();
let mut early = Vec::new();
let mut arrivals = 0usize;
let mut peer_control = None;
std::future::poll_fn(|cx| {
loop {
let mut over = false;
while !over {
match web_transport_trait::poll::Session::poll_accept_uni(&mut conn, cx) {
Poll::Ready(Ok(recv)) => {
arrivals += 1;
pending.push(PendingUni::new(recv));
over = arrivals > HANDSHAKE_STREAMS;
}
Poll::Ready(Err(err)) => return Poll::Ready(Err(err)),
Poll::Pending => break,
}
}
let mut progressed = false;
let mut index = 0;
while index < pending.len() {
let Poll::Ready(class) = pending[index].poll_classify(cx) else {
index += 1;
continue;
};
let stream = pending.swap_remove(index);
progressed = true;
match class {
UniClass::Control => {
peer_control = Some(stream.recv);
return Poll::Ready(Ok(()));
}
UniClass::Qpack => held.push(stream.recv),
UniClass::Web(session) => early.push((session, stream.recv)),
UniClass::Unknown => {}
}
}
if over {
return Poll::Ready(Err(Error::Web("too many streams before the control stream".into())));
}
if !progressed {
return Poll::Pending;
}
}
})
.await?;
let mut peer_control = peer_control.expect("the loop only ends with a control stream");
let settings = read_settings(&mut peer_control).await?;
if settings.supports_webtransport() == 0 {
return Err(Error::Web("peer does not support WebTransport".into()));
}
held.push(peer_control);
let (send, mut recv) = accept_bi(&mut conn).await?;
let request = read_connect(&mut recv).await?;
Ok(Self {
guard: Guard::new(conn.clone()),
conn,
request,
send,
recv,
held,
control,
early,
pending,
})
}
pub fn url(&self) -> &url::Url {
&self.request.url
}
pub fn protocols(&self) -> &[String] {
&self.request.protocols
}
pub async fn respond(mut self, response: Response) -> Result<Session, Error> {
let Response { protocol } = response;
let mut encoded = proto::ConnectResponse::OK;
if let Some(protocol) = &protocol {
if !self.request.protocols.iter().any(|offered| offered == protocol) {
return Err(Error::Web(format!("subprotocol {protocol:?} was not offered")));
}
encoded = encoded.with_protocol(protocol);
}
let mut buf = Vec::new();
encoded.encode(&mut buf).map_err(|err| Error::Web(err.to_string()))?;
write_all(&mut self.send, &buf).await?;
self.guard.disarm();
Ok(Session::establish(self, protocol))
}
pub async fn ok(self) -> Result<Session, Error> {
self.respond(Response::default()).await
}
pub async fn reject(mut self, reason: Rejected) -> Result<(), Error> {
let response = proto::ConnectResponse::new(reason.status());
let mut buf = Vec::new();
response.encode(&mut buf).map_err(|err| Error::Web(err.to_string()))?;
write_all(&mut self.send, &buf).await?;
web_transport_trait::poll::SendStream::finish(&mut self.send)?;
let mut deadline = self.conn.owner().after(CLOSE_GRACE);
let send = &mut self.send;
kio::wait(|waiter| {
let mut cx = Context::from_waker(waiter.waker());
if web_transport_trait::poll::SendStream::poll_closed(send, &mut cx).is_ready() {
return Poll::Ready(());
}
deadline.poll(waiter)
})
.await;
self.guard.disarm();
self.conn.close_code(H3_NO_ERROR, "");
Ok(())
}
}
#[derive(Clone, Debug, Default)]
pub struct Response {
protocol: Option<String>,
}
impl Response {
pub fn with_protocol(mut self, protocol: impl Into<String>) -> Self {
self.protocol = Some(protocol.into());
self
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
#[non_exhaustive]
pub enum Rejected {
Unauthorized,
Forbidden,
NotFound,
BadRequest,
Unavailable,
}
impl Rejected {
fn status(self) -> http::StatusCode {
match self {
Self::Unauthorized => http::StatusCode::UNAUTHORIZED,
Self::Forbidden => http::StatusCode::FORBIDDEN,
Self::NotFound => http::StatusCode::NOT_FOUND,
Self::BadRequest => http::StatusCode::BAD_REQUEST,
Self::Unavailable => http::StatusCode::SERVICE_UNAVAILABLE,
}
}
}
impl std::fmt::Debug for Request {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("Request").field("url", &self.request.url).finish()
}
}
struct Web {
session_id: u64,
header_uni: Bytes,
header_bi: Bytes,
header_datagram: Bytes,
state: RefCell<State>,
closed: RefCell<Option<Error>>,
}
struct State {
pending_uni: Vec<PendingUni>,
pending_bi: Vec<PendingBi>,
ready_uni: VecDeque<super::RecvStream>,
ready_bi: VecDeque<(super::SendStream, super::RecvStream)>,
held_recv: Vec<super::RecvStream>,
_control_send: Option<super::SendStream>,
connect_send: Option<super::SendStream>,
}
pub struct Session {
conn: Connection,
web: Option<Rc<Web>>,
protocol: Option<String>,
}
impl Session {
pub fn raw(conn: Connection) -> Self {
let protocol = web_transport_trait::poll::Session::protocol(&conn).map(str::to_owned);
Self {
conn,
web: None,
protocol,
}
}
fn establish(request: Request, protocol: Option<String>) -> Self {
let Request {
conn,
guard: _,
request: _,
send,
recv,
held,
control,
early,
pending,
} = request;
let session_id = send.id();
let mut header_uni = Vec::new();
encode_varint(STREAM_WEBTRANSPORT, &mut header_uni);
encode_varint(session_id, &mut header_uni);
let mut header_bi = Vec::new();
encode_varint(FRAME_WEBTRANSPORT, &mut header_bi);
encode_varint(session_id, &mut header_bi);
let mut header_datagram = Vec::new();
encode_varint(session_id, &mut header_datagram);
let ready_uni = early
.into_iter()
.filter_map(|(session, recv)| (session == session_id).then_some(recv))
.collect();
let web = Rc::new(Web {
session_id,
header_uni: header_uni.into(),
header_bi: header_bi.into(),
header_datagram: header_datagram.into(),
state: RefCell::new(State {
pending_uni: pending,
pending_bi: Vec::new(),
ready_uni,
ready_bi: VecDeque::new(),
held_recv: held,
_control_send: Some(control),
connect_send: Some(send),
}),
closed: RefCell::new(None),
});
let capsules = web.clone();
let capsule_conn = conn.clone();
conn.owner()
.spawn(async move { read_capsules(capsules, capsule_conn, recv).await });
Self {
conn,
web: Some(web),
protocol,
}
}
fn map_err(&self, err: Error) -> Error {
match self.web {
Some(_) => unmap_err(err),
None => err,
}
}
}
impl Clone for Session {
fn clone(&self) -> Self {
Self {
conn: self.conn.clone(),
web: self.web.clone(),
protocol: self.protocol.clone(),
}
}
}
impl std::fmt::Debug for Session {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("Session")
.field("web", &self.web.is_some())
.field("protocol", &self.protocol)
.finish()
}
}
impl web_transport_trait::poll::Session for Session {
type SendStream = SendStream;
type RecvStream = RecvStream;
type Error = Error;
fn poll_accept_uni(&mut self, cx: &mut Context<'_>) -> Poll<Result<Self::RecvStream, Self::Error>> {
let Some(web) = self.web.clone() else {
let inner = ready!(web_transport_trait::poll::Session::poll_accept_uni(&mut self.conn, cx))?;
return Poll::Ready(Ok(RecvStream { inner, web: false }));
};
loop {
if let Some(inner) = web.state.borrow_mut().ready_uni.pop_front() {
return Poll::Ready(Ok(RecvStream { inner, web: true }));
}
loop {
match web_transport_trait::poll::Session::poll_accept_uni(&mut self.conn, cx) {
Poll::Ready(Ok(recv)) => web.state.borrow_mut().pending_uni.push(PendingUni::new(recv)),
Poll::Ready(Err(err)) => return Poll::Ready(Err(unmap_err(err))),
Poll::Pending => break,
}
}
if !classify_uni(&web, cx) {
return Poll::Pending;
}
}
}
fn poll_accept_bi(
&mut self,
cx: &mut Context<'_>,
) -> Poll<Result<(Self::SendStream, Self::RecvStream), Self::Error>> {
let Some(web) = self.web.clone() else {
let (send, recv) = ready!(web_transport_trait::poll::Session::poll_accept_bi(&mut self.conn, cx))?;
return Poll::Ready(Ok((
SendStream {
inner: send,
prefix: Bytes::new(),
finishing: false,
web: false,
},
RecvStream {
inner: recv,
web: false,
},
)));
};
loop {
if let Some((send, recv)) = web.state.borrow_mut().ready_bi.pop_front() {
return Poll::Ready(Ok((
SendStream {
inner: send,
prefix: Bytes::new(),
finishing: false,
web: true,
},
RecvStream { inner: recv, web: true },
)));
}
loop {
match web_transport_trait::poll::Session::poll_accept_bi(&mut self.conn, cx) {
Poll::Ready(Ok((send, recv))) => web.state.borrow_mut().pending_bi.push(PendingBi::new(send, recv)),
Poll::Ready(Err(err)) => return Poll::Ready(Err(unmap_err(err))),
Poll::Pending => break,
}
}
if !classify_bi(&web, cx) {
return Poll::Pending;
}
}
}
fn poll_open_uni(&mut self, cx: &mut Context<'_>) -> Poll<Result<Self::SendStream, Self::Error>> {
let inner = ready!(web_transport_trait::poll::Session::poll_open_uni(&mut self.conn, cx))
.map_err(|err| self.map_err(err))?;
let prefix = self.web.as_ref().map(|web| web.header_uni.clone()).unwrap_or_default();
Poll::Ready(Ok(SendStream {
inner,
prefix,
finishing: false,
web: self.web.is_some(),
}))
}
fn poll_open_bi(
&mut self,
cx: &mut Context<'_>,
) -> Poll<Result<(Self::SendStream, Self::RecvStream), Self::Error>> {
let (send, recv) = ready!(web_transport_trait::poll::Session::poll_open_bi(&mut self.conn, cx))
.map_err(|err| self.map_err(err))?;
let prefix = self.web.as_ref().map(|web| web.header_bi.clone()).unwrap_or_default();
Poll::Ready(Ok((
SendStream {
inner: send,
prefix,
finishing: false,
web: self.web.is_some(),
},
RecvStream {
inner: recv,
web: self.web.is_some(),
},
)))
}
fn poll_send_datagram(&mut self, cx: &mut Context<'_>, payload: &[u8]) -> Poll<Result<(), Self::Error>> {
let Some(web) = &self.web else {
return web_transport_trait::poll::Session::poll_send_datagram(&mut self.conn, cx, payload);
};
let mut framed = Vec::with_capacity(web.header_datagram.len() + payload.len());
framed.extend_from_slice(&web.header_datagram);
framed.extend_from_slice(payload);
web_transport_trait::poll::Session::poll_send_datagram(&mut self.conn, cx, &framed).map_err(unmap_err)
}
fn poll_recv_datagram(&mut self, cx: &mut Context<'_>) -> Poll<Result<Bytes, Self::Error>> {
let Some(web) = self.web.clone() else {
return web_transport_trait::poll::Session::poll_recv_datagram(&mut self.conn, cx);
};
loop {
let datagram = ready!(web_transport_trait::poll::Session::poll_recv_datagram(
&mut self.conn,
cx
))
.map_err(unmap_err)?;
let mut peek: &[u8] = &datagram;
match decode_varint(&mut peek) {
Some(id) if id == web.session_id => {
let start = datagram.len() - peek.len();
return Poll::Ready(Ok(datagram.slice(start..)));
}
_ => tracing::debug!("dropping a datagram for an unknown session"),
}
}
}
fn max_datagram_size(&self) -> usize {
let inner = web_transport_trait::poll::Session::max_datagram_size(&self.conn);
match &self.web {
Some(web) => inner.saturating_sub(web.header_datagram.len()),
None => inner,
}
}
fn protocol(&self) -> Option<&str> {
self.protocol.as_deref()
}
fn close(&mut self, code: u32, reason: &str) {
let Some(web) = &self.web else {
return web_transport_trait::poll::Session::close(&mut self.conn, code, reason);
};
{
let mut closed = web.closed.borrow_mut();
if closed.is_some() {
return;
}
*closed = Some(Error::App {
code: u64::from(code),
reason: reason.to_string(),
});
}
let connect_send = web.state.borrow_mut().connect_send.take();
let http3 = proto::error_to_http3(code);
let Some(mut send) = connect_send else {
self.conn.close_code(http3, reason);
return;
};
let capsule = proto::Capsule::CloseWebTransportSession {
code,
reason: reason.to_string(),
};
let mut payload = Vec::new();
capsule.encode(&mut payload);
let mut frame = Vec::new();
encode_varint(FRAME_DATA, &mut frame);
encode_varint(payload.len() as u64, &mut frame);
frame.extend_from_slice(&payload);
let mut deadline = self.conn.owner().after(CLOSE_GRACE);
let reason = reason.to_string();
let mut conn = self.conn.clone();
self.conn.owner().spawn(async move {
let mut offset = 0;
kio::wait(|waiter| {
let mut cx = Context::from_waker(waiter.waker());
if web_transport_trait::poll::Session::poll_closed(&mut conn, &mut cx).is_ready() {
return Poll::Ready(());
}
while offset < frame.len() {
match web_transport_trait::poll::SendStream::poll_write(&mut send, &mut cx, &frame[offset..]) {
Poll::Ready(Ok(n)) => offset += n,
Poll::Ready(Err(_)) => return Poll::Ready(()),
Poll::Pending => break,
}
if offset == frame.len() {
let _ = web_transport_trait::poll::SendStream::finish(&mut send);
}
}
deadline.poll(waiter)
})
.await;
conn.close_code(http3, &reason);
});
}
fn poll_closed(&mut self, cx: &mut Context<'_>) -> Poll<Self::Error> {
let err = ready!(web_transport_trait::poll::Session::poll_closed(&mut self.conn, cx));
let Some(web) = &self.web else {
return Poll::Ready(err);
};
if let Some(recorded) = web.closed.borrow().clone() {
return Poll::Ready(recorded);
}
Poll::Ready(unmap_err(err))
}
fn stats(&self) -> impl web_transport_trait::Stats {
web_transport_trait::poll::Session::stats(&self.conn)
}
}
pub struct SendStream {
inner: super::SendStream,
prefix: Bytes,
finishing: bool,
web: bool,
}
impl SendStream {
fn map(&self, err: Error) -> Error {
match self.web {
true => unmap_err(err),
false => err,
}
}
}
impl web_transport_trait::poll::SendStream for SendStream {
type Error = Error;
fn poll_write(&mut self, cx: &mut Context<'_>, buf: &[u8]) -> Poll<Result<usize, Self::Error>> {
if self.finishing {
return Poll::Ready(Err(Error::Quic("stream already finished".to_string())));
}
while !self.prefix.is_empty() {
let n = ready!(web_transport_trait::poll::SendStream::poll_write(
&mut self.inner,
cx,
&self.prefix
))
.map_err(|err| self.map(err))?;
self.prefix.advance(n);
}
match web_transport_trait::poll::SendStream::poll_write(&mut self.inner, cx, buf) {
Poll::Ready(Err(err)) => Poll::Ready(Err(self.map(err))),
other => other,
}
}
fn poll_write_buf<B: Buf>(&mut self, cx: &mut Context<'_>, buf: &mut B) -> Poll<Result<usize, Self::Error>> {
if self.finishing {
return Poll::Ready(Err(Error::Quic("stream already finished".to_string())));
}
while !self.prefix.is_empty() {
ready!(web_transport_trait::poll::SendStream::poll_write_buf(
&mut self.inner,
cx,
&mut self.prefix
))
.map_err(|err| self.map(err))?;
}
match web_transport_trait::poll::SendStream::poll_write_buf(&mut self.inner, cx, buf) {
Poll::Ready(Err(err)) => Poll::Ready(Err(self.map(err))),
other => other,
}
}
fn set_priority(&mut self, order: u8) {
web_transport_trait::poll::SendStream::set_priority(&mut self.inner, order);
}
fn finish(&mut self) -> Result<(), Self::Error> {
if !self.prefix.is_empty() {
let n = self.inner.try_write(&self.prefix);
self.prefix.advance(n);
if !self.prefix.is_empty() {
self.finishing = true;
return Ok(());
}
}
web_transport_trait::poll::SendStream::finish(&mut self.inner).map_err(|err| self.map(err))
}
fn reset(&mut self, code: u32) {
match self.web {
true => self.inner.reset_code(proto::error_to_http3(code)),
false => web_transport_trait::poll::SendStream::reset(&mut self.inner, code),
}
}
fn poll_closed(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
while self.finishing {
match ready!(web_transport_trait::poll::SendStream::poll_write(
&mut self.inner,
cx,
&self.prefix
)) {
Ok(n) => self.prefix.advance(n),
Err(err) => return Poll::Ready(Err(self.map(err))),
}
if self.prefix.is_empty() {
self.finishing = false;
if let Err(err) = web_transport_trait::poll::SendStream::finish(&mut self.inner) {
return Poll::Ready(Err(self.map(err)));
}
}
}
match web_transport_trait::poll::SendStream::poll_closed(&mut self.inner, cx) {
Poll::Ready(Err(err)) => Poll::Ready(Err(self.map(err))),
other => other,
}
}
}
impl Drop for SendStream {
fn drop(&mut self) {
if self.finishing {
let n = self.inner.try_write(&self.prefix);
self.prefix.advance(n);
if self.prefix.is_empty() {
let _ = web_transport_trait::poll::SendStream::finish(&mut self.inner);
}
}
if self.web && !self.inner.ended() {
self.inner.reset_code(proto::error_to_http3(0));
}
}
}
impl std::fmt::Debug for SendStream {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
self.inner.fmt(f)
}
}
pub struct RecvStream {
inner: super::RecvStream,
web: bool,
}
impl RecvStream {
fn map(&self, err: Error) -> Error {
match self.web {
true => unmap_err(err),
false => err,
}
}
}
impl web_transport_trait::poll::RecvStream for RecvStream {
type Error = Error;
fn poll_read(&mut self, cx: &mut Context<'_>, dst: &mut [u8]) -> Poll<Result<Option<usize>, Self::Error>> {
match web_transport_trait::poll::RecvStream::poll_read(&mut self.inner, cx, dst) {
Poll::Ready(Err(err)) => Poll::Ready(Err(self.map(err))),
other => other,
}
}
fn stop(&mut self, code: u32) {
match self.web {
true => self.inner.stop_code(proto::error_to_http3(code)),
false => web_transport_trait::poll::RecvStream::stop(&mut self.inner, code),
}
}
fn poll_closed(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
match web_transport_trait::poll::RecvStream::poll_closed(&mut self.inner, cx) {
Poll::Ready(Err(err)) => Poll::Ready(Err(self.map(err))),
other => other,
}
}
}
impl Drop for RecvStream {
fn drop(&mut self) {
if self.web && !self.inner.ended() {
self.inner.stop_code(proto::error_to_http3(0));
}
}
}
impl std::fmt::Debug for RecvStream {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
self.inner.fmt(f)
}
}
fn unmap_err(err: Error) -> Error {
fn unmap(code: u64) -> Option<u64> {
proto::error_from_http3(code).map(u64::from)
}
match err {
Error::Reset(code) => match unmap(code) {
Some(code) => Error::Reset(code),
None => Error::Http3 {
code,
reason: String::new(),
},
},
Error::Stop(code) => match unmap(code) {
Some(code) => Error::Stop(code),
None => Error::Http3 {
code,
reason: String::new(),
},
},
Error::App { code, reason } => match unmap(code) {
Some(code) => Error::App { code, reason },
None => Error::Http3 { code, reason },
},
other => other,
}
}
#[derive(Default)]
struct VarRead {
buf: [u8; 8],
have: usize,
}
enum VarPoll {
Value(u64),
End,
}
impl VarRead {
fn poll(&mut self, cx: &mut Context<'_>, recv: &mut super::RecvStream) -> Poll<VarPoll> {
loop {
let need = match self.have {
0 => 1,
_ => 1usize << (self.buf[0] >> 6),
};
if self.have >= need {
let mut value = u64::from(self.buf[0] & 0x3f);
for byte in &self.buf[1..need] {
value = (value << 8) | u64::from(*byte);
}
return Poll::Ready(VarPoll::Value(value));
}
match ready!(web_transport_trait::poll::RecvStream::poll_read(
recv,
cx,
&mut self.buf[self.have..need]
)) {
Ok(Some(n)) => self.have += n,
Ok(None) | Err(_) => return Poll::Ready(VarPoll::End),
}
}
}
}
fn encode_varint(value: u64, buf: &mut Vec<u8>) {
if value < 1 << 6 {
buf.push(value as u8);
} else if value < 1 << 14 {
buf.extend_from_slice(&((value as u16) | 0x4000).to_be_bytes());
} else if value < 1 << 30 {
buf.extend_from_slice(&((value as u32) | 0x8000_0000).to_be_bytes());
} else {
buf.extend_from_slice(&(value | 0xc000_0000_0000_0000).to_be_bytes());
}
}
fn decode_varint(buf: &mut &[u8]) -> Option<u64> {
let first = *buf.first()?;
let len = 1usize << (first >> 6);
if buf.len() < len {
return None;
}
let mut value = u64::from(first & 0x3f);
for byte in &buf[1..len] {
value = (value << 8) | u64::from(*byte);
}
*buf = &buf[len..];
Some(value)
}
enum UniClass {
Control,
Qpack,
Web(u64),
Unknown,
}
struct PendingUni {
recv: super::RecvStream,
typ: VarRead,
session: VarRead,
typ_value: Option<u64>,
}
impl PendingUni {
fn new(recv: super::RecvStream) -> Self {
Self {
recv,
typ: VarRead::default(),
session: VarRead::default(),
typ_value: None,
}
}
fn poll_classify(&mut self, cx: &mut Context<'_>) -> Poll<UniClass> {
let typ = match self.typ_value {
Some(typ) => typ,
None => match ready!(self.typ.poll(cx, &mut self.recv)) {
VarPoll::Value(typ) => {
self.typ_value = Some(typ);
typ
}
VarPoll::End => return Poll::Ready(UniClass::Unknown),
},
};
match typ {
STREAM_WEBTRANSPORT => match ready!(self.session.poll(cx, &mut self.recv)) {
VarPoll::Value(session) => Poll::Ready(UniClass::Web(session)),
VarPoll::End => Poll::Ready(UniClass::Unknown),
},
STREAM_CONTROL => Poll::Ready(UniClass::Control),
STREAM_QPACK_ENCODER | STREAM_QPACK_DECODER => Poll::Ready(UniClass::Qpack),
_ => {
tracing::debug!(typ, "ignoring an unknown unidirectional stream");
Poll::Ready(UniClass::Unknown)
}
}
}
}
fn classify_uni(web: &Rc<Web>, cx: &mut Context<'_>) -> bool {
let pending = std::mem::take(&mut web.state.borrow_mut().pending_uni);
let mut keep = Vec::new();
let mut progressed = false;
for mut stream in pending {
match stream.poll_classify(cx) {
Poll::Pending => keep.push(stream),
Poll::Ready(UniClass::Web(session)) if session == web.session_id => {
web.state.borrow_mut().ready_uni.push_back(stream.recv);
progressed = true;
}
Poll::Ready(UniClass::Control | UniClass::Qpack) => {
let mut state = web.state.borrow_mut();
if state.held_recv.len() < HELD_STREAMS {
state.held_recv.push(stream.recv);
}
}
Poll::Ready(_) => {}
}
}
let mut state = web.state.borrow_mut();
state.pending_uni.extend(keep);
progressed
}
struct PendingBi {
send: super::SendStream,
recv: super::RecvStream,
typ: VarRead,
session: VarRead,
typ_value: Option<u64>,
}
impl PendingBi {
fn new(send: super::SendStream, recv: super::RecvStream) -> Self {
Self {
send,
recv,
typ: VarRead::default(),
session: VarRead::default(),
typ_value: None,
}
}
fn poll_classify(&mut self, cx: &mut Context<'_>) -> Poll<Option<u64>> {
let typ = match self.typ_value {
Some(typ) => typ,
None => match ready!(self.typ.poll(cx, &mut self.recv)) {
VarPoll::Value(typ) => {
self.typ_value = Some(typ);
typ
}
VarPoll::End => return Poll::Ready(None),
},
};
if typ != FRAME_WEBTRANSPORT {
tracing::debug!(typ, "ignoring an unknown bidirectional stream");
return Poll::Ready(None);
}
match ready!(self.session.poll(cx, &mut self.recv)) {
VarPoll::Value(session) => Poll::Ready(Some(session)),
VarPoll::End => Poll::Ready(None),
}
}
}
fn classify_bi(web: &Rc<Web>, cx: &mut Context<'_>) -> bool {
let pending = std::mem::take(&mut web.state.borrow_mut().pending_bi);
let mut keep = Vec::new();
let mut progressed = false;
for mut stream in pending {
match stream.poll_classify(cx) {
Poll::Pending => keep.push(stream),
Poll::Ready(Some(session)) if session == web.session_id => {
web.state.borrow_mut().ready_bi.push_back((stream.send, stream.recv));
progressed = true;
}
Poll::Ready(_) => {}
}
}
let mut state = web.state.borrow_mut();
state.pending_bi.extend(keep);
progressed
}
async fn open_uni(conn: &mut Connection) -> Result<super::SendStream, Error> {
std::future::poll_fn(|cx| web_transport_trait::poll::Session::poll_open_uni(conn, cx)).await
}
async fn accept_bi(conn: &mut Connection) -> Result<(super::SendStream, super::RecvStream), Error> {
std::future::poll_fn(|cx| web_transport_trait::poll::Session::poll_accept_bi(conn, cx)).await
}
async fn write_all(send: &mut super::SendStream, mut buf: &[u8]) -> Result<(), Error> {
while !buf.is_empty() {
let n = std::future::poll_fn(|cx| web_transport_trait::poll::SendStream::poll_write(send, cx, buf)).await?;
buf = &buf[n..];
}
Ok(())
}
async fn read_some(recv: &mut super::RecvStream, buf: &mut BytesMut) -> Result<bool, Error> {
let mut chunk = [0u8; 4096];
let n = std::future::poll_fn(|cx| web_transport_trait::poll::RecvStream::poll_read(recv, cx, &mut chunk)).await?;
match n {
Some(n) => {
if buf.len() + n > MESSAGE_LIMIT {
return Err(Error::Web("an HTTP/3 message exceeded the buffer limit".into()));
}
buf.extend_from_slice(&chunk[..n]);
Ok(true)
}
None => Ok(false),
}
}
async fn read_settings(recv: &mut super::RecvStream) -> Result<proto::Settings, Error> {
let mut buf = BytesMut::new();
buf.extend_from_slice(&[STREAM_CONTROL as u8]);
loop {
let mut peek: &[u8] = &buf;
match proto::Settings::decode(&mut peek) {
Ok(settings) => return Ok(settings),
Err(proto::SettingsError::UnexpectedEnd) => {}
Err(err) => return Err(Error::Web(err.to_string())),
}
if !read_some(recv, &mut buf).await? {
return Err(Error::Web("control stream ended before SETTINGS".into()));
}
}
}
async fn read_connect(recv: &mut super::RecvStream) -> Result<proto::ConnectRequest, Error> {
let mut buf = BytesMut::new();
loop {
let mut peek: &[u8] = &buf;
match proto::ConnectRequest::decode(&mut peek) {
Ok(request) => return Ok(request),
Err(proto::ConnectError::UnexpectedEnd) => {}
Err(err) => return Err(Error::Web(err.to_string())),
}
if !read_some(recv, &mut buf).await? {
return Err(Error::Web("stream ended before the CONNECT request".into()));
}
}
}
async fn read_capsules(web: Rc<Web>, conn: Connection, mut recv: super::RecvStream) {
let mut capsules = Capsules::default();
let capsule = loop {
match capsules.take() {
Ok(Some(proto::Capsule::CloseWebTransportSession { code, reason })) => {
break Some((code, reason));
}
Ok(Some(_)) => continue,
Ok(None) => {}
Err(err) => {
tracing::debug!(%err, "failed to parse a capsule on the CONNECT stream");
break None;
}
}
match read_some(&mut recv, &mut capsules.frames).await {
Ok(true) => {}
Ok(false) | Err(_) => break None,
}
};
match capsule {
Some((code, reason)) => {
web.closed.borrow_mut().get_or_insert(Error::App {
code: u64::from(code),
reason: reason.clone(),
});
conn.close_code(proto::error_to_http3(code), &reason);
}
None => conn.close_code(proto::error_to_http3(0), ""),
}
}
#[derive(Default)]
struct Capsules {
frames: BytesMut,
body: BytesMut,
}
impl Capsules {
fn take(&mut self) -> Result<Option<proto::Capsule>, Error> {
self.demux()?;
let mut peek: &[u8] = &self.body;
match proto::Capsule::decode(&mut peek) {
Ok(capsule) => {
let consumed = self.body.len() - peek.len();
self.body.advance(consumed);
Ok(Some(capsule))
}
Err(proto::CapsuleError::UnexpectedEnd | proto::CapsuleError::VarInt(_)) => Ok(None),
Err(err) => Err(Error::Web(err.to_string())),
}
}
fn demux(&mut self) -> Result<(), Error> {
loop {
let mut peek: &[u8] = &self.frames;
let Some(typ) = decode_varint(&mut peek) else {
return Ok(());
};
let Some(len) = decode_varint(&mut peek) else {
return Ok(());
};
let len = usize::try_from(len).map_err(|_| Error::Web("oversized HTTP/3 frame".into()))?;
if peek.len() < len {
return Ok(());
}
let header = self.frames.len() - peek.len();
if typ == FRAME_DATA {
if self.body.len() + len > MESSAGE_LIMIT {
return Err(Error::Web("a capsule exceeded the buffer limit".into()));
}
self.body.extend_from_slice(&peek[..len]);
}
self.frames.advance(header + len);
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn an_unmappable_code_is_not_an_application_error() {
use web_transport_trait::Error as _;
let app = unmap_err(Error::App {
code: proto::error_to_http3(7),
reason: "seven".into(),
});
assert!(
matches!(&app, Error::App { code: 7, reason } if reason == "seven"),
"got {app:?}"
);
assert_eq!(app.session_error(), Some((7, "seven".to_string())));
let h3 = unmap_err(Error::App {
code: 0x100,
reason: "done".into(),
});
assert!(
matches!(&h3, Error::Http3 { code: 0x100, reason } if reason == "done"),
"got {h3:?}"
);
assert_eq!(h3.session_error(), None, "not a code the peer's application chose");
let reset = unmap_err(Error::Reset(0x100));
assert!(matches!(reset, Error::Http3 { code: 0x100, .. }), "got {reset:?}");
assert_eq!(reset.stream_error(), None, "not a MoQ stream code");
assert!(matches!(
unmap_err(Error::Stop(proto::error_to_http3(3))),
Error::Stop(3)
));
}
fn frame(typ: u64, payload: &[u8]) -> Vec<u8> {
let mut buf = Vec::new();
encode_varint(typ, &mut buf);
encode_varint(payload.len() as u64, &mut buf);
buf.extend_from_slice(payload);
buf
}
fn close_capsule(code: u32, reason: &str) -> Vec<u8> {
let mut buf = Vec::new();
proto::Capsule::CloseWebTransportSession {
code,
reason: reason.to_string(),
}
.encode(&mut buf);
buf
}
#[test]
fn a_capsule_spans_data_frames() {
let capsule = close_capsule(42, "split");
let (head, tail) = capsule.split_at(capsule.len() / 2);
let mut capsules = Capsules::default();
capsules.frames.extend_from_slice(&frame(FRAME_DATA, head));
assert!(capsules.take().expect("parse").is_none(), "half a capsule is not one");
capsules.frames.extend_from_slice(&frame(FRAME_DATA, tail));
let capsule = capsules.take().expect("parse").expect("the second half completes it");
assert_eq!(
capsule,
proto::Capsule::CloseWebTransportSession {
code: 42,
reason: "split".to_string()
}
);
}
#[test]
fn one_frame_carries_several_capsules() {
let mut payload = close_capsule(1, "first");
payload.extend_from_slice(&close_capsule(2, "second"));
let mut capsules = Capsules::default();
capsules.frames.extend_from_slice(&frame(FRAME_DATA, &payload));
for (code, reason) in [(1, "first"), (2, "second")] {
let capsule = capsules.take().expect("parse").expect("a whole capsule");
assert_eq!(
capsule,
proto::Capsule::CloseWebTransportSession {
code,
reason: reason.to_string()
}
);
}
assert!(capsules.take().expect("parse").is_none(), "only two were written");
}
#[test]
fn a_non_data_frame_is_skipped() {
let capsule = close_capsule(7, "after");
let (head, tail) = capsule.split_at(1);
let mut capsules = Capsules::default();
capsules.frames.extend_from_slice(&frame(FRAME_DATA, head));
capsules.frames.extend_from_slice(&frame(0x07, b"\x00"));
capsules.frames.extend_from_slice(&frame(FRAME_DATA, tail));
let capsule = capsules.take().expect("parse").expect("a whole capsule");
assert_eq!(
capsule,
proto::Capsule::CloseWebTransportSession {
code: 7,
reason: "after".to_string()
}
);
}
#[test]
fn a_capsule_stream_is_bounded() {
let mut header = Vec::new();
encode_varint(0x2843, &mut header);
encode_varint(65536, &mut header);
let mut capsules = Capsules::default();
capsules.frames.extend_from_slice(&frame(FRAME_DATA, &header));
let chunk = vec![0u8; 8 * 1024];
let err = loop {
match capsules.take() {
Ok(None) => {}
Ok(Some(capsule)) => panic!("the payload never arrived, got {capsule:?}"),
Err(err) => break err,
}
capsules.frames.extend_from_slice(&frame(FRAME_DATA, &chunk));
};
assert!(matches!(err, Error::Web(_)), "refused with {err}");
assert!(capsules.body.len() <= MESSAGE_LIMIT, "the buffer stayed bounded");
}
}