use crate::{announce, frame, group, origin, track};
use std::{sync::Arc, task::Poll, time::Duration};
use bytes::Buf;
use futures::{FutureExt, StreamExt, stream::FuturesUnordered};
use web_transport_trait::Stats;
use crate::{
AsPath, Error, Origin, OriginList,
coding::{Encode, Stream, Writer},
lite::{
self,
priority::{Priority, PriorityHandle, PriorityQueue},
},
util::{MaybeBoxedExt, MaybeSendBox, TaskSet},
};
use super::Version;
struct WatchedRoute {
consumer: crate::broadcast::Consumer,
demand: crate::broadcast::Demand,
path: crate::PathOwned,
sent: Option<SentRoute>,
idle_at: Option<web_async::time::Instant>,
}
#[derive(Clone, PartialEq, Eq)]
struct SentRoute {
hops: OriginList,
cost: lite::RouteCost,
}
pub(super) struct PublisherConfig<S: web_transport_trait::Session> {
pub session: S,
pub origin: origin::Consumer,
pub version: Version,
}
pub(super) struct Publisher<S: web_transport_trait::Session> {
session: S,
origin: origin::Consumer,
self_origin: Origin,
priority: PriorityQueue,
version: Version,
}
impl<S: web_transport_trait::Session> Publisher<S> {
pub fn new(config: PublisherConfig<S>) -> Self {
let self_origin = *config.origin;
Self {
session: config.session,
origin: config.origin,
self_origin,
priority: Default::default(),
version: config.version,
}
}
pub async fn run(self) -> Result<(), Error> {
let this = Arc::new(self);
let mut tasks = TaskSet::owned();
loop {
let stream = tasks.drive(Stream::accept(&this.session, this.version)).await?;
let this = this.clone();
tasks.push(async move {
if let Err(err) = this.handle(stream).await {
tracing::warn!(%err, "control stream error");
}
});
}
}
async fn handle(&self, mut stream: Stream<S, Version>) -> Result<(), Error> {
let kind = stream.reader.decode().await?;
match kind {
lite::ControlType::Announce => self.recv_announce(stream).await,
lite::ControlType::Subscribe => self.recv_subscribe(stream).await,
lite::ControlType::Fetch => self.recv_fetch(stream).await,
lite::ControlType::Track => self.recv_track(stream).await,
lite::ControlType::Probe => {
self.recv_probe(stream).await;
Ok(())
}
lite::ControlType::Goaway => {
tracing::info!("received goaway stream");
Ok(())
}
lite::ControlType::Session => Err(Error::UnexpectedStream),
}
}
async fn recv_probe(&self, mut stream: Stream<S, Version>) {
match Self::run_probe(&self.session, &mut stream, self.version).await {
Ok(()) => {
tracing::debug!("probe stream closed");
}
Err(err) => {
tracing::warn!(%err, "probe stream error");
stream.writer.abort(&err);
}
}
}
async fn run_probe(session: &S, stream: &mut Stream<S, Version>, _version: Version) -> Result<(), Error> {
const PROBE_INTERVAL: Duration = Duration::from_millis(100);
const PROBE_MAX_AGE: Duration = Duration::from_secs(10);
const PROBE_MAX_DELTA: f64 = 0.25;
let mut last_sent: Option<(u64, web_async::time::Instant)> = None;
let mut interval = web_async::time::interval(PROBE_INTERVAL);
loop {
let closed = {
let mut closed = std::pin::pin!(stream.reader.closed());
let mut tick = std::pin::pin!(interval.tick());
kio::wait(|waiter| {
if let Poll::Ready(res) = waiter.poll_future(closed.as_mut()) {
return Poll::Ready(Some(res));
}
waiter.poll_future(tick.as_mut()).map(|_| None)
})
.await
};
if let Some(res) = closed {
return res;
}
let Some(bitrate) = session.stats().estimated_send_rate() else {
continue;
};
let should_send = match last_sent {
None => true,
Some((0, _)) => bitrate > 0,
Some((prev, at)) => {
let elapsed = at.elapsed().as_secs_f64();
let t = elapsed.clamp(PROBE_INTERVAL.as_secs_f64(), PROBE_MAX_AGE.as_secs_f64());
let range = PROBE_MAX_AGE.as_secs_f64() - PROBE_INTERVAL.as_secs_f64();
let threshold = PROBE_MAX_DELTA * (PROBE_MAX_AGE.as_secs_f64() - t) / range;
let change = (bitrate as f64 - prev as f64).abs() / prev as f64;
change >= threshold
}
};
if should_send {
let rtt = session.stats().rtt().map(|d| d.as_millis() as u64);
stream.writer.encode(&lite::Probe { bitrate, rtt }).await?;
last_sent = Some((bitrate, web_async::time::Instant::now()));
}
}
}
pub async fn recv_announce(&self, mut stream: Stream<S, Version>) -> Result<(), Error> {
let interest = stream.reader.decode::<lite::AnnounceRequest>().await?;
let prefix = interest.prefix.to_owned();
let exclude_hop = interest.exclude_hop;
let origin = self
.origin
.scope(&[prefix.as_path()])
.unwrap_or_else(|| self.origin.empty());
let mut announced = origin.announced();
if let Err(err) = Self::run_announce(
&mut stream,
&origin,
&mut announced,
&prefix,
self.self_origin,
exclude_hop,
self.version,
)
.await
{
match &err {
Error::Cancel | Error::Transport(_) => {
tracing::debug!(prefix = %origin.absolute(prefix), "announcing cancelled");
}
err => {
tracing::warn!(%err, prefix = %origin.absolute(prefix), "announcing error");
}
}
stream.writer.abort(&err);
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
async fn run_announce(
stream: &mut Stream<S, Version>,
origin: &origin::Consumer,
announced: &mut announce::Consumer,
prefix: impl AsPath,
self_origin: Origin,
exclude_hop: u64,
version: Version,
) -> Result<(), Error> {
let prefix = prefix.as_path();
let mut next_announce_id: u64 = 0;
let mut announce_ids: std::collections::HashMap<crate::PathOwned, u64> = std::collections::HashMap::new();
let mut watched: std::collections::HashMap<crate::PathOwned, WatchedRoute> = std::collections::HashMap::new();
match version {
Version::Lite01 | Version::Lite02 => {
let mut init = Vec::new();
while let Some(crate::announce::Update { path, broadcast }) = announced.try_next() {
let suffix = path
.strip_prefix(&prefix)
.expect("origin returned invalid path")
.to_owned();
let absolute = origin.absolute(&path).to_owned();
if broadcast.is_some() {
tracing::debug!(broadcast = %absolute, "announce");
if !init.contains(&suffix) {
init.push(suffix);
}
} else {
tracing::debug!(broadcast = %absolute, "unannounce");
init.retain(|p| p != &suffix);
}
}
let announce_init = lite::AnnounceInit { suffixes: init };
stream.writer.encode(&announce_init).await?;
}
_ if version.has_announce_ok() => {
let mut initial: Vec<(crate::PathOwned, SentRoute)> = Vec::new();
while let Some(crate::announce::Update { path, broadcast }) = announced.try_next() {
let suffix = path
.strip_prefix(&prefix)
.expect("origin returned invalid path")
.to_owned();
let absolute = origin.absolute(&path).to_owned();
match broadcast {
Some(broadcast) => {
let route = broadcast.route();
let hops = route.hops.clone();
let demand = broadcast.demand();
let cost = Self::outgoing_cost(version, &demand, &route);
watched.insert(
suffix.clone(),
WatchedRoute {
consumer: broadcast.clone(),
demand,
path: path.clone(),
sent: None,
idle_at: None,
},
);
if exclude_hop != 0 && hops.iter().any(|h| h.id() == exclude_hop) {
continue;
}
if hops.contains(&self_origin) {
continue;
}
tracing::debug!(broadcast = %absolute, "announce");
initial.retain(|(s, _)| s != &suffix);
initial.push((suffix, SentRoute { hops, cost }));
}
None => {
tracing::debug!(broadcast = %absolute, "unannounce");
watched.remove(&suffix);
initial.retain(|(s, _)| s != &suffix);
}
}
}
let ok = lite::AnnounceOk {
origin: self_origin,
active: initial.len() as u64,
};
let mut buf = bytes::BytesMut::new();
ok.encode(&mut buf, version)?;
for (suffix, route) in &initial {
if version.has_announce_id() {
announce_ids.insert(suffix.clone(), next_announce_id);
next_announce_id += 1;
}
if let Some(entry) = watched.get_mut(suffix) {
entry.sent = Some(route.clone());
}
lite::AnnounceBroadcast::Active {
suffix: suffix.as_path(),
hops: route.hops.clone(),
cost: route.cost,
}
.encode(&mut buf, version)?;
}
let mut buf = buf.freeze();
stream.writer.write_all(&mut buf).await?;
}
_ => {
}
}
enum Op {
Announce(Option<crate::announce::Update>),
Route(crate::PathOwned, Result<crate::broadcast::Route, Error>),
Idle(crate::PathOwned),
Linger,
}
const COST_LINGER: Duration = Duration::from_secs(5);
loop {
let deadline = watched
.values()
.filter_map(|entry| entry.idle_at)
.min()
.map(|at| at + COST_LINGER);
let op = {
let mut closed = std::pin::pin!(stream.reader.closed());
let mut linger = std::pin::pin!(async move {
match deadline {
Some(at) => {
web_async::time::sleep(at.saturating_duration_since(web_async::time::Instant::now())).await
}
None => std::future::pending().await,
}
});
let mut fired: Option<web_async::time::Instant> = None;
kio::wait(|waiter| {
if let Poll::Ready(res) = waiter.poll_future(closed.as_mut()) {
return Poll::Ready(Err(res));
}
if let Poll::Ready(next) = announced.poll_next(waiter) {
return Poll::Ready(Ok(Op::Announce(next)));
}
if fired.is_none() && waiter.poll_future(linger.as_mut()).is_ready() {
fired = Some(web_async::time::Instant::now());
}
for (suffix, entry) in watched.iter_mut() {
if let Poll::Ready(res) = entry.consumer.poll_route_changed(waiter) {
return Poll::Ready(Ok(Op::Route(suffix.clone(), res)));
}
if !version.has_route_cost() {
continue;
}
let Some(sent) = &entry.sent else { continue };
if sent.cost != lite::RouteCost(0) {
if let Poll::Ready(Ok(())) = entry.demand.poll_used(waiter) {
return Poll::Ready(Ok(Op::Route(suffix.clone(), Ok(entry.consumer.route()))));
}
continue;
}
match entry.idle_at {
Some(_) if entry.demand.is_used() => entry.idle_at = None,
Some(at) if fired.is_some_and(|now| now >= at + COST_LINGER) => {
entry.idle_at = None;
return Poll::Ready(Ok(Op::Route(suffix.clone(), Ok(entry.consumer.route()))));
}
Some(_) => {
let _ = entry.demand.poll_used(waiter);
continue;
}
None => {}
}
if let Poll::Ready(Ok(())) = entry.demand.poll_unused(waiter) {
return Poll::Ready(Ok(Op::Idle(suffix.clone())));
}
}
match fired {
Some(_) => Poll::Ready(Ok(Op::Linger)),
None => Poll::Pending,
}
})
.await
};
let op = match op {
Ok(op) => op,
Err(res) => return res,
};
match op {
Op::Announce(None) => {
stream.writer.finish()?;
return stream.writer.closed().await;
}
Op::Announce(Some(crate::announce::Update { path, broadcast })) => {
let suffix = path
.strip_prefix(&prefix)
.expect("origin returned invalid path")
.to_owned();
let absolute = origin.absolute(&path).to_owned();
match broadcast {
Some(active) => {
let route = active.route();
let demand = active.demand();
if lite::restart_supported(version) {
watched.insert(
suffix.clone(),
WatchedRoute {
consumer: active.clone(),
demand: demand.clone(),
path: path.clone(),
sent: None,
idle_at: None,
},
);
}
let Some(hops) =
Self::prepare_active_hops(&route.hops, self_origin, exclude_hop, version, &absolute)
else {
continue;
};
let cost = Self::outgoing_cost(version, &demand, &route);
tracing::debug!(broadcast = %absolute, "announce");
if version.has_announce_id() {
let prev = announce_ids.insert(suffix.clone(), next_announce_id);
debug_assert!(prev.is_none(), "announce id still assigned for a new announce");
next_announce_id += 1;
}
if let Some(entry) = watched.get_mut(&suffix) {
entry.sent = Some(SentRoute {
hops: hops.clone(),
cost,
});
}
stream
.writer
.encode(&lite::AnnounceBroadcast::Active { suffix, hops, cost })
.await?;
}
None => {
tracing::debug!(broadcast = %absolute, "unannounce");
let retracted = watched.remove(&suffix).is_some_and(|entry| entry.sent.is_none());
if version.has_announce_id() {
if let Some(id) = announce_ids.remove(&suffix) {
stream.writer.encode(&lite::AnnounceBroadcast::EndedId { id }).await?;
}
} else if !retracted {
stream
.writer
.encode(&lite::AnnounceBroadcast::Ended {
suffix,
hops: OriginList::new(),
})
.await?;
}
}
}
}
Op::Route(suffix, res) => {
let Ok(route) = res else {
watched.remove(&suffix);
continue;
};
let Some(entry) = watched.get_mut(&suffix) else {
continue;
};
entry.idle_at = None;
let absolute = origin.absolute(&entry.path).to_owned();
let cost = Self::outgoing_cost(version, &entry.demand, &route);
let hops = Self::prepare_active_hops(&route.hops, self_origin, exclude_hop, version, &absolute)
.map(|hops| SentRoute { hops, cost });
let sent = entry.sent.clone();
match (hops, sent) {
(Some(route), Some(sent)) if route == sent => {}
(Some(route), Some(_)) => {
tracing::debug!(broadcast = %absolute, "reannounce");
if version.has_announce_id() {
let Some(id) = announce_ids.get(&suffix).copied() else {
debug_assert!(false, "announced path without an announce id");
tracing::warn!(broadcast = %absolute, "restart without an announce id; skipping");
continue;
};
entry.sent = Some(route.clone());
stream
.writer
.encode(&lite::AnnounceBroadcast::Restart {
id,
hops: route.hops,
cost: route.cost,
})
.await?;
} else {
entry.sent = Some(route.clone());
stream
.writer
.encode(&lite::AnnounceBroadcast::Active {
suffix,
hops: route.hops,
cost: route.cost,
})
.await?;
}
}
(Some(route), None) => {
tracing::debug!(broadcast = %absolute, "announce");
if version.has_announce_id() {
announce_ids.insert(suffix.clone(), next_announce_id);
next_announce_id += 1;
}
entry.sent = Some(route.clone());
stream
.writer
.encode(&lite::AnnounceBroadcast::Active {
suffix,
hops: route.hops,
cost: route.cost,
})
.await?;
}
(None, Some(_)) => {
tracing::debug!(broadcast = %absolute, "unannounce (filtered route)");
entry.sent = None;
if version.has_announce_id() {
if let Some(id) = announce_ids.remove(&suffix) {
stream.writer.encode(&lite::AnnounceBroadcast::EndedId { id }).await?;
}
} else {
stream
.writer
.encode(&lite::AnnounceBroadcast::Ended {
suffix,
hops: OriginList::new(),
})
.await?;
}
}
(None, None) => {}
}
}
Op::Idle(suffix) => {
if let Some(entry) = watched.get_mut(&suffix) {
entry.idle_at = Some(web_async::time::Instant::now());
}
}
Op::Linger => {}
}
}
}
fn outgoing_cost(
version: Version,
demand: &crate::broadcast::Demand,
route: &crate::broadcast::Route,
) -> lite::RouteCost {
if !version.has_route_cost() {
return lite::RouteCost::default();
}
match demand.is_used() {
true => lite::RouteCost(0),
false => lite::RouteCost(route.cost),
}
}
fn prepare_active_hops(
hops: &OriginList,
self_origin: Origin,
exclude_hop: u64,
version: Version,
absolute: &crate::Path,
) -> Option<OriginList> {
if exclude_hop != 0 && hops.iter().any(|h| h.id() == exclude_hop) {
tracing::debug!(broadcast = %absolute, %exclude_hop, "skipping announce per peer's exclude_hop");
return None;
}
if hops.contains(&self_origin) {
tracing::debug!(broadcast = %absolute, "skipping reflected announce");
return None;
}
let mut hops = hops.clone();
if !version.has_announce_ok() && hops.push(self_origin).is_err() {
tracing::warn!(broadcast = %absolute, "dropping announce; hop chain at MAX_HOPS (possible loop)");
return None;
}
Some(hops)
}
pub async fn recv_track(&self, mut stream: Stream<S, Version>) -> Result<(), Error> {
if !self.version.has_track_stream() {
return Err(Error::UnexpectedStream);
}
let request = stream.reader.decode::<lite::Track>().await?;
let track = request.track.clone();
let absolute = self.origin.absolute(&request.broadcast).to_owned();
tracing::debug!(broadcast = %absolute, %track, "track info requested");
if let Err(err) = self.run_track_info(&mut stream, &request).await {
match &err {
Error::Cancel | Error::Transport(_) => {
tracing::debug!(broadcast = %absolute, %track, "track info cancelled")
}
err => tracing::warn!(broadcast = %absolute, %track, %err, "track info error"),
}
stream.writer.abort(&err);
}
Ok(())
}
async fn run_track_info(&self, stream: &mut Stream<S, Version>, request: &lite::Track<'_>) -> Result<(), Error> {
let broadcast = self.origin.request_broadcast(&request.broadcast).await?;
let info = broadcast.track(&request.track)?.info().await?;
stream
.writer
.encode(&lite::TrackInfo {
priority: info.priority,
ordered: info.ordered,
latency_max: info.latency_max,
timescale: info.timescale,
})
.await?;
stream.writer.finish()?;
stream.writer.closed().await
}
pub async fn recv_subscribe(&self, mut stream: Stream<S, Version>) -> Result<(), Error> {
let subscribe = stream.reader.decode::<lite::Subscribe>().await?;
let id = subscribe.id;
let track = subscribe.track.clone();
let absolute = self.origin.absolute(&subscribe.broadcast).to_owned();
tracing::info!(%id, broadcast = %absolute, %track, "subscribed started");
let broadcast = self.origin.request_broadcast(&subscribe.broadcast);
if let Err(err) = Self::run_subscribe(
self.session.clone(),
&mut stream,
&subscribe,
broadcast,
self.priority.clone(),
self.version,
)
.await
{
match &err {
Error::Cancel | Error::Transport(_) => {
tracing::info!(%id, broadcast = %absolute, %track, "subscribed cancelled")
}
err => {
tracing::warn!(%id, broadcast = %absolute, %track, %err, "subscribed error")
}
}
stream.writer.abort(&err);
} else {
tracing::info!(%id, broadcast = %absolute, %track, "subscribed complete")
}
Ok(())
}
async fn run_subscribe(
session: S,
stream: &mut Stream<S, Version>,
subscribe: &lite::Subscribe<'_>,
broadcast: kio::Pending<origin::Requesting>,
priority: PriorityQueue,
version: Version,
) -> Result<(), Error> {
let subscription = crate::track::Subscription {
priority: subscribe.priority,
ordered: subscribe.ordered,
latency_max: subscribe.max_latency,
group_start: subscribe.start_group,
group_end: subscribe.end_group,
};
let broadcast = broadcast.await?;
let track_consumer = broadcast.track(&subscribe.track)?;
let track = track_consumer.subscribe(subscription).await?;
let timescale = if version.has_track_stream() {
Some(track.info().timescale)
} else {
None
};
if !version.has_track_stream() {
let info = lite::SubscribeOk {
priority: subscribe.priority,
ordered: false,
max_latency: std::time::Duration::ZERO,
start_group: None,
end_group: None,
};
stream.writer.encode(&lite::SubscribeResponse::Ok(info)).await?;
}
let track_priority_tx = kio::Producer::new(subscribe.priority);
let sub = Subscription {
session,
id: subscribe.id,
track_name: Arc::from(track.name()),
priority,
track_priority: track_priority_tx.consume(),
track_priority_seen: subscribe.priority,
version,
timescale,
};
sub.run_track(
track,
subscribe.start_group,
subscribe.end_group,
&mut stream.reader,
&mut stream.writer,
&track_priority_tx,
)
.await?;
stream.writer.finish()?;
stream.writer.closed().await
}
pub async fn recv_fetch(&self, mut stream: Stream<S, Version>) -> Result<(), Error> {
if !self.version.has_track_stream() {
return Err(Error::UnexpectedStream);
}
let fetch = stream.reader.decode::<lite::Fetch>().await?;
let track = fetch.track.clone();
let group = fetch.group;
let absolute = self.origin.absolute(&fetch.broadcast).to_owned();
tracing::info!(broadcast = %absolute, %track, %group, "fetch started");
let broadcast = self.origin.request_broadcast(&fetch.broadcast);
if let Err(err) = Self::run_fetch(&mut stream, &fetch, broadcast, self.version).await {
match &err {
Error::Cancel | Error::Transport(_) => {
tracing::info!(broadcast = %absolute, %track, %group, "fetch cancelled")
}
err => tracing::warn!(broadcast = %absolute, %track, %group, %err, "fetch error"),
}
stream.writer.abort(&err);
} else {
tracing::info!(broadcast = %absolute, %track, %group, "fetch complete");
}
Ok(())
}
async fn run_fetch(
stream: &mut Stream<S, Version>,
fetch: &lite::Fetch<'_>,
broadcast: kio::Pending<origin::Requesting>,
version: Version,
) -> Result<(), Error> {
let broadcast = broadcast.await?;
let track = broadcast.track(&fetch.track)?;
let mut group = track
.fetch_group(
fetch.group,
group::Fetch {
priority: fetch.priority,
},
)
.await?;
let timescale = if version.has_track_stream() {
Some(group.timescale())
} else {
None
};
let mut prev_ts: u64 = 0;
while let Some(mut frame) = group.next_frame().await? {
write_fetch_frame(&mut stream.writer, &mut frame, timescale, &mut prev_ts).await?;
}
stream.writer.finish()?;
stream.writer.closed().await
}
}
#[cfg(test)]
mod test {
use super::*;
use crate::{Timestamp, broadcast};
fn track_producer(name: impl Into<Arc<str>>) -> track::Producer {
track::Producer::new(Arc::new(broadcast::Info::default()), name, None)
}
#[tokio::test]
async fn recv_next_drains_datagram_before_finished() {
let mut producer = track_producer("test");
let mut subscriber = producer.subscribe(None);
producer
.append_datagram(Timestamp::from_millis(1).unwrap(), &b"last"[..])
.unwrap();
producer.finish().unwrap();
match recv_next(&mut subscriber, true, false).await.unwrap() {
Recv::Datagram(datagram) => assert_eq!(&datagram.payload[..], b"last"),
_ => panic!("expected datagram before finished"),
}
match recv_next(&mut subscriber, true, false).await.unwrap() {
Recv::Finished => {}
_ => panic!("expected finished after datagram"),
}
}
#[tokio::test]
async fn recv_next_reports_future_boundary_before_finished() {
let mut producer = track_producer("test");
let mut subscriber = producer.subscribe(None);
producer.create_group(group::Info { sequence: 5 }).unwrap();
producer.finish_at(7).unwrap();
match recv_next(&mut subscriber, false, true).await.unwrap() {
Recv::Group(group) => assert_eq!(group.sequence, 5),
_ => panic!("expected group 5"),
}
match recv_next(&mut subscriber, false, true).await.unwrap() {
Recv::Boundary(group) => assert_eq!(group, 7),
_ => panic!("expected the future boundary"),
}
producer.create_group(group::Info { sequence: 6 }).unwrap();
match recv_next(&mut subscriber, false, false).await.unwrap() {
Recv::Group(group) => assert_eq!(group.sequence, 6),
_ => panic!("expected group 6"),
}
match recv_next(&mut subscriber, false, false).await.unwrap() {
Recv::Finished => {}
_ => panic!("expected finished once the boundary is reached"),
}
}
}
#[cfg(test)]
mod announce_test {
use super::*;
use crate::coding::{Decode, Reader};
use std::sync::Mutex;
#[derive(Debug, Clone, Default)]
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)]
struct SinkSend {
writes: Arc<Mutex<Vec<u8>>>,
}
impl web_transport_trait::SendStream for SinkSend {
type Error = SinkError;
async fn write(&mut self, buf: &[u8]) -> Result<usize, Self::Error> {
self.writes.lock().unwrap().extend_from_slice(buf);
Ok(buf.len())
}
fn set_priority(&mut self, _order: u8) {}
fn finish(&mut self) -> Result<(), Self::Error> {
Ok(())
}
fn reset(&mut self, _code: u32) {}
async fn closed(&mut self) -> Result<(), Self::Error> {
std::future::pending().await
}
}
struct PendingRecv;
impl web_transport_trait::RecvStream for PendingRecv {
type Error = SinkError;
async fn read(&mut self, _dst: &mut [u8]) -> Result<Option<usize>, Self::Error> {
std::future::pending().await
}
fn stop(&mut self, _code: u32) {}
async fn closed(&mut self) -> Result<(), Self::Error> {
std::future::pending().await
}
}
#[derive(Clone)]
struct SinkSession;
impl web_transport_trait::Session for SinkSession {
type SendStream = SinkSend;
type RecvStream = PendingRecv;
type Error = SinkError;
async fn accept_uni(&self) -> Result<Self::RecvStream, Self::Error> {
std::future::pending().await
}
async fn accept_bi(&self) -> Result<(Self::SendStream, Self::RecvStream), Self::Error> {
std::future::pending().await
}
async fn open_bi(&self) -> Result<(Self::SendStream, Self::RecvStream), Self::Error> {
std::future::pending().await
}
async fn open_uni(&self) -> Result<Self::SendStream, Self::Error> {
std::future::pending().await
}
fn send_datagram(&self, _payload: bytes::Bytes) -> Result<(), Self::Error> {
Ok(())
}
async fn recv_datagram(&self) -> Result<bytes::Bytes, Self::Error> {
std::future::pending().await
}
fn max_datagram_size(&self) -> usize {
0
}
fn protocol(&self) -> Option<&str> {
None
}
fn close(&self, _code: u32, _reason: &str) {}
async fn closed(&self) -> Self::Error {
std::future::pending().await
}
fn stats(&self) -> impl web_transport_trait::Stats {
SinkStats
}
}
struct SinkStats;
impl web_transport_trait::Stats for SinkStats {
fn estimated_send_rate(&self) -> Option<u64> {
None
}
}
const VERSION: Version = Version::Lite06Wip;
const COLD: u64 = 7;
struct Wire {
writes: Arc<Mutex<Vec<u8>>>,
cursor: usize,
}
impl Wire {
fn pending(&self) -> Vec<u8> {
self.writes.lock().unwrap()[self.cursor..].to_vec()
}
fn take_ok(&mut self) -> lite::AnnounceOk {
let buf = self.pending();
let mut slice = &buf[..];
let ok = lite::AnnounceOk::decode(&mut slice, VERSION).expect("announce ok");
self.cursor += buf.len() - slice.len();
ok
}
fn take_announces(&mut self) -> Vec<lite::AnnounceBroadcast<'static>> {
let buf = self.pending();
let mut slice = &buf[..];
let mut msgs = Vec::new();
while !slice.is_empty() {
msgs.push(own(
lite::AnnounceBroadcast::decode(&mut slice, VERSION).expect("announce message")
));
}
self.cursor += buf.len();
msgs
}
fn assert_quiet(&self) {
let pending = self.pending();
assert!(pending.is_empty(), "unexpected wire bytes: {pending:?}");
}
}
fn own(msg: lite::AnnounceBroadcast<'_>) -> lite::AnnounceBroadcast<'static> {
match msg {
lite::AnnounceBroadcast::Active { suffix, hops, cost } => lite::AnnounceBroadcast::Active {
suffix: suffix.to_owned(),
hops,
cost,
},
lite::AnnounceBroadcast::Ended { suffix, hops } => lite::AnnounceBroadcast::Ended {
suffix: suffix.to_owned(),
hops,
},
lite::AnnounceBroadcast::EndedId { id } => lite::AnnounceBroadcast::EndedId { id },
lite::AnnounceBroadcast::Restart { id, hops, cost } => lite::AnnounceBroadcast::Restart { id, hops, cost },
}
}
struct Harness {
origin: origin::Producer,
source: crate::broadcast::Producer,
downstream: crate::broadcast::Consumer,
wire: Wire,
task: tokio::task::JoinHandle<Result<(), Error>>,
}
impl Harness {
fn assert_idle(&self) {
self.wire.assert_quiet();
assert!(!self.task.is_finished(), "the announce loop ended unexpectedly");
}
async fn announce(&mut self, name: &str) -> (crate::broadcast::Producer, track::Consumer) {
let source = self
.origin
.create_broadcast(name, crate::broadcast::Route::new().with_cost(COLD).with_announce(true))
.unwrap();
let downstream = self.origin.consume().announced_broadcast(name).await.unwrap();
let track = downstream.track("video").unwrap();
settle().await;
match self.wire.take_announces().as_slice() {
[
lite::AnnounceBroadcast::Active { cost: first, .. },
lite::AnnounceBroadcast::Restart { cost: second, .. },
] => {
assert_eq!(*first, lite::RouteCost(COLD));
assert_eq!(*second, lite::RouteCost(0));
}
[lite::AnnounceBroadcast::Active { cost, .. }] => assert_eq!(*cost, lite::RouteCost(0)),
other => panic!("expected {name} to announce, got {other:?}"),
}
(source, track)
}
}
async fn settle() {
tokio::time::sleep(Duration::from_millis(1)).await;
}
async fn harness(demand: bool) -> (Harness, Option<track::Consumer>) {
let origin = Origin::new(1).unwrap().produce();
let source = origin
.create_broadcast(
"cam",
crate::broadcast::Route::new().with_cost(COLD).with_announce(true),
)
.unwrap();
let downstream = origin.consume().announced_broadcast("cam").await.unwrap();
let track = demand.then(|| downstream.track("video").unwrap());
let writes = Arc::new(Mutex::new(Vec::new()));
let consumer = origin.consume();
let mut stream = Stream::<SinkSession, Version> {
writer: Writer::new(SinkSend { writes: writes.clone() }, VERSION),
reader: Reader::new(PendingRecv, VERSION),
};
let task = tokio::spawn(async move {
let mut announced = consumer.announced();
let self_origin = *consumer;
Publisher::<SinkSession>::run_announce(&mut stream, &consumer, &mut announced, "", self_origin, 0, VERSION)
.await
});
settle().await;
let mut wire = Wire { writes, cursor: 0 };
assert_eq!(wire.take_ok().active, 1, "expected one initial announce");
let expected = if demand { 0 } else { COLD };
match wire.take_announces().as_slice() {
[lite::AnnounceBroadcast::Active { cost, .. }] => assert_eq!(*cost, lite::RouteCost(expected)),
other => panic!("expected the initial announce, got {other:?}"),
}
(
Harness {
origin,
source,
downstream,
wire,
task,
},
track,
)
}
#[tokio::test(start_paused = true)]
async fn drain_defers_the_cold_restore() {
let (h, track) = harness(true).await;
drop(track);
settle().await;
h.assert_idle();
tokio::time::sleep(Duration::from_secs(3)).await;
h.assert_idle();
}
#[tokio::test(start_paused = true)]
async fn demand_return_cancels_the_restore() {
let (mut h, track) = harness(true).await;
drop(track);
tokio::time::sleep(Duration::from_secs(3)).await;
let track = h.downstream.track("video").unwrap();
tokio::time::sleep(Duration::from_secs(1)).await;
drop(track);
tokio::time::sleep(Duration::from_secs(2)).await;
h.assert_idle();
tokio::time::sleep(Duration::from_secs(4)).await;
match h.wire.take_announces().as_slice() {
[lite::AnnounceBroadcast::Restart { id: 0, cost, .. }] => assert_eq!(*cost, lite::RouteCost(COLD)),
other => panic!("expected the restore on the fresh deadline, got {other:?}"),
}
}
#[tokio::test(start_paused = true)]
async fn staggered_lingers_restore_independently() {
let (mut h, first) = harness(true).await;
let (_second_source, second) = h.announce("cam2").await;
drop(first);
tokio::time::sleep(Duration::from_secs(2)).await;
h.assert_idle();
drop(second);
tokio::time::sleep(Duration::from_secs(4)).await;
match h.wire.take_announces().as_slice() {
[lite::AnnounceBroadcast::Restart { id: 0, cost, .. }] => assert_eq!(*cost, lite::RouteCost(COLD)),
other => panic!("expected only the first restore, got {other:?}"),
}
tokio::time::sleep(Duration::from_secs(2)).await;
match h.wire.take_announces().as_slice() {
[lite::AnnounceBroadcast::Restart { id: 1, cost, .. }] => assert_eq!(*cost, lite::RouteCost(COLD)),
other => panic!("expected the second restore, got {other:?}"),
}
}
#[tokio::test(start_paused = true)]
async fn linger_expiry_restores_the_cold_cost() {
let (mut h, track) = harness(true).await;
drop(track);
tokio::time::sleep(Duration::from_secs(6)).await;
match h.wire.take_announces().as_slice() {
[lite::AnnounceBroadcast::Restart { id: 0, cost, .. }] => assert_eq!(*cost, lite::RouteCost(COLD)),
other => panic!("expected one cold-cost restart, got {other:?}"),
}
tokio::time::sleep(Duration::from_secs(30)).await;
h.assert_idle();
}
#[tokio::test(start_paused = true)]
async fn route_change_supersedes_the_linger() {
let (mut h, track) = harness(true).await;
drop(track);
tokio::time::sleep(Duration::from_secs(3)).await;
h.wire.assert_quiet();
let hops = OriginList::try_from(vec![Origin::new(9).unwrap()]).unwrap();
h.source
.set_route(
crate::broadcast::Route::new()
.with_hops(hops.clone())
.with_cost(COLD)
.with_announce(true),
)
.unwrap();
settle().await;
match h.wire.take_announces().as_slice() {
[
lite::AnnounceBroadcast::Restart {
id: 0,
hops: sent,
cost,
},
] => {
assert_eq!(sent, &hops);
assert_eq!(*cost, lite::RouteCost(COLD));
}
other => panic!("expected the failover restart, got {other:?}"),
}
tokio::time::sleep(Duration::from_secs(30)).await;
h.assert_idle();
}
#[tokio::test(start_paused = true)]
async fn demand_reprices_warm_immediately() {
let (mut h, _) = harness(false).await;
let _track = h.downstream.track("video").unwrap();
settle().await;
match h.wire.take_announces().as_slice() {
[lite::AnnounceBroadcast::Restart { id: 0, cost, .. }] => assert_eq!(*cost, lite::RouteCost(0)),
other => panic!("expected the warm restart, got {other:?}"),
}
}
}
async fn encode_frame_timing<W: web_transport_trait::SendStream>(
writer: &mut Writer<W, Version>,
frame: &frame::Consumer,
timescale: Option<crate::Timescale>,
prev_ts: &mut u64,
) -> Result<(), Error> {
if timescale.is_none() {
return Ok(());
}
let ts = frame.timestamp.value();
encode_zigzag_delta(writer, ts, prev_ts).await?;
Ok(())
}
async fn encode_zigzag_delta<W: web_transport_trait::SendStream>(
writer: &mut Writer<W, Version>,
curr: u64,
prev: &mut u64,
) -> Result<(), Error> {
let delta: i64 = (curr as i128 - *prev as i128)
.try_into()
.map_err(|_| Error::BoundsExceeded(crate::coding::BoundsExceeded))?;
let zz = crate::coding::VarInt::from_zigzag(delta).map_err(crate::coding::EncodeError::from)?;
writer.encode(&zz).await?;
*prev = curr;
Ok(())
}
async fn write_fetch_frame<W: web_transport_trait::SendStream>(
writer: &mut Writer<W, Version>,
frame: &mut frame::Consumer,
timescale: Option<crate::Timescale>,
prev_ts: &mut u64,
) -> Result<(), Error> {
encode_frame_timing(writer, frame, timescale, prev_ts).await?;
writer.encode(&frame.size).await?;
while let Some(chunk) = frame.read_chunk().await? {
writer.write_chunk(chunk).await?;
}
Ok(())
}
#[allow(clippy::large_enum_variant)]
enum Recv {
Group(group::Consumer),
Datagram(crate::Datagram),
Boundary(u64),
Finished,
}
fn poll_recv_next(
track: &mut track::Subscriber,
datagrams: bool,
emit_boundary: bool,
waiter: &kio::Waiter,
) -> Poll<Result<Recv, Error>> {
{
let mut groups_finished = false;
match track.poll_next_group(waiter) {
Poll::Ready(Ok(Some(group))) => return Poll::Ready(Ok(Recv::Group(group))),
Poll::Ready(Ok(None)) => groups_finished = true,
Poll::Ready(Err(err)) => return Poll::Ready(Err(err)),
Poll::Pending => {}
}
if datagrams {
match track.poll_recv_datagram(waiter) {
Poll::Ready(Ok(Some(datagram))) => return Poll::Ready(Ok(Recv::Datagram(datagram))),
Poll::Ready(Ok(None)) => {}
Poll::Ready(Err(err)) => return Poll::Ready(Err(err)),
Poll::Pending => {}
}
}
if emit_boundary && let Poll::Ready(res) = track.poll_finished(waiter) {
return Poll::Ready(res.map(Recv::Boundary));
}
if groups_finished {
return Poll::Ready(Ok(Recv::Finished));
}
Poll::Pending
}
}
#[cfg(test)]
async fn recv_next(track: &mut track::Subscriber, datagrams: bool, emit_boundary: bool) -> Result<Recv, Error> {
kio::wait(|waiter| poll_recv_next(track, datagrams, emit_boundary, waiter)).await
}
#[derive(Clone)]
struct Subscription<S: web_transport_trait::Session> {
session: S,
id: u64,
track_name: Arc<str>,
priority: PriorityQueue,
track_priority: kio::Consumer<u8>,
track_priority_seen: u8,
version: Version,
timescale: Option<crate::Timescale>,
}
impl<S: web_transport_trait::Session> Subscription<S> {
async fn run_track(
mut self,
mut track: track::Subscriber,
start_group: Option<u64>,
initial_end_group: Option<u64>,
reader: &mut crate::coding::Reader<S::RecvStream, Version>,
writer: &mut Writer<S::SendStream, Version>,
track_priority_tx: &kio::Producer<u8>,
) -> Result<(), Error> {
let mut tasks: FuturesUnordered<MaybeSendBox<'static, ()>> = FuturesUnordered::new();
if let Some(start_group) = start_group.or_else(|| track.latest()) {
track.start_at(start_group);
}
track.end_at(initial_end_group);
let emit_range = self.version.has_track_stream();
let mut start_sent = false;
let mut end_sent = false;
let datagrams = self.version.has_datagrams() && self.session.max_datagram_size() > 0;
#[allow(clippy::large_enum_variant)]
enum Event {
Recv(Result<Recv, Error>),
Update(Result<Option<lite::SubscribeUpdate>, Error>),
}
loop {
let event = {
let emit_boundary = emit_range && !end_sent;
let mut update = std::pin::pin!(reader.decode_maybe::<lite::SubscribeUpdate>());
kio::wait(|waiter| {
let mut cx = std::task::Context::from_waker(waiter.waker());
while let Poll::Ready(Some(())) = tasks.poll_next_unpin(&mut cx) {}
if let Poll::Ready(upd) = waiter.poll_future(update.as_mut()) {
return Poll::Ready(Event::Update(upd));
}
if let Poll::Ready(res) = poll_recv_next(&mut track, datagrams, emit_boundary, waiter) {
return Poll::Ready(Event::Recv(res));
}
Poll::Pending
})
.await
};
match event {
Event::Recv(res) => match res? {
Recv::Group(group) => {
if emit_range && !start_sent {
start_sent = true;
writer
.encode(&lite::SubscribeResponse::Start(lite::SubscribeStart {
group: group.sequence,
}))
.await?;
}
self.queue_serve(group, &mut tasks);
}
Recv::Datagram(datagram) => self.serve_datagram(datagram),
Recv::Boundary(group) => {
end_sent = true;
writer
.encode(&lite::SubscribeResponse::End(lite::SubscribeEnd { group }))
.await?;
}
Recv::Finished => {
while tasks.next().await.is_some() {}
return Ok(());
}
},
Event::Update(upd) => {
let Some(upd) = upd? else {
return Ok(());
};
if let Ok(mut value) = track_priority_tx.write() {
*value = upd.priority;
}
let _ = track.update(crate::track::Subscription {
priority: upd.priority,
ordered: upd.ordered,
latency_max: upd.max_latency,
group_start: upd.start_group,
group_end: upd.end_group,
..Default::default()
});
if let Some(start_group) = upd.start_group {
track.start_at(start_group);
}
track.end_at(upd.end_group);
}
}
}
}
fn queue_serve(&mut self, group: group::Consumer, tasks: &mut FuturesUnordered<MaybeSendBox<'static, ()>>) {
let sequence = group.sequence;
tracing::debug!(subscribe = self.id, track = %self.track_name, sequence, "serving group");
let current_priority = self.track_priority_current();
let handle = self.priority.insert(Priority::new(current_priority, sequence));
let fut = self.clone().serve_group(sequence, handle, group);
tasks.push(fut.map(|_| ()).maybe_boxed());
}
async fn serve_group(
mut self,
sequence: u64,
mut priority: PriorityHandle,
mut group: group::Consumer,
) -> Result<(), Error> {
let msg = lite::Group {
subscribe: self.id,
sequence,
};
let stream = self.session.open_uni().await.map_err(Error::from_transport)?;
let mut stream = Writer::new(stream, self.version);
stream.set_priority(priority.current());
stream.encode(&lite::DataType::Group).await?;
stream.encode(&msg).await?;
let mut prev_ts: u64 = 0;
while let Some(frame) = self.next_frame(&mut stream, &mut priority, &mut group).await? {
self.serve_frame(&mut stream, &mut priority, frame, &mut prev_ts)
.await?;
}
stream.finish()?;
stream.closed().await?;
tracing::debug!(sequence, "finished group");
Ok(())
}
fn serve_datagram(&self, datagram: crate::Datagram) {
let body = lite::Datagram {
subscribe: self.id,
sequence: datagram.sequence,
timestamp: datagram.timestamp.value(),
payload: datagram.payload,
};
let Ok(body) = body.encode_bytes(self.version) else {
return;
};
let max = self.session.max_datagram_size();
if body.len() > max {
tracing::debug!(
sequence = datagram.sequence,
size = body.len(),
max,
"dropping datagram larger than the transport limit"
);
return;
}
let _ = self.session.send_datagram(body);
}
async fn serve_frame(
&mut self,
stream: &mut Writer<S::SendStream, Version>,
priority: &mut PriorityHandle,
mut frame: frame::Consumer,
prev_ts: &mut u64,
) -> Result<(), Error> {
encode_frame_timing(stream, &frame, self.timescale, prev_ts).await?;
stream.encode(&frame.size).await?;
while let Some(chunk) = self.read_chunk(stream, priority, &mut frame).await? {
self.write_chunk(stream, priority, chunk).await?;
}
Ok(())
}
async fn next_frame(
&mut self,
stream: &mut Writer<S::SendStream, Version>,
priority: &mut PriorityHandle,
group: &mut group::Consumer,
) -> Result<Option<frame::Consumer>, Error> {
Self::serve_step(
stream,
priority,
&self.track_priority,
&mut self.track_priority_seen,
|waiter| group.poll_next_frame(waiter),
)
.await
}
async fn read_chunk(
&mut self,
stream: &mut Writer<S::SendStream, Version>,
priority: &mut PriorityHandle,
frame: &mut frame::Consumer,
) -> Result<Option<bytes::Bytes>, Error> {
Self::serve_step(
stream,
priority,
&self.track_priority,
&mut self.track_priority_seen,
|waiter| frame.poll_read_chunk(waiter),
)
.await
}
async fn serve_step<T>(
stream: &mut Writer<S::SendStream, Version>,
priority: &mut PriorityHandle,
track_priority: &kio::Consumer<u8>,
track_priority_seen: &mut u8,
mut work: impl FnMut(&kio::Waiter) -> Poll<Result<T, Error>>,
) -> Result<T, Error> {
enum Event<T> {
Closed,
Work(Result<T, Error>),
Priority(u8),
TrackPriority(u8),
}
loop {
let event = {
let mut closed = std::pin::pin!(stream.closed());
let seen = *track_priority_seen;
kio::wait(|waiter| {
if waiter.poll_future(closed.as_mut()).is_ready() {
return Poll::Ready(Event::Closed);
}
if let Poll::Ready(res) = work(waiter) {
return Poll::Ready(Event::Work(res));
}
if let Poll::Ready(new_pri) = priority.poll_next(waiter) {
return Poll::Ready(Event::Priority(new_pri));
}
match track_priority.poll(waiter, |value| {
if **value != seen {
Poll::Ready(**value)
} else {
Poll::Pending
}
}) {
Poll::Ready(Ok(value)) => Poll::Ready(Event::TrackPriority(value)),
Poll::Ready(Err(_)) | Poll::Pending => Poll::Pending,
}
})
.await
};
match event {
Event::Closed => return Err(Error::Cancel),
Event::Work(res) => return res,
Event::Priority(new_pri) => stream.set_priority(new_pri),
Event::TrackPriority(new_track) => {
*track_priority_seen = new_track;
priority.set_track(new_track);
}
}
}
}
fn track_priority_current(&mut self) -> u8 {
self.track_priority_seen = *self.track_priority.read();
self.track_priority_seen
}
async fn write_chunk(
&mut self,
stream: &mut Writer<S::SendStream, Version>,
priority: &mut PriorityHandle,
mut chunk: bytes::Bytes,
) -> Result<(), Error> {
while chunk.has_remaining() {
self.apply_priority(stream, priority);
stream.write(&mut chunk).await?;
}
Ok(())
}
fn apply_priority(&mut self, stream: &mut Writer<S::SendStream, Version>, priority: &mut PriorityHandle) {
let track_priority = self.track_priority_current();
priority.set_track(track_priority);
stream.set_priority(priority.current());
}
}