use crate::runtime::Timers as _;
use crate::{SessionError, announce, frame, group, origin, track};
use std::{
collections::HashMap,
ops::Bound,
sync::Arc,
task::{Context, Poll, ready},
time::Duration,
};
use web_transport_trait::Stats;
use crate::{
Error, Hop, Hops,
coding::{Encode, Stream, Writer},
lite::{
self,
priority::{Priority, PriorityHandle, PriorityQueue},
},
};
use super::Version;
pub(super) struct PublisherConfig<S: crate::transport::poll::Session> {
pub runtime: crate::time::Clock,
pub session: S,
pub origin: origin::Consumer,
pub version: Version,
pub peer_setup: super::PeerSetup,
pub goaway: crate::goaway::Protocol,
pub peer_hop: Option<Hop>,
}
struct Shared<S: crate::transport::poll::Session> {
session: S,
origin: origin::Consumer,
self_origin: Hop,
peer_setup: super::PeerSetup,
peer_hop: Option<Hop>,
serving: std::sync::OnceLock<origin::Consumer>,
priority: PriorityQueue,
version: Version,
goaway: crate::goaway::Protocol,
}
const MAX_SAFE_AGE_MS: u64 = (1_u64 << 53) - 1;
fn serving_max_age(version: Version, requested: Duration) -> Duration {
match version.carries_max_age() {
true => requested,
false => Duration::from_millis(MAX_SAFE_AGE_MS),
}
}
fn position_cursor(track: &mut track::Subscriber, version: Version, start_group: Option<u64>) {
if version.resolves_start() || start_group.is_some() {
return;
}
if let Some(latest) = track.latest() {
track.start_at(latest);
}
}
impl<S: crate::transport::poll::Session> Shared<S> {
fn poll_serving_origin(&self, waiter: &kio::Waiter) -> Poll<origin::Consumer> {
if let Some(origin) = self.serving.get() {
return Poll::Ready(origin.clone());
}
let declared = match self.version.has_setup_stream() {
true => ready!(self.peer_setup.poll_hop(waiter)),
false => None,
};
let origin = match declared.or(self.peer_hop) {
Some(peer) => self.origin.clone().excluding(peer),
None => self.origin.clone(),
};
Poll::Ready(self.serving.get_or_init(|| origin).clone())
}
}
pub(super) struct Publisher<S: crate::transport::poll::Session> {
shared: Arc<Shared<S>>,
runtime: crate::time::Clock,
accept: S,
children: kio::Tasks<Control<S>>,
}
impl<S: crate::transport::poll::Session> Publisher<S> {
pub fn new(config: PublisherConfig<S>) -> Self {
let self_origin = config.origin.hop();
let accept = config.session.clone();
Self {
shared: Arc::new(Shared {
session: config.session,
origin: config.origin,
self_origin,
peer_setup: config.peer_setup,
peer_hop: config.peer_hop,
serving: std::sync::OnceLock::new(),
priority: Default::default(),
version: config.version,
goaway: config.goaway,
}),
runtime: config.runtime,
accept,
children: kio::Tasks::new(),
}
}
}
impl<S> Publisher<S>
where
S: crate::transport::poll::Session,
{
pub fn poll(&mut self, waiter: &kio::Waiter) -> Poll<Result<(), Error>> {
let _ = self.children.poll(waiter);
let mut cx = Context::from_waker(waiter.waker());
loop {
match Stream::poll_accept(&mut self.accept, self.shared.version, &mut cx) {
Poll::Ready(Ok(stream)) => {
self.children.push(Control {
shared: self.shared.clone(),
runtime: self.runtime.clone(),
state: ControlState::Start { stream },
});
}
Poll::Ready(Err(err)) => return Poll::Ready(Err(err)),
Poll::Pending => break,
}
}
let _ = self.children.poll(waiter);
Poll::Pending
}
}
#[cfg(test)]
impl<S: crate::transport::poll::Session> Publisher<S> {
async fn run_announce(
stream: &mut Stream<S, Version>,
origin: &origin::Consumer,
announced: &mut announce::Consumer,
prefix: impl crate::AsPath,
self_origin: Hop,
version: Version,
) -> Result<(), Error> {
let mut run = AnnounceRun::new(prefix.as_path().to_owned(), self_origin, version);
kio::wait(|waiter| run.poll(stream, origin, announced, waiter)).await
}
}
struct Control<S: crate::transport::poll::Session> {
shared: Arc<Shared<S>>,
runtime: crate::time::Clock,
state: ControlState<S>,
}
#[allow(clippy::large_enum_variant)]
enum ControlState<S: crate::transport::poll::Session> {
Start {
stream: Stream<S, Version>,
},
Announce(AnnounceServe<S>),
Subscribe(SubscribeServe<S>),
Fetch(FetchServe<S>),
TrackInfo(TrackInfoServe<S>),
Probe(ProbeServe<S>),
Goaway {
stream: Stream<S, Version>,
},
Done,
}
impl<S: crate::transport::poll::Session> kio::Task for Control<S> {
fn poll(&mut self, waiter: &kio::Waiter) -> Poll<()> {
if let Err(err) = ready!(self.poll_serve(waiter)) {
tracing::warn!(%err, "control stream error");
}
Poll::Ready(())
}
}
impl<S: crate::transport::poll::Session> Control<S> {
fn poll_serve(&mut self, waiter: &kio::Waiter) -> Poll<Result<(), Error>> {
loop {
match &mut self.state {
ControlState::Start { stream } => {
let mut cx = Context::from_waker(waiter.waker());
let kind = ready!(stream.reader.poll_decode::<lite::ControlType>(&mut cx))?;
let ControlState::Start { stream } = std::mem::replace(&mut self.state, ControlState::Done) else {
unreachable!()
};
self.state = match kind {
lite::ControlType::Announce => {
ControlState::Announce(AnnounceServe::new(self.shared.clone(), stream))
}
lite::ControlType::Subscribe => {
ControlState::Subscribe(SubscribeServe::new(self.shared.clone(), stream))
}
lite::ControlType::Fetch => ControlState::Fetch(FetchServe::new(self.shared.clone(), stream)?),
lite::ControlType::Track => {
ControlState::TrackInfo(TrackInfoServe::new(self.shared.clone(), stream)?)
}
lite::ControlType::Probe => {
ControlState::Probe(ProbeServe::new(self.shared.clone(), self.runtime.clone(), stream))
}
lite::ControlType::Goaway => ControlState::Goaway { stream },
lite::ControlType::Session => return Poll::Ready(Err(Error::UnexpectedStream)),
};
}
ControlState::Announce(serve) => return serve.poll(waiter),
ControlState::Subscribe(serve) => return serve.poll(waiter),
ControlState::Fetch(serve) => return serve.poll(waiter),
ControlState::TrackInfo(serve) => return serve.poll(waiter),
ControlState::Probe(serve) => return serve.poll(waiter),
ControlState::Goaway { stream } => {
let mut cx = Context::from_waker(waiter.waker());
let msg = ready!(stream.reader.poll_decode::<lite::Goaway>(&mut cx))?;
tracing::info!(uri = %msg.uri, "received goaway");
let uri = msg.uri.into_owned();
let goaway = crate::goaway::Goaway {
uri: uri.clone(),
timeout: None,
};
if let Err(err) = self.shared.goaway.record(goaway) {
tracing::warn!(%uri, "duplicate GOAWAY received; closing session");
self.shared
.session
.clone()
.close(SessionError::from(&err).to_code(), &err.to_string());
return Poll::Ready(Err(err));
}
return Poll::Ready(Ok(()));
}
ControlState::Done => return Poll::Ready(Ok(())),
}
}
}
}
struct ProbeServe<S: crate::transport::poll::Session> {
shared: Arc<Shared<S>>,
runtime: crate::time::Clock,
stream: Option<Stream<S, Version>>,
last_sent: Option<(lite::Probe, crate::runtime::Instant)>,
next_probe: crate::runtime::Deadline<crate::time::Clock>,
}
impl<S: crate::transport::poll::Session> ProbeServe<S> {
const PROBE_INTERVAL: Duration = Duration::from_millis(100);
const PROBE_MAX_AGE: Duration = Duration::from_secs(10);
const PROBE_MAX_DELTA: f64 = 0.25;
const PROBE_RTT_DELTA: f64 = 0.25;
fn moved(prev: Option<u64>, next: Option<u64>, threshold: f64) -> bool {
match (prev, next) {
(None, None) => false,
(Some(prev), Some(next)) => {
if prev == 0 {
return next != 0;
}
(next as f64 - prev as f64).abs() / prev as f64 >= threshold
}
_ => true,
}
}
fn bitrate_threshold(elapsed: Duration) -> f64 {
let t = elapsed
.as_secs_f64()
.clamp(Self::PROBE_INTERVAL.as_secs_f64(), Self::PROBE_MAX_AGE.as_secs_f64());
let range = Self::PROBE_MAX_AGE.as_secs_f64() - Self::PROBE_INTERVAL.as_secs_f64();
Self::PROBE_MAX_DELTA * (Self::PROBE_MAX_AGE.as_secs_f64() - t) / range
}
fn new(shared: Arc<Shared<S>>, runtime: crate::time::Clock, stream: Stream<S, Version>) -> Self {
Self {
shared,
stream: Some(stream),
last_sent: None,
next_probe: crate::runtime::Deadline::at(&runtime, runtime.now()),
runtime,
}
}
fn poll(&mut self, waiter: &kio::Waiter) -> Poll<Result<(), Error>> {
match ready!(self.poll_probe(waiter)) {
Ok(()) => tracing::debug!("probe stream closed"),
Err(err) => {
tracing::warn!(%err, "probe stream error");
self.stream.take().expect("stream present").writer.abort(&err);
}
}
Poll::Ready(Ok(()))
}
fn poll_probe(&mut self, waiter: &kio::Waiter) -> Poll<Result<(), Error>> {
let stream = self.stream.as_mut().expect("stream present");
let mut cx = Context::from_waker(waiter.waker());
loop {
ready!(stream.writer.poll_flush(&mut cx))?;
if let Poll::Ready(res) = stream.reader.poll_closed(&mut cx) {
return Poll::Ready(res);
}
ready!(self.next_probe.poll(waiter));
let next = self
.next_probe
.deadline()
.and_then(|at| at.checked_add(Self::PROBE_INTERVAL));
self.next_probe.set(next);
let report = {
let stats = self.shared.session.stats();
lite::Probe {
bitrate: stats.estimated_send_rate(),
rtt: self
.shared
.version
.has_probe_rtt()
.then(|| stats.rtt().map(|d| d.as_millis() as u64))
.flatten(),
}
};
if report.bitrate.is_none() && report.rtt.is_none() {
let retracts = self
.last_sent
.as_ref()
.is_some_and(|(prev, _)| prev.bitrate.is_some() || prev.rtt.is_some());
if !retracts {
continue;
}
}
let should_send = match &self.last_sent {
None => true,
Some((prev, at)) => {
let elapsed = self.runtime.now().duration_since(*at);
elapsed >= Self::PROBE_MAX_AGE
|| Self::moved(prev.bitrate, report.bitrate, Self::bitrate_threshold(elapsed))
|| Self::moved(prev.rtt, report.rtt, Self::PROBE_RTT_DELTA)
}
};
if should_send {
stream.writer.buffer(&report)?;
self.last_sent = Some((report, self.runtime.now()));
}
}
}
}
struct AnnounceServe<S: crate::transport::poll::Session> {
shared: Arc<Shared<S>>,
stream: Option<Stream<S, Version>>,
state: AnnounceState,
}
#[allow(clippy::large_enum_variant)]
enum AnnounceState {
Decode,
ExcludeHop { prefix: crate::PathOwned, hidden: bool },
Run {
origin: origin::Consumer,
announced: announce::Consumer,
run: AnnounceRun,
},
}
impl<S: crate::transport::poll::Session> AnnounceServe<S> {
fn new(shared: Arc<Shared<S>>, stream: Stream<S, Version>) -> Self {
Self {
shared,
stream: Some(stream),
state: AnnounceState::Decode,
}
}
fn poll(&mut self, waiter: &kio::Waiter) -> Poll<Result<(), Error>> {
loop {
match &mut self.state {
AnnounceState::Decode => {
let stream = self.stream.as_mut().expect("stream present");
let mut cx = Context::from_waker(waiter.waker());
let interest = ready!(stream.reader.poll_decode::<lite::AnnounceRequest>(&mut cx))?;
let prefix = interest.prefix.to_owned();
let hidden = interest.hidden;
let assigned = self.shared.peer_hop.map(|origin| origin.id()).unwrap_or(0);
if self.shared.version.has_exclude_hop() {
let exclude_hop = match interest.exclude_hop {
0 => assigned,
id => id,
};
self.start(prefix, exclude_hop, hidden);
} else if self.shared.version.has_setup_stream() {
self.state = AnnounceState::ExcludeHop { prefix, hidden };
} else {
self.start(prefix, assigned, hidden);
}
}
AnnounceState::ExcludeHop { prefix, hidden } => {
let assigned = self.shared.peer_hop.map(|origin| origin.id()).unwrap_or(0);
let exclude_hop = ready!(self.shared.peer_setup.poll_hop(waiter))
.map(|origin| origin.id())
.unwrap_or(assigned);
let (prefix, hidden) = (prefix.clone(), *hidden);
self.start(prefix, exclude_hop, hidden);
}
AnnounceState::Run { origin, announced, run } => {
let stream = self.stream.as_mut().expect("stream present");
let res = ready!(run.poll(stream, origin, announced, waiter));
if let Err(err) = res {
match &err {
Error::Cancel
| Error::Stream(crate::StreamError::Cancel)
| Error::Session(crate::SessionError::Cancel)
| Error::Transport(_) => {
tracing::debug!(prefix = %origin.absolute(&run.prefix), "announcing cancelled");
}
err => {
tracing::warn!(%err, prefix = %origin.absolute(&run.prefix), "announcing error");
}
}
self.stream.take().expect("stream present").writer.abort(&err);
}
return Poll::Ready(Ok(()));
}
}
}
}
fn start(&mut self, prefix: crate::PathOwned, exclude_hop: u64, hidden: bool) {
let scope = crate::Pattern::subtree(prefix.as_str())
.map(crate::Patterns::from)
.unwrap_or_default();
let origin = self
.shared
.origin
.scope("", &scope)
.unwrap_or_else(|_| self.shared.origin.empty());
let origin = origin.excluding(Hop::new(exclude_hop).unwrap_or(Hop::UNKNOWN));
let origin = match hidden {
true => origin.with_hidden(true),
false => origin,
};
let announced = origin.announced();
let run = AnnounceRun::new(prefix, self.shared.self_origin, self.shared.version);
self.state = AnnounceState::Run { origin, announced, run };
}
}
struct AnnounceRun {
prefix: crate::PathOwned,
self_origin: Hop,
version: Version,
next_announce_id: u64,
live: HashMap<crate::PathOwned, Option<u64>>,
phase: AnnouncePhase,
}
enum AnnouncePhase {
Init,
Running,
Closing,
}
impl AnnounceRun {
fn new(prefix: crate::PathOwned, self_origin: Hop, version: Version) -> Self {
Self {
prefix,
self_origin,
version,
next_announce_id: 0,
live: HashMap::new(),
phase: AnnouncePhase::Init,
}
}
fn suffix(&self, update: &announce::Update) -> crate::PathOwned {
update
.prefix
.strip_prefix(&self.prefix)
.expect("origin returned a route outside the requested prefix")
.to_owned()
}
fn outgoing(&self, route: &crate::origin::Route, absolute: &crate::Path) -> Option<(Hops, crate::origin::Cost)> {
let mut hops = route.hops.clone();
if self.self_origin != Hop::UNKNOWN && hops.contains(&self.self_origin) {
tracing::debug!(route = %absolute, "dropping reflected route");
return None;
}
if !self.version.has_announce_ok() && hops.push(self.self_origin).is_err() {
tracing::warn!(route = %absolute, "dropping announce; hop chain at MAX_HOPS (possible loop)");
return None;
}
let cost = match self.version.has_route_cost() {
true => route.cost.clamped(),
false => crate::origin::Cost::UNKNOWN,
};
Some((hops, cost))
}
fn assign_id(&mut self) -> Option<u64> {
if !self.version.has_announce_id() {
return None;
}
let id = self.next_announce_id;
self.next_announce_id += 1;
Some(id)
}
fn retract<S: crate::transport::poll::Session>(
&mut self,
stream: &mut Stream<S, Version>,
suffix: crate::PathOwned,
absolute: &crate::Path,
) -> Result<(), Error> {
let Some(id) = self.live.remove(&suffix) else {
return Ok(());
};
tracing::debug!(route = %absolute, "unannounce");
match id {
Some(id) => stream.writer.buffer(&lite::AnnounceBroadcast::EndedId { id })?,
None => stream.writer.buffer(&lite::AnnounceBroadcast::Ended {
suffix,
hops: Hops::new(),
})?,
}
Ok(())
}
fn init<S: crate::transport::poll::Session>(
&mut self,
stream: &mut Stream<S, Version>,
origin: &origin::Consumer,
announced: &mut announce::Consumer,
) -> Result<(), Error> {
match self.version {
Version::Lite01 | Version::Lite02 => {
let mut init = Vec::new();
while let Some(update) = announced.try_next() {
let absolute = origin.absolute(&update.prefix);
let suffix = self.suffix(&update);
if update.kind.is_active() {
if self.outgoing(&update.route, &absolute).is_none() {
continue;
}
tracing::debug!(route = %absolute, "announce");
if !init.contains(&suffix) {
init.push(suffix);
}
} else {
tracing::debug!(route = %absolute, "unannounce");
init.retain(|p| p != &suffix);
}
}
let announce_init = lite::AnnounceInit { suffixes: init };
stream.writer.buffer(&announce_init)?;
}
_ if self.version.has_announce_ok() => {
let mut initial: Vec<(crate::PathOwned, Hops, crate::origin::Cost)> = Vec::new();
while let Some(update) = announced.try_next() {
let absolute = origin.absolute(&update.prefix);
let suffix = self.suffix(&update);
if update.kind.is_active() {
let Some((hops, cost)) = self.outgoing(&update.route, &absolute) else {
continue;
};
tracing::debug!(route = %absolute, "announce");
initial.retain(|(s, ..)| s != &suffix);
initial.push((suffix, hops, cost));
} else {
tracing::debug!(route = %absolute, "unannounce");
initial.retain(|(s, ..)| s != &suffix);
}
}
let ok = lite::AnnounceOk {
origin: self.self_origin,
active: initial.len() as u64,
};
stream.writer.buffer(&ok)?;
for (suffix, hops, cost) in initial {
let id = self.assign_id();
self.live.insert(suffix.clone(), id);
stream
.writer
.buffer(&lite::AnnounceBroadcast::Active { suffix, hops, cost })?;
}
}
_ => {
}
}
Ok(())
}
fn poll<S: crate::transport::poll::Session>(
&mut self,
stream: &mut Stream<S, Version>,
origin: &origin::Consumer,
announced: &mut announce::Consumer,
waiter: &kio::Waiter,
) -> Poll<Result<(), Error>> {
let mut cx = Context::from_waker(waiter.waker());
if matches!(self.phase, AnnouncePhase::Init) {
self.init(stream, origin, announced)?;
self.phase = AnnouncePhase::Running;
}
loop {
ready!(stream.writer.poll_flush(&mut cx))?;
if matches!(self.phase, AnnouncePhase::Closing) {
return stream.writer.poll_closed(&mut cx);
}
if let Poll::Ready(res) = stream.reader.poll_closed(&mut cx) {
return Poll::Ready(res);
}
let Poll::Ready(next) = announced.poll_next(waiter) else {
return Poll::Pending;
};
let Some(update) = next else {
stream.writer.finish()?;
self.phase = AnnouncePhase::Closing;
continue;
};
let absolute = origin.absolute(&update.prefix);
let suffix = self.suffix(&update);
if !update.kind.is_active() {
self.retract(stream, suffix, &absolute)?;
continue;
}
match self.outgoing(&update.route, &absolute) {
Some((hops, cost)) => match self.live.get(&suffix) {
Some(&id) if lite::restart_supported(self.version) => {
tracing::debug!(route = %absolute, "reannounce");
match id {
Some(id) => stream
.writer
.buffer(&lite::AnnounceBroadcast::Restart { id, hops, cost })?,
None => stream
.writer
.buffer(&lite::AnnounceBroadcast::Active { suffix, hops, cost })?,
}
}
Some(_) => {}
None => {
tracing::debug!(route = %absolute, "announce");
let id = self.assign_id();
self.live.insert(suffix.clone(), id);
stream
.writer
.buffer(&lite::AnnounceBroadcast::Active { suffix, hops, cost })?;
}
},
None => self.retract(stream, suffix, &absolute)?,
}
}
}
}
struct TrackInfoServe<S: crate::transport::poll::Session> {
shared: Arc<Shared<S>>,
stream: Option<Stream<S, Version>>,
state: TrackInfoState,
absolute: crate::PathOwned,
track: String,
}
enum TrackInfoState {
Decode,
Hop {
msg: lite::Track<'static>,
},
Request {
msg: lite::Track<'static>,
requesting: origin::Requesting,
},
Query {
querying: track::Querying,
},
Finish {
finished: bool,
},
}
impl<S: crate::transport::poll::Session> TrackInfoServe<S> {
fn new(shared: Arc<Shared<S>>, stream: Stream<S, Version>) -> Result<Self, Error> {
if !shared.version.has_track_stream() {
return Err(Error::UnexpectedStream);
}
Ok(Self {
shared,
stream: Some(stream),
state: TrackInfoState::Decode,
absolute: Default::default(),
track: Default::default(),
})
}
fn poll(&mut self, waiter: &kio::Waiter) -> Poll<Result<(), Error>> {
match ready!(self.poll_serve(waiter)) {
Ok(()) => Poll::Ready(Ok(())),
Err(err) if matches!(self.state, TrackInfoState::Decode) => Poll::Ready(Err(err)),
Err(err) => {
match &err {
Error::Cancel
| Error::Stream(crate::StreamError::Cancel)
| Error::Session(crate::SessionError::Cancel)
| Error::Transport(_) => {
tracing::debug!(broadcast = %self.absolute, track = %self.track, "track info cancelled")
}
err => {
tracing::warn!(broadcast = %self.absolute, track = %self.track, %err, "track info error")
}
}
self.stream.take().expect("stream present").writer.abort(&err);
Poll::Ready(Ok(()))
}
}
}
fn poll_serve(&mut self, waiter: &kio::Waiter) -> Poll<Result<(), Error>> {
loop {
match &mut self.state {
TrackInfoState::Decode => {
let stream = self.stream.as_mut().expect("stream present");
let mut cx = Context::from_waker(waiter.waker());
let msg = ready!(stream.reader.poll_decode::<lite::Track>(&mut cx))?;
self.absolute = self.shared.origin.absolute(&msg.broadcast).to_owned();
self.track = msg.track.to_string();
tracing::debug!(broadcast = %self.absolute, track = %self.track, "track info requested");
self.state = TrackInfoState::Hop { msg };
}
TrackInfoState::Hop { .. } => {
let origin = ready!(self.shared.poll_serving_origin(waiter));
let TrackInfoState::Hop { msg } = std::mem::replace(&mut self.state, TrackInfoState::Decode) else {
unreachable!()
};
let requesting = origin.request_broadcast(&msg.broadcast).into_inner();
self.state = TrackInfoState::Request { msg, requesting };
}
TrackInfoState::Request { msg, requesting } => {
let broadcast = ready!(requesting.poll_ok(waiter))?;
let querying = broadcast.track(&msg.track)?.query().into_inner();
self.state = TrackInfoState::Query { querying };
}
TrackInfoState::Query { querying } => {
let info = ready!(querying.poll_ok(waiter))?;
let stream = self.stream.as_mut().expect("stream present");
stream.writer.buffer(&lite::TrackInfo {
priority: info.priority,
max_age: info.max_age,
timescale: info.timescale,
})?;
self.state = TrackInfoState::Finish { finished: false };
}
TrackInfoState::Finish { finished } => {
let stream = self.stream.as_mut().expect("stream present");
let mut cx = Context::from_waker(waiter.waker());
if !*finished {
ready!(stream.writer.poll_flush(&mut cx))?;
stream.writer.finish()?;
*finished = true;
}
return stream.writer.poll_closed(&mut cx);
}
}
}
}
}
struct SubscribeServe<S: crate::transport::poll::Session> {
shared: Arc<Shared<S>>,
stream: Option<Stream<S, Version>>,
state: SubscribeState<S>,
id: u64,
absolute: crate::PathOwned,
track: String,
}
enum SubscribeState<S: crate::transport::poll::Session> {
Decode,
Hop {
msg: lite::Subscribe<'static>,
},
Request {
msg: lite::Subscribe<'static>,
requesting: origin::Requesting,
},
Confirm {
msg: lite::Subscribe<'static>,
subscribing: track::Subscribing,
},
Run(Box<TrackRun<S>>),
Drain {
children: kio::Tasks<GroupServe<S>>,
},
Finish {
finished: bool,
},
}
impl<S: crate::transport::poll::Session> SubscribeServe<S> {
fn new(shared: Arc<Shared<S>>, stream: Stream<S, Version>) -> Self {
Self {
shared,
stream: Some(stream),
state: SubscribeState::Decode,
id: 0,
absolute: Default::default(),
track: Default::default(),
}
}
fn poll(&mut self, waiter: &kio::Waiter) -> Poll<Result<(), Error>> {
match ready!(self.poll_serve(waiter)) {
Ok(()) => {
tracing::info!(id = self.id, broadcast = %self.absolute, track = %self.track, "subscribed complete");
Poll::Ready(Ok(()))
}
Err(err) if matches!(self.state, SubscribeState::Decode) => Poll::Ready(Err(err)),
Err(err) => {
match &err {
Error::Cancel
| Error::Stream(crate::StreamError::Cancel)
| Error::Session(crate::SessionError::Cancel)
| Error::Transport(_) => {
tracing::info!(id = self.id, broadcast = %self.absolute, track = %self.track, "subscribed cancelled")
}
err => {
tracing::warn!(id = self.id, broadcast = %self.absolute, track = %self.track, %err, "subscribed error")
}
}
self.stream.take().expect("stream present").writer.abort(&err);
Poll::Ready(Ok(()))
}
}
}
fn poll_serve(&mut self, waiter: &kio::Waiter) -> Poll<Result<(), Error>> {
loop {
match &mut self.state {
SubscribeState::Decode => {
let stream = self.stream.as_mut().expect("stream present");
let mut cx = Context::from_waker(waiter.waker());
let msg = ready!(stream.reader.poll_decode::<lite::Subscribe>(&mut cx))?;
self.id = msg.id;
self.absolute = self.shared.origin.absolute(&msg.broadcast).to_owned();
self.track = msg.track.to_string();
tracing::info!(id = self.id, broadcast = %self.absolute, track = %self.track, "subscribed started");
self.state = SubscribeState::Hop { msg };
}
SubscribeState::Hop { .. } => {
let origin = ready!(self.shared.poll_serving_origin(waiter));
let SubscribeState::Hop { msg } = std::mem::replace(&mut self.state, SubscribeState::Decode) else {
unreachable!()
};
let requesting = origin.request_broadcast(&msg.broadcast).into_inner();
self.state = SubscribeState::Request { msg, requesting };
}
SubscribeState::Request { requesting, .. } => {
let broadcast = ready!(requesting.poll_ok(waiter))?;
let SubscribeState::Request { msg, .. } =
std::mem::replace(&mut self.state, SubscribeState::Decode)
else {
unreachable!()
};
let subscription = crate::track::Subscription {
priority: msg.priority,
max_age: serving_max_age(self.shared.version, msg.max_age),
..Bounds::from(&msg).positions()
};
let track_consumer = broadcast.track(&msg.track)?;
let subscribing = track_consumer.subscribe(subscription).into_inner();
self.state = SubscribeState::Confirm { msg, subscribing };
}
SubscribeState::Confirm { subscribing, .. } => {
let track = ready!(subscribing.poll_ok(waiter))?;
let SubscribeState::Confirm { msg, .. } =
std::mem::replace(&mut self.state, SubscribeState::Decode)
else {
unreachable!()
};
let stream = self.stream.as_mut().expect("stream present");
let timescale = if self.shared.version.has_track_stream() {
Some(track.info().timescale)
} else {
None
};
if !self.shared.version.has_track_stream() {
let info = lite::SubscribeOk {
priority: msg.priority,
max_age: Duration::ZERO,
start_group: None,
end_group: None,
};
stream.writer.buffer(&lite::SubscribeResponse::Ok(info))?;
}
let track_priority_tx = kio::Producer::new(msg.priority);
let sub = Subscription {
session: self.shared.session.clone(),
id: msg.id,
track_name: Arc::from(track.name()),
priority: self.shared.priority.clone(),
track_priority: track_priority_tx.consume(),
track_priority_seen: msg.priority,
version: self.shared.version,
timescale,
};
let run = TrackRun::new(sub, track, Bounds::from(&msg), track_priority_tx);
self.state = SubscribeState::Run(Box::new(run));
}
SubscribeState::Run(run) => {
let stream = self.stream.as_mut().expect("stream present");
match ready!(run.poll(stream, waiter))? {
TrackEnd::PeerFin => self.state = SubscribeState::Finish { finished: false },
TrackEnd::Finished => {
let SubscribeState::Run(run) = std::mem::replace(&mut self.state, SubscribeState::Decode)
else {
unreachable!()
};
self.state = SubscribeState::Drain { children: run.children };
}
}
}
SubscribeState::Drain { children } => {
ready!(children.poll(waiter));
self.state = SubscribeState::Finish { finished: false };
}
SubscribeState::Finish { finished } => {
let stream = self.stream.as_mut().expect("stream present");
let mut cx = Context::from_waker(waiter.waker());
if !*finished {
ready!(stream.writer.poll_flush(&mut cx))?;
stream.writer.finish()?;
*finished = true;
}
return stream.writer.poll_closed(&mut cx);
}
}
}
}
}
struct FetchServe<S: crate::transport::poll::Session> {
shared: Arc<Shared<S>>,
stream: Option<Stream<S, Version>>,
state: FetchState,
absolute: crate::PathOwned,
track: String,
group: u64,
}
#[allow(clippy::large_enum_variant)]
enum FetchState {
Decode,
Hop {
msg: lite::Fetch<'static>,
},
Request {
msg: lite::Fetch<'static>,
requesting: origin::Requesting,
},
Fetch {
msg: lite::Fetch<'static>,
fetching: track::Fetching,
},
Serve {
group: group::Consumer,
timescale: Option<crate::Timescale>,
prev_ts: u64,
frame: Option<frame::Consumer>,
chunk: Option<bytes::Bytes>,
batch: Box<frame::Buffer>,
batch_pos: usize,
},
Finish {
finished: bool,
},
}
impl<S: crate::transport::poll::Session> FetchServe<S> {
fn new(shared: Arc<Shared<S>>, stream: Stream<S, Version>) -> Result<Self, Error> {
if !shared.version.has_track_stream() {
return Err(Error::UnexpectedStream);
}
Ok(Self {
shared,
stream: Some(stream),
state: FetchState::Decode,
absolute: Default::default(),
track: Default::default(),
group: 0,
})
}
fn poll(&mut self, waiter: &kio::Waiter) -> Poll<Result<(), Error>> {
match ready!(self.poll_serve(waiter)) {
Ok(()) => {
tracing::info!(broadcast = %self.absolute, track = %self.track, group = %self.group, "fetch complete");
Poll::Ready(Ok(()))
}
Err(err) if matches!(self.state, FetchState::Decode) => Poll::Ready(Err(err)),
Err(err) => {
match &err {
Error::Cancel
| Error::Stream(crate::StreamError::Cancel)
| Error::Session(crate::SessionError::Cancel)
| Error::Transport(_) => {
tracing::info!(broadcast = %self.absolute, track = %self.track, group = %self.group, "fetch cancelled")
}
err => {
tracing::warn!(broadcast = %self.absolute, track = %self.track, group = %self.group, %err, "fetch error")
}
}
self.stream.take().expect("stream present").writer.abort(&err);
Poll::Ready(Ok(()))
}
}
}
fn poll_serve(&mut self, waiter: &kio::Waiter) -> Poll<Result<(), Error>> {
loop {
match &mut self.state {
FetchState::Decode => {
let stream = self.stream.as_mut().expect("stream present");
let mut cx = Context::from_waker(waiter.waker());
let msg = ready!(stream.reader.poll_decode::<lite::Fetch>(&mut cx))?;
self.absolute = self.shared.origin.absolute(&msg.broadcast).to_owned();
self.track = msg.track.to_string();
self.group = msg.group;
tracing::info!(broadcast = %self.absolute, track = %self.track, group = %self.group, "fetch started");
self.state = FetchState::Hop { msg };
}
FetchState::Hop { .. } => {
let origin = ready!(self.shared.poll_serving_origin(waiter));
let FetchState::Hop { msg } = std::mem::replace(&mut self.state, FetchState::Decode) else {
unreachable!()
};
let requesting = origin.request_broadcast(&msg.broadcast).into_inner();
self.state = FetchState::Request { msg, requesting };
}
FetchState::Request { requesting, .. } => {
let broadcast = ready!(requesting.poll_ok(waiter))?;
let FetchState::Request { msg, .. } = std::mem::replace(&mut self.state, FetchState::Decode) else {
unreachable!()
};
let track = broadcast.track(&msg.track)?;
let fetching = track
.fetch_group(
msg.group,
group::Fetch {
priority: msg.priority,
frame_start: msg.start_frame,
..Default::default()
},
)
.into_inner();
self.state = FetchState::Fetch { msg, fetching };
}
FetchState::Fetch { msg, fetching } => {
let mut group = ready!(kio::Pollable::poll(fetching, waiter))?;
if group.index() != msg.start_frame {
return Poll::Ready(Err(Error::Lagged));
}
group.end_at(msg.end_frame.map_or(Bound::Unbounded, Bound::Included));
let timescale = if self.shared.version.has_track_stream() {
Some(group.timescale())
} else {
None
};
self.state = FetchState::Serve {
group,
timescale,
prev_ts: 0,
frame: None,
chunk: None,
batch: Box::new(frame::Buffer::new()),
batch_pos: 0,
};
}
FetchState::Serve {
group,
timescale,
prev_ts,
frame,
chunk,
batch,
batch_pos,
} => {
let stream = self.stream.as_mut().expect("stream present");
let mut cx = Context::from_waker(waiter.waker());
loop {
ready!(stream.writer.poll_flush(&mut cx))?;
if let Some(pending) = chunk {
ready!(stream.writer.poll_write(&mut cx, pending))?;
if !bytes::Buf::has_remaining(pending) {
*chunk = None;
}
} else if let Some(pending) = frame {
match ready!(pending.poll_read_chunk(waiter))? {
Some(next) => *chunk = Some(next),
None => *frame = None,
}
} else if *batch_pos < batch.len() {
let batched = &mut batch.filled_mut()[*batch_pos];
buffer_frame_info(
&mut stream.writer,
batched.timestamp,
batched.payload.len() as u64,
*timescale,
prev_ts,
)?;
let payload = std::mem::take(&mut batched.payload);
if !payload.is_empty() {
*chunk = Some(payload);
}
*batch_pos += 1;
group.keep_alive();
} else {
match group.poll_read_frames(waiter, batch) {
Poll::Ready(Ok(count)) if count > 0 => {
*batch_pos = 0;
continue;
}
Poll::Ready(Ok(_)) => break,
Poll::Ready(Err(err)) => return Poll::Ready(Err(err)),
Poll::Pending => {}
}
match ready!(group.poll_next_frame(waiter))? {
Some(next) => {
buffer_frame_info(
&mut stream.writer,
next.timestamp,
next.size,
*timescale,
prev_ts,
)?;
*frame = Some(next);
}
None => break,
}
}
}
self.state = FetchState::Finish { finished: false };
}
FetchState::Finish { finished } => {
let stream = self.stream.as_mut().expect("stream present");
let mut cx = Context::from_waker(waiter.waker());
if !*finished {
ready!(stream.writer.poll_flush(&mut cx))?;
stream.writer.finish()?;
*finished = true;
}
return stream.writer.poll_closed(&mut cx);
}
}
}
}
}
#[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)
}
#[test]
fn a_pre06_wire_is_pinned_to_the_live_edge() {
use futures::FutureExt;
let producer = track_producer("test");
for second in 0..3 {
let mut group = producer.append_group().unwrap();
group
.write_frame(Timestamp::from_millis(second * 1000).unwrap(), b"x".to_vec())
.unwrap();
group.finish().unwrap();
}
let served = |version| {
producer.subscribe(
track::Subscription::default().with_max_age(serving_max_age(version, std::time::Duration::ZERO)),
)
};
let drain = |subscriber: &mut track::Subscriber| {
let mut sequences = Vec::new();
while let Some(Ok(Some(group))) = subscriber.recv_group().now_or_never() {
sequences.push(group.sequence);
}
sequences
};
let mut unpinned = served(Version::Lite01);
assert_eq!(
drain(&mut unpinned),
vec![0, 1, 2],
"the unbounded budget alone resolves to the whole cache"
);
let mut legacy = served(Version::Lite01);
position_cursor(&mut legacy, Version::Lite01, None);
assert_eq!(drain(&mut legacy), vec![2]);
let mut tolerant =
producer.subscribe(track::Subscription::default().with_max_age(std::time::Duration::from_secs(5)));
position_cursor(&mut tolerant, Version::Lite05, None);
assert_eq!(drain(&mut tolerant), vec![2]);
let mut declared =
producer.subscribe(track::Subscription::default().with_max_age(std::time::Duration::from_secs(5)));
position_cursor(&mut declared, Version::Lite06, None);
assert_eq!(drain(&mut declared), vec![0, 1, 2]);
}
#[tokio::test]
async fn start_floor_suppresses_late_lower_arrivals() {
use futures::FutureExt;
let mut producer = track_producer("test");
let mut subscriber = producer.subscribe(None);
let write = |producer: &mut track::Producer, sequence: u64| {
let mut group = producer.create_group(crate::group::Info { sequence }).unwrap();
group
.write_frame(Timestamp::from_millis(1).unwrap(), b"x".to_vec())
.unwrap();
group.finish().unwrap();
};
write(&mut producer, 7);
match recv_next(&mut subscriber, false, false).await.unwrap() {
Recv::Group(group) => {
assert_eq!(group.sequence, 7);
subscriber.start_at(group.sequence);
}
_ => panic!("expected the first group"),
}
write(&mut producer, 5);
assert!(
recv_next(&mut subscriber, false, false).now_or_never().is_none(),
"a group below the resolved start must be suppressed"
);
}
#[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"),
}
}
#[tokio::test]
async fn recv_next_serves_late_arrival_after_newer_group() {
use futures::FutureExt;
let mut producer = track_producer("test");
let mut subscriber = producer.subscribe(track::Subscription::default().with_max_age(Duration::from_secs(5)));
producer.create_group(group::Info { sequence: 2 }).unwrap();
match recv_next(&mut subscriber, false, false).await.unwrap() {
Recv::Group(group) => assert_eq!(group.sequence, 2),
_ => panic!("expected group 2"),
}
producer.create_group(group::Info { sequence: 1 }).unwrap();
match recv_next(&mut subscriber, false, false).now_or_never() {
Some(Ok(Recv::Group(group))) => assert_eq!(group.sequence, 1),
Some(_) => panic!("expected the late-arriving group"),
None => panic!("the late-arriving group was skipped"),
}
producer.finish_at(3).unwrap();
match recv_next(&mut subscriber, false, false).await.unwrap() {
Recv::Finished => {}
_ => panic!("expected finished"),
}
}
}
#[cfg(test)]
mod announce_test {
use super::*;
use crate::coding::{Decode, Reader};
use crate::lite::test_transport::*;
use crate::model::ProduceTest;
use std::sync::Mutex;
type TestPublisher = Publisher<SinkSession>;
const VERSION: Version = Version::Lite06;
fn pub_hops() -> Hops {
Hops::try_from(vec![Hop::new(9).unwrap()]).unwrap()
}
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 },
lite::AnnounceBroadcast::Skipped => lite::AnnounceBroadcast::Skipped,
}
}
struct Harness {
origin: origin::Producer,
announcement: crate::model::AnnounceProducer,
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 settle() {
tokio::time::sleep(Duration::from_millis(1)).await;
}
async fn harness() -> Harness {
let origin = Hop::new(1).unwrap().produce();
let announcement = origin
.announce(
"cam",
crate::origin::Route::default().with_hops(pub_hops()).with_cost(7),
)
.unwrap();
let log = Log::default();
let writes = log.writes.clone();
let consumer = origin.consume();
let mut stream = Stream::<SinkSession, Version> {
writer: Writer::new(SinkSend::new(log), VERSION),
reader: Reader::new(PendingRecv, VERSION),
};
let task = tokio::spawn(async move {
let mut announced = consumer.announced();
let self_origin = consumer.hop();
TestPublisher::run_announce(&mut stream, &consumer, &mut announced, "", self_origin, VERSION).await
});
settle().await;
let mut wire = Wire { writes, cursor: 0 };
assert_eq!(wire.take_ok().active, 1, "expected one initial announce");
match wire.take_announces().as_slice() {
[lite::AnnounceBroadcast::Active { suffix, hops, cost }] => {
assert_eq!(suffix.as_str(), "cam");
assert_eq!(hops, &pub_hops());
assert_eq!(*cost, crate::origin::Cost::new(7));
}
other => panic!("expected the initial announce, got {other:?}"),
}
Harness {
origin,
announcement,
wire,
task,
}
}
#[tokio::test(start_paused = true)]
async fn announce_and_retract() {
let mut h = harness().await;
let late = h
.origin
.announce("mic", crate::origin::Route::default().with_hops(pub_hops()))
.unwrap();
settle().await;
match h.wire.take_announces().as_slice() {
[lite::AnnounceBroadcast::Active { suffix, .. }] => assert_eq!(suffix.as_str(), "mic"),
other => panic!("expected an announce, got {other:?}"),
}
drop(late);
settle().await;
match h.wire.take_announces().as_slice() {
[lite::AnnounceBroadcast::EndedId { id: 1 }] => {}
other => panic!("expected the retraction, got {other:?}"),
}
h.assert_idle();
}
#[tokio::test(start_paused = true)]
async fn update_restarts_in_place() {
let mut h = harness().await;
h.announcement
.update(crate::origin::Route::default().with_hops(pub_hops()).with_cost(3))
.unwrap();
settle().await;
match h.wire.take_announces().as_slice() {
[lite::AnnounceBroadcast::Restart { id: 0, hops, cost }] => {
assert_eq!(hops, &pub_hops());
assert_eq!(*cost, crate::origin::Cost::new(3));
}
other => panic!("expected a restart, got {other:?}"),
}
h.assert_idle();
}
#[tokio::test(start_paused = true)]
async fn identical_update_is_quiet() {
let h = harness().await;
h.announcement
.update(crate::origin::Route::default().with_hops(pub_hops()).with_cost(7))
.unwrap();
settle().await;
h.assert_idle();
}
#[tokio::test(start_paused = true)]
async fn excluded_routes_are_filtered() {
let peer = Hop::new(42).unwrap();
let origin = Hop::new(1).unwrap().produce();
let tainted = Hops::try_from(vec![peer]).unwrap();
let _tainted = origin
.announce("echoed", crate::origin::Route::default().with_hops(tainted))
.unwrap();
let _clean = origin
.announce("local", crate::origin::Route::default().with_hops(pub_hops()))
.unwrap();
let log = Log::default();
let writes = log.writes.clone();
let consumer = origin.consume().excluding(peer);
let mut stream = Stream::<SinkSession, Version> {
writer: Writer::new(SinkSend::new(log), VERSION),
reader: Reader::new(PendingRecv, VERSION),
};
let task = tokio::spawn(async move {
let mut announced = consumer.announced();
let self_origin = consumer.hop();
TestPublisher::run_announce(&mut stream, &consumer, &mut announced, "", self_origin, VERSION).await
});
settle().await;
let mut wire = Wire { writes, cursor: 0 };
assert_eq!(wire.take_ok().active, 1, "only the clean route is announced");
match wire.take_announces().as_slice() {
[lite::AnnounceBroadcast::Active { suffix, .. }] => assert_eq!(suffix.as_str(), "local"),
other => panic!("expected the clean announce, got {other:?}"),
}
task.abort();
}
#[tokio::test(start_paused = true)]
async fn late_route_emits_announce_start() {
let mut h = harness().await;
let _late = h
.origin
.announce("mic", crate::origin::Route::default().with_hops(pub_hops()))
.unwrap();
settle().await;
match h.wire.take_announces().as_slice() {
[lite::AnnounceBroadcast::Active { suffix, .. }] => assert_eq!(suffix.as_str(), "mic"),
other => panic!("expected ANNOUNCE_START, got {other:?}"),
}
h.assert_idle();
}
#[tokio::test(start_paused = true)]
async fn cost_clamps_to_the_wire_ceiling() {
let mut h = harness().await;
h.announcement
.update(
crate::origin::Route::default()
.with_hops(pub_hops())
.with_cost(u64::MAX),
)
.unwrap();
settle().await;
match h.wire.take_announces().as_slice() {
[lite::AnnounceBroadcast::Restart { cost, .. }] => {
assert_eq!(*cost, crate::origin::Cost::MAX);
}
other => panic!("expected a clamped restart, got {other:?}"),
}
}
}
fn buffer_frame_info<W: crate::transport::poll::SendStream>(
writer: &mut Writer<W, Version>,
timestamp: crate::Timestamp,
size: u64,
timescale: Option<crate::Timescale>,
prev_ts: &mut u64,
) -> Result<(), Error> {
if timescale.is_some() {
buffer_zigzag_delta(writer, timestamp.value(), prev_ts)?;
}
writer.buffer(&size)?;
Ok(())
}
fn buffer_zigzag_delta<W: crate::transport::poll::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.buffer(&zz)?;
*prev = curr;
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_recv_group(waiter)? {
Poll::Ready(Some(group)) => return Poll::Ready(Ok(Recv::Group(group))),
Poll::Ready(None) => groups_finished = true,
Poll::Pending => {}
}
if datagrams {
match track.poll_recv_datagram(waiter)? {
Poll::Ready(Some(datagram)) => return Poll::Ready(Ok(Recv::Datagram(datagram))),
Poll::Ready(None) => {}
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
}
fn position_group(group: &mut group::Consumer, start: Option<(u64, u64)>, end: Option<(u64, u64)>) -> bool {
let expected = match start {
Some((sequence, frame)) if sequence == group.sequence => frame,
_ => 0,
};
group.start_at(expected);
if group.index() != expected {
return false;
}
if let Some((sequence, frame)) = end
&& sequence == group.sequence
{
group.end_at(Bound::Included(frame));
}
true
}
struct Bounds {
start_group: Option<u64>,
start_frame: u64,
end_group: Option<u64>,
end_frame: Option<u64>,
}
impl Bounds {
fn positions(&self) -> crate::track::Subscription {
let mut sub = crate::track::Subscription::default();
if let Some(group) = self.start_group {
sub = sub.with_start(track::Position {
group,
frame: self.start_frame,
});
}
if let Some(group) = self.end_group {
sub = sub.with_end(match self.end_frame {
Some(frame) => track::Position::after(group, frame),
None => track::Position::after_group(group),
});
}
sub
}
fn start_frame(&self) -> Option<(u64, u64)> {
self.start_group
.map(|group| (group, self.start_frame))
.filter(|(_, frame)| *frame != 0)
}
fn end_frame(&self) -> Option<(u64, u64)> {
self.end_group.zip(self.end_frame)
}
}
impl From<&lite::Subscribe<'_>> for Bounds {
fn from(msg: &lite::Subscribe<'_>) -> Self {
Self {
start_group: msg.start_group,
start_frame: msg.start_frame,
end_group: msg.end_group,
end_frame: msg.end_frame,
}
}
}
impl From<&lite::SubscribeUpdate> for Bounds {
fn from(msg: &lite::SubscribeUpdate) -> Self {
Self {
start_group: msg.start_group,
start_frame: msg.start_frame,
end_group: msg.end_group,
end_frame: msg.end_frame,
}
}
}
#[derive(Clone)]
struct Subscription<S: crate::transport::poll::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: crate::transport::poll::Session> Subscription<S> {
fn serve_datagram(&mut 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);
}
fn track_priority_current(&mut self) -> u8 {
self.track_priority_seen = *self.track_priority.read();
self.track_priority_seen
}
#[cfg(test)]
async fn serve_group(
self,
sequence: u64,
frame_start: u64,
priority: PriorityHandle,
group: group::Consumer,
) -> Result<(), Error> {
let mut serve = Box::new(GroupServe::new(self, sequence, frame_start, priority, group));
kio::wait(move |waiter| serve.poll_serve(waiter)).await
}
}
enum TrackEnd {
PeerFin,
Finished,
}
struct TrackRun<S: crate::transport::poll::Session> {
ctx: Subscription<S>,
track: track::Subscriber,
track_priority_tx: kio::Producer<u8>,
start_frame: Option<(u64, u64)>,
end_frame: Option<(u64, u64)>,
emit_range: bool,
start_sent: bool,
end_sent: bool,
datagrams: bool,
children: kio::Tasks<GroupServe<S>>,
}
impl<S: crate::transport::poll::Session> TrackRun<S> {
fn new(
ctx: Subscription<S>,
mut track: track::Subscriber,
bounds: Bounds,
track_priority_tx: kio::Producer<u8>,
) -> Self {
position_cursor(&mut track, ctx.version, bounds.start_group);
track.end_at(bounds.end_group.map_or(Bound::Unbounded, Bound::Included));
let emit_range = ctx.version.has_track_stream();
let datagrams = ctx.version.has_datagrams() && ctx.session.max_datagram_size() > 0;
Self {
start_frame: bounds.start_frame(),
end_frame: bounds.end_frame(),
ctx,
track,
track_priority_tx,
emit_range,
start_sent: false,
end_sent: false,
datagrams,
children: kio::Tasks::new(),
}
}
fn poll(&mut self, stream: &mut Stream<S, Version>, waiter: &kio::Waiter) -> Poll<Result<TrackEnd, Error>> {
let mut cx = Context::from_waker(waiter.waker());
loop {
ready!(stream.writer.poll_flush(&mut cx))?;
let _ = self.children.poll(waiter);
if let Poll::Ready(upd) = stream.reader.poll_decode_maybe::<lite::SubscribeUpdate>(&mut cx) {
let Some(upd) = upd? else {
return Poll::Ready(Ok(TrackEnd::PeerFin));
};
if let Ok(mut value) = self.track_priority_tx.write() {
*value = upd.priority;
}
let bounds = Bounds::from(&upd);
let _ = self.track.update(crate::track::Subscription {
priority: upd.priority,
max_age: serving_max_age(self.ctx.version, upd.max_age),
..bounds.positions()
});
if let Some(start_group) = upd.start_group {
self.track.start_at(start_group);
}
self.track
.end_at(upd.end_group.map_or(Bound::Unbounded, Bound::Included));
self.start_frame = bounds.start_frame();
self.end_frame = bounds.end_frame();
continue;
}
let emit_boundary = self.emit_range && !self.end_sent;
if let Poll::Ready(res) = poll_recv_next(&mut self.track, self.datagrams, emit_boundary, waiter) {
match res? {
Recv::Group(mut group) => {
let sequence = group.sequence;
if !position_group(&mut group, self.start_frame, self.end_frame) {
tracing::debug!(subscribe = self.ctx.id, track = %self.ctx.track_name, sequence, "skipping group with a missing head");
continue;
}
if self.emit_range && !self.start_sent {
self.start_sent = true;
stream
.writer
.buffer(&lite::SubscribeResponse::Start(lite::SubscribeStart {
group: sequence,
}))?;
self.track.start_at(sequence);
}
let frame_start = group.index();
tracing::debug!(subscribe = self.ctx.id, track = %self.ctx.track_name, sequence, "serving group");
let current_priority = self.ctx.track_priority_current();
let handle = self
.ctx
.priority
.insert(Priority::new(current_priority, self.ctx.id, sequence));
self.children
.push(GroupServe::new(self.ctx.clone(), sequence, frame_start, handle, group));
}
Recv::Datagram(datagram) => self.ctx.serve_datagram(datagram),
Recv::Boundary(group) => {
self.end_sent = true;
stream
.writer
.buffer(&lite::SubscribeResponse::End(lite::SubscribeEnd { group }))?;
}
Recv::Finished => return Poll::Ready(Ok(TrackEnd::Finished)),
}
continue;
}
return Poll::Pending;
}
}
}
struct GroupServe<S: crate::transport::poll::Session> {
ctx: Subscription<S>,
priority: PriorityHandle,
group: group::Consumer,
sequence: u64,
frame_start: u64,
prev_ts: u64,
state: GroupState<S>,
}
#[allow(clippy::large_enum_variant)]
enum GroupState<S: crate::transport::poll::Session> {
Open,
Serve {
writer: Writer<S::SendStream, Version>,
frame: Option<frame::Consumer>,
chunk: Option<bytes::Bytes>,
batch: Box<frame::Buffer>,
batch_pos: usize,
},
Closed {
writer: Writer<S::SendStream, Version>,
},
Done,
}
impl<S: crate::transport::poll::Session> GroupServe<S> {
fn new(
ctx: Subscription<S>,
sequence: u64,
frame_start: u64,
priority: PriorityHandle,
group: group::Consumer,
) -> Self {
Self {
ctx,
priority,
group,
sequence,
frame_start,
prev_ts: 0,
state: GroupState::Open,
}
}
fn poll_serve(&mut self, waiter: &kio::Waiter) -> Poll<Result<(), Error>> {
loop {
match &mut self.state {
GroupState::Open => {
if self.group.poll_expired(waiter) {
self.state = GroupState::Done;
return Poll::Ready(Err(Error::Old));
}
let mut cx = Context::from_waker(waiter.waker());
let stream = match ready!(self.ctx.session.poll_open_uni(&mut cx)) {
Ok(stream) => stream,
Err(err) => {
self.state = GroupState::Done;
return Poll::Ready(Err(Error::from_transport(err)));
}
};
let mut writer = Writer::new(stream, self.ctx.version);
writer.set_priority(self.priority.send_order());
let msg = lite::Group {
subscribe: self.ctx.id,
sequence: self.sequence,
frame_start: self.frame_start,
};
if let Err(err) = writer.buffer(&lite::DataType::Group).and_then(|()| writer.buffer(&msg)) {
self.state = GroupState::Done;
writer.abort(&err);
return Poll::Ready(Err(err));
}
self.state = GroupState::Serve {
writer,
frame: None,
chunk: None,
batch: Box::new(frame::Buffer::new()),
batch_pos: 0,
};
}
GroupState::Serve {
writer,
frame,
chunk,
batch,
batch_pos,
} => {
let mut cx = Context::from_waker(waiter.waker());
while let Poll::Ready(rank) = self.priority.poll_next(waiter) {
writer.set_priority(PriorityHandle::send_order_of(rank));
}
let seen = self.ctx.track_priority_seen;
if let Poll::Ready(Ok(value)) = self.ctx.track_priority.poll(waiter, |value| {
if **value != seen {
Poll::Ready(**value)
} else {
Poll::Pending
}
}) {
self.ctx.track_priority_seen = value;
let rank = self.priority.set_track(value);
writer.set_priority(PriorityHandle::send_order_of(rank));
}
let outcome = 'serve: {
if writer.poll_closed(&mut cx).is_ready() {
break 'serve Err(Error::Cancel);
}
loop {
match writer.poll_flush(&mut cx) {
Poll::Ready(Ok(())) => {}
Poll::Ready(Err(err)) => break 'serve Err(err),
Poll::Pending => {
if self.group.poll_expired_while_pending(waiter, true) {
break 'serve Err(Error::Old);
}
return Poll::Pending;
}
}
if let Some(pending) = chunk {
match writer.poll_write(&mut cx, pending) {
Poll::Ready(Ok(_)) => {
if !bytes::Buf::has_remaining(pending) {
*chunk = None;
}
}
Poll::Ready(Err(err)) => break 'serve Err(err),
Poll::Pending => {
if self.group.poll_expired_while_pending(waiter, true) {
break 'serve Err(Error::Old);
}
return Poll::Pending;
}
}
} else if let Some(pending) = frame {
match pending.poll_read_chunk(waiter) {
Poll::Ready(Ok(Some(next))) => *chunk = Some(next),
Poll::Ready(Ok(None)) => *frame = None,
Poll::Ready(Err(err)) => break 'serve Err(err),
Poll::Pending => return Poll::Pending,
}
} else if *batch_pos < batch.len() {
let batched = &mut batch.filled_mut()[*batch_pos];
let buffered = buffer_frame_info(
writer,
batched.timestamp,
batched.payload.len() as u64,
self.ctx.timescale,
&mut self.prev_ts,
);
if let Err(err) = buffered {
break 'serve Err(err);
}
let payload = std::mem::take(&mut batched.payload);
if !payload.is_empty() {
*chunk = Some(payload);
}
*batch_pos += 1;
self.group.keep_alive();
} else {
match self.group.poll_read_frames(waiter, batch) {
Poll::Ready(Ok(count)) if count > 0 => {
*batch_pos = 0;
continue;
}
Poll::Ready(Ok(_)) => break 'serve Ok(()),
Poll::Ready(Err(err)) => break 'serve Err(err),
Poll::Pending => {}
}
match self.group.poll_next_frame(waiter) {
Poll::Ready(Ok(Some(next))) => {
let buffered = buffer_frame_info(
writer,
next.timestamp,
next.size,
self.ctx.timescale,
&mut self.prev_ts,
);
if let Err(err) = buffered {
break 'serve Err(err);
}
*frame = Some(next);
}
Poll::Ready(Ok(None)) => break 'serve Ok(()),
Poll::Ready(Err(err)) => break 'serve Err(err),
Poll::Pending => return Poll::Pending,
}
}
}
};
let GroupState::Serve { writer, .. } = std::mem::replace(&mut self.state, GroupState::Done) else {
unreachable!()
};
match outcome {
Ok(()) => {
let mut writer = writer;
match writer.finish() {
Ok(()) => self.state = GroupState::Closed { writer },
Err(err) => return Poll::Ready(Err(err)),
}
}
Err(err) => {
writer.abort(&err);
return Poll::Ready(Err(err));
}
}
}
GroupState::Closed { writer } => {
let mut cx = Context::from_waker(waiter.waker());
let res = ready!(writer.poll_close(&mut cx));
self.state = GroupState::Done;
return Poll::Ready(res.map(|()| {
tracing::debug!(sequence = self.sequence, "finished group");
}));
}
GroupState::Done => return Poll::Ready(Ok(())),
}
}
}
}
impl<S: crate::transport::poll::Session> kio::Task for GroupServe<S> {
fn poll(&mut self, waiter: &kio::Waiter) -> Poll<()> {
ready!(self.poll_serve(waiter)).map(|()| ()).unwrap_or(());
Poll::Ready(())
}
}
#[cfg(all(test, not(loom)))]
mod serve_group_test {
use super::*;
use crate::lite::test_transport::*;
use crate::{Timestamp, broadcast};
use futures::FutureExt;
#[test]
fn bounds_convert_to_positions() {
let whole = Bounds {
start_group: None,
start_frame: 0,
end_group: Some(5),
end_frame: None,
};
assert_eq!(whole.positions().end, Some(track::Position::group(6)));
let capped = Bounds {
end_frame: Some(2),
..whole
};
assert_eq!(capped.positions().end, Some(track::Position { group: 5, frame: 3 }));
let started = Bounds {
start_group: Some(5),
start_frame: 3,
end_group: None,
end_frame: None,
};
let positions = started.positions();
assert_eq!(positions.start, Some(track::Position { group: 5, frame: 3 }));
assert_eq!(positions.end, None);
let orphan = Bounds {
start_group: None,
start_frame: 3,
end_group: None,
end_frame: Some(7),
};
assert_eq!((orphan.positions().start, orphan.positions().end), (None, None));
}
#[test]
fn position_group_skips_a_missing_head() {
let track = track::Producer::new(Arc::new(broadcast::Info::default()), "video", None);
let mut group = track.create_group(group::Info { sequence: 3 }).unwrap();
group.start_at(5).unwrap();
group.write_frame(Timestamp::ZERO, b"tail".to_vec()).unwrap();
let mut consumer = group.consume();
assert!(!position_group(&mut consumer, None, None));
let mut consumer = group.consume();
assert!(position_group(&mut consumer, Some((3, 5)), None));
assert_eq!(consumer.index(), 5);
let mut consumer = group.consume();
assert!(!position_group(&mut consumer, Some((3, 2)), None));
let mut other = track.create_group(group::Info { sequence: 4 }).unwrap();
other.write_frame(Timestamp::ZERO, b"whole".to_vec()).unwrap();
let mut consumer = other.consume();
assert!(position_group(&mut consumer, Some((3, 5)), None));
assert_eq!(consumer.index(), 0);
}
#[test]
fn position_group_caps_the_end_group() {
let track = track::Producer::new(Arc::new(broadcast::Info::default()), "video", None);
let mut group = track.create_group(group::Info { sequence: 7 }).unwrap();
for i in 0..4u8 {
group.write_frame(Timestamp::ZERO, vec![i]).unwrap();
}
group.finish().unwrap();
let mut consumer = group.consume();
assert!(position_group(&mut consumer, None, Some((7, 1))));
assert_eq!(
consumer.read_frame().now_or_never().unwrap().unwrap().unwrap().payload[0],
0
);
assert_eq!(
consumer.read_frame().now_or_never().unwrap().unwrap().unwrap().payload[0],
1
);
assert!(
consumer.read_frame().now_or_never().unwrap().unwrap().is_none(),
"capped"
);
let mut consumer = group.consume();
assert!(position_group(&mut consumer, None, Some((8, 1))));
for i in 0..4u8 {
assert_eq!(
consumer.read_frame().now_or_never().unwrap().unwrap().unwrap().payload[0],
i
);
}
}
#[tokio::test]
async fn resets_with_the_abort_code() {
let log = Log::default();
let session = SinkSession::new(log.clone());
let track_priority = kio::Producer::new(0u8);
let subscription = Subscription {
session,
id: 0,
track_name: "test".into(),
priority: PriorityQueue::default(),
track_priority: track_priority.consume(),
track_priority_seen: 0,
version: Version::Lite06,
timescale: Some(crate::Timescale::default()),
};
let track = track::Producer::new(Arc::new(broadcast::Info::default()), "test", None);
let mut group = track.create_group(group::Info { sequence: 0 }).unwrap();
group
.write_frame(Timestamp::from_millis(0).unwrap(), b"hello".as_slice())
.unwrap();
let handle = subscription.priority.insert(Priority::new(0, 0, 0));
let mut serve = std::pin::pin!(subscription.serve_group(0, 0, handle, group.consume()));
assert!(futures::poll!(serve.as_mut()).is_pending());
group.abort(Error::Old).unwrap();
assert!(matches!(serve.await, Err(Error::Old)));
assert_eq!(log.resets(), vec![crate::StreamError::Old.to_code()]);
}
#[tokio::test]
async fn blocked_transport_write_expires_with_the_group() {
tokio::time::pause();
let gate = kio::Producer::new(false);
let session = SinkSession::gated_uni(gate.consume());
let log = session.log.clone();
let track_priority = kio::Producer::new(0u8);
let subscription = Subscription {
session,
id: 0,
track_name: "test".into(),
priority: PriorityQueue::default(),
track_priority: track_priority.consume(),
track_priority_seen: 0,
version: Version::Lite06,
timescale: Some(crate::Timescale::default()),
};
let track = track::Producer::new(Arc::new(broadcast::Info::default()), "test", None);
let mut subscriber = track.subscribe(None);
let mut old = track.append_group().unwrap();
old.write_frame(Timestamp::ZERO, b"old".as_slice()).unwrap();
old.finish().unwrap();
let group = subscriber.recv_group().await.unwrap().expect("old group");
let handle = subscription.priority.insert(Priority::new(0, 0, 0));
let mut serve = std::pin::pin!(subscription.serve_group(0, 0, handle, group));
assert!(
futures::poll!(serve.as_mut()).is_pending(),
"transport write is blocked"
);
tokio::time::advance(Duration::from_secs(1)).await;
let mut edge = track.append_group().unwrap();
edge.write_frame(Timestamp::from_millis(1000).unwrap(), b"edge".as_slice())
.unwrap();
edge.finish().unwrap();
assert!(matches!(serve.await, Err(Error::Old)));
assert_eq!(log.resets(), vec![crate::StreamError::Old.to_code()]);
}
#[tokio::test]
async fn blocked_final_transport_chunk_expires_with_the_group() {
tokio::time::pause();
let gate = kio::Producer::new(true);
let session = SinkSession::gated_uni(gate.consume());
let log = session.log.clone();
let track_priority = kio::Producer::new(0u8);
let subscription = Subscription {
session,
id: 0,
track_name: "test".into(),
priority: PriorityQueue::default(),
track_priority: track_priority.consume(),
track_priority_seen: 0,
version: Version::Lite06,
timescale: Some(crate::Timescale::default()),
};
let track = track::Producer::new(Arc::new(broadcast::Info::default()), "test", None);
let mut subscriber = track.subscribe(None);
let mut old = track.append_group().unwrap();
let mut frame = old
.create_frame(frame::Info {
timestamp: Timestamp::ZERO,
size: 2,
})
.unwrap();
frame.write(b"a".as_slice()).unwrap();
let group = subscriber.recv_group().await.unwrap().expect("old group");
let handle = subscription.priority.insert(Priority::new(0, 0, 0));
let mut serve = std::pin::pin!(subscription.serve_group(0, 0, handle, group));
assert!(
futures::poll!(serve.as_mut()).is_pending(),
"waiting for the final byte"
);
let Ok(mut open) = gate.write() else {
panic!("transport gate closed");
};
*open = false;
drop(open);
frame.write(b"b".as_slice()).unwrap();
frame.finish().unwrap();
old.finish().unwrap();
assert!(
futures::poll!(serve.as_mut()).is_pending(),
"the final byte is transport-blocked"
);
tokio::time::advance(Duration::from_secs(1)).await;
let mut edge = track.append_group().unwrap();
edge.write_frame(Timestamp::from_millis(1000).unwrap(), b"edge".as_slice())
.unwrap();
edge.finish().unwrap();
assert!(matches!(serve.await, Err(Error::Old)));
assert_eq!(log.resets(), vec![crate::StreamError::Old.to_code()]);
}
#[tokio::test]
async fn blocked_transport_open_expires_with_the_group() {
tokio::time::pause();
let gate = kio::Producer::new(false);
let session = SinkSession::gated_open_uni(gate.consume());
let track_priority = kio::Producer::new(0u8);
let subscription = Subscription {
session,
id: 0,
track_name: "test".into(),
priority: PriorityQueue::default(),
track_priority: track_priority.consume(),
track_priority_seen: 0,
version: Version::Lite06,
timescale: Some(crate::Timescale::default()),
};
let track = track::Producer::new(Arc::new(broadcast::Info::default()), "test", None);
let mut subscriber = track.subscribe(None);
let mut old = track.append_group().unwrap();
old.write_frame(Timestamp::ZERO, b"old".as_slice()).unwrap();
old.finish().unwrap();
let group = subscriber.recv_group().await.unwrap().expect("old group");
let handle = subscription.priority.insert(Priority::new(0, 0, 0));
let mut serve = std::pin::pin!(subscription.serve_group(0, 0, handle, group));
assert!(
futures::poll!(serve.as_mut()).is_pending(),
"stream credit is exhausted"
);
tokio::time::advance(Duration::from_secs(1)).await;
let mut edge = track.append_group().unwrap();
edge.write_frame(Timestamp::from_millis(1000).unwrap(), b"edge".as_slice())
.unwrap();
edge.finish().unwrap();
assert!(matches!(serve.await, Err(Error::Old)));
}
#[test]
fn a_version_without_the_field_serves_a_non_dropping_budget() {
for version in [Version::Lite01, Version::Lite02] {
let max_age = serving_max_age(version, Duration::ZERO);
assert!(
max_age >= Duration::from_secs(86_400),
"{version:?} must not be served as real time: {max_age:?}"
);
}
assert_eq!(serving_max_age(Version::Lite05, Duration::ZERO), Duration::ZERO);
assert_eq!(
serving_max_age(Version::Lite05, Duration::from_secs(3)),
Duration::from_secs(3)
);
}
#[tokio::test]
async fn completed_group_does_not_reset() {
let log = Log::default();
let session = SinkSession::new(log.clone());
let track_priority = kio::Producer::new(0u8);
let subscription = Subscription {
session,
id: 0,
track_name: "test".into(),
priority: PriorityQueue::default(),
track_priority: track_priority.consume(),
track_priority_seen: 0,
version: Version::Lite06,
timescale: Some(crate::Timescale::default()),
};
let track = track::Producer::new(Arc::new(broadcast::Info::default()), "test", None);
let mut group = track.create_group(group::Info { sequence: 0 }).unwrap();
group
.write_frame(Timestamp::from_millis(0).unwrap(), b"hello".as_slice())
.unwrap();
let consumer = group.consume();
group.finish().unwrap();
let handle = subscription.priority.insert(Priority::new(0, 0, 0));
subscription.serve_group(0, 0, handle, consumer).await.unwrap();
assert_eq!(log.resets(), Vec::<u32>::new(), "clean completion must not reset");
let priorities = log.priorities();
assert!(!priorities.is_empty(), "the group stream must set a priority");
assert!(
priorities.iter().all(|&p| p == 255),
"rank 0 must reach the transport as send order 255: {priorities:?}",
);
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::lite::test_transport::SinkSession;
use crate::model::ProduceTest;
#[tokio::test(start_paused = true)]
async fn serving_origin_falls_back_to_the_assigned_identity() {
let assigned = crate::Hop::new(777).unwrap();
let upstream = crate::Hop::new(778).unwrap();
let origin = crate::origin::Config::new(crate::Hop::new(1).unwrap()).produce();
let mut echoed_hops = Hops::new();
echoed_hops.push(crate::Hop::UNKNOWN).unwrap();
let _echoed = origin
.dynamic(
"echoed",
crate::origin::Route::default()
.with_hops(echoed_hops)
.with_via(assigned),
)
.unwrap();
let mut local_hops = Hops::new();
local_hops.push(upstream).unwrap();
let _local = origin
.dynamic("local", crate::origin::Route::default().with_hops(local_hops))
.unwrap();
let peer_setup = crate::lite::PeerSetup::default();
peer_setup.set(crate::lite::Setup::default());
let (_, goaway) = crate::goaway::Handle::new(true);
let publisher = Publisher::new(PublisherConfig {
runtime: crate::time::Clock::tokio(),
session: SinkSession::new(Default::default()),
origin: origin.consume(),
version: Version::Lite06,
peer_setup,
goaway,
peer_hop: Some(assigned),
});
let serving = kio::wait(|waiter| publisher.shared.poll_serving_origin(waiter)).await;
use futures::FutureExt;
assert!(
serving.request_broadcast("echoed/x").now_or_never().unwrap().is_err(),
"served the peer its own route"
);
assert!(
serving.request_broadcast("local/x").now_or_never().is_none(),
"withheld an independent route"
);
}
#[tokio::test]
async fn announce_init_applies_route_selection() {
let assigned = crate::Hop::new(777).unwrap();
let clean_publisher = crate::Hop::new(778).unwrap();
let self_origin = crate::Hop::new(1).unwrap();
let origin = crate::origin::Config::new(self_origin).produce();
let mut tainted_hops = Hops::new();
tainted_hops.push(crate::Hop::UNKNOWN).unwrap();
let _tainted = origin
.announce(
"echoed",
crate::origin::Route::default()
.with_hops(tainted_hops)
.with_via(assigned),
)
.unwrap();
let mut clean_hops = Hops::new();
clean_hops.push(clean_publisher).unwrap();
let _clean = origin
.announce("local", crate::origin::Route::default().with_hops(clean_hops))
.unwrap();
let gate = kio::Producer::new(true);
let session = SinkSession::gated_bi(gate.consume());
let log = session.log.clone();
let mut stream = Stream::open(&mut session.clone(), Version::Lite01).await.unwrap();
let consumer = origin.consume().excluding(assigned);
let mut announced = consumer.announced();
let mut run = std::pin::pin!(Publisher::<SinkSession>::run_announce(
&mut stream,
&consumer,
&mut announced,
crate::Path::new(""),
self_origin,
Version::Lite01,
));
assert!(futures::poll!(run.as_mut()).is_pending());
let writes = log.writes.lock().unwrap();
assert!(
writes.windows(b"local".len()).any(|w| w == b"local"),
"clean broadcast in ANNOUNCE_INIT"
);
assert!(
!writes.windows(b"echoed".len()).any(|w| w == b"echoed"),
"echoed broadcast filtered from ANNOUNCE_INIT"
);
}
fn decode_probes(bytes: &[u8]) -> Vec<lite::Probe> {
decode_probes_version(bytes, Version::Lite05)
}
fn decode_probes_version(bytes: &[u8], version: Version) -> Vec<lite::Probe> {
use crate::coding::Decode as _;
let mut slice = bytes;
let mut out = Vec::new();
while bytes::Buf::remaining(&slice) > 0 {
out.push(lite::Probe::decode(&mut slice, version).unwrap());
}
out
}
fn probe_server(
session: SinkSession,
stream: Stream<SinkSession, Version>,
version: Version,
) -> ProbeServe<SinkSession> {
let origin = Hop::random().produce();
let (_, goaway) = crate::goaway::Handle::new(true);
let publisher = Publisher::new(PublisherConfig {
runtime: crate::time::Clock::tokio(),
session,
origin: origin.consume(),
version,
peer_setup: crate::lite::PeerSetup::default(),
goaway,
peer_hop: None,
});
ProbeServe::new(publisher.shared.clone(), publisher.runtime.clone(), stream)
}
async fn probe_writes(stats: crate::lite::test_transport::SinkStats) -> Vec<lite::Probe> {
probe_writes_version(stats, Version::Lite05).await
}
async fn probe_writes_version(stats: crate::lite::test_transport::SinkStats, version: Version) -> Vec<lite::Probe> {
let gate = kio::Producer::new(true);
let mut session = SinkSession::gated_bi(gate.consume()).with_stats(stats);
let log = session.log.clone();
let stream = Stream::open(&mut session, version).await.unwrap();
let mut server = probe_server(session, stream, version);
let mut run = std::pin::pin!(kio::wait(|waiter| server.poll_probe(waiter)));
assert!(futures::poll!(run.as_mut()).is_pending());
tokio::time::sleep(Duration::from_millis(150)).await;
assert!(futures::poll!(run.as_mut()).is_pending());
let writes = log.writes.lock().unwrap().clone();
decode_probes_version(&writes, version)
}
#[tokio::test(start_paused = true)]
async fn reports_rtt_without_a_bitrate() {
let stats = crate::lite::test_transport::SinkStats::default().with_rtt(std::time::Duration::from_millis(40));
let probes = probe_writes(stats).await;
assert_eq!(probes.len(), 1, "expected exactly one report");
assert_eq!(probes[0].rtt, Some(40));
assert_eq!(probes[0].bitrate, None, "unknown bitrate, not a measured zero");
}
#[tokio::test(start_paused = true)]
async fn reports_bitrate_without_an_rtt() {
let stats = crate::lite::test_transport::SinkStats::default().with_send_rate(1_000_000);
let probes = probe_writes(stats).await;
assert_eq!(probes.len(), 1);
assert_eq!(probes[0].bitrate, Some(1_000_000));
assert_eq!(probes[0].rtt, None);
}
#[tokio::test(start_paused = true)]
async fn reports_nothing_when_nothing_is_measurable() {
let probes = probe_writes(crate::lite::test_transport::SinkStats::default()).await;
assert!(probes.is_empty(), "expected no report, got {probes:?}");
}
#[tokio::test(start_paused = true)]
async fn lite03_sends_nothing_for_an_rtt_only_report() {
let stats = crate::lite::test_transport::SinkStats::default().with_rtt(std::time::Duration::from_millis(40));
let probes = probe_writes_version(stats, Version::Lite03).await;
assert!(probes.is_empty(), "expected no report on lite-03, got {probes:?}");
}
#[tokio::test(start_paused = true)]
async fn lite03_reports_the_bitrate() {
let stats = crate::lite::test_transport::SinkStats::default()
.with_send_rate(1_000_000)
.with_rtt(std::time::Duration::from_millis(40));
let probes = probe_writes_version(stats, Version::Lite03).await;
assert_eq!(probes.len(), 1);
assert_eq!(probes[0].bitrate, Some(1_000_000));
assert_eq!(probes[0].rtt, None, "lite-03 carries no RTT field");
}
#[tokio::test(start_paused = true)]
async fn a_bitrate_going_unknown_is_retracted_once() {
let gate = kio::Producer::new(true);
let stats = crate::lite::test_transport::SinkStats::default().with_send_rate(1_000_000);
let mut session = SinkSession::gated_bi(gate.consume()).with_stats(stats);
let log = session.log.clone();
let stream = Stream::open(&mut session, Version::Lite05).await.unwrap();
let mut server = probe_server(session.clone(), stream, Version::Lite05);
let mut run = std::pin::pin!(kio::wait(|waiter| server.poll_probe(waiter)));
assert!(futures::poll!(run.as_mut()).is_pending());
tokio::time::sleep(Duration::from_millis(150)).await;
assert!(futures::poll!(run.as_mut()).is_pending());
session.set_stats(crate::lite::test_transport::SinkStats::default());
for _ in 0..3 {
tokio::time::sleep(Duration::from_secs(11)).await;
assert!(futures::poll!(run.as_mut()).is_pending());
}
let writes = log.writes.lock().unwrap().clone();
let probes = decode_probes(&writes);
assert_eq!(
probes.len(),
2,
"the measurement then one retraction, not a repeating 'unknown': {probes:?}"
);
assert_eq!(probes[0].bitrate, Some(1_000_000));
assert_eq!(probes[1].bitrate, None);
}
}