use std::{fmt, sync::Arc};
use snafu::ResultExt;
use tokio::{io::AsyncWriteExt, sync::mpsc};
use super::{
error::{
Closed, DatagramError, OpenSnafu, OpenStreamError, UnsupportedSnafu, WriteHeaderSnafu,
},
protocol::Registry,
};
use crate::{
codec::{BoxReadStream, BoxWriteStream, EncodeExt, SinkWriter},
quic::{self},
varint::VarInt,
};
pub(super) type RoutedBiStream = (BoxReadStream, BoxWriteStream);
pub(super) type RoutedUniStream = BoxReadStream;
pub struct WebTransportSession {
session_id: VarInt,
bidi_rx: tokio::sync::Mutex<mpsc::Receiver<RoutedBiStream>>,
uni_rx: tokio::sync::Mutex<mpsc::Receiver<RoutedUniStream>>,
conn: Arc<dyn quic::DynManageStream>,
registry: Registry,
}
impl WebTransportSession {
pub(super) fn new(
session_id: VarInt,
bidi_rx: mpsc::Receiver<RoutedBiStream>,
uni_rx: mpsc::Receiver<RoutedUniStream>,
conn: Arc<dyn quic::DynManageStream>,
registry: Registry,
) -> Self {
Self {
session_id,
bidi_rx: tokio::sync::Mutex::new(bidi_rx),
uni_rx: tokio::sync::Mutex::new(uni_rx),
conn,
registry,
}
}
pub fn session_id(&self) -> VarInt {
self.session_id
}
pub async fn open_bi(&self) -> Result<(BoxReadStream, BoxWriteStream), OpenStreamError> {
let (reader, writer) = self.conn.open_bi().await.context(OpenSnafu)?;
let writer = write_header(writer, super::WT_BIDI_SIGNAL, self.session_id).await?;
Ok((reader, writer))
}
pub async fn open_uni(&self) -> Result<BoxWriteStream, OpenStreamError> {
let writer = self.conn.open_uni().await.context(OpenSnafu)?;
let writer = write_header(writer, super::WT_UNI_SIGNAL, self.session_id).await?;
Ok(writer)
}
pub async fn accept_bi(&self) -> Result<(BoxReadStream, BoxWriteStream), Closed> {
self.bidi_rx.lock().await.recv().await.ok_or(Closed)
}
pub async fn accept_uni(&self) -> Result<BoxReadStream, Closed> {
self.uni_rx.lock().await.recv().await.ok_or(Closed)
}
pub async fn send_datagram(&self, _data: &[u8]) -> Result<(), DatagramError> {
UnsupportedSnafu.fail()
}
pub async fn recv_datagram(&self) -> Result<Vec<u8>, DatagramError> {
UnsupportedSnafu.fail()
}
}
async fn write_header(
writer: BoxWriteStream,
signal: VarInt,
session_id: VarInt,
) -> Result<BoxWriteStream, OpenStreamError> {
let mut codec_writer = SinkWriter::new(writer);
codec_writer
.encode_one(signal)
.await
.map_err(quic::StreamError::from)
.context(WriteHeaderSnafu)?;
codec_writer
.encode_one(session_id)
.await
.map_err(quic::StreamError::from)
.context(WriteHeaderSnafu)?;
AsyncWriteExt::flush(&mut codec_writer)
.await
.map_err(quic::StreamError::from)
.context(WriteHeaderSnafu)?;
Ok(codec_writer.into_inner())
}
impl fmt::Debug for WebTransportSession {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("WebTransportSession")
.field("session_id", &self.session_id)
.finish()
}
}
impl Drop for WebTransportSession {
fn drop(&mut self) {
if let Ok(mut registry) = self.registry.lock() {
registry.remove(&self.session_id);
}
}
}
#[cfg(test)]
mod tests {
use super::*;
const fn assert_send_sync<T: Send + Sync>() {}
const _: () = assert_send_sync::<WebTransportSession>();
}