use std::{
sync::{
Arc, Mutex,
atomic::{AtomicUsize, Ordering},
},
task::{Context, Poll},
};
use web_transport_trait::poll;
#[derive(Debug, Clone, Default)]
pub struct SinkError;
impl std::fmt::Display for SinkError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(f, "sink transport error")
}
}
impl std::error::Error for SinkError {}
impl web_transport_trait::Error for SinkError {
fn session_error(&self) -> Option<(u32, String)> {
Some((0, "closed".to_string()))
}
}
#[derive(Clone, Default)]
pub struct Log {
pub writes: Arc<Mutex<Vec<u8>>>,
pub resets: Arc<Mutex<Vec<u32>>>,
stops: Arc<Mutex<Vec<u32>>>,
closes: Arc<Mutex<Vec<(u32, String)>>>,
bi_opens: Arc<AtomicUsize>,
priorities: Arc<Mutex<Vec<u8>>>,
}
impl Log {
pub fn resets(&self) -> Vec<u32> {
self.resets.lock().unwrap().clone()
}
pub fn stops(&self) -> Vec<u32> {
self.stops.lock().unwrap().clone()
}
pub fn priorities(&self) -> Vec<u8> {
self.priorities.lock().unwrap().clone()
}
pub fn closes(&self) -> Vec<(u32, String)> {
self.closes.lock().unwrap().clone()
}
pub fn bi_opens(&self) -> usize {
self.bi_opens.load(Ordering::Relaxed)
}
}
pub struct SinkSend {
pub log: Log,
gate: Option<kio::Consumer<bool>>,
park: kio::Park,
finished: bool,
}
impl SinkSend {
pub fn new(log: Log) -> Self {
Self {
log,
gate: None,
park: kio::Park::default(),
finished: false,
}
}
pub fn gated(log: Log, gate: kio::Consumer<bool>) -> Self {
Self {
log,
gate: Some(gate),
park: kio::Park::default(),
finished: false,
}
}
}
impl poll::SendStream for SinkSend {
type Error = SinkError;
fn poll_write(&mut self, cx: &mut Context<'_>, buf: &[u8]) -> Poll<Result<usize, Self::Error>> {
if let Some(gate) = &self.gate {
let waiter = self.park.hold(cx);
match gate.poll(waiter, |open| if **open { Poll::Ready(()) } else { Poll::Pending }) {
Poll::Ready(Ok(())) => {}
Poll::Ready(Err(_)) | Poll::Pending => return Poll::Pending,
}
}
self.log.writes.lock().unwrap().extend_from_slice(buf);
Poll::Ready(Ok(buf.len()))
}
fn set_priority(&mut self, order: u8) {
self.log.priorities.lock().unwrap().push(order);
}
fn finish(&mut self) -> Result<(), Self::Error> {
self.finished = true;
Ok(())
}
fn reset(&mut self, code: u32) {
self.log.resets.lock().unwrap().push(code);
}
fn poll_closed(&mut self, _cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
match self.finished {
true => Poll::Ready(Ok(())),
false => Poll::Pending,
}
}
}
pub struct PendingRecv;
impl poll::RecvStream for PendingRecv {
type Error = SinkError;
fn poll_read(&mut self, _cx: &mut Context<'_>, _dst: &mut [u8]) -> Poll<Result<Option<usize>, Self::Error>> {
Poll::Pending
}
fn stop(&mut self, _code: u32) {}
fn poll_closed(&mut self, _cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
Poll::Pending
}
}
#[derive(Debug, Clone, Default)]
pub struct ResetError;
impl std::fmt::Display for ResetError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(f, "stream reset by peer (code 0)")
}
}
impl std::error::Error for ResetError {}
impl web_transport_trait::Error for ResetError {
fn session_error(&self) -> Option<(u32, String)> {
None
}
fn stream_error(&self) -> Option<u32> {
Some(0)
}
}
pub struct DeadRecv;
impl poll::RecvStream for DeadRecv {
type Error = ResetError;
fn poll_read(&mut self, _cx: &mut Context<'_>, _dst: &mut [u8]) -> Poll<Result<Option<usize>, Self::Error>> {
Poll::Ready(Err(ResetError))
}
fn stop(&mut self, _code: u32) {}
fn poll_closed(&mut self, _cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
Poll::Ready(Err(ResetError))
}
}
#[derive(Clone)]
pub struct DeadStreamSession {
pub log: Log,
unis: Arc<Mutex<usize>>,
bis: Arc<Mutex<usize>>,
}
impl DeadStreamSession {
pub fn unis(count: usize) -> Self {
Self {
log: Log::default(),
unis: Arc::new(Mutex::new(count)),
bis: Arc::new(Mutex::new(0)),
}
}
pub fn bis(count: usize) -> Self {
Self {
log: Log::default(),
unis: Arc::new(Mutex::new(0)),
bis: Arc::new(Mutex::new(count)),
}
}
fn take(counter: &Mutex<usize>) -> bool {
let mut remaining = counter.lock().unwrap();
match *remaining {
0 => false,
_ => {
*remaining -= 1;
true
}
}
}
}
impl poll::Session for DeadStreamSession {
type SendStream = SinkSend;
type RecvStream = DeadRecv;
type Error = SinkError;
fn poll_accept_uni(&mut self, _cx: &mut Context<'_>) -> Poll<Result<Self::RecvStream, Self::Error>> {
match Self::take(&self.unis) {
true => Poll::Ready(Ok(DeadRecv)),
false => Poll::Pending,
}
}
fn poll_accept_bi(&mut self, _cx: &mut Context<'_>) -> Poll<Result<poll::BiStreams<Self>, Self::Error>> {
match Self::take(&self.bis) {
true => Poll::Ready(Ok((SinkSend::new(self.log.clone()), DeadRecv))),
false => Poll::Pending,
}
}
fn poll_open_bi(&mut self, _cx: &mut Context<'_>) -> Poll<Result<poll::BiStreams<Self>, Self::Error>> {
self.log.bi_opens.fetch_add(1, Ordering::Relaxed);
Poll::Ready(Ok((SinkSend::new(self.log.clone()), DeadRecv)))
}
fn poll_open_uni(&mut self, _cx: &mut Context<'_>) -> Poll<Result<Self::SendStream, Self::Error>> {
Poll::Ready(Ok(SinkSend::new(self.log.clone())))
}
fn poll_send_datagram(&mut self, _cx: &mut Context<'_>, _payload: &[u8]) -> Poll<Result<(), Self::Error>> {
Poll::Ready(Ok(()))
}
fn poll_recv_datagram(&mut self, _cx: &mut Context<'_>) -> Poll<Result<bytes::Bytes, Self::Error>> {
Poll::Pending
}
fn max_datagram_size(&self) -> usize {
0
}
fn protocol(&self) -> Option<&str> {
None
}
fn close(&mut self, code: u32, reason: &str) {
self.log.closes.lock().unwrap().push((code, reason.to_owned()));
}
fn poll_closed(&mut self, _cx: &mut Context<'_>) -> Poll<Self::Error> {
Poll::Pending
}
fn stats(&self) -> impl web_transport_trait::Stats {
SinkStats::default()
}
}
#[derive(Clone, Default)]
pub struct SinkSession {
pub log: Log,
bi_gate: Option<kio::Consumer<bool>>,
accept_gate: Option<kio::Consumer<bool>>,
uni_gate: Option<kio::Consumer<bool>>,
uni_open_gate: Option<kio::Consumer<bool>>,
uni_open_park: kio::Park,
protocol: Option<&'static str>,
stats: Arc<Mutex<SinkStats>>,
}
impl SinkSession {
pub fn new(log: Log) -> Self {
Self {
log,
bi_gate: None,
accept_gate: None,
uni_gate: None,
uni_open_gate: None,
uni_open_park: kio::Park::default(),
protocol: None,
stats: Arc::new(Mutex::new(SinkStats::default())),
}
}
pub fn with_stats(self, stats: SinkStats) -> Self {
self.set_stats(stats);
self
}
pub fn set_stats(&self, stats: SinkStats) {
*self.stats.lock().unwrap() = stats;
}
pub fn with_protocol(mut self, protocol: &'static str) -> Self {
self.protocol = Some(protocol);
self
}
pub fn gated_bi(gate: kio::Consumer<bool>) -> Self {
Self {
log: Log::default(),
bi_gate: Some(gate),
accept_gate: None,
uni_gate: None,
uni_open_gate: None,
uni_open_park: kio::Park::default(),
protocol: None,
stats: Arc::new(Mutex::new(SinkStats::default())),
}
}
pub fn accepted_bi(gate: kio::Consumer<bool>) -> Self {
Self {
log: Log::default(),
bi_gate: None,
accept_gate: Some(gate),
uni_gate: None,
uni_open_gate: None,
uni_open_park: kio::Park::default(),
protocol: None,
stats: Arc::new(Mutex::new(SinkStats::default())),
}
}
pub fn gated_uni(gate: kio::Consumer<bool>) -> Self {
Self {
log: Log::default(),
bi_gate: None,
accept_gate: None,
uni_gate: Some(gate),
uni_open_gate: None,
uni_open_park: kio::Park::default(),
protocol: None,
stats: Arc::new(Mutex::new(SinkStats::default())),
}
}
pub fn gated_open_uni(gate: kio::Consumer<bool>) -> Self {
Self {
log: Log::default(),
bi_gate: None,
accept_gate: None,
uni_gate: None,
uni_open_gate: Some(gate),
uni_open_park: kio::Park::default(),
protocol: None,
stats: Arc::new(Mutex::new(SinkStats::default())),
}
}
}
impl poll::Session for SinkSession {
type SendStream = SinkSend;
type RecvStream = PendingRecv;
type Error = SinkError;
fn poll_accept_uni(&mut self, _cx: &mut Context<'_>) -> Poll<Result<Self::RecvStream, Self::Error>> {
Poll::Pending
}
fn poll_accept_bi(&mut self, _cx: &mut Context<'_>) -> Poll<Result<poll::BiStreams<Self>, Self::Error>> {
let Some(gate) = self.accept_gate.clone() else {
return Poll::Pending;
};
let send = SinkSend {
log: self.log.clone(),
gate: Some(gate),
park: kio::Park::default(),
finished: false,
};
Poll::Ready(Ok((send, PendingRecv)))
}
fn poll_open_bi(&mut self, _cx: &mut Context<'_>) -> Poll<Result<poll::BiStreams<Self>, Self::Error>> {
let Some(gate) = self.bi_gate.clone() else {
return Poll::Pending;
};
self.log.bi_opens.fetch_add(1, Ordering::Relaxed);
let send = SinkSend {
log: self.log.clone(),
gate: Some(gate),
park: kio::Park::default(),
finished: false,
};
Poll::Ready(Ok((send, PendingRecv)))
}
fn poll_open_uni(&mut self, cx: &mut Context<'_>) -> Poll<Result<Self::SendStream, Self::Error>> {
if let Some(gate) = &self.uni_open_gate {
let waiter = self.uni_open_park.hold(cx);
match gate.poll(waiter, |open| (**open).then_some(()).map_or(Poll::Pending, Poll::Ready)) {
Poll::Ready(Ok(())) => {}
Poll::Ready(Err(_)) | Poll::Pending => return Poll::Pending,
}
}
Poll::Ready(Ok(match &self.uni_gate {
Some(gate) => SinkSend::gated(self.log.clone(), gate.clone()),
None => SinkSend::new(self.log.clone()),
}))
}
fn poll_send_datagram(&mut self, _cx: &mut Context<'_>, _payload: &[u8]) -> Poll<Result<(), Self::Error>> {
Poll::Ready(Ok(()))
}
fn poll_recv_datagram(&mut self, _cx: &mut Context<'_>) -> Poll<Result<bytes::Bytes, Self::Error>> {
Poll::Pending
}
fn max_datagram_size(&self) -> usize {
0
}
fn protocol(&self) -> Option<&str> {
self.protocol
}
fn close(&mut self, code: u32, reason: &str) {
self.log.closes.lock().unwrap().push((code, reason.to_owned()));
}
fn poll_closed(&mut self, _cx: &mut Context<'_>) -> Poll<Self::Error> {
Poll::Pending
}
fn stats(&self) -> impl web_transport_trait::Stats {
*self.stats.lock().unwrap()
}
}
#[derive(Default, Clone, Copy)]
pub struct SinkStats {
pub estimated_send_rate: Option<u64>,
pub rtt: Option<std::time::Duration>,
}
impl SinkStats {
pub fn with_send_rate(mut self, rate: u64) -> Self {
self.estimated_send_rate = Some(rate);
self
}
pub fn with_rtt(mut self, rtt: std::time::Duration) -> Self {
self.rtt = Some(rtt);
self
}
}
impl web_transport_trait::Stats for SinkStats {
fn estimated_send_rate(&self) -> Option<u64> {
self.estimated_send_rate
}
fn rtt(&self) -> Option<std::time::Duration> {
self.rtt
}
}
pub struct ScriptedRecv {
script: Arc<Mutex<Vec<u8>>>,
eof: bool,
log: Log,
}
impl poll::RecvStream for ScriptedRecv {
type Error = SinkError;
fn poll_read(&mut self, _cx: &mut Context<'_>, dst: &mut [u8]) -> Poll<Result<Option<usize>, Self::Error>> {
let take = {
let mut script = self.script.lock().unwrap();
if script.is_empty() {
0
} else {
let take = dst.len().min(script.len());
dst[..take].copy_from_slice(&script[..take]);
script.drain(..take);
take
}
};
match take {
0 if self.eof => Poll::Ready(Ok(None)),
0 => Poll::Pending,
take => Poll::Ready(Ok(Some(take))),
}
}
fn stop(&mut self, code: u32) {
self.log.stops.lock().unwrap().push(code);
}
fn poll_closed(&mut self, _cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
Poll::Pending
}
}
#[derive(Clone)]
pub struct ScriptedSession {
pub log: Log,
eof: bool,
script: Arc<Mutex<Vec<u8>>>,
queue: Option<Arc<Mutex<std::collections::VecDeque<Vec<u8>>>>>,
open_gate: Option<kio::Consumer<bool>>,
park: kio::Park,
incoming_unis: Arc<Mutex<std::collections::VecDeque<Vec<u8>>>>,
}
impl ScriptedSession {
pub fn new(script: Vec<u8>) -> Self {
Self {
log: Log::default(),
eof: false,
script: Arc::new(Mutex::new(script)),
queue: None,
open_gate: None,
park: kio::Park::default(),
incoming_unis: Arc::new(Mutex::new(std::collections::VecDeque::new())),
}
}
pub fn with_incoming_unis(mut self, scripts: Vec<Vec<u8>>) -> Self {
self.incoming_unis = Arc::new(Mutex::new(scripts.into_iter().collect()));
self
}
pub fn eof(script: Vec<u8>) -> Self {
Self {
eof: true,
..Self::new(script)
}
}
pub fn per_stream(scripts: Vec<Vec<u8>>) -> Self {
Self {
queue: Some(Arc::new(Mutex::new(scripts.into_iter().collect()))),
..Self::new(Vec::new())
}
}
pub fn per_stream_eof(scripts: Vec<Vec<u8>>) -> Self {
Self {
eof: true,
..Self::per_stream(scripts)
}
}
pub fn gated_open(scripts: Vec<Vec<u8>>, gate: kio::Consumer<bool>) -> Self {
Self {
open_gate: Some(gate),
..Self::per_stream(scripts)
}
}
}
impl poll::Session for ScriptedSession {
type SendStream = SinkSend;
type RecvStream = ScriptedRecv;
type Error = SinkError;
fn poll_accept_uni(&mut self, _cx: &mut Context<'_>) -> Poll<Result<Self::RecvStream, Self::Error>> {
let Some(script) = self.incoming_unis.lock().unwrap().pop_front() else {
return Poll::Pending;
};
Poll::Ready(Ok(ScriptedRecv {
script: Arc::new(Mutex::new(script)),
eof: self.eof,
log: self.log.clone(),
}))
}
fn poll_accept_bi(&mut self, _cx: &mut Context<'_>) -> Poll<Result<poll::BiStreams<Self>, Self::Error>> {
Poll::Pending
}
fn poll_open_bi(&mut self, cx: &mut Context<'_>) -> Poll<Result<poll::BiStreams<Self>, Self::Error>> {
if let Some(gate) = self.open_gate.clone() {
let waiter = self.park.hold(cx);
if gate
.poll(waiter, |open| if **open { Poll::Ready(()) } else { Poll::Pending })
.is_pending()
{
return Poll::Pending;
}
}
self.log.bi_opens.fetch_add(1, Ordering::Relaxed);
let script = match &self.queue {
Some(queue) => Arc::new(Mutex::new(queue.lock().unwrap().pop_front().unwrap_or_default())),
None => self.script.clone(),
};
Poll::Ready(Ok((
SinkSend::new(self.log.clone()),
ScriptedRecv {
script,
eof: self.eof,
log: self.log.clone(),
},
)))
}
fn poll_open_uni(&mut self, _cx: &mut Context<'_>) -> Poll<Result<Self::SendStream, Self::Error>> {
Poll::Ready(Ok(SinkSend::new(self.log.clone())))
}
fn poll_send_datagram(&mut self, _cx: &mut Context<'_>, _payload: &[u8]) -> Poll<Result<(), Self::Error>> {
Poll::Ready(Ok(()))
}
fn poll_recv_datagram(&mut self, _cx: &mut Context<'_>) -> Poll<Result<bytes::Bytes, Self::Error>> {
Poll::Pending
}
fn max_datagram_size(&self) -> usize {
0
}
fn protocol(&self) -> Option<&str> {
None
}
fn close(&mut self, code: u32, reason: &str) {
self.log.closes.lock().unwrap().push((code, reason.to_owned()));
}
fn poll_closed(&mut self, _cx: &mut Context<'_>) -> Poll<Self::Error> {
Poll::Pending
}
fn stats(&self) -> impl web_transport_trait::Stats {
SinkStats::default()
}
}