use std::{
sync::{Arc, Mutex},
task::{Context, Poll},
};
use bytes::Bytes;
use web_transport_trait::poll;
#[derive(Debug, Clone)]
pub struct MockError {
code: Option<u32>,
reason: String,
}
impl MockError {
fn closed() -> Self {
Self {
code: Some(0),
reason: "session closed".into(),
}
}
fn stream_reset(code: u32) -> Self {
Self {
code: Some(code),
reason: "stream reset".into(),
}
}
}
impl std::fmt::Display for MockError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(f, "mock transport: {}", self.reason)
}
}
impl std::error::Error for MockError {}
impl web_transport_trait::Error for MockError {
fn session_error(&self) -> Option<(u32, String)> {
self.code.map(|c| (c, self.reason.clone()))
}
fn stream_error(&self) -> Option<u32> {
self.code
}
}
enum StreamChunk {
Data(Bytes),
Fin,
Reset(u32),
}
#[derive(Default)]
struct ClosedSignal {
result: Mutex<Option<Result<(), MockError>>>,
waiters: kio::Fan,
}
impl ClosedSignal {
fn set(&self, result: Result<(), MockError>) {
let mut slot = self.result.lock().unwrap();
if slot.is_none() {
*slot = Some(result);
self.waiters.wake();
}
}
}
pub struct MockSendStream {
tx: Option<kio::Queue<StreamChunk>>,
closed: Arc<ClosedSignal>,
park: kio::Park,
}
impl poll::SendStream for MockSendStream {
type Error = MockError;
fn poll_write(&mut self, _cx: &mut Context<'_>, buf: &[u8]) -> Poll<Result<usize, Self::Error>> {
let Some(tx) = self.tx.as_ref() else {
return Poll::Ready(Err(MockError::closed()));
};
match tx.try_push(StreamChunk::Data(Bytes::copy_from_slice(buf))) {
Ok(()) => Poll::Ready(Ok(buf.len())),
Err(_) => Poll::Ready(Err(MockError::closed())),
}
}
fn set_priority(&mut self, _order: u8) {}
fn finish(&mut self) -> Result<(), Self::Error> {
if let Some(tx) = self.tx.take() {
let _ = tx.try_push(StreamChunk::Fin);
}
Ok(())
}
fn reset(&mut self, code: u32) {
if let Some(tx) = self.tx.take() {
let _ = tx.try_push(StreamChunk::Reset(code));
}
}
fn poll_closed(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
self.closed.waiters.register(self.park.hold(cx));
match self.closed.result.lock().unwrap().clone() {
Some(result) => Poll::Ready(result),
None => Poll::Pending,
}
}
}
impl Drop for MockSendStream {
fn drop(&mut self) {
if let Some(tx) = self.tx.take() {
let _ = tx.try_push(StreamChunk::Fin);
}
}
}
pub struct MockRecvStream {
rx: kio::Queue<StreamChunk>,
buf: Bytes,
done: bool,
closed: Arc<ClosedSignal>,
park: kio::Park,
}
impl MockRecvStream {
fn poll_chunk(&mut self, cx: &mut Context<'_>) -> Poll<Option<StreamChunk>> {
let waiter = self.park.hold(cx);
match self.rx.poll_pop(waiter) {
Poll::Ready(Ok(chunk)) => Poll::Ready(Some(chunk)),
Poll::Ready(Err(_)) => Poll::Ready(None),
Poll::Pending => Poll::Pending,
}
}
}
impl poll::RecvStream for MockRecvStream {
type Error = MockError;
fn poll_read(&mut self, cx: &mut Context<'_>, dst: &mut [u8]) -> Poll<Result<Option<usize>, Self::Error>> {
if self.done {
return Poll::Ready(Ok(None));
}
if !self.buf.is_empty() {
let n = dst.len().min(self.buf.len());
dst[..n].copy_from_slice(&self.buf[..n]);
self.buf = self.buf.slice(n..);
return Poll::Ready(Ok(Some(n)));
}
match std::task::ready!(self.poll_chunk(cx)) {
Some(StreamChunk::Data(data)) => {
let n = dst.len().min(data.len());
dst[..n].copy_from_slice(&data[..n]);
if n < data.len() {
self.buf = data.slice(n..);
}
Poll::Ready(Ok(Some(n)))
}
Some(StreamChunk::Fin) | None => {
self.done = true;
Poll::Ready(Ok(None))
}
Some(StreamChunk::Reset(code)) => {
self.done = true;
Poll::Ready(Err(MockError::stream_reset(code)))
}
}
}
fn stop(&mut self, _code: u32) {
self.closed.set(Ok(()));
self.done = true;
}
fn poll_closed(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
if self.done {
return Poll::Ready(Ok(()));
}
loop {
match std::task::ready!(self.poll_chunk(cx)) {
Some(StreamChunk::Data(_)) => {}
Some(StreamChunk::Fin) | None => {
self.done = true;
return Poll::Ready(Ok(()));
}
Some(StreamChunk::Reset(code)) => {
self.done = true;
return Poll::Ready(Err(MockError::stream_reset(code)));
}
}
}
}
}
impl Drop for MockRecvStream {
fn drop(&mut self) {
self.closed.set(Ok(()));
self.rx.close();
}
}
fn new_stream_pair() -> (MockSendStream, MockRecvStream) {
let queue = kio::Queue::new();
let closed = Arc::new(ClosedSignal::default());
let send = MockSendStream {
tx: Some(queue.clone()),
closed: closed.clone(),
park: kio::Park::default(),
};
let recv = MockRecvStream {
rx: queue,
buf: Bytes::new(),
done: false,
closed,
park: kio::Park::default(),
};
(send, recv)
}
#[derive(Default)]
struct ConnectionState {
close_state: Mutex<Option<(u32, String)>>,
waiters: kio::Fan,
}
struct SessionSide {
bidi: kio::Queue<(MockSendStream, MockRecvStream)>,
uni: kio::Queue<MockRecvStream>,
peer_bidi: kio::Queue<(MockSendStream, MockRecvStream)>,
peer_uni: kio::Queue<MockRecvStream>,
datagrams: kio::Queue<Bytes>,
peer_datagrams: kio::Queue<Bytes>,
protocol: Option<&'static str>,
conn: Arc<ConnectionState>,
}
#[derive(Clone)]
pub struct MockSession {
side: Arc<SessionSide>,
accept_uni: kio::Park,
accept_bi: kio::Park,
datagram: kio::Park,
closed: kio::Park,
}
impl poll::Session for MockSession {
type SendStream = MockSendStream;
type RecvStream = MockRecvStream;
type Error = MockError;
fn poll_accept_uni(&mut self, cx: &mut Context<'_>) -> Poll<Result<Self::RecvStream, Self::Error>> {
let waiter = self.accept_uni.hold(cx);
self.side.conn.waiters.register(waiter);
if let Poll::Ready(res) = self.side.uni.poll_pop(waiter) {
return Poll::Ready(res.map_err(|_| self.close_error()));
}
let closed = self.side.conn.close_state.lock().unwrap().is_some();
match closed {
true => Poll::Ready(Err(self.close_error())),
false => Poll::Pending,
}
}
fn poll_accept_bi(
&mut self,
cx: &mut Context<'_>,
) -> Poll<Result<(Self::SendStream, Self::RecvStream), Self::Error>> {
let waiter = self.accept_bi.hold(cx);
self.side.conn.waiters.register(waiter);
if let Poll::Ready(res) = self.side.bidi.poll_pop(waiter) {
return Poll::Ready(res.map_err(|_| self.close_error()));
}
let closed = self.side.conn.close_state.lock().unwrap().is_some();
match closed {
true => Poll::Ready(Err(self.close_error())),
false => Poll::Pending,
}
}
fn poll_open_bi(
&mut self,
_cx: &mut Context<'_>,
) -> Poll<Result<(Self::SendStream, Self::RecvStream), Self::Error>> {
let (our_send, peer_recv) = new_stream_pair();
let (peer_send, our_recv) = new_stream_pair();
match self.side.peer_bidi.try_push((peer_send, peer_recv)) {
Ok(()) => Poll::Ready(Ok((our_send, our_recv))),
Err(_) => Poll::Ready(Err(self.close_error())),
}
}
fn poll_open_uni(&mut self, _cx: &mut Context<'_>) -> Poll<Result<Self::SendStream, Self::Error>> {
let (our_send, peer_recv) = new_stream_pair();
match self.side.peer_uni.try_push(peer_recv) {
Ok(()) => Poll::Ready(Ok(our_send)),
Err(_) => Poll::Ready(Err(self.close_error())),
}
}
fn poll_send_datagram(&mut self, _cx: &mut Context<'_>, payload: &[u8]) -> Poll<Result<(), Self::Error>> {
match self.side.peer_datagrams.try_push(Bytes::copy_from_slice(payload)) {
Ok(()) => Poll::Ready(Ok(())),
Err(_) => Poll::Ready(Err(self.close_error())),
}
}
fn poll_recv_datagram(&mut self, cx: &mut Context<'_>) -> Poll<Result<Bytes, Self::Error>> {
let waiter = self.datagram.hold(cx);
self.side.conn.waiters.register(waiter);
if let Poll::Ready(res) = self.side.datagrams.poll_pop(waiter) {
return Poll::Ready(res.map_err(|_| self.close_error()));
}
let closed = self.side.conn.close_state.lock().unwrap().is_some();
match closed {
true => Poll::Ready(Err(self.close_error())),
false => Poll::Pending,
}
}
fn max_datagram_size(&self) -> usize {
1200
}
fn protocol(&self) -> Option<&str> {
self.side.protocol
}
fn close(&mut self, code: u32, reason: &str) {
let mut state = self.side.conn.close_state.lock().unwrap();
if state.is_none() {
*state = Some((code, reason.to_string()));
drop(state);
self.side.conn.waiters.wake();
}
}
fn poll_closed(&mut self, cx: &mut Context<'_>) -> Poll<Self::Error> {
self.side.conn.waiters.register(self.closed.hold(cx));
let state = self.side.conn.close_state.lock().unwrap().clone();
match state {
Some((code, reason)) => Poll::Ready(MockError {
code: Some(code),
reason,
}),
None => Poll::Pending,
}
}
fn stats(&self) -> impl web_transport_trait::Stats {
web_transport_trait::StatsUnavailable
}
}
impl MockSession {
fn close_error(&self) -> MockError {
self.side
.conn
.close_state
.lock()
.unwrap()
.as_ref()
.map(|(code, reason)| MockError {
code: Some(*code),
reason: reason.clone(),
})
.unwrap_or_else(MockError::closed)
}
}
pub fn create_mock_session_pair(protocol: Option<&'static str>) -> (MockSession, MockSession) {
let conn = Arc::new(ConnectionState::default());
let c2s_bidi = kio::Queue::new();
let c2s_uni = kio::Queue::new();
let s2c_bidi = kio::Queue::new();
let s2c_uni = kio::Queue::new();
let c2s_datagrams = kio::Queue::new();
let s2c_datagrams = kio::Queue::new();
let client_side = Arc::new(SessionSide {
bidi: s2c_bidi.clone(),
uni: s2c_uni.clone(),
peer_bidi: c2s_bidi.clone(),
peer_uni: c2s_uni.clone(),
datagrams: s2c_datagrams.clone(),
peer_datagrams: c2s_datagrams.clone(),
protocol,
conn: conn.clone(),
});
let server_side = Arc::new(SessionSide {
bidi: c2s_bidi,
uni: c2s_uni,
peer_bidi: s2c_bidi,
peer_uni: s2c_uni,
datagrams: c2s_datagrams,
peer_datagrams: s2c_datagrams,
protocol,
conn,
});
let new = |side| MockSession {
side,
accept_uni: kio::Park::default(),
accept_bi: kio::Park::default(),
datagram: kio::Park::default(),
closed: kio::Park::default(),
};
(new(client_side), new(server_side))
}