use crate::runtime::Timers as _;
use crate::{frame, group, origin, track};
use std::{
collections::HashMap,
sync::{Arc, atomic},
task::{Poll, ready},
time::Duration,
};
use crate::{
AsPath, Error, Path, PathOwned, Timescale, Timestamp, bandwidth,
coding::{Decode, Reader, Stream},
lite,
track::{Position, Subscription},
};
use super::Version;
use crate::tail::{self, Settle, Tail};
use kio::Lock;
pub(super) struct SubscriberConfig<S: crate::transport::poll::Session> {
pub runtime: crate::time::Clock,
pub session: S,
pub origin: origin::Producer,
pub recv_bandwidth: Option<bandwidth::Producer>,
pub version: Version,
pub peer_setup: super::PeerSetup,
pub peer_hop: Option<crate::Hop>,
pub cost: Option<u64>,
pub going_away: crate::goaway::GoingAway,
}
#[derive(Clone)]
pub(super) struct Subscriber<S: crate::transport::poll::Session> {
runtime: crate::time::Clock,
session: S,
origin: origin::Producer,
recv_bandwidth: Option<bandwidth::Producer>,
self_origin: crate::Hop,
session_origin: crate::Hop,
subscribes: Lock<HashMap<u64, TrackEntry>>,
next_id: Arc<atomic::AtomicU64>,
version: Version,
peer_setup: super::PeerSetup,
cost: Option<u64>,
sources: kio::Queue<(PathOwned, crate::broadcast::Dynamic)>,
going_away: crate::goaway::GoingAway,
}
#[derive(Clone)]
struct TrackEntry {
producer: track::Producer,
timescale: Option<Timescale>,
tail: kio::Producer<Tail>,
}
impl<S: crate::transport::poll::Session> Subscriber<S> {
pub fn new(config: SubscriberConfig<S>) -> Self {
let self_origin = config.origin.hop();
Self {
session: config.session,
runtime: config.runtime,
origin: config.origin,
recv_bandwidth: config.recv_bandwidth,
self_origin,
session_origin: config.peer_hop.unwrap_or(crate::Hop::UNKNOWN),
subscribes: Default::default(),
next_id: Default::default(),
version: config.version,
peer_setup: config.peer_setup,
cost: config.cost,
sources: kio::Queue::new(),
going_away: config.going_away,
}
}
fn check_going_away(&self) -> Result<(), Error> {
if self.going_away.is_set() {
return Err(Error::GoingAway);
}
Ok(())
}
fn poll_link_cost(&self, waiter: &kio::Waiter) -> Poll<u64> {
if !self.version.has_route_cost() {
return Poll::Ready(0);
}
match self.cost {
Some(cost) => Poll::Ready(cost),
None => self
.peer_setup
.poll_cost(waiter)
.map(|cost| cost.unwrap_or(super::DEFAULT_COST)),
}
}
fn handle_announce(
&mut self,
prefix: &PathOwned,
announce: lite::AnnounceBroadcast<'_>,
run: &mut PrefixRun,
) -> Result<(), Error> {
match announce {
lite::AnnounceBroadcast::Active { suffix, hops, cost } => {
let path = prefix.join(&suffix);
if self.version.has_announce_id() {
run.announced_by_id.insert(run.next_announce_id, path.clone());
run.next_announce_id += 1;
}
if lite::restart_supported(self.version)
&& !self.version.has_announce_id()
&& run.announced.contains(&path)
{
self.restart_announce(
path,
hops,
cost,
run.link_cost,
run.responder_origin,
&mut run.announced,
)?;
} else {
self.start_announce(
path,
hops,
cost,
run.link_cost,
run.responder_origin,
&mut run.announced,
)?;
}
}
lite::AnnounceBroadcast::Ended { suffix, .. } => {
let path = prefix.join(&suffix);
tracing::debug!(broadcast = %self.log_path(&path), "unannounced");
run.announced.retire(&path);
}
lite::AnnounceBroadcast::EndedId { id } => {
let Some(path) = run.announced_by_id.remove(&id) else {
return Err(Error::ProtocolViolation);
};
tracing::debug!(broadcast = %self.log_path(&path), "unannounced");
run.announced.retire(&path);
}
lite::AnnounceBroadcast::Restart { id, hops, cost } => {
let Some(path) = run.announced_by_id.get(&id).cloned() else {
return Err(Error::ProtocolViolation);
};
self.restart_announce(
path,
hops,
cost,
run.link_cost,
run.responder_origin,
&mut run.announced,
)?;
}
lite::AnnounceBroadcast::Skipped => {}
}
Ok(())
}
fn start_announce(
&mut self,
path: PathOwned,
mut hops: crate::Hops,
cost: crate::origin::Cost,
link_cost: u64,
responder_origin: Option<crate::Hop>,
announced: &mut Announced,
) -> Result<bool, Error> {
if announced.contains(&path) {
return Err(Error::ProtocolViolation);
}
announced.reserve(path.clone());
if let Some(responder) = responder_origin {
if responder != crate::Hop::UNKNOWN && hops.contains(&responder) {
tracing::debug!(route = %self.log_path(&path), "dropping announce reflected by its sender");
return Ok(false);
}
if hops.push(responder).is_err() {
tracing::warn!(
route = %self.log_path(&path),
"dropping announce; hop chain at MAX_HOPS (possible loop)",
);
return Ok(false);
}
}
if hops.contains(&self.self_origin) {
tracing::debug!(route = %self.log_path(&path), "dropping reflected announce");
return Ok(false);
}
if hops.is_empty() {
hops.push(crate::Hop::UNKNOWN)
.expect("an empty hop chain always has room for one entry, and repeats nothing");
}
tracing::debug!(route = %self.log_path(&path), hops = hops.len(), "announce");
let route = self.announced_route(hops, cost, link_cost, responder_origin);
let Ok(dynamic) = self.origin.dynamic(&path, route.clone()) else {
return Ok(false);
};
announced.attach(path, AnnouncedRoute::new(route, dynamic));
Ok(true)
}
fn announced_route(
&self,
hops: crate::Hops,
cost: crate::origin::Cost,
link_cost: u64,
responder: Option<crate::Hop>,
) -> crate::origin::Route {
let mut route = crate::origin::Route::default()
.with_hops(hops)
.with_cost(cost.charged(link_cost))
.with_via(self.via(responder));
if self.going_away.is_set() {
route.cost = crate::origin::Cost::DRAIN;
}
route
}
fn via(&self, responder: Option<crate::Hop>) -> crate::Hop {
responder
.filter(|hop| *hop != crate::Hop::UNKNOWN)
.unwrap_or(self.session_origin)
}
fn restart_announce(
&mut self,
path: PathOwned,
mut hops: crate::Hops,
cost: crate::origin::Cost,
link_cost: u64,
responder_origin: Option<crate::Hop>,
announced: &mut Announced,
) -> Result<bool, Error> {
let reflected = match responder_origin {
Some(responder) => {
(responder != crate::Hop::UNKNOWN && hops.contains(&responder))
|| hops.push(responder).is_err()
|| hops.contains(&self.self_origin)
}
None => hops.contains(&self.self_origin),
};
if reflected {
tracing::debug!(route = %self.log_path(&path), "dropping reflected restart");
announced.declined(path);
return Ok(false);
}
if hops.is_empty() {
hops.push(crate::Hop::UNKNOWN)
.expect("an empty hop chain always has room for one entry, and repeats nothing");
}
tracing::debug!(route = %self.log_path(&path), hops = hops.len(), "restart");
let metadata = self.announced_route(hops, cost, link_cost, responder_origin);
if let Some(entry) = announced.attached(&path) {
entry.update(metadata);
return Ok(true);
}
let Ok(dynamic) = self.origin.dynamic(&path, metadata.clone()) else {
announced.declined(path);
return Ok(false);
};
announced.attach(path, AnnouncedRoute::new(metadata, dynamic));
Ok(true)
}
fn remove_subscribe(&self, id: u64) {
self.subscribes.lock().remove(&id);
}
fn route_datagram(&self, payload: bytes::Bytes) -> Result<(), Error> {
let mut buf = payload;
let dg = lite::Datagram::decode(&mut buf, self.version)?;
let mut subscribes = self.subscribes.lock();
let Some(entry) = subscribes.get_mut(&dg.subscribe) else {
return Ok(());
};
let scale = entry.timescale.unwrap_or_default();
let timestamp =
Timestamp::new(dg.timestamp, scale).map_err(|_| Error::BoundsExceeded(crate::coding::BoundsExceeded))?;
if let Ok(mut tail) = entry.tail.write() {
tail.account(dg.sequence..dg.sequence.saturating_add(1));
}
entry.producer.insert_datagram(dg.sequence, timestamp, dg.payload)?;
Ok(())
}
fn log_path(&self, path: impl AsPath) -> Path<'_> {
self.origin.root().join(path)
}
}
struct SubscriptionCleanup(Lock<HashMap<u64, TrackEntry>>);
impl Drop for SubscriptionCleanup {
fn drop(&mut self) {
for (_, entry) in self.0.lock().drain() {
let _ = entry.producer.abort(Error::Cancel);
}
}
}
pub(super) struct SubscriberDriver<S: crate::transport::poll::Session> {
subscriber: Subscriber<S>,
_cleanup: SubscriptionCleanup,
prefixes: Vec<AnnouncePrefix<S>>,
uni: UniAccept<S>,
bandwidth: Option<RecvBandwidth<S>>,
datagrams: Option<DatagramRecv<S>>,
sources: kio::Tasks<SourceServe<S>>,
}
impl<S: crate::transport::poll::Session> SubscriberDriver<S> {
pub fn new(subscriber: Subscriber<S>) -> Self {
let prefixes = crate::model::interest_prefixes(&subscriber.origin.allowed())
.into_iter()
.map(|prefix| AnnouncePrefix::new(subscriber.clone(), prefix))
.collect();
Self {
prefixes,
_cleanup: SubscriptionCleanup(subscriber.subscribes.clone()),
uni: UniAccept::new(subscriber.clone()),
bandwidth: Some(RecvBandwidth::new(subscriber.clone())),
datagrams: Some(DatagramRecv::new(subscriber.clone())),
sources: kio::Tasks::new(),
subscriber,
}
}
pub fn poll(&mut self, waiter: &kio::Waiter) -> Poll<Result<(), Error>> {
let mut i = 0;
while i < self.prefixes.len() {
match self.prefixes[i].poll(waiter) {
Poll::Ready(Ok(())) => {
self.prefixes.swap_remove(i);
}
Poll::Ready(Err(err)) => return Poll::Ready(Err(err)),
Poll::Pending => i += 1,
}
}
if let Poll::Ready(res) = self.uni.poll(waiter) {
return Poll::Ready(res);
}
if let Some(bandwidth) = &mut self.bandwidth {
match bandwidth.poll(waiter) {
Poll::Ready(Err(err)) => return Poll::Ready(Err(err)),
Poll::Ready(Ok(())) => self.bandwidth = None,
Poll::Pending => {}
}
}
if let Some(datagrams) = &mut self.datagrams {
match datagrams.poll(waiter) {
Poll::Ready(Err(err)) => return Poll::Ready(Err(err)),
Poll::Ready(Ok(())) => self.datagrams = None,
Poll::Pending => {}
}
}
while let Poll::Ready(Ok((path, dynamic))) = self.subscriber.sources.poll_pop(waiter) {
self.sources
.push(SourceServe::new(self.subscriber.clone(), path, dynamic));
}
let _ = self.sources.poll(waiter);
Poll::Pending
}
}
struct UniAccept<S: crate::transport::poll::Session> {
subscriber: Subscriber<S>,
accept: S,
children: kio::Tasks<UniServe<S>>,
}
impl<S: crate::transport::poll::Session> UniAccept<S> {
fn new(subscriber: Subscriber<S>) -> Self {
let accept = subscriber.session.clone();
Self {
subscriber,
accept,
children: kio::Tasks::new(),
}
}
fn poll(&mut self, waiter: &kio::Waiter) -> Poll<Result<(), Error>> {
let _ = self.children.poll(waiter);
let mut cx = std::task::Context::from_waker(waiter.waker());
loop {
match self.accept.poll_accept_uni(&mut cx) {
Poll::Ready(Ok(stream)) => {
self.children.push(UniServe {
subscriber: self.subscriber.clone(),
state: UniState::Start {
reader: Reader::new(stream, self.subscriber.version),
},
});
}
Poll::Ready(Err(err)) => return Poll::Ready(Err(Error::from_transport(err))),
Poll::Pending => break,
}
}
let _ = self.children.poll(waiter);
Poll::Pending
}
}
struct UniServe<S: crate::transport::poll::Session> {
subscriber: Subscriber<S>,
state: UniState<S>,
}
#[allow(clippy::large_enum_variant)]
enum UniState<S: crate::transport::poll::Session> {
Start {
reader: Reader<S::RecvStream, Version>,
},
Setup {
reader: Reader<S::RecvStream, Version>,
},
Group(GroupRecv<S>),
Done,
}
impl<S: crate::transport::poll::Session> kio::Task for UniServe<S> {
fn poll(&mut self, waiter: &kio::Waiter) -> Poll<()> {
if let Err(err) = ready!(self.poll_serve(waiter)) {
tracing::debug!(%err, "error running uni stream");
}
Poll::Ready(())
}
}
impl<S: crate::transport::poll::Session> UniServe<S> {
fn abort(&mut self, err: &Error) {
match &mut self.state {
UniState::Start { reader } | UniState::Setup { reader } => reader.abort(err),
UniState::Group(recv) => recv.reader.abort(err),
UniState::Done => {}
}
}
fn poll_serve(&mut self, waiter: &kio::Waiter) -> Poll<Result<(), Error>> {
loop {
match &mut self.state {
UniState::Start { reader } => {
let mut cx = std::task::Context::from_waker(waiter.waker());
let kind = ready!(reader.poll_decode::<lite::DataType>(&mut cx))?;
let UniState::Start { reader } = std::mem::replace(&mut self.state, UniState::Done) else {
unreachable!()
};
self.state = match kind {
lite::DataType::Group => UniState::Group(GroupRecv::new(self.subscriber.clone(), reader)),
lite::DataType::Setup => UniState::Setup { reader },
};
}
UniState::Setup { reader } => {
if !self.subscriber.version.has_setup_stream() {
let err = Error::UnexpectedStream;
self.abort(&err);
return Poll::Ready(Ok(()));
}
let mut cx = std::task::Context::from_waker(waiter.waker());
let res = ready!(reader.poll_decode::<lite::Setup>(&mut cx));
match res {
Ok(setup) => {
tracing::debug!(?setup, "received peer setup");
self.subscriber.peer_setup.set(setup);
return Poll::Ready(Ok(()));
}
Err(err) => {
self.abort(&err);
return Poll::Ready(Ok(()));
}
}
}
UniState::Group(recv) => {
let res = ready!(recv.poll_serve(waiter));
if let Err(err) = res {
self.abort(&err);
}
return Poll::Ready(Ok(()));
}
UniState::Done => return Poll::Ready(Ok(())),
}
}
}
}
struct GroupRecv<S: crate::transport::poll::Session> {
subscriber: Subscriber<S>,
reader: Reader<S::RecvStream, Version>,
state: GroupRecvState,
}
#[allow(clippy::large_enum_variant)]
enum GroupRecvState {
Header,
Serve {
group: crate::recv::Group,
track: track::Producer,
ingest: FrameIngest,
},
Done,
}
impl<S: crate::transport::poll::Session> GroupRecv<S> {
fn new(subscriber: Subscriber<S>, reader: Reader<S::RecvStream, Version>) -> Self {
Self {
subscriber,
reader,
state: GroupRecvState::Header,
}
}
fn poll_serve(&mut self, waiter: &kio::Waiter) -> Poll<Result<(), Error>> {
loop {
match &mut self.state {
GroupRecvState::Header => {
let mut cx = std::task::Context::from_waker(waiter.waker());
let hdr = ready!(self.reader.poll_decode::<lite::Group>(&mut cx))?;
let (group, track, timescale) = {
let mut subs = self.subscriber.subscribes.lock();
let entry = subs.get_mut(&hdr.subscribe).ok_or(Error::Cancel)?;
if let Ok(mut tail) = entry.tail.write() {
tail.open(hdr.sequence);
}
let group_info = group::Info { sequence: hdr.sequence };
let mut group = entry.producer.create_group(group_info)?;
group.start_at(hdr.frame_start)?;
(group, entry.producer.clone(), entry.timescale)
};
self.state = GroupRecvState::Serve {
group: crate::recv::Group::new(group),
track,
ingest: FrameIngest::new(self.subscriber.runtime.clone(), timescale),
};
}
GroupRecvState::Serve { group, track, ingest } => {
let res = 'serve: {
if let Poll::Ready(err) = track.poll_closed(waiter) {
break 'serve Err(err);
}
if let Poll::Ready(err) = group.poll_closed(waiter) {
break 'serve Err(err);
}
match ingest.poll(&mut self.reader, group, waiter) {
Poll::Ready(res) => break 'serve res,
Poll::Pending => return Poll::Pending,
}
};
let GroupRecvState::Serve { group, .. } = std::mem::replace(&mut self.state, GroupRecvState::Done)
else {
unreachable!()
};
match res {
Ok(()) => {
let _ = group.finish();
}
Err(err @ (Error::Cancel | Error::Stream(crate::StreamError::Cancel))) => {
let _ = group.abort(err);
}
Err(err) => {
tracing::debug!(%err, group = %group.sequence, "group error");
let _ = group.abort(err.clone());
return Poll::Ready(Err(err));
}
}
return Poll::Ready(Ok(()));
}
GroupRecvState::Done => return Poll::Ready(Ok(())),
}
}
}
}
struct FrameIngest {
runtime: crate::time::Clock,
timescale: Option<Timescale>,
prev_ts: u64,
phase: IngestPhase,
}
enum IngestPhase {
Timing,
Size { timestamp: Option<Timestamp> },
Payload { frame: frame::ProducerOwned },
}
impl FrameIngest {
fn new(runtime: crate::time::Clock, timescale: Option<Timescale>) -> Self {
Self {
timescale,
prev_ts: 0,
phase: IngestPhase::Timing,
runtime,
}
}
fn poll<R: crate::transport::poll::RecvStream>(
&mut self,
reader: &mut Reader<R, Version>,
group: &mut group::Producer,
waiter: &kio::Waiter,
) -> Poll<Result<(), Error>> {
let mut cx = std::task::Context::from_waker(waiter.waker());
loop {
match &mut self.phase {
IngestPhase::Timing => {
let Some(scale) = self.timescale else {
self.phase = IngestPhase::Size { timestamp: None };
continue;
};
let Some(zz) = ready!(reader.poll_decode_maybe::<crate::coding::VarInt>(&mut cx))? else {
return Poll::Ready(Ok(()));
};
let next: u64 = (self.prev_ts as i128 + zz.to_zigzag() as i128)
.try_into()
.map_err(|_| Error::BoundsExceeded(crate::coding::BoundsExceeded))?;
self.prev_ts = next;
let timestamp = Timestamp::new(next, scale)
.map_err(|_| Error::BoundsExceeded(crate::coding::BoundsExceeded))?;
self.phase = IngestPhase::Size {
timestamp: Some(timestamp),
};
}
IngestPhase::Size { timestamp } => {
let Some(size) = ready!(reader.poll_decode_maybe::<u64>(&mut cx))? else {
return Poll::Ready(Ok(()));
};
let timestamp = timestamp.unwrap_or_else(|| Timestamp::from(self.runtime.now()));
let frame = group.create_frame_owned(frame::Info { size, timestamp })?;
self.phase = IngestPhase::Payload { frame };
}
IngestPhase::Payload { frame } => {
let failed = ready!(reader.poll_read_frame(&mut cx, frame)).err();
let IngestPhase::Payload { frame } = std::mem::replace(&mut self.phase, IngestPhase::Timing) else {
unreachable!()
};
match failed {
None => frame.finish()?,
Some(err) => {
let _ = frame.abort(err.clone());
return Poll::Ready(Err(err));
}
}
}
}
}
}
}
struct DatagramRecv<S: crate::transport::poll::Session> {
subscriber: Subscriber<S>,
recv: S,
enabled: bool,
}
impl<S: crate::transport::poll::Session> DatagramRecv<S> {
fn new(subscriber: Subscriber<S>) -> Self {
let recv = subscriber.session.clone();
let enabled = subscriber.version.has_datagrams() && recv.max_datagram_size() > 0;
Self {
subscriber,
recv,
enabled,
}
}
fn poll(&mut self, waiter: &kio::Waiter) -> Poll<Result<(), Error>> {
if !self.enabled {
return Poll::Ready(Ok(()));
}
let mut cx = std::task::Context::from_waker(waiter.waker());
loop {
let payload = ready!(self.recv.poll_recv_datagram(&mut cx)).map_err(Error::from_transport)?;
if let Err(err) = self.subscriber.route_datagram(payload) {
tracing::debug!(%err, "dropping datagram");
}
}
}
}
struct RecvBandwidth<S: crate::transport::poll::Session> {
subscriber: Subscriber<S>,
state: BandwidthState<S>,
}
#[allow(clippy::large_enum_variant)]
enum BandwidthState<S: crate::transport::poll::Session> {
Gate,
WaitUsed,
Probing(ProbeStream<S>),
}
impl<S: crate::transport::poll::Session> RecvBandwidth<S> {
fn new(subscriber: Subscriber<S>) -> Self {
Self {
subscriber,
state: BandwidthState::Gate,
}
}
fn poll(&mut self, waiter: &kio::Waiter) -> Poll<Result<(), Error>> {
loop {
match &mut self.state {
BandwidthState::Gate => {
if self.subscriber.recv_bandwidth.is_none() {
return Poll::Ready(Ok(()));
}
if self.subscriber.version.has_setup_stream()
&& ready!(self.subscriber.peer_setup.poll_probe_level(waiter)) < lite::ProbeLevel::Report
{
tracing::debug!("peer does not support probing; skipping probe stream");
return Poll::Ready(Ok(()));
}
self.state = BandwidthState::WaitUsed;
}
BandwidthState::WaitUsed => {
let bandwidth = self.subscriber.recv_bandwidth.as_ref().expect("gated above");
match ready!(bandwidth.poll_used(waiter)) {
Ok(()) => self.state = BandwidthState::Probing(ProbeStream::new(&self.subscriber)),
Err(_) => return Poll::Ready(Ok(())),
}
}
BandwidthState::Probing(probe) => {
let bandwidth = self.subscriber.recv_bandwidth.as_ref().expect("gated above");
match bandwidth.poll_unused(waiter) {
Poll::Ready(Ok(())) => {
self.state = BandwidthState::WaitUsed;
continue;
}
Poll::Ready(Err(_)) => return Poll::Ready(Ok(())),
Poll::Pending => {}
}
match ready!(probe.poll(waiter)) {
Ok(()) => tracing::debug!("probe stream closed"),
Err(err) => tracing::warn!(%err, "probe stream error"),
}
return Poll::Ready(Ok(()));
}
}
}
}
}
struct ProbeStream<S: crate::transport::poll::Session> {
subscriber: Subscriber<S>,
session: S,
state: ProbeState<S>,
}
enum ProbeState<S: crate::transport::poll::Session> {
Open,
Send { stream: Stream<S, Version> },
Read { stream: Stream<S, Version> },
}
impl<S: crate::transport::poll::Session> ProbeStream<S> {
fn new(subscriber: &Subscriber<S>) -> Self {
Self {
subscriber: subscriber.clone(),
session: subscriber.session.clone(),
state: ProbeState::Open,
}
}
fn poll(&mut self, waiter: &kio::Waiter) -> Poll<Result<(), Error>> {
let mut cx = std::task::Context::from_waker(waiter.waker());
loop {
match &mut self.state {
ProbeState::Open => {
if self.subscriber.going_away.is_set() {
return Poll::Ready(Ok(()));
}
let mut stream = ready!(Stream::poll_open(&mut self.session, self.subscriber.version, &mut cx))?;
stream.writer.buffer(&lite::ControlType::Probe)?;
self.state = ProbeState::Send { stream };
}
ProbeState::Send { stream } => {
ready!(stream.writer.poll_flush(&mut cx))?;
let ProbeState::Send { stream } = std::mem::replace(&mut self.state, ProbeState::Open) else {
unreachable!()
};
self.state = ProbeState::Read { stream };
}
ProbeState::Read { stream } => {
let bandwidth = self.subscriber.recv_bandwidth.as_ref().expect("gated by RecvBandwidth");
loop {
let Some(probe) = ready!(stream.reader.poll_decode_maybe::<lite::Probe>(&mut cx))? else {
return Poll::Ready(Ok(()));
};
bandwidth.set(probe.bitrate.map(bandwidth::Rate::from_bps))?;
}
}
}
}
}
}
struct AnnouncePrefix<S: crate::transport::poll::Session> {
subscriber: Subscriber<S>,
prefix: PathOwned,
state: PrefixState<S>,
}
enum PrefixState<S: crate::transport::poll::Session> {
Open,
Send { stream: Stream<S, Version> },
ReadOk { stream: Stream<S, Version> },
Cost {
stream: Stream<S, Version>,
responder_origin: Option<crate::Hop>,
},
ReadInit { stream: Stream<S, Version>, run: PrefixRun },
Run { stream: Stream<S, Version>, run: PrefixRun },
}
struct PrefixRun {
responder_origin: Option<crate::Hop>,
link_cost: u64,
announced: Announced,
next_announce_id: u64,
announced_by_id: HashMap<u64, PathOwned>,
}
impl<S: crate::transport::poll::Session> AnnouncePrefix<S> {
fn new(subscriber: Subscriber<S>, prefix: PathOwned) -> Self {
Self {
subscriber,
prefix,
state: PrefixState::Open,
}
}
fn poll(&mut self, waiter: &kio::Waiter) -> Poll<Result<(), Error>> {
let mut cx = std::task::Context::from_waker(waiter.waker());
loop {
match &mut self.state {
PrefixState::Open => {
self.subscriber.check_going_away()?;
let mut stream = ready!(Stream::poll_open(
&mut self.subscriber.session,
self.subscriber.version,
&mut cx
))?;
stream.writer.buffer(&lite::ControlType::Announce)?;
stream.writer.buffer(&lite::AnnounceRequest {
prefix: self.prefix.as_path(),
exclude_hop: self.subscriber.self_origin.id(),
hidden: true,
})?;
self.state = PrefixState::Send { stream };
}
PrefixState::Send { stream } => {
ready!(stream.writer.poll_flush(&mut cx))?;
let PrefixState::Send { stream } = std::mem::replace(&mut self.state, PrefixState::Open) else {
unreachable!()
};
self.state = match self.subscriber.version.has_announce_ok() {
true => PrefixState::ReadOk { stream },
false => PrefixState::Cost {
stream,
responder_origin: None,
},
};
}
PrefixState::ReadOk { stream } => {
let ok = ready!(stream.reader.poll_decode::<lite::AnnounceOk>(&mut cx))?;
let origin = ok.origin;
let PrefixState::ReadOk { stream } = std::mem::replace(&mut self.state, PrefixState::Open) else {
unreachable!()
};
self.state = PrefixState::Cost {
stream,
responder_origin: Some(origin),
};
}
PrefixState::Cost { .. } => {
let link_cost = ready!(self.subscriber.poll_link_cost(waiter));
let PrefixState::Cost {
stream,
responder_origin,
} = std::mem::replace(&mut self.state, PrefixState::Open)
else {
unreachable!()
};
let run = PrefixRun {
responder_origin,
link_cost,
announced: Announced::default(),
next_announce_id: 0,
announced_by_id: HashMap::new(),
};
self.state = match self.subscriber.version {
Version::Lite01 | Version::Lite02 => PrefixState::ReadInit { stream, run },
_ => PrefixState::Run { stream, run },
};
}
PrefixState::ReadInit { stream, run } => {
let msg = ready!(stream.reader.poll_decode::<lite::AnnounceInit>(&mut cx))?;
for suffix in msg.suffixes {
let path = self.prefix.join(&suffix);
self.subscriber.start_announce(
path,
crate::Hops::new(),
crate::origin::Cost::UNKNOWN,
0,
run.responder_origin,
&mut run.announced,
)?;
}
let PrefixState::ReadInit { stream, run } = std::mem::replace(&mut self.state, PrefixState::Open)
else {
unreachable!()
};
self.state = PrefixState::Run { stream, run };
}
PrefixState::Run { stream, run } => {
if self.subscriber.going_away.poll(waiter).is_ready() {
run.announced.drain();
}
loop {
match stream.reader.poll_decode_maybe::<lite::AnnounceBroadcast>(&mut cx) {
Poll::Ready(Ok(Some(announce))) => {
self.subscriber.handle_announce(&self.prefix, announce, run)?;
}
Poll::Ready(Ok(None)) => {
stream.writer.finish().ok();
return Poll::Ready(Ok(()));
}
Poll::Ready(Err(err)) => return Poll::Ready(Err(err)),
Poll::Pending => break,
}
}
run.announced.poll_serve(&self.subscriber, waiter);
return Poll::Pending;
}
}
}
}
}
struct SourceServe<S: crate::transport::poll::Session> {
subscriber: Subscriber<S>,
path: PathOwned,
dynamic: crate::broadcast::Dynamic,
closed: S,
tracks: kio::Tasks<TrackServeRun<S>>,
ended: bool,
}
impl<S: crate::transport::poll::Session> SourceServe<S> {
fn new(subscriber: Subscriber<S>, path: PathOwned, dynamic: crate::broadcast::Dynamic) -> Self {
let closed = subscriber.session.clone();
Self {
subscriber,
path,
dynamic,
closed,
tracks: kio::Tasks::new(),
ended: false,
}
}
}
impl<S: crate::transport::poll::Session> kio::Task for SourceServe<S> {
fn poll(&mut self, waiter: &kio::Waiter) -> Poll<()> {
let _ = self.tracks.poll(waiter);
let mut cx = std::task::Context::from_waker(waiter.waker());
loop {
if self.closed.poll_closed(&mut cx).is_ready() {
return Poll::Ready(());
}
if self.ended {
return self.tracks.poll(waiter);
}
match self.dynamic.poll_requested_track(waiter) {
Poll::Ready(Ok(request)) => {
let serve = TrackServe {
subscriber: self.subscriber.clone(),
path: self.path.clone(),
name: request.name().to_string(),
};
self.tracks.push(TrackServeRun::new(serve, request));
}
Poll::Ready(Err(err)) => {
tracing::debug!(%err, "source closed");
self.ended = true;
}
Poll::Pending => break,
}
}
let _ = self.tracks.poll(waiter);
Poll::Pending
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::coding::{Decode, Encode};
use crate::lite::test_transport::SinkSession;
use crate::model::ProduceTest;
use futures::FutureExt;
const VERSION: Version = Version::Lite05;
#[test]
fn unsubscribe_drops_the_datagram_and_releases_the_producer() {
let origin = origin::Config::new(crate::Hop::new(1).unwrap()).produce();
let subscriber = Subscriber::new(SubscriberConfig {
runtime: crate::time::Clock::tokio(),
session: SinkSession::default(),
origin,
recv_bandwidth: None,
version: VERSION,
peer_setup: Default::default(),
peer_hop: None,
cost: None,
going_away: Default::default(),
});
let broadcast = crate::broadcast::Info::new().produce();
let producer = broadcast.create_track("datagrams", None).unwrap();
let mut received = producer.subscribe(None);
subscriber.subscribes.lock().insert(
7,
TrackEntry {
producer,
timescale: Some(Timescale::default()),
tail: Default::default(),
},
);
let payload = |sequence| {
lite::Datagram {
subscribe: 7,
sequence,
timestamp: sequence,
payload: bytes::Bytes::from_static(b"x"),
}
.encode_bytes(VERSION)
.unwrap()
};
subscriber.route_datagram(payload(1)).unwrap();
assert_eq!(
received
.recv_datagram()
.now_or_never()
.unwrap()
.unwrap()
.unwrap()
.sequence,
1
);
subscriber.remove_subscribe(7);
subscriber.route_datagram(payload(2)).unwrap();
assert!(
matches!(received.recv_datagram().now_or_never(), Some(Err(Error::Dropped))),
"the track outlived its subscription"
);
}
#[tokio::test]
async fn establish_sends_one_registered_subscribe() {
let gate = kio::Producer::new(false);
let session = SinkSession::gated_bi(gate.consume());
let origin = origin::Config::new(crate::Hop::new(1).unwrap()).produce();
let subscriber = Subscriber::new(SubscriberConfig {
runtime: crate::time::Clock::tokio(),
session: session.clone(),
origin,
recv_bandwidth: None,
version: VERSION,
peer_setup: Default::default(),
peer_hop: None,
cost: None,
going_away: Default::default(),
});
let subscribes = subscriber.subscribes.clone();
let serve = TrackServe {
subscriber,
path: Path::new("room/host").to_owned(),
name: "catalog.json".to_string(),
};
let broadcast = crate::broadcast::Info::new().produce();
let mut producer = broadcast.create_track("catalog.json", None).unwrap();
let mut sub = Sub::None;
let mut establish = std::pin::pin!(serve.establish(
&mut producer,
&mut sub,
Subscription::default(),
Some(Timescale::default()),
));
assert!(futures::poll!(establish.as_mut()).is_pending());
assert_eq!(session.log.bi_opens(), 1);
assert!(subscribes.lock().contains_key(&0), "registered before the wire");
let Ok(mut open) = gate.write() else {
panic!("gate closed")
};
*open = true;
drop(open);
establish.await.unwrap();
assert_eq!(session.log.bi_opens(), 1);
let writes = session.log.writes.lock().unwrap().clone();
let mut wire = writes.as_slice();
assert_eq!(
lite::ControlType::decode(&mut wire, VERSION).unwrap(),
lite::ControlType::Subscribe
);
let msg = lite::Subscribe::decode(&mut wire, VERSION).unwrap();
assert_eq!(msg.id, 0);
assert_eq!(msg.track, "catalog.json");
assert!(wire.is_empty(), "a second SUBSCRIBE trailed the first");
}
struct Harness {
serve: TrackServe<SinkSession>,
session: SinkSession,
producer: track::Producer,
_broadcast: crate::broadcast::Producer,
_gate: kio::Producer<bool>,
}
impl Harness {
fn new(version: Version) -> Self {
let gate = kio::Producer::new(true);
let session = SinkSession::gated_bi(gate.consume());
let origin = origin::Config::new(crate::Hop::new(1).unwrap()).produce();
let subscriber = Subscriber::new(SubscriberConfig {
runtime: crate::time::Clock::tokio(),
session: session.clone(),
origin,
recv_bandwidth: None,
version,
peer_setup: Default::default(),
peer_hop: None,
cost: None,
going_away: Default::default(),
});
let broadcast = crate::broadcast::Info::new().produce();
let producer = broadcast.create_track("catalog.json", None).unwrap();
Self {
serve: TrackServe {
subscriber,
path: Path::new("room/host").to_owned(),
name: "catalog.json".to_string(),
},
session,
producer,
_broadcast: broadcast,
_gate: gate,
}
}
fn wire(&self) -> Vec<u8> {
self.session.log.writes.lock().unwrap().clone()
}
}
fn mid_group_demand() -> Subscription {
Subscription::default()
.with_start(Position { group: 5, frame: 3 })
.with_end(Position::after(5, 7))
}
#[tokio::test]
async fn frame_bounds_widen_for_an_older_peer() {
let mut h = Harness::new(Version::Lite05);
let mut sub = Sub::None;
h.serve
.handle_subscription(
&mut h.producer,
&mut sub,
Some(mid_group_demand()),
true,
Some(Timescale::default()),
)
.await
.expect("an older peer must not fail the subscribe");
let wire = h.wire();
let mut wire = wire.as_slice();
assert_eq!(
lite::ControlType::decode(&mut wire, Version::Lite05).unwrap(),
lite::ControlType::Subscribe
);
let msg = lite::Subscribe::decode(&mut wire, Version::Lite05).unwrap();
assert_eq!((msg.start_group, msg.end_group), (Some(5), Some(5)));
assert_eq!((msg.start_frame, msg.end_frame), (0, None));
}
#[tokio::test]
async fn frame_bounds_survive_on_a_lite06_peer() {
let mut h = Harness::new(Version::Lite06);
let mut sub = Sub::None;
h.serve
.handle_subscription(
&mut h.producer,
&mut sub,
Some(mid_group_demand()),
true,
Some(Timescale::default()),
)
.await
.unwrap();
let wire = h.wire();
let mut wire = wire.as_slice();
assert_eq!(
lite::ControlType::decode(&mut wire, Version::Lite06).unwrap(),
lite::ControlType::Subscribe
);
let msg = lite::Subscribe::decode(&mut wire, Version::Lite06).unwrap();
assert_eq!((msg.start_frame, msg.end_frame), (3, Some(7)));
}
#[test]
fn wire_bounds_convert_the_exclusive_end() {
let bounds = WireBounds::new(None, Some(Position::group(6)));
assert_eq!((bounds.end_group, bounds.end_frame), (Some(5), None));
let bounds = WireBounds::new(None, Some(Position { group: 5, frame: 3 }));
assert_eq!((bounds.end_group, bounds.end_frame), (Some(5), Some(2)));
let bounds = WireBounds::new(None, None);
assert_eq!((bounds.end_group, bounds.end_frame), (None, None));
let bounds = WireBounds::new(Some(Position { group: 5, frame: 3 }), None);
assert_eq!((bounds.start_group, bounds.start_frame), (Some(5), 3));
}
#[test]
fn wire_bounds_match_the_builders() {
let whole = Subscription::default().with_end(Position::after_group(5));
let bounds = WireBounds::new(whole.start, whole.end);
assert_eq!((bounds.end_group, bounds.end_frame), (Some(5), None));
let capped = Subscription::default().with_end(Position::after(5, 2));
let bounds = WireBounds::new(capped.start, capped.end);
assert_eq!((bounds.end_group, bounds.end_frame), (Some(5), Some(2)));
let started = Subscription::default().with_start(Position { group: 5, frame: 3 });
let bounds = WireBounds::new(started.start, started.end);
assert_eq!((bounds.start_group, bounds.start_frame), (Some(5), 3));
}
#[tokio::test]
async fn an_empty_range_opens_no_subscription() {
let mut h = Harness::new(Version::Lite06);
let mut sub = Sub::None;
let empty = Subscription::default().with_end(Position::group(0));
h.serve
.handle_subscription(&mut h.producer, &mut sub, Some(empty), true, Some(Timescale::default()))
.await
.unwrap();
assert!(matches!(sub, Sub::None), "must not open a subscription");
assert!(h.wire().is_empty(), "nothing reached the wire");
}
#[tokio::test]
async fn an_empty_range_cancels_a_live_subscription() {
let mut h = Harness::new(Version::Lite06);
let mut sub = Sub::None;
h.serve
.handle_subscription(
&mut h.producer,
&mut sub,
Some(Subscription::default()),
true,
Some(Timescale::default()),
)
.await
.unwrap();
assert!(matches!(sub, Sub::Active(_)), "the first subscriber opens one");
let established = h.wire().len();
let empty = Subscription::default().with_end(Position::group(0));
h.serve
.handle_subscription(&mut h.producer, &mut sub, Some(empty), true, Some(Timescale::default()))
.await
.unwrap();
assert!(matches!(sub, Sub::None), "the upstream must be canceled");
assert_eq!(h.wire().len(), established, "no SUBSCRIBE_UPDATE claiming group 0");
}
#[tokio::test]
async fn a_nonzero_empty_range_opens_no_subscription() {
let mut h = Harness::new(Version::Lite06);
let mut sub = Sub::None;
let empty = Subscription::default()
.with_start(Position::group(5))
.with_end(Position::group(5));
h.serve
.handle_subscription(&mut h.producer, &mut sub, Some(empty), true, Some(Timescale::default()))
.await
.unwrap();
assert!(matches!(sub, Sub::None), "must not open a subscription");
assert!(h.wire().is_empty(), "nothing reached the wire");
}
#[tokio::test]
async fn a_nonzero_empty_range_cancels_a_live_subscription() {
let mut h = Harness::new(Version::Lite06);
let mut sub = Sub::None;
h.serve
.handle_subscription(
&mut h.producer,
&mut sub,
Some(Subscription::default()),
true,
Some(Timescale::default()),
)
.await
.unwrap();
assert!(matches!(sub, Sub::Active(_)), "the first subscriber opens one");
let established = h.wire().len();
let empty = Subscription::default()
.with_start(Position::group(5))
.with_end(Position::group(5));
h.serve
.handle_subscription(&mut h.producer, &mut sub, Some(empty), true, Some(Timescale::default()))
.await
.unwrap();
assert!(matches!(sub, Sub::None), "the upstream must be canceled");
assert_eq!(h.wire().len(), established, "no SUBSCRIBE_UPDATE inverting the range");
}
#[tokio::test]
async fn frame_bounds_widen_outward_at_the_last_group() {
let h = Harness::new(Version::Lite05);
let mut subscription = Subscription::default()
.with_start(Position {
group: u64::MAX,
frame: 1,
})
.with_end(Position::after(u64::MAX, 5));
h.serve.widen_frame_bounds(&mut subscription);
assert_eq!(subscription.start, Some(Position::group(u64::MAX)));
assert_eq!(subscription.end, None);
}
#[tokio::test]
async fn frame_bounds_widen_on_update() {
let mut h = Harness::new(Version::Lite05);
let mut sub = Sub::None;
h.serve
.handle_subscription(
&mut h.producer,
&mut sub,
Some(Subscription::default()),
true,
Some(Timescale::default()),
)
.await
.unwrap();
let established = h.wire().len();
h.serve
.handle_subscription(
&mut h.producer,
&mut sub,
Some(mid_group_demand()),
true,
Some(Timescale::default()),
)
.await
.expect("a downstream frame offset must not tear down an older upstream");
let wire = h.wire();
let mut wire = &wire[established..];
let msg = lite::SubscribeUpdate::decode(&mut wire, Version::Lite05).unwrap();
assert_eq!((msg.start_group, msg.start_frame), (Some(5), 0));
assert_eq!((msg.end_group, msg.end_frame), (Some(5), None));
}
#[tokio::test]
async fn buffered_start_applies_iff_demand_matches() {
let mut h = Harness::new(Version::Lite05);
let mut sub = Sub::None;
let demand = |group: u64| Some(Subscription::default().with_start(Position::group(group)));
let applies = |sub: &Sub<SinkSession>| matches!(sub, Sub::Active(active) if active.start == active.requested);
h.serve
.handle_subscription(&mut h.producer, &mut sub, demand(3), true, Some(Timescale::default()))
.await
.unwrap();
assert!(applies(&sub), "a fresh subscription accepts its START");
h.serve
.handle_subscription(
&mut h.producer,
&mut sub,
Some(
Subscription::default()
.with_start(Position::group(3))
.with_end(Position::group(9)),
),
true,
Some(Timescale::default()),
)
.await
.unwrap();
assert!(applies(&sub), "an unmoved start keeps the START applicable");
h.serve
.handle_subscription(&mut h.producer, &mut sub, demand(8), true, Some(Timescale::default()))
.await
.unwrap();
assert!(!applies(&sub), "a moved start must invalidate a buffered START");
h.serve
.handle_subscription(&mut h.producer, &mut sub, demand(3), true, Some(Timescale::default()))
.await
.unwrap();
assert!(applies(&sub), "demand returning restores the START");
}
#[tokio::test]
async fn updates_move_the_declared_floor_both_ways() {
let mut h = Harness::new(Version::Lite05);
let mut sub = Sub::None;
let demand = |group: u64| Some(Subscription::default().with_start(Position::group(group)));
h.serve
.handle_subscription(&mut h.producer, &mut sub, demand(5), true, Some(Timescale::default()))
.await
.unwrap();
assert_eq!(h.producer.start_sequence(), Some(5));
h.serve
.handle_subscription(&mut h.producer, &mut sub, demand(8), true, Some(Timescale::default()))
.await
.unwrap();
assert_eq!(h.producer.start_sequence(), Some(8));
h.serve
.handle_subscription(&mut h.producer, &mut sub, demand(3), true, Some(Timescale::default()))
.await
.unwrap();
assert_eq!(h.producer.start_sequence(), Some(3));
h.serve
.handle_subscription(
&mut h.producer,
&mut sub,
Some(Subscription::default()),
true,
Some(Timescale::default()),
)
.await
.unwrap();
assert_eq!(h.producer.start_sequence(), None);
}
#[tokio::test]
async fn a_double_announce_is_an_error_even_when_reflected() {
let assigned = crate::Hop::new(777).unwrap();
let origin = origin::Config::new(crate::Hop::new(1).unwrap()).produce();
let mut subscriber = Subscriber::new(SubscriberConfig {
runtime: crate::time::Clock::tokio(),
session: SinkSession::new(Default::default()),
origin,
recv_bandwidth: None,
version: VERSION,
peer_setup: Default::default(),
cost: None,
peer_hop: Some(assigned),
going_away: Default::default(),
});
let path = Path::new("room/host").to_owned();
let mut announced = Announced::default();
assert!(
subscriber
.start_announce(
path.clone(),
crate::Hops::new(),
crate::origin::Cost::default(),
0,
Some(assigned),
&mut announced,
)
.unwrap()
);
let mut reflected = crate::Hops::new();
reflected.push(assigned).unwrap();
assert!(
matches!(
subscriber.start_announce(
path.clone(),
reflected,
crate::origin::Cost::default(),
0,
Some(assigned),
&mut announced,
),
Err(Error::ProtocolViolation)
),
"the double announce must be reported, not silently dropped",
);
}
#[tokio::test]
async fn every_declined_announce_is_recorded() {
let assigned = crate::Hop::new(777).unwrap();
let origin = origin::Config::new(crate::Hop::new(1).unwrap()).produce();
let mut subscriber = Subscriber::new(SubscriberConfig {
runtime: crate::time::Clock::tokio(),
session: SinkSession::new(Default::default()),
origin,
recv_bandwidth: None,
version: Version::Lite03,
peer_setup: Default::default(),
cost: None,
peer_hop: Some(assigned),
going_away: Default::default(),
});
let path = Path::new("room/host").to_owned();
let mut announced = Announced::default();
let hops = crate::Hops::try_from(vec![crate::Hop::new(1).unwrap()]).unwrap();
assert!(
!subscriber
.start_announce(
path.clone(),
hops,
crate::origin::Cost::default(),
0,
None,
&mut announced
)
.unwrap(),
"a chain that already names this session must be declined",
);
let mut fresh = crate::Hops::new();
fresh.push(crate::Hop::new(7).unwrap()).unwrap();
assert!(
matches!(
subscriber.start_announce(
path.clone(),
fresh,
crate::origin::Cost::default(),
0,
None,
&mut announced
),
Err(Error::ProtocolViolation)
),
"a declined announce still holds its path, so a second start for it is a violation",
);
}
#[tokio::test]
async fn a_dropped_announce_still_holds_its_path() {
let assigned = crate::Hop::new(777).unwrap();
let origin = origin::Config::new(crate::Hop::new(1).unwrap()).produce();
let consumer = origin.consume();
let mut subscriber = Subscriber::new(SubscriberConfig {
runtime: crate::time::Clock::tokio(),
session: SinkSession::new(Default::default()),
origin,
recv_bandwidth: None,
version: VERSION,
peer_setup: Default::default(),
cost: None,
peer_hop: Some(assigned),
going_away: Default::default(),
});
let path = Path::new("room/host").to_owned();
let mut announced = Announced::default();
let mut reflected = crate::Hops::new();
reflected.push(assigned).unwrap();
assert!(
!subscriber
.start_announce(
path.clone(),
reflected,
crate::origin::Cost::default(),
0,
Some(assigned),
&mut announced,
)
.unwrap(),
"a chain naming its own sender must be dropped",
);
assert!(consumer.get_broadcast("room/host").is_none());
let mut hops = crate::Hops::new();
hops.push(crate::Hop::new(7).unwrap()).unwrap();
let err = subscriber
.start_announce(
path.clone(),
hops,
crate::origin::Cost::default(),
0,
Some(assigned),
&mut announced,
)
.expect_err("a second start for the peer-owned path must be rejected");
assert!(matches!(err, Error::ProtocolViolation));
}
#[tokio::test]
async fn an_announce_reflected_by_its_sender_is_dropped() {
let assigned = crate::Hop::new(777).unwrap();
let origin = origin::Config::new(crate::Hop::new(1).unwrap()).produce();
let consumer = origin.consume();
let mut subscriber = Subscriber::new(SubscriberConfig {
runtime: crate::time::Clock::tokio(),
session: SinkSession::new(Default::default()),
origin,
recv_bandwidth: None,
version: VERSION,
peer_setup: Default::default(),
cost: None,
peer_hop: Some(assigned),
going_away: Default::default(),
});
let mut hops = crate::Hops::new();
hops.push(assigned).unwrap();
let mut announced = Announced::default();
let accepted = subscriber
.start_announce(
Path::new("room/host").to_owned(),
hops,
crate::origin::Cost::default(),
0,
Some(assigned),
&mut announced,
)
.unwrap();
assert!(!accepted, "a chain naming its own sender must not become a route");
tokio::time::sleep(std::time::Duration::from_millis(1)).await;
assert!(consumer.get_broadcast("room/host").is_none());
}
#[tokio::test]
async fn a_reflected_announce_does_not_displace_the_local_front() {
let relay = crate::Hop::new(1).unwrap();
let origin = origin::Config::new(relay).produce();
let assigned = crate::Hop::new(777).unwrap();
let local = origin.publish("room/host", origin::Route::default()).unwrap();
let track = local.create_track("video", None).unwrap();
let mut group = track.append_group().unwrap();
group.write_frame(crate::Timestamp::ZERO, b"local".as_ref()).unwrap();
group.finish().unwrap();
let peer = origin.consume().excluding(assigned);
let resolved = peer.request_broadcast("room/host").await.expect("resolves");
let mut sub = resolved
.track("video")
.unwrap()
.subscribe(None)
.await
.expect("subscribe");
let mut group = sub.recv_group().await.expect("recv group").expect("track ended early");
assert_eq!(
&group.read_frame().await.expect("read frame").expect("frame").payload[..],
b"local"
);
let mut subscriber = Subscriber::new(SubscriberConfig {
runtime: crate::time::Clock::tokio(),
session: SinkSession::new(Default::default()),
origin: origin.clone(),
recv_bandwidth: None,
version: VERSION,
peer_setup: Default::default(),
cost: None,
peer_hop: Some(assigned),
going_away: Default::default(),
});
let mut hops = crate::Hops::new();
hops.push(relay).unwrap();
let mut announced = Announced::default();
let accepted = subscriber
.start_announce(
Path::new("room/host").to_owned(),
hops,
crate::origin::Cost::default(),
0,
Some(assigned),
&mut announced,
)
.unwrap();
assert!(!accepted, "an announce that already names this origin must be dropped");
let still = peer
.request_broadcast("room/host")
.await
.expect("the local front keeps serving");
assert!(
still.is_clone(&resolved),
"the reflected announce must not replace the local front"
);
let mut group = track.append_group().unwrap();
group.write_frame(crate::Timestamp::ZERO, b"still".as_ref()).unwrap();
group.finish().unwrap();
let mut group = sub.recv_group().await.expect("recv group").expect("track ended early");
assert_eq!(
&group.read_frame().await.expect("read frame").expect("frame").payload[..],
b"still"
);
}
#[tokio::test]
async fn assigned_peer_hop_attributes_announces() {
let session = SinkSession::new(Default::default());
let assigned = crate::Hop::new(777).unwrap();
let origin = origin::Config::new(crate::Hop::new(1).unwrap()).produce();
let consumer = origin.consume();
let mut subscriber = Subscriber::new(SubscriberConfig {
runtime: crate::time::Clock::tokio(),
session,
origin,
recv_bandwidth: None,
version: VERSION,
peer_setup: Default::default(),
peer_hop: Some(assigned),
cost: None,
going_away: Default::default(),
});
let mut announced = Announced::default();
let accepted = subscriber
.start_announce(
Path::new("room/host").to_owned(),
crate::Hops::new(),
crate::origin::Cost::UNKNOWN,
0,
None,
&mut announced,
)
.unwrap();
assert!(accepted);
let mut cursor = consumer.announced();
let route = cursor.assert_next_active("room/host");
let hops: Vec<_> = route.hops.iter().copied().collect();
assert_eq!(hops, vec![crate::Hop::UNKNOWN]);
assert!(route.is_anonymous());
let mut hidden = consumer.excluding(assigned).announced();
hidden.assert_next_wait();
}
#[tokio::test]
async fn lite03_placeholders_stay_anonymous() {
let assigned = crate::Hop::new(777).unwrap();
let origin = origin::Config::new(crate::Hop::new(1).unwrap()).produce();
let consumer = origin.consume();
let mut subscriber = Subscriber::new(SubscriberConfig {
runtime: crate::time::Clock::tokio(),
session: SinkSession::new(Default::default()),
origin,
recv_bandwidth: None,
version: Version::Lite03,
peer_setup: Default::default(),
cost: None,
peer_hop: Some(assigned),
going_away: Default::default(),
});
let hops = crate::Hops::try_from(vec![crate::Hop::UNKNOWN, crate::Hop::UNKNOWN]).unwrap();
let mut announced = Announced::default();
assert!(
subscriber
.start_announce(
Path::new("room/host").to_owned(),
hops,
crate::origin::Cost::UNKNOWN,
0,
None,
&mut announced,
)
.unwrap()
);
let mut cursor = consumer.announced();
let route = cursor.assert_next_active("room/host");
let hops: Vec<_> = route.hops.iter().copied().collect();
assert_eq!(hops, vec![crate::Hop::UNKNOWN, crate::Hop::UNKNOWN]);
assert!(route.is_anonymous());
}
#[tokio::test]
async fn absent_peer_hop_stamps_unknown() {
let (mut subscriber, consumer) = restart_subscriber(SinkSession::new(Default::default()));
let mut announced = Announced::default();
subscriber
.start_announce(
Path::new("room/host").to_owned(),
crate::Hops::new(),
crate::origin::Cost::UNKNOWN,
0,
None,
&mut announced,
)
.unwrap();
let mut cursor = consumer.announced();
let route = cursor.assert_next_active("room/host");
let hops: Vec<_> = route.hops.iter().copied().collect();
assert_eq!(hops, vec![crate::Hop::UNKNOWN]);
}
fn restart_subscriber(session: SinkSession) -> (Subscriber<SinkSession>, crate::origin::Consumer) {
let origin = origin::Config::new(crate::Hop::new(1).unwrap()).produce();
let consumer = origin.consume();
let subscriber = Subscriber::new(SubscriberConfig {
runtime: crate::time::Clock::tokio(),
session,
origin,
recv_bandwidth: None,
version: VERSION,
peer_setup: Default::default(),
cost: None,
peer_hop: None,
going_away: Default::default(),
});
(subscriber, consumer)
}
#[tokio::test]
async fn restart_updates_the_route_in_place() {
let (mut subscriber, consumer) = restart_subscriber(SinkSession::new(Default::default()));
let mut announced = Announced::default();
let path = Path::new("room/host").to_owned();
subscriber
.start_announce(
path.clone(),
crate::Hops::new(),
crate::origin::Cost::UNKNOWN,
0,
Some(crate::Hop::new(7).unwrap()),
&mut announced,
)
.unwrap();
let mut cursor = consumer.announced();
cursor.assert_next_active("room/host");
subscriber
.restart_announce(
path.clone(),
crate::Hops::new(),
crate::origin::Cost::new(5),
0,
Some(crate::Hop::new(7).unwrap()),
&mut announced,
)
.unwrap();
let route = cursor.assert_next_active("room/host");
assert_eq!(route.cost, crate::origin::Cost::new(5).charged(0));
}
#[tokio::test(start_paused = true)]
async fn a_lost_announce_stream_retracts_the_route() {
let origin = crate::origin::Config::new(crate::Hop::new(1).unwrap()).produce();
let consumer = origin.consume();
let mut subscriber = Subscriber::new(SubscriberConfig {
runtime: crate::time::Clock::tokio(),
session: SinkSession::new(Default::default()),
origin,
recv_bandwidth: None,
version: VERSION,
peer_setup: Default::default(),
cost: None,
peer_hop: None,
going_away: Default::default(),
});
let path = Path::new("room/host").to_owned();
let hops = crate::Hops::try_from(vec![crate::Hop::new(7).unwrap()]).unwrap();
let mut announced = Announced::default();
subscriber
.start_announce(
path.clone(),
hops,
crate::origin::Cost::default(),
1,
None,
&mut announced,
)
.unwrap();
let mut cursor = consumer.announced();
cursor.assert_next_active("room/host");
drop(announced);
cursor.assert_next_ended("room/host");
let hops = crate::Hops::try_from(vec![crate::Hop::new(7).unwrap()]).unwrap();
let mut announced = Announced::default();
subscriber
.start_announce(
path.clone(),
hops,
crate::origin::Cost::default(),
1,
None,
&mut announced,
)
.unwrap();
cursor.assert_next_active("room/host");
assert!(announced.contains(&path.clone()), "the announce was not recorded");
announced.retire(&path.clone());
cursor.assert_next_ended("room/host");
}
}
struct WireBounds {
start_group: Option<u64>,
start_frame: u64,
end_group: Option<u64>,
end_frame: Option<u64>,
}
impl WireBounds {
fn new(start: Option<Position>, end: Option<Position>) -> Self {
debug_assert!(
end.is_none_or(|end| end > start.unwrap_or_default()),
"an empty range cannot be encoded; it should have been dropped as no demand"
);
let (end_group, end_frame) = match end {
Some(end) if end.frame == 0 => (Some(end.group.saturating_sub(1)), None),
Some(end) => (Some(end.group), Some(end.frame - 1)),
None => (None, None),
};
Self {
start_group: start.map(|start| start.group),
start_frame: start.map_or(0, |start| start.frame),
end_group,
end_frame,
}
}
}
struct SubStream<S: crate::transport::poll::Session> {
stream: Stream<S, Version>,
id: u64,
max_age: Duration,
start: Option<Position>,
priority: u8,
requested: Option<Position>,
tail: kio::Producer<Tail>,
served: Option<u64>,
end: Option<u64>,
}
impl<S: crate::transport::poll::Session> SubStream<S> {
fn owed(&self, requested_end: Option<u64>) -> Option<std::ops::Range<u64>> {
let end = self.end?;
let end = requested_end.map_or(end, |requested| requested.min(end));
Some(self.served.unwrap_or(end)..end)
}
}
enum Sub<S: crate::transport::poll::Session> {
None,
Active(SubStream<S>),
}
#[derive(Default)]
struct Announced(HashMap<PathOwned, Option<AnnouncedRoute>>);
impl Announced {
fn contains(&self, path: &PathOwned) -> bool {
self.0.contains_key(path)
}
fn attach(&mut self, path: PathOwned, route: AnnouncedRoute) {
self.0.insert(path, Some(route));
}
fn declined(&mut self, path: PathOwned) {
if let Some(Some(route)) = self.0.insert(path, None) {
route.finish();
}
}
fn reserve(&mut self, path: PathOwned) {
debug_assert!(!self.0.contains_key(&path), "reserved a prefix already advertised");
self.0.insert(path, None);
}
fn attached(&mut self, path: &PathOwned) -> Option<&mut AnnouncedRoute> {
self.0.get_mut(path)?.as_mut()
}
fn retire(&mut self, path: &PathOwned) {
if let Some(Some(route)) = self.0.remove(path) {
route.finish();
}
}
fn poll_serve<S: crate::transport::poll::Session>(&mut self, subscriber: &Subscriber<S>, waiter: &kio::Waiter) {
let root = subscriber.origin.root().to_owned();
for entry in self.0.values_mut().flatten() {
while let Poll::Ready(Ok(request)) = entry.dynamic.poll_requested_broadcast(waiter) {
let Some(path) = request.path().strip_prefix(&root) else {
continue;
};
let path = path.to_owned();
let source = subscriber.origin.create_source(&path);
let _ = subscriber.sources.try_push((path.clone(), source.dynamic()));
request.accept(&source);
entry
.sources
.insert(path, crate::model::broadcast::SourceGuard::new(source));
}
}
}
fn drain(&mut self) {
for entry in self.0.values_mut().flatten() {
entry.drain();
}
}
}
struct AnnouncedRoute {
route: crate::origin::Route,
dynamic: crate::origin::Dynamic,
sources: HashMap<PathOwned, crate::model::broadcast::SourceGuard>,
drained: bool,
}
impl AnnouncedRoute {
fn new(route: crate::origin::Route, dynamic: crate::origin::Dynamic) -> Self {
Self {
route,
dynamic,
sources: HashMap::new(),
drained: false,
}
}
fn finish(self) {
for (_, source) in self.sources {
source.finish();
}
}
fn update(&mut self, route: crate::origin::Route) {
self.route = route.clone();
self.drained = false;
let _ = self.dynamic.update(route);
}
fn drain(&mut self) {
if self.drained {
return;
}
self.drained = true;
let mut route = self.route.clone();
route.cost = crate::origin::Cost::DRAIN;
let _ = self.dynamic.update(route);
}
}
enum ServeEnd {
Finished,
GiveBack(Error),
Idle,
}
#[derive(Clone)]
struct TrackServe<S: crate::transport::poll::Session> {
subscriber: Subscriber<S>,
path: PathOwned,
name: String,
}
impl<S: crate::transport::poll::Session> TrackServe<S> {
fn widen_frame_bounds(&self, subscription: &mut Subscription) {
if self.subscriber.version.has_frame_bounds() {
return;
}
let start = subscription.start.map(|start| Position::group(start.group));
let end = subscription.end.and_then(|end| match end.frame {
0 => Some(end),
_ => Position::after_group(end.group),
});
if (start, end) != (subscription.start, subscription.end) {
tracing::debug!(
track = %self.name,
version = ?self.subscriber.version,
"widening frame bounds to whole groups for an older peer"
);
}
subscription.start = start;
subscription.end = end;
}
fn begin_subscription(
&self,
producer: &mut track::Producer,
sub: &mut Sub<S>,
pref: Option<Subscription>,
supports_update: bool,
timescale: Option<Timescale>,
) -> Result<Begin<S>, Error> {
let pref = pref.filter(|sub| sub.end.is_none_or(|end| end > sub.start.unwrap_or_default()));
match pref {
Some(mut subscription) => {
self.widen_frame_bounds(&mut subscription);
match sub {
Sub::None => {
Ok(Begin::Establish(self.prepare_establish(
producer,
subscription,
timescale,
)))
}
Sub::Active(active) => {
let start_moved = active.start != subscription.start;
active.priority = subscription.priority;
active.max_age = subscription.max_age;
active.start = subscription.start;
if supports_update {
if start_moved {
let _ = producer.start_at(active.start.map(|start| start.group));
}
buffer_update(active, subscription.end)?;
}
Ok(Begin::None)
}
}
}
None => {
if let Sub::Active(active) = sub {
self.subscriber.remove_subscribe(active.id);
let _ = active.stream.writer.finish();
tracing::info!(track = %self.name, "subscribe canceled (idle)");
*sub = Sub::None;
}
Ok(Begin::None)
}
}
}
fn prepare_establish(
&self,
producer: &mut track::Producer,
subscription: Subscription,
timescale: Option<Timescale>,
) -> Establish<S> {
let id = self.subscriber.next_id.fetch_add(1, atomic::Ordering::Relaxed);
let floor = subscription.start.map(|start| start.group);
let _ = match self.subscriber.version.resolves_start() {
true => producer.request_start(floor),
false => producer.start_at(floor),
};
tracing::info!(id, broadcast = %self.subscriber.log_path(&self.path), track = %self.name, "subscribe started");
let tail = kio::Producer::new(Tail::default());
self.subscriber.subscribes.lock().insert(
id,
TrackEntry {
producer: producer.clone(),
timescale,
tail: tail.clone(),
},
);
let session = self.subscriber.session.clone();
Establish {
serve: self.clone(),
closed: session.clone(),
session,
id,
subscription,
tail,
state: EstablishState::Open,
}
}
#[cfg(test)]
async fn establish(
&self,
producer: &mut track::Producer,
sub: &mut Sub<S>,
subscription: Subscription,
timescale: Option<Timescale>,
) -> Result<(), Error> {
let mut est = Box::new(self.prepare_establish(producer, subscription, timescale));
let id = est.id;
match kio::wait(move |waiter| est.poll(waiter)).await {
Ok(active) => {
*sub = Sub::Active(active);
Ok(())
}
Err(err) => {
self.subscriber.remove_subscribe(id);
Err(err)
}
}
}
#[cfg(test)]
async fn handle_subscription(
&self,
producer: &mut track::Producer,
sub: &mut Sub<S>,
pref: Option<Subscription>,
supports_update: bool,
timescale: Option<Timescale>,
) -> Result<(), Error> {
match self.begin_subscription(producer, sub, pref, supports_update, timescale)? {
Begin::Establish(est) => {
let mut est = Box::new(est);
let id = est.id;
match kio::wait(move |waiter| est.poll(waiter)).await {
Ok(active) => *sub = Sub::Active(active),
Err(err) => {
self.subscriber.remove_subscribe(id);
return Err(err);
}
}
}
Begin::None => {
if let Sub::Active(active) = sub {
std::future::poll_fn(|cx| active.stream.writer.poll_flush(cx)).await?;
}
}
}
Ok(())
}
}
#[allow(clippy::large_enum_variant)]
enum Begin<S: crate::transport::poll::Session> {
None,
Establish(Establish<S>),
}
fn buffer_update<S: crate::transport::poll::Session>(
active: &mut SubStream<S>,
end: Option<Position>,
) -> Result<(), Error> {
let bounds = WireBounds::new(active.start, end);
active.stream.writer.buffer(&lite::SubscribeUpdate {
priority: active.priority,
max_age: active.max_age,
start_group: bounds.start_group,
end_group: bounds.end_group,
start_frame: bounds.start_frame,
end_frame: bounds.end_frame,
})
}
struct Establish<S: crate::transport::poll::Session> {
serve: TrackServe<S>,
session: S,
closed: S,
id: u64,
subscription: Subscription,
tail: kio::Producer<Tail>,
state: EstablishState<S>,
}
enum EstablishState<S: crate::transport::poll::Session> {
Open,
Send {
stream: Stream<S, Version>,
},
WaitOk {
stream: Stream<S, Version>,
},
}
impl<S: crate::transport::poll::Session> Establish<S> {
fn poll(&mut self, waiter: &kio::Waiter) -> Poll<Result<SubStream<S>, Error>> {
let mut cx = std::task::Context::from_waker(waiter.waker());
loop {
match &mut self.state {
EstablishState::Open => {
self.serve.subscriber.check_going_away()?;
let mut stream = ready!(Stream::poll_open(
&mut self.session,
self.serve.subscriber.version,
&mut cx
))?;
let bounds = WireBounds::new(self.subscription.start, self.subscription.end);
let msg = lite::Subscribe {
id: self.id,
broadcast: self.serve.path.as_path(),
track: self.serve.name.as_str().into(),
priority: self.subscription.priority,
max_age: self.subscription.max_age,
start_group: bounds.start_group,
end_group: bounds.end_group,
start_frame: bounds.start_frame,
end_frame: bounds.end_frame,
};
stream.writer.buffer(&lite::ControlType::Subscribe)?;
stream.writer.buffer(&msg)?;
self.state = EstablishState::Send { stream };
}
EstablishState::Send { stream } => {
ready!(stream.writer.poll_flush(&mut cx))?;
let EstablishState::Send { stream } = std::mem::replace(&mut self.state, EstablishState::Open)
else {
unreachable!()
};
if !self.serve.subscriber.version.has_track_stream() {
self.state = EstablishState::WaitOk { stream };
continue;
}
return Poll::Ready(Ok(self.activate(stream)));
}
EstablishState::WaitOk { stream } => {
if self.closed.poll_closed(&mut cx).is_ready() {
return Poll::Ready(Err(Error::Dropped));
}
let resp = ready!(stream.reader.poll_decode::<lite::SubscribeResponse>(&mut cx))?;
if !matches!(resp, lite::SubscribeResponse::Ok(_)) {
return Poll::Ready(Err(Error::ProtocolViolation));
}
let EstablishState::WaitOk { stream } = std::mem::replace(&mut self.state, EstablishState::Open)
else {
unreachable!()
};
return Poll::Ready(Ok(self.activate(stream)));
}
}
}
}
fn activate(&self, stream: Stream<S, Version>) -> SubStream<S> {
SubStream {
stream,
id: self.id,
max_age: self.subscription.max_age,
start: self.subscription.start,
priority: self.subscription.priority,
requested: self.subscription.start,
tail: self.tail.clone(),
served: None,
end: None,
}
}
}
struct TrackServeRun<S: crate::transport::poll::Session> {
serve: TrackServe<S>,
state: TrackRunState<S>,
}
#[allow(clippy::large_enum_variant)]
enum TrackRunState<S: crate::transport::poll::Session> {
Info {
request: Option<track::Request>,
info: TrackInfoFetch<S>,
},
Serve(ServeLoop<S>),
Done,
}
impl<S: crate::transport::poll::Session> TrackServeRun<S> {
fn new(serve: TrackServe<S>, request: track::Request) -> Self {
let state = if serve.subscriber.version.has_track_stream() {
TrackRunState::Info {
request: Some(request),
info: TrackInfoFetch::new(&serve),
}
} else {
let info = track::Info::default().with_max_age(serve.subscriber.origin.default_max_age());
TrackRunState::Serve(ServeLoop::new(&serve, request, info, None))
};
Self { serve, state }
}
}
impl<S: crate::transport::poll::Session> kio::Task for TrackServeRun<S> {
fn poll(&mut self, waiter: &kio::Waiter) -> Poll<()> {
loop {
match &mut self.state {
TrackRunState::Info { request, info } => {
let res = ready!(info.poll_fetch(&self.serve, waiter));
let request = request.take().expect("request pending");
match res {
Ok(info) => {
let timescale = Some(info.timescale);
self.state = TrackRunState::Serve(ServeLoop::new(&self.serve, request, info, timescale));
}
Err(err) => {
tracing::warn!(broadcast = %self.serve.subscriber.log_path(&self.serve.path), track = %self.serve.name, %err, "track info failed");
request.reject(err);
self.state = TrackRunState::Done;
return Poll::Ready(());
}
}
}
TrackRunState::Serve(serve_loop) => {
let teardown = ready!(serve_loop.poll(&self.serve, waiter));
let TrackRunState::Serve(mut serve_loop) = std::mem::replace(&mut self.state, TrackRunState::Done)
else {
unreachable!()
};
match teardown {
ServeEnd::Idle => match serve_loop.serving.abort_unused(Error::Cancel) {
Ok(()) => {
tracing::debug!(broadcast = %self.serve.subscriber.log_path(&self.serve.path), track = %self.serve.name, "track released (idle)");
}
Err(used) => {
serve_loop.serving = used;
self.state = TrackRunState::Serve(serve_loop);
continue;
}
},
ServeEnd::Finished => {
let _ = serve_loop.serving.finish();
}
ServeEnd::GiveBack(err) => {
let _ = serve_loop.serving.abort(err);
}
}
if let Sub::Active(active) = &mut serve_loop.sub {
self.serve.subscriber.remove_subscribe(active.id);
let _ = active.stream.writer.finish();
}
return Poll::Ready(());
}
TrackRunState::Done => return Poll::Ready(()),
}
}
}
}
struct TrackInfoFetch<S: crate::transport::poll::Session> {
session: S,
closed: S,
state: TrackInfoState<S>,
}
enum TrackInfoState<S: crate::transport::poll::Session> {
Open,
Send { stream: Stream<S, Version> },
Read { stream: Stream<S, Version> },
}
impl<S: crate::transport::poll::Session> TrackInfoFetch<S> {
fn new(serve: &TrackServe<S>) -> Self {
let session = serve.subscriber.session.clone();
Self {
closed: session.clone(),
session,
state: TrackInfoState::Open,
}
}
fn poll_fetch(&mut self, serve: &TrackServe<S>, waiter: &kio::Waiter) -> Poll<Result<track::Info, Error>> {
let mut cx = std::task::Context::from_waker(waiter.waker());
loop {
match &mut self.state {
TrackInfoState::Open => {
serve.subscriber.check_going_away()?;
let mut stream = ready!(Stream::poll_open(&mut self.session, serve.subscriber.version, &mut cx))?;
stream.writer.buffer(&lite::ControlType::Track)?;
stream.writer.buffer(&lite::Track {
broadcast: serve.path.as_path(),
track: serve.name.as_str().into(),
})?;
self.state = TrackInfoState::Send { stream };
}
TrackInfoState::Send { stream } => {
ready!(stream.writer.poll_flush(&mut cx))?;
let TrackInfoState::Send { stream } = std::mem::replace(&mut self.state, TrackInfoState::Open)
else {
unreachable!()
};
self.state = TrackInfoState::Read { stream };
}
TrackInfoState::Read { stream } => {
if self.closed.poll_closed(&mut cx).is_ready() {
return Poll::Ready(Err(Error::Dropped));
}
let info = ready!(stream.reader.poll_decode::<lite::TrackInfo>(&mut cx))?;
let _ = stream.writer.finish();
let model = track::Info::default()
.with_timescale(info.timescale)
.with_max_age(info.max_age)
.with_priority(info.priority);
return Poll::Ready(Ok(model));
}
}
}
}
}
struct ServeLoop<S: crate::transport::poll::Session> {
serving: track::Producer,
dynamic: track::Dynamic,
sub: Sub<S>,
fetches: kio::Tasks<FetchServeRun<S>>,
closed: S,
supports_update: bool,
supports_fetch: bool,
timescale: Option<Timescale>,
mode: ServeMode<S>,
}
#[allow(clippy::large_enum_variant)]
enum ServeMode<S: crate::transport::poll::Session> {
Select,
Establish(Establish<S>),
Tail {
settle: Settle,
owed: Option<std::ops::Range<u64>>,
},
}
impl<S: crate::transport::poll::Session> ServeLoop<S> {
fn new(serve: &TrackServe<S>, request: track::Request, info: track::Info, timescale: Option<Timescale>) -> Self {
let dynamic = request.dynamic();
let request = match serve.subscriber.version.resolves_start() {
true => request.resolving_start(),
false => request,
};
let serving = request.accept(info);
Self {
serving,
dynamic,
sub: Sub::None,
fetches: kio::Tasks::new(),
closed: serve.subscriber.session.clone(),
supports_update: !matches!(serve.subscriber.version, Version::Lite01 | Version::Lite02),
supports_fetch: serve.subscriber.version.has_track_stream(),
timescale,
mode: ServeMode::Select,
}
}
fn poll(&mut self, serve: &TrackServe<S>, waiter: &kio::Waiter) -> Poll<ServeEnd> {
loop {
match &mut self.mode {
ServeMode::Establish(est) => {
let res = ready!(est.poll(waiter));
let id = est.id;
self.mode = ServeMode::Select;
match res {
Ok(active) => self.sub = Sub::Active(active),
Err(err) => {
serve.subscriber.remove_subscribe(id);
return Poll::Ready(ServeEnd::GiveBack(err));
}
}
}
ServeMode::Tail { settle, owed } => {
let _ = self.fetches.poll(waiter);
if settle
.poll(waiter, |tail| owed.clone().is_some_and(|owed| tail.covers(owed)))
.is_ready()
{
return Poll::Ready(ServeEnd::Finished);
}
if self.fetches.is_empty() && self.serving.poll_unused(waiter).is_ready() {
return Poll::Ready(ServeEnd::Idle);
}
let mut cx = std::task::Context::from_waker(waiter.waker());
if self.closed.poll_closed(&mut cx).is_ready() {
return Poll::Ready(ServeEnd::GiveBack(Error::Dropped));
}
return Poll::Pending;
}
ServeMode::Select => {
let mut cx = std::task::Context::from_waker(waiter.waker());
if let Sub::Active(active) = &mut self.sub {
match active.stream.writer.poll_flush(&mut cx) {
Poll::Ready(Ok(())) => {}
Poll::Ready(Err(err)) => {
serve.subscriber.remove_subscribe(active.id);
self.sub = Sub::None;
return Poll::Ready(ServeEnd::GiveBack(err));
}
Poll::Pending => return Poll::Pending,
}
}
match self.dynamic.poll_requested_group(waiter) {
Poll::Ready(Ok(req)) => {
if self.supports_fetch {
self.fetches
.push(FetchServeRun::new(serve.clone(), req, self.timescale));
} else {
req.reject(Error::Version);
}
continue;
}
Poll::Ready(Err(_)) => return Poll::Ready(ServeEnd::GiveBack(Error::Dropped)),
Poll::Pending => {}
}
match self.serving.poll_subscription_changed(waiter) {
Poll::Ready(Ok(pref)) => {
match serve.begin_subscription(
&mut self.serving,
&mut self.sub,
pref,
self.supports_update,
self.timescale,
) {
Ok(Begin::Establish(est)) => self.mode = ServeMode::Establish(est),
Ok(Begin::None) => {}
Err(err) => return Poll::Ready(ServeEnd::GiveBack(err)),
}
continue;
}
Poll::Ready(Err(_)) => return Poll::Ready(ServeEnd::GiveBack(Error::Dropped)),
Poll::Pending => {}
}
let _ = self.fetches.poll(waiter);
if self.fetches.is_empty() && self.serving.poll_unused(waiter).is_ready() {
return Poll::Ready(ServeEnd::Idle);
}
if let Sub::Active(active) = &mut self.sub
&& let Poll::Ready(res) = active
.stream
.reader
.poll_decode_maybe::<lite::SubscribeResponse>(&mut cx)
{
match res {
Ok(Some(msg)) => {
match &msg {
lite::SubscribeResponse::End(end) => {
if let Err(err) = self.serving.finish_at(end.group) {
tracing::warn!(track = %serve.name, group = end.group, %err, "invalid subscribe end");
}
active.end = Some(end.group);
}
lite::SubscribeResponse::Start(start) => {
if active.start == active.requested {
let _ = self.serving.start_at(start.group);
}
active.served = Some(start.group);
}
lite::SubscribeResponse::Drop(dropped) => {
if let Ok(mut tail) = active.tail.write() {
tail.account(dropped.start..dropped.end.saturating_add(1));
}
}
lite::SubscribeResponse::Ok(_) => {
tracing::debug!(track = %serve.name, ?msg, "subscribe response")
}
}
continue;
}
Ok(None) => {
tracing::info!(broadcast = %serve.subscriber.log_path(&serve.path), track = %serve.name, "subscribe complete");
let subscription = self.serving.subscription();
let requested_end = subscription.as_ref().and_then(|sub| sub.end).map(|end| match end
.frame
{
0 => end.group,
_ => end.group.saturating_add(1),
});
let grace = subscription
.map(|sub| sub.max_age)
.filter(|max_age| !max_age.is_zero())
.unwrap_or(tail::GRACE);
self.mode = ServeMode::Tail {
settle: Settle::new(&serve.subscriber.runtime, active.tail.consume(), grace),
owed: active.owed(requested_end),
};
continue;
}
Err(err) => {
tracing::warn!(broadcast = %serve.subscriber.log_path(&serve.path), track = %serve.name, %err, "subscribe error");
return Poll::Ready(ServeEnd::GiveBack(err));
}
}
}
if self.closed.poll_closed(&mut cx).is_ready() {
return Poll::Ready(ServeEnd::GiveBack(Error::Dropped));
}
return Poll::Pending;
}
}
}
}
}
struct FetchServeRun<S: crate::transport::poll::Session> {
serve: TrackServe<S>,
session: S,
timescale: Option<Timescale>,
group: u64,
state: FetchRunState<S>,
}
enum FetchRunState<S: crate::transport::poll::Session> {
Open {
request: Option<group::Request>,
},
Send {
request: Option<group::Request>,
stream: Stream<S, Version>,
frame_start: u64,
},
Ingest {
stream: Stream<S, Version>,
producer: group::Producer,
ingest: FrameIngest,
},
Done,
}
impl<S: crate::transport::poll::Session> FetchServeRun<S> {
fn new(serve: TrackServe<S>, request: group::Request, timescale: Option<Timescale>) -> Self {
let session = serve.subscriber.session.clone();
let group = request.sequence();
Self {
serve,
session,
timescale,
group,
state: FetchRunState::Open { request: Some(request) },
}
}
}
impl<S: crate::transport::poll::Session> kio::Task for FetchServeRun<S> {
fn poll(&mut self, waiter: &kio::Waiter) -> Poll<()> {
let mut cx = std::task::Context::from_waker(waiter.waker());
loop {
match &mut self.state {
FetchRunState::Open { request } => {
tracing::info!(broadcast = %self.serve.subscriber.log_path(&self.serve.path), track = %self.serve.name, group = self.group, "fetch started");
if self.serve.subscriber.going_away.is_set() {
request.take().expect("request pending").reject(Error::GoingAway);
self.state = FetchRunState::Done;
return Poll::Ready(());
}
let mut stream = match ready!(Stream::poll_open(
&mut self.session,
self.serve.subscriber.version,
&mut cx
)) {
Ok(stream) => stream,
Err(err) => {
tracing::warn!(track = %self.serve.name, %err, "fetch stream open failed");
request.take().expect("request pending").reject(err);
self.state = FetchRunState::Done;
return Poll::Ready(());
}
};
let request = request.take().expect("request pending");
let frame_start = match self.serve.subscriber.version.has_frame_bounds() {
true => request.frame_start(),
false => 0,
};
let msg = lite::Fetch {
broadcast: self.serve.path.as_path(),
track: self.serve.name.as_str().into(),
priority: request.priority(),
group: self.group,
start_frame: frame_start,
end_frame: None,
};
let buffered = stream
.writer
.buffer(&lite::ControlType::Fetch)
.and_then(|()| stream.writer.buffer(&msg));
if let Err(err) = buffered {
stream.writer.abort(&err);
request.reject(err);
self.state = FetchRunState::Done;
return Poll::Ready(());
}
self.state = FetchRunState::Send {
request: Some(request),
stream,
frame_start,
};
}
FetchRunState::Send { stream, .. } => {
if let Err(err) = ready!(stream.writer.poll_flush(&mut cx)) {
let FetchRunState::Send { request, stream, .. } =
std::mem::replace(&mut self.state, FetchRunState::Done)
else {
unreachable!()
};
stream.writer.abort(&err);
request.expect("request pending").reject(err);
return Poll::Ready(());
}
let FetchRunState::Send {
request,
stream,
frame_start,
} = std::mem::replace(&mut self.state, FetchRunState::Done)
else {
unreachable!()
};
let request = request.expect("request pending");
let group_info = track::Info::default()
.with_timescale(self.timescale.unwrap_or_default())
.with_max_age(self.serve.subscriber.origin.default_max_age());
let mut producer = match request.accept(group_info) {
Ok(producer) => producer,
Err(err) => {
tracing::debug!(track = %self.serve.name, group = self.group, %err, "fetch not served");
stream.writer.abort(&err);
return Poll::Ready(());
}
};
if let Err(err) = producer.start_at(frame_start) {
stream.writer.abort(&err);
let _ = producer.abort(err);
return Poll::Ready(());
}
self.state = FetchRunState::Ingest {
stream,
producer,
ingest: FrameIngest::new(self.serve.subscriber.runtime.clone(), self.timescale),
};
}
FetchRunState::Ingest {
stream,
producer,
ingest,
} => {
let res = ready!(ingest.poll(&mut stream.reader, producer, waiter));
let FetchRunState::Ingest { producer, .. } =
std::mem::replace(&mut self.state, FetchRunState::Done)
else {
unreachable!()
};
match res {
Ok(()) => {
let producer = producer;
let _ = producer.finish();
}
Err(err) => {
let _ = producer.abort(err);
}
}
return Poll::Ready(());
}
FetchRunState::Done => return Poll::Ready(()),
}
}
}
}