use std::{sync::Arc, task::Poll, time::Duration};
use web_transport_trait::Stats as _;
use crate::{Error, SessionError, Version, bandwidth, goaway};
#[derive(Clone)]
struct Close {
code: u32,
reason: String,
}
struct StatsState {
sample: Stats,
demanded: bool,
}
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
#[non_exhaustive]
pub struct Stats {
pub rtt: Option<Duration>,
pub estimated_send_rate: Option<bandwidth::Rate>,
pub estimated_recv_rate: Option<bandwidth::Rate>,
pub bytes_sent: Option<u64>,
pub bytes_received: Option<u64>,
pub bytes_lost: Option<u64>,
pub packets_sent: Option<u64>,
pub packets_received: Option<u64>,
pub packets_lost: Option<u64>,
}
#[derive(Clone)]
pub struct Session {
close: kio::Producer<Option<Close>>,
closed: kio::Consumer<Option<Error>>,
stats: kio::Shared<StatsState>,
version: Version,
send_bandwidth: Option<bandwidth::Consumer>,
recv_bandwidth: Option<bandwidth::Consumer>,
goaway: Arc<goaway::Handle>,
}
impl Session {
pub fn version(&self) -> Version {
self.version
}
pub fn send_bandwidth(&self) -> Option<bandwidth::Consumer> {
self.send_bandwidth.clone()
}
pub fn recv_bandwidth(&self) -> Option<bandwidth::Consumer> {
self.recv_bandwidth.clone()
}
pub fn stats(&self) -> Stats {
let mut stats = {
let mut state = self.stats.lock();
if !state.demanded {
state.demanded = true;
}
state.sample
};
stats.estimated_recv_rate = self.recv_bandwidth.as_ref().and_then(bandwidth::Consumer::peek);
stats
}
pub fn abort(&self, err: Error) {
if let Ok(mut close) = self.close.write()
&& close.is_none()
{
*close = Some(Close {
code: SessionError::from(&err).to_code(),
reason: err.to_string(),
});
}
}
pub async fn closed(&self) -> Error {
match self
.closed
.wait(|state| match &**state {
Some(err) => Poll::Ready(err.clone()),
None => Poll::Pending,
})
.await
{
Ok(err) => err,
Err(kio::Closed) => Error::Cancel,
}
}
pub fn drain(&self) -> goaway::Producer {
self.goaway.producer()
}
pub fn draining(&self) -> goaway::Consumer {
self.goaway.consumer()
}
}
impl Session {
pub(super) fn new<S>(
runtime: crate::time::Clock,
session: S,
version: Version,
recv_bandwidth: Option<bandwidth::Consumer>,
protocol: crate::driver::Protocol<S>,
goaway: goaway::Handle,
) -> (Self, crate::Driver<S>)
where
S: crate::transport::poll::Session,
{
let sample = snapshot(&session);
let (send_bandwidth, send_producer) = if sample.estimated_send_rate.is_some() {
let producer = bandwidth::Producer::new();
(Some(producer.consume()), Some(producer))
} else {
(None, None)
};
let close = kio::Producer::new(None);
let closed = kio::Producer::new(None);
let closed_consumer = closed.consume();
let stats = kio::Shared::new(StatsState {
sample,
demanded: false,
});
let supervisor = Supervisor {
runtime: runtime.clone(),
closed_watch: session.clone(),
session,
close: Some(close.consume()),
closed,
stats: stats.clone(),
send_bandwidth: send_producer,
mode: SamplerMode::Idle,
};
let session = Self {
close,
closed: closed_consumer,
stats,
version,
send_bandwidth,
recv_bandwidth,
goaway: Arc::new(goaway),
};
let driver = crate::Driver::new(
runtime.clone(),
crate::driver::State {
protocol,
supervisor: Some(supervisor),
result: None,
},
);
(session, driver)
}
}
pub(crate) struct Supervisor<S> {
runtime: crate::time::Clock,
session: S,
closed_watch: S,
close: Option<kio::Consumer<Option<Close>>>,
closed: kio::Producer<Option<Error>>,
stats: kio::Shared<StatsState>,
send_bandwidth: Option<bandwidth::Producer>,
mode: SamplerMode,
}
enum SamplerMode {
Idle,
Polling {
deadline: crate::runtime::Deadline<crate::time::Clock>,
},
}
impl<S: crate::transport::poll::Session> Supervisor<S> {
const POLL_INTERVAL: Duration = Duration::from_millis(100);
pub(crate) fn poll(&mut self, waiter: &kio::Waiter) -> Poll<()> {
let mut cx = std::task::Context::from_waker(waiter.waker());
if let Poll::Ready(err) = self.closed_watch.poll_closed(&mut cx) {
self.stats.lock().sample = snapshot(&self.session);
if let Ok(mut closed) = self.closed.write() {
*closed = Some(Error::from_transport(err));
}
return Poll::Ready(());
}
if let Some(close) = &self.close {
let request = match close.poll(waiter, |state| match &**state {
Some(request) => Poll::Ready(request.clone()),
None => Poll::Pending,
}) {
Poll::Ready(Ok(request)) => Some(request),
Poll::Ready(Err(last)) => Some(last.clone().unwrap_or_else(|| Close {
code: SessionError::Cancel.to_code(),
reason: "dropped".to_string(),
})),
Poll::Pending => None,
};
if let Some(request) = request {
self.session.close(request.code, &request.reason);
self.close = None;
}
}
self.poll_sampler(waiter);
Poll::Pending
}
fn sample(&mut self) {
let sample = snapshot(&self.session);
if let Some(producer) = &self.send_bandwidth {
if producer.set(sample.estimated_send_rate).is_err() {
self.send_bandwidth = None;
}
}
let mut stats = self.stats.lock();
stats.sample = sample;
stats.demanded = false;
drop(stats);
self.mode = SamplerMode::Polling {
deadline: crate::runtime::Deadline::after(&self.runtime, Self::POLL_INTERVAL),
};
}
fn poll_sampler(&mut self, waiter: &kio::Waiter) {
loop {
match &mut self.mode {
SamplerMode::Idle => {
let mut demanded = match &self.send_bandwidth {
Some(producer) => match producer.poll_used(waiter) {
Poll::Ready(Ok(())) => true,
Poll::Ready(Err(_)) => {
self.send_bandwidth = None;
false
}
Poll::Pending => false,
},
None => false,
};
demanded |= self
.stats
.poll(waiter, |state| match state.demanded {
true => Poll::Ready(()),
false => Poll::Pending,
})
.is_ready();
if !demanded {
return;
}
self.sample();
}
SamplerMode::Polling { deadline } => {
if deadline.poll(waiter).is_pending() {
return;
}
let used = self.send_bandwidth.as_ref().is_some_and(bandwidth::Producer::is_used);
if !used && !self.stats.read().demanded {
self.mode = SamplerMode::Idle;
continue;
}
self.sample();
}
}
}
}
}
fn snapshot<S: crate::transport::poll::Session>(session: &S) -> Stats {
let stats = session.stats();
Stats {
rtt: stats.rtt(),
estimated_send_rate: stats.estimated_send_rate().map(bandwidth::Rate::from_bps),
bytes_sent: stats.bytes_sent(),
bytes_received: stats.bytes_received(),
bytes_lost: stats.bytes_lost(),
packets_sent: stats.packets_sent(),
packets_received: stats.packets_received(),
packets_lost: stats.packets_lost(),
..Default::default()
}
}
const _: () = {
const fn assert_send_sync<T: Send + Sync>() {}
assert_send_sync::<Session>();
};