use std::{
collections::{HashMap, hash_map::Entry},
task::Poll,
time::Duration,
};
use crate::{
Error, Path, PathOwned, Timescale, broadcast,
coding::{Reader, Stream},
frame, group,
ietf::{self, Control, FilterType, GroupOrder, RequestId},
origin, track,
util::{MaybeBoxedExt, MaybeSendBox, TaskSet, Tasks},
};
use super::{Message, Version, cluster};
use web_async::Lock;
const TRACK_ALIAS_TIMEOUT: Duration = Duration::from_secs(1);
type TrackAliases = kio::Producer<HashMap<u64, RequestId>>;
fn insert_track_alias(aliases: &TrackAliases, alias: u64, request_id: RequestId) -> Result<(), Error> {
let mut aliases = aliases.write().map_err(|_| Error::Dropped)?;
match aliases.entry(alias) {
Entry::Occupied(entry) if *entry.get() == request_id => Ok(()),
Entry::Occupied(_) => Err(Error::Duplicate),
Entry::Vacant(entry) => {
entry.insert(request_id);
Ok(())
}
}
}
fn is_protocol_violation(err: &Error) -> bool {
matches!(
err,
Error::Decode(_)
| Error::Encode(_)
| Error::BoundsExceeded(_)
| Error::WrongSize
| Error::TooManyParameters
| Error::ProtocolViolation
| Error::UnexpectedMessage
| Error::UnexpectedStream
)
}
fn remove_track_alias(aliases: &TrackAliases, alias: u64, request_id: RequestId) {
let Ok(mut aliases) = aliases.write() else {
return;
};
if aliases.get(&alias) == Some(&request_id) {
aliases.remove(&alias);
}
}
#[derive(Default)]
struct State {
subscribes: HashMap<RequestId, TrackState>,
aliases: TrackAliases,
broadcasts: HashMap<PathOwned, BroadcastState>,
}
struct TrackState {
producer: track::Producer,
alias: Option<u64>,
timescale: Option<Timescale>,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
enum Detach {
Graceful,
Abrupt,
}
struct BroadcastState {
producer: crate::model::broadcast::SourceGuard,
count: usize,
publisher: Option<crate::Origin>,
}
struct Advertised {
route: broadcast::Route,
publisher: Option<crate::Origin>,
}
#[derive(Clone)]
pub(super) struct Subscriber<S: web_transport_trait::Session> {
session: S,
origin: origin::Producer,
control: Control,
session_origin: crate::Origin,
self_origin: crate::Origin,
peer_setup: cluster::PeerSetup,
cost: Option<u64>,
state: Lock<State>,
tasks: Tasks,
version: Version,
}
async fn resolve_track_alias(aliases: kio::Consumer<HashMap<u64, RequestId>>, alias: u64) -> Result<RequestId, Error> {
let mut timeout = kio::time::Deadline::after(TRACK_ALIAS_TIMEOUT);
kio::wait(|waiter| {
let resolved = aliases.poll(waiter, |aliases| match aliases.get(&alias) {
Some(request_id) => Poll::Ready(*request_id),
None => Poll::Pending,
});
if let Poll::Ready(result) = resolved {
return Poll::Ready(result.map_err(|_| Error::Dropped));
}
if timeout.poll(waiter).is_ready() {
return Poll::Ready(Err(Error::NotFound));
}
Poll::Pending
})
.await
}
impl<S: web_transport_trait::Session> Subscriber<S> {
#[allow(clippy::too_many_arguments)]
pub fn new(
session: S,
origin: origin::Producer,
control: Control,
peer_origin: Option<crate::Origin>,
peer_setup: cluster::PeerSetup,
self_origin: crate::Origin,
cost: Option<u64>,
version: Version,
tasks: Tasks,
) -> Self {
Self {
session,
origin,
control,
session_origin: peer_origin.unwrap_or(crate::Origin::UNKNOWN),
self_origin,
peer_setup,
cost,
state: Default::default(),
tasks,
version,
}
}
pub(super) async fn peer(&self) -> cluster::Peer {
match cluster::supported(self.version) {
true => self.peer_setup.get().await,
false => cluster::Peer::default(),
}
}
fn session_route(&self, peer: &cluster::Peer) -> broadcast::Route {
let mut hops = crate::OriginList::new();
hops.push(self.session_origin)
.expect("an empty hop chain has room for one entry");
broadcast::Route::new()
.with_hops(hops)
.with_cost(cluster::link_cost(self.cost, peer))
.with_announce(true)
}
fn route(&self, advert: Option<&cluster::Advert>, peer: &cluster::Peer) -> Option<Advertised> {
let Some(advert) = advert else {
return Some(Advertised {
route: self.session_route(peer),
publisher: None,
});
};
if advert.loops(self.self_origin) {
return None;
}
Some(Advertised {
route: advert.route(cluster::link_cost(self.cost, peer)),
publisher: Some(
advert
.hops
.hops()
.iter()
.next()
.copied()
.unwrap_or(crate::Origin::UNKNOWN),
),
})
}
fn register_alias(&self, request_id: RequestId, alias: u64) -> Result<(), Error> {
let mut state = self.state.lock();
if !state.subscribes.contains_key(&request_id) {
return Err(Error::NotFound);
}
insert_track_alias(&state.aliases, alias, request_id)?;
state.subscribes.get_mut(&request_id).unwrap().alias = Some(alias);
Ok(())
}
fn remove_subscribe(&self, request_id: RequestId) -> Option<TrackState> {
let mut state = self.state.lock();
let track = state.subscribes.remove(&request_id)?;
if let Some(alias) = track.alias {
remove_track_alias(&state.aliases, alias, request_id);
}
Some(track)
}
pub fn subscribe_prefixes(&self) -> Vec<PathOwned> {
self.origin.allowed().map(|p| p.to_owned()).collect()
}
pub async fn run_subscribe_namespace<T: web_transport_trait::Session>(
&mut self,
mut stream: Stream<T, Version>,
prefix: PathOwned,
) -> Result<(), Error> {
let request_id = self.control.next_request_id().await?;
match self.version {
Version::Draft14 | Version::Draft15 | Version::Draft16 | Version::Draft17 => {
let msg = ietf::SubscribeNamespaceLegacy {
request_id,
namespace: prefix.clone(),
subscribe_options: 0x01, };
stream.writer.encode(&ietf::SubscribeNamespaceLegacy::ID).await?;
stream.writer.encode(&msg).await?;
}
_ => {
let msg = ietf::SubscribeNamespace {
request_id,
namespace: prefix.clone(),
};
stream.writer.encode(&ietf::SubscribeNamespace::ID).await?;
stream.writer.encode(&msg).await?;
}
}
tracing::debug!(%prefix, "subscribe_namespace sent");
let type_id: u64 = stream.reader.decode().await?;
let size: u16 = stream.reader.decode().await?;
let mut data = stream.reader.read_exact(size as usize).await?;
match type_id {
ietf::SubscribeNamespaceOk::ID if self.version == Version::Draft14 => {
let _msg = ietf::SubscribeNamespaceOk::decode_msg(&mut data, self.version)?;
}
ietf::RequestOk::ID => {
let _msg = ietf::RequestOk::decode_msg(&mut data, self.version)?;
}
ietf::SubscribeNamespaceError::ID if self.version == Version::Draft14 => {
let msg = ietf::SubscribeNamespaceError::decode_msg(&mut data, self.version)?;
tracing::warn!(error_code = %msg.error_code, reason = %msg.reason_phrase, "subscribe_namespace error");
return Err(Error::Cancel);
}
ietf::RequestError::ID => {
let msg = ietf::RequestError::decode_msg(&mut data, self.version)?;
tracing::warn!(error_code = %msg.error_code, reason = %msg.reason_phrase, "subscribe_namespace error");
return Err(Error::Cancel);
}
_ => return Err(Error::UnexpectedMessage),
}
tracing::debug!(%prefix, "subscribe_namespace ok");
let peer = self.peer().await;
let mut live: std::collections::HashSet<PathOwned> = std::collections::HashSet::new();
let res = self.run_namespace_entries(&mut stream, &prefix, &peer, &mut live).await;
for path in live {
let _ = self.stop_announce(path, Detach::Abrupt);
}
res
}
async fn run_namespace_entries<T: web_transport_trait::Session>(
&mut self,
stream: &mut Stream<T, Version>,
prefix: &PathOwned,
peer: &cluster::Peer,
live: &mut std::collections::HashSet<PathOwned>,
) -> Result<(), Error> {
loop {
let type_id: u64 = match stream.reader.decode_maybe().await? {
Some(id) => id,
None => break, };
let size: u16 = stream.reader.decode().await?;
let mut data = stream.reader.read_exact(size as usize).await?;
match type_id {
ietf::Namespace::ID => {
let msg = ietf::Namespace::decode_body(&mut data, self.version, peer.negotiated())?;
if !data.is_empty() {
return Err(Error::WrongSize);
}
let path = prefix.join(&msg.suffix);
let Some(advert) = self.route(msg.cluster.as_ref(), peer) else {
tracing::debug!(%path, "dropping reflected namespace");
if live.remove(&path) {
let _ = self.stop_announce(path, Detach::Graceful);
}
continue;
};
tracing::debug!(%path, hops = advert.route.hops.len(), cost = advert.route.cost, "namespace");
if live.contains(&path) {
self.update_announce(path, advert)?;
} else {
self.start_announce(path.clone(), advert)?;
live.insert(path);
}
}
ietf::NamespaceDone::ID => {
let msg = ietf::NamespaceDone::decode_msg(&mut data, self.version)?;
let path = prefix.join(&msg.suffix);
tracing::debug!(%path, "namespace_done");
if live.remove(&path) {
let _ = self.stop_announce(path, Detach::Graceful);
}
}
_ => {
tracing::warn!(type_id, "unexpected message on subscribe_namespace stream");
return Err(Error::UnexpectedMessage);
}
}
}
Ok(())
}
pub fn handle_stream(
&mut self,
id: u64,
mut data: bytes::Bytes,
stream: Stream<S, Version>,
peer: cluster::Peer,
) -> Result<MaybeSendBox<'static, ()>, Error> {
let mut this = self.clone();
let task = match id {
ietf::Publish::ID => {
let msg = ietf::Publish::decode_msg(&mut data, this.version)?;
if !data.is_empty() {
return Err(Error::WrongSize);
}
tracing::debug!(message = ?msg, "received publish");
async move {
if let Err(err) = this.run_publish_stream(stream, msg).await {
tracing::debug!(%err, "publish stream error");
}
}
.maybe_boxed()
}
ietf::PublishNamespace::ID => {
let msg = ietf::PublishNamespace::decode_body(&mut data, this.version, peer.negotiated())?;
if !data.is_empty() {
return Err(Error::WrongSize);
}
tracing::debug!(message = ?msg, "received publish_namespace");
async move {
if let Err(err) = this.run_publish_namespace_stream(stream, msg, peer).await {
if is_protocol_violation(&err) {
this.session.close(err.to_code(), err.to_string().as_ref());
}
tracing::debug!(%err, "publish_namespace stream error");
}
}
.maybe_boxed()
}
_ => {
tracing::warn!(id, "unexpected bidi stream type for subscriber");
return Err(Error::UnexpectedStream);
}
};
Ok(task)
}
async fn run_publish_namespace_stream(
&mut self,
mut stream: Stream<S, Version>,
msg: ietf::PublishNamespace<'_>,
peer: cluster::Peer,
) -> Result<(), Error> {
let request_id = msg.request_id;
let path = msg.track_namespace.to_owned();
let Some(advert) = self.route(msg.cluster.as_ref(), &peer) else {
tracing::debug!(%path, "dropping reflected publish_namespace");
self.write_error(&mut stream, request_id, 400, "route loops through this relay")
.await?;
let _ = stream.writer.finish();
let _ = stream.writer.closed().await;
return Ok(());
};
match self.start_announce(path.clone(), advert) {
Ok(_) => {
if let Err(err) = self.write_ok(&mut stream, request_id).await {
let _ = self.stop_announce(path, Detach::Graceful);
return Err(err);
}
}
Err(err) => {
self.write_error(&mut stream, request_id, 400, &err.to_string()).await?;
let _ = stream.writer.finish();
let _ = stream.writer.closed().await;
return Ok(());
}
}
let mut attached = true;
let res = self
.run_publish_namespace_updates(&mut stream, &path, request_id, peer, &mut attached)
.await;
if attached {
let detach = match res.is_ok() {
true => Detach::Graceful,
false => Detach::Abrupt,
};
self.stop_announce(path, detach)?;
}
res
}
fn terminal_publish_namespace(&self, type_id: u64) -> bool {
match self.version {
Version::Draft14 | Version::Draft15 | Version::Draft16 => type_id == ietf::PublishNamespaceDone::ID,
_ => false,
}
}
async fn run_publish_namespace_updates(
&mut self,
stream: &mut Stream<S, Version>,
path: &PathOwned,
request_id: RequestId,
peer: cluster::Peer,
attached: &mut bool,
) -> Result<(), Error> {
loop {
let type_id: u64 = match stream.reader.decode_maybe().await? {
Some(id) => id,
None => return Ok(()),
};
let terminal = self.terminal_publish_namespace(type_id);
if type_id != ietf::PublishNamespace::ID && !terminal {
tracing::warn!(type_id, "unexpected message on publish_namespace stream");
return Err(Error::UnexpectedMessage);
}
let size: u16 = stream.reader.decode().await?;
let mut data = stream.reader.read_exact(size as usize).await?;
if terminal {
ietf::PublishNamespaceDone::decode_msg(&mut data, self.version)?;
if !data.is_empty() {
return Err(Error::WrongSize);
}
tracing::debug!(%path, "publish_namespace_done");
return Ok(());
}
let msg = ietf::PublishNamespace::decode_body(&mut data, self.version, peer.negotiated())?;
if !data.is_empty() {
return Err(Error::WrongSize);
}
if msg.request_id != request_id || msg.track_namespace.as_str() != path.as_str() {
tracing::warn!(%path, "publish_namespace update does not match its stream");
return Err(Error::ProtocolViolation);
}
let Some(advert) = self.route(msg.cluster.as_ref(), &peer) else {
if std::mem::take(attached) {
tracing::debug!(%path, "publish_namespace now loops back; detaching");
let _ = self.stop_announce(path.clone(), Detach::Graceful);
}
continue;
};
tracing::debug!(%path, hops = advert.route.hops.len(), cost = advert.route.cost, "publish_namespace update");
match *attached {
true => self.update_announce(path.clone(), advert)?,
false => {
self.start_announce(path.clone(), advert)?;
*attached = true;
}
}
}
}
async fn run_publish_stream(
&mut self,
mut stream: Stream<S, Version>,
msg: ietf::Publish<'_>,
) -> Result<(), Error> {
tracing::debug!(broadcast = %msg.track_namespace, track = %msg.track_name, "rejecting publish");
self.write_publish_error(&mut stream, msg.request_id, 400, "PUBLISH is not supported")
.await?;
let _ = stream.writer.finish();
Ok(())
}
async fn write_ok(&self, stream: &mut Stream<S, Version>, request_id: RequestId) -> Result<(), Error> {
match self.version {
Version::Draft14 => {
stream.writer.encode(&ietf::PublishNamespaceOk::ID).await?;
stream.writer.encode(&ietf::PublishNamespaceOk { request_id }).await?;
}
Version::Draft15 | Version::Draft16 => {
stream.writer.encode(&ietf::RequestOk::ID).await?;
stream
.writer
.encode(&ietf::RequestOk {
request_id: Some(request_id),
})
.await?;
}
_ => {
stream.writer.encode(&ietf::RequestOk::ID).await?;
stream.writer.encode(&ietf::RequestOk { request_id: None }).await?;
}
}
Ok(())
}
async fn write_error(
&self,
stream: &mut Stream<S, Version>,
request_id: RequestId,
error_code: u64,
reason: &str,
) -> Result<(), Error> {
match self.version {
Version::Draft14 => {
stream.writer.encode(&ietf::PublishNamespaceError::ID).await?;
stream
.writer
.encode(&ietf::PublishNamespaceError {
request_id,
error_code,
reason_phrase: reason.into(),
})
.await?;
}
Version::Draft15 | Version::Draft16 => {
stream.writer.encode(&ietf::RequestError::ID).await?;
stream
.writer
.encode(&ietf::RequestError {
request_id: Some(request_id),
error_code,
reason_phrase: reason.into(),
retry_interval: 0,
})
.await?;
}
_ => {
stream.writer.encode(&ietf::RequestError::ID).await?;
stream
.writer
.encode(&ietf::RequestError {
request_id: None,
error_code,
reason_phrase: reason.into(),
retry_interval: 0,
})
.await?;
}
}
Ok(())
}
async fn write_publish_error(
&self,
stream: &mut Stream<S, Version>,
request_id: RequestId,
error_code: u64,
reason: &str,
) -> Result<(), Error> {
match self.version {
Version::Draft14 => {
stream.writer.encode(&ietf::PublishError::ID).await?;
stream
.writer
.encode(&ietf::PublishError {
request_id,
error_code,
reason_phrase: reason.into(),
})
.await?;
}
Version::Draft15 | Version::Draft16 => {
stream.writer.encode(&ietf::RequestError::ID).await?;
stream
.writer
.encode(&ietf::RequestError {
request_id: Some(request_id),
error_code,
reason_phrase: reason.into(),
retry_interval: 0,
})
.await?;
}
_ => {
stream.writer.encode(&ietf::RequestError::ID).await?;
stream
.writer
.encode(&ietf::RequestError {
request_id: None,
error_code,
reason_phrase: reason.into(),
retry_interval: 0,
})
.await?;
}
}
Ok(())
}
fn start_announce(&mut self, path: PathOwned, advert: Advertised) -> Result<broadcast::Producer, Error> {
let mut state = self.state.lock();
let existing = state.broadcasts.contains_key(&path);
let producer = self.attach(&mut state, path.clone(), advert)?;
if existing && let Some(entry) = state.broadcasts.get_mut(&path) {
entry.count += 1;
}
Ok(producer)
}
fn update_announce(&mut self, path: PathOwned, advert: Advertised) -> Result<(), Error> {
let mut state = self.state.lock();
if !state.broadcasts.contains_key(&path) {
return Err(Error::NotFound);
}
self.attach(&mut state, path, advert)?;
Ok(())
}
fn attach(&self, state: &mut State, path: PathOwned, advert: Advertised) -> Result<broadcast::Producer, Error> {
let publisher = advert.publisher;
let mut carried = None;
if let Entry::Occupied(entry) = state.broadcasts.entry(path.clone())
&& let (Some(old), Some(new)) = (entry.get().publisher, publisher)
&& (old != new || new == crate::Origin::UNKNOWN)
{
tracing::debug!(broadcast = %self.origin.absolute(&path), "publisher changed; replacing the source");
carried = Some(entry.get().count);
entry.remove().producer.finish();
}
match state.broadcasts.entry(path.clone()) {
Entry::Occupied(entry) => {
let mut producer = entry.get().producer.producer();
producer.set_route(advert.route)?;
Ok(producer)
}
Entry::Vacant(entry) => {
let route = advert.route;
let broadcast = self.origin.create_broadcast(&path, route)?;
let dynamic = broadcast.dynamic();
entry.insert(BroadcastState {
producer: crate::model::broadcast::SourceGuard::new(broadcast.clone()),
count: carried.unwrap_or(1),
publisher,
});
tracing::debug!(broadcast = %self.origin.absolute(&path), "announce");
let this = self.clone();
self.tasks.push(async move {
if let Err(err) = this.run_broadcast(path, dynamic).await {
tracing::debug!(%err, "error running broadcast");
}
});
Ok(broadcast)
}
}
}
fn stop_announce(&mut self, path: PathOwned, detach: Detach) -> Result<(), Error> {
let mut state = self.state.lock();
match state.broadcasts.entry(path.clone()) {
Entry::Occupied(mut entry) => {
entry.get_mut().count -= 1;
if entry.get().count == 0 {
tracing::debug!(broadcast = %self.origin.absolute(&path), ?detach, "unannounced");
let producer = entry.remove().producer;
match detach {
Detach::Graceful => producer.finish(),
Detach::Abrupt => drop(producer),
}
}
}
Entry::Vacant(_) => return Err(Error::NotFound),
};
Ok(())
}
async fn run_broadcast(&self, path: Path<'_>, mut broadcast: broadcast::Dynamic) -> Result<(), Error> {
let mut subscribes = TaskSet::owned();
loop {
let next = subscribes
.drive(async {
let mut closed = std::pin::pin!(self.session.closed());
kio::wait(|waiter| {
if waiter.poll_future(closed.as_mut()).is_ready() {
return Poll::Ready(None);
}
broadcast.poll_requested_track(waiter).map(Some)
})
.await
})
.await;
let request = match next {
Some(Ok(request)) => request,
Some(Err(err)) => {
tracing::debug!(%err, "broadcast closed");
break;
}
None => break,
};
let mut this = self.clone();
let path = path.to_owned();
let broadcast = broadcast.clone();
subscribes.push(async move {
this.run_subscribe(path, broadcast, request).await;
});
}
Ok(())
}
async fn run_subscribe(
&mut self,
broadcast_path: Path<'_>,
broadcast: broadcast::Dynamic,
request: track::Request,
) {
let info = track::Info::default().with_timescale(crate::Timescale::MICRO);
let mut track = request.accept(info);
let request_id = match self.control.next_request_id().await {
Ok(id) => id,
Err(err) => {
let _ = track.abort(err);
return;
}
};
let mut stream = match Stream::open(&self.session, self.version).await {
Ok(s) => s,
Err(err) => {
tracing::debug!(%err, "failed to open subscribe stream");
let _ = track.abort(err);
return;
}
};
{
let mut state = self.state.lock();
state.subscribes.insert(
request_id,
TrackState {
producer: track.clone(),
alias: None,
timescale: None,
},
);
}
if let Err(err) = self
.write_subscribe(&mut stream, request_id, &broadcast_path, &track)
.await
{
tracing::debug!(%err, "failed to write subscribe");
self.remove_subscribe(request_id);
let _ = track.abort(err);
return;
}
tracing::info!(broadcast = %self.origin.absolute(&broadcast_path), track = %track.name(), "subscribe started");
match self.read_subscribe_response(&mut stream).await {
Ok(Some((alias, timescale))) => {
if let Some(timescale) = timescale {
let mut state = self.state.lock();
if let Some(track) = state.subscribes.get_mut(&request_id) {
track.timescale = Some(timescale);
}
}
if let Err(err) = self.register_alias(request_id, alias) {
self.session.close(err.to_code(), err.to_string().as_ref());
self.remove_subscribe(request_id);
let _ = track.abort(err);
return;
}
}
Ok(None) => {}
Err(err) => {
tracing::debug!(%err, "subscribe response error");
self.remove_subscribe(request_id);
let _ = track.abort(err);
return;
}
};
enum End {
Unused,
BroadcastClosed(Error),
StreamClosed(Result<(), Error>),
}
let end = {
let mut closed = std::pin::pin!(stream.reader.closed());
kio::wait(|waiter| {
if track.poll_unused(waiter).is_ready() {
return Poll::Ready(End::Unused);
}
if let Poll::Ready(err) = broadcast.poll_closed(waiter) {
return Poll::Ready(End::BroadcastClosed(err));
}
waiter.poll_future(closed.as_mut()).map(End::StreamClosed)
})
.await
};
match end {
End::Unused => {
tracing::info!(broadcast = %self.origin.absolute(&broadcast_path), track = %track.name(), "subscribe cancelled");
let _ = track.abort(Error::Cancel);
}
End::BroadcastClosed(err) => {
tracing::info!(broadcast = %self.origin.absolute(&broadcast_path), track = %track.name(), "broadcast closed");
let _ = track.abort(err);
}
End::StreamClosed(res) => match res {
Ok(()) => {
tracing::info!(broadcast = %self.origin.absolute(&broadcast_path), track = %track.name(), "subscribe complete");
let _ = track.finish();
}
Err(err) => {
tracing::debug!(%err, "subscribe stream closed with error");
let _ = track.abort(err);
}
},
}
self.remove_subscribe(request_id);
stream.writer.finish().ok();
}
async fn write_subscribe(
&self,
stream: &mut Stream<S, Version>,
request_id: RequestId,
broadcast: &Path<'_>,
track: &track::Producer,
) -> Result<(), Error> {
stream.writer.encode(&ietf::Subscribe::ID).await?;
stream
.writer
.encode(&ietf::Subscribe {
request_id,
track_namespace: broadcast.to_owned(),
track_name: track.name().into(),
subscriber_priority: track.subscription().map(|s| s.priority).unwrap_or(0),
group_order: GroupOrder::Descending,
filter_type: FilterType::LargestObject,
})
.await?;
Ok(())
}
async fn read_subscribe_response(
&self,
stream: &mut Stream<S, Version>,
) -> Result<Option<(u64, Option<Timescale>)>, Error> {
let type_id: u64 = stream.reader.decode().await?;
let size: u16 = stream.reader.decode().await?;
let mut data = stream.reader.read_exact(size as usize).await?;
match type_id {
ietf::SubscribeOk::ID => {
let msg = ietf::SubscribeOk::decode_msg(&mut data, self.version)?;
tracing::debug!(message = ?msg, "received subscribe ok");
Ok(Some((msg.track_alias, msg.timescale)))
}
ietf::SubscribeError::ID if self.version == Version::Draft14 => {
let msg = ietf::SubscribeError::decode_msg(&mut data, self.version)?;
tracing::warn!(message = ?msg, "subscribe error");
Err(Error::Cancel)
}
ietf::RequestError::ID => {
let msg = ietf::RequestError::decode_msg(&mut data, self.version)?;
tracing::warn!(message = ?msg, "request error");
Err(Error::Cancel)
}
_ => Err(Error::UnexpectedMessage),
}
}
pub async fn recv_group(&mut self, stream: &mut Reader<S::RecvStream, Version>) -> Result<(), Error> {
let group: ietf::GroupHeader = stream.decode().await?;
if group.sub_group_id != 0 {
tracing::warn!(sub_group_id = %group.sub_group_id, "subgroup ID is not supported, dropping stream");
return Err(Error::Unsupported);
}
let aliases = self.state.lock().aliases.consume();
let request_id = resolve_track_alias(aliases, group.track_alias).await.inspect_err(|_| {
tracing::warn!(track_alias = %group.track_alias, "unknown track alias");
})?;
let (mut producer, track, timescale) = {
let mut state = self.state.lock();
let track = state.subscribes.get_mut(&request_id).ok_or(Error::NotFound)?;
let group_info = group::Info {
sequence: group.group_id,
};
let producer = track.producer.create_group(group_info)?;
(producer, track.producer.clone(), track.timescale)
};
let res = {
let mut serve = std::pin::pin!(self.run_group(group, stream, producer.clone(), timescale));
kio::wait(|waiter| {
if let Poll::Ready(err) = track.poll_closed(waiter) {
return Poll::Ready(Err(err));
}
if let Poll::Ready(err) = producer.poll_closed(waiter) {
return Poll::Ready(Err(err));
}
waiter.poll_future(serve.as_mut())
})
.await
};
match res {
Err(Error::Cancel) => {
let _ = producer.abort(Error::Cancel);
}
Err(err) => {
tracing::debug!(%err, group = %producer.sequence, "group error");
let _ = producer.abort(err);
}
_ => {
let _ = producer.finish();
}
}
Ok(())
}
async fn run_group(
&mut self,
group: ietf::GroupHeader,
stream: &mut Reader<S::RecvStream, Version>,
mut producer: group::Producer,
timescale: Option<Timescale>,
) -> Result<(), Error> {
while let Some(id_delta) = stream.decode_maybe::<u64>().await? {
if id_delta != 0 {
tracing::warn!(id_delta = %id_delta, "object ID delta is not supported, dropping stream");
return Err(Error::Unsupported);
}
let timestamp = match (group.flags.has_extensions, timescale) {
(true, Some(timescale)) => {
let size: usize = stream.decode().await?;
let mut ext = stream.read_exact(size).await?;
ietf::decode_object_time(&mut ext, timescale, self.version)?
}
(true, None) => {
let size: usize = stream.decode().await?;
stream.read_exact(size).await?;
None
}
(false, _) => None,
};
let size: u64 = stream.decode().await?;
if size == 0 {
let status: u64 = stream.decode().await?;
if status == 0 {
let timestamp = timestamp.unwrap_or_else(crate::Timestamp::now);
let frame = producer.create_frame(frame::Info { size: 0, timestamp })?;
frame.finish()?;
} else if status == 3 && !group.flags.has_end {
break;
} else {
return Err(Error::Unsupported);
}
} else {
let timestamp = timestamp.unwrap_or_else(crate::Timestamp::now);
let mut frame = producer.create_frame(frame::Info { size, timestamp })?;
if let Err(err) = self.run_frame(stream, &mut frame).await {
let _ = frame.abort(err.clone());
return Err(err);
}
frame.finish()?;
}
}
Ok(())
}
async fn run_frame(
&mut self,
stream: &mut Reader<S::RecvStream, Version>,
frame: &mut frame::Producer<'_>,
) -> Result<(), Error> {
while frame.remaining() > 0 {
match stream.read_chunk(frame.remaining()).await? {
Some(chunk) if !chunk.is_empty() => {
frame.write(chunk)?;
}
_ => return Err(Error::WrongSize),
}
}
Ok(())
}
}
#[cfg(test)]
mod tests {
use futures::poll;
use super::*;
#[tokio::test(start_paused = true)]
async fn track_alias_waits_for_control_message() {
let aliases = TrackAliases::default();
let pending = resolve_track_alias(aliases.consume(), 7);
tokio::pin!(pending);
assert!(poll!(&mut pending).is_pending());
insert_track_alias(&aliases, 7, RequestId(11)).unwrap();
assert_eq!(pending.await.unwrap(), RequestId(11));
}
#[tokio::test(start_paused = true)]
async fn unknown_track_alias_times_out() {
let aliases = TrackAliases::default();
assert!(matches!(
resolve_track_alias(aliases.consume(), 7).await,
Err(Error::NotFound)
));
}
async fn settle() {
tokio::time::sleep(std::time::Duration::from_millis(1)).await;
}
fn occurrences(log: &crate::lite::test_transport::Log, needle: &[u8]) -> usize {
let writes = log.writes.lock().unwrap();
writes.windows(needle.len()).filter(|window| *window == needle).count()
}
#[tokio::test]
async fn a_rooted_subscriber_asks_for_its_scope_not_its_root() {
let origin = crate::origin::Info::new(crate::Origin::new(1).unwrap()).produce();
let scoped = origin
.with_root("rootns")
.and_then(|rooted| rooted.scope(&[crate::Path::new("cam")]))
.expect("scope the origin");
let gate = kio::Producer::new(true);
let session = crate::lite::test_transport::SinkSession::gated_bi(gate.consume());
let log = session.log.clone();
let (tasks, _task_set) = crate::util::TaskSet::new();
let mut subscriber = Subscriber::new(
session.clone(),
scoped,
Control::new(None, false),
None,
cluster::PeerSetup::default(),
crate::Origin::new(1).unwrap(),
None,
Version::Draft16,
tasks,
);
assert_eq!(
subscriber.subscribe_prefixes(),
vec![crate::Path::new("cam").to_owned()],
"one SUBSCRIBE_NAMESPACE per permitted prefix, relative to the root",
);
let stream = Stream::open(&session, Version::Draft16).await.unwrap();
let mut run = std::pin::pin!(subscriber.run_subscribe_namespace(stream, crate::Path::new("cam").to_owned()));
assert!(futures::poll!(run.as_mut()).is_pending());
assert_eq!(occurrences(&log, b"cam"), 1, "asked the peer for our scope");
assert_eq!(occurrences(&log, b"rootns"), 0, "asked the peer for our local root");
}
async fn namespace_response(version: Version, suffix: &str) -> Vec<u8> {
let log = crate::lite::test_transport::Log::default();
let mut writer = crate::coding::Writer::new(crate::lite::test_transport::SinkSend::new(log.clone()), version);
writer.encode(&ietf::RequestOk::ID).await.unwrap();
writer.encode(&ietf::RequestOk { request_id: None }).await.unwrap();
writer.encode(&ietf::Namespace::ID).await.unwrap();
writer
.encode(&ietf::Namespace {
suffix: crate::Path::new(suffix),
cluster: None,
})
.await
.unwrap();
let writes = log.writes.lock().unwrap();
writes.clone()
}
#[tokio::test]
async fn a_rooted_subscriber_mounts_a_reply_under_its_root_once() {
const VERSION: Version = Version::Draft18;
let origin = crate::origin::Info::new(crate::Origin::new(1).unwrap()).produce();
let consumer = origin.consume();
let scoped = origin
.with_root("rootns")
.and_then(|rooted| rooted.scope(&[crate::Path::new("cam")]))
.expect("scope the origin");
let session = crate::lite::test_transport::ScriptedSession::new(namespace_response(VERSION, "x.hang").await);
let (tasks, _task_set) = crate::util::TaskSet::new();
let peer_setup = cluster::PeerSetup::default();
peer_setup.set(cluster::Peer::default());
let mut subscriber = Subscriber::new(
session.clone(),
scoped,
Control::new(None, false),
None,
peer_setup,
crate::Origin::new(1).unwrap(),
None,
VERSION,
tasks,
);
let prefix = subscriber.subscribe_prefixes().pop().expect("one prefix");
let stream = Stream::open(&session, VERSION).await.unwrap();
let mut run = std::pin::pin!(subscriber.run_subscribe_namespace(stream, prefix));
for _ in 0..100 {
let _ = futures::poll!(run.as_mut());
if consumer.get_broadcast("rootns/cam/x.hang").is_some() {
break;
}
settle().await;
}
assert!(
consumer.get_broadcast("rootns/cam/x.hang").is_some(),
"the reply mounts under the root once",
);
assert!(
consumer.get_broadcast("rootns/rootns/cam/x.hang").is_none(),
"the root was applied twice",
);
}
#[test]
fn removing_old_track_does_not_remove_reused_alias() {
let aliases = TrackAliases::default();
insert_track_alias(&aliases, 7, RequestId(11)).unwrap();
remove_track_alias(&aliases, 7, RequestId(13));
assert_eq!(aliases.read().get(&7), Some(&RequestId(11)));
}
#[tokio::test]
async fn assigned_peer_origin_attributes_announces() {
let session = crate::lite::test_transport::SinkSession::new(Default::default());
let assigned = crate::Origin::new(777).unwrap();
let origin = crate::origin::Info::new(crate::Origin::new(1).unwrap()).produce();
let consumer = origin.consume();
let (tasks, _task_set) = crate::util::TaskSet::new();
let mut subscriber = Subscriber::new(
session,
origin,
Control::new(None, false),
Some(assigned),
cluster::PeerSetup::default(),
crate::Origin::new(1).unwrap(),
None,
Version::Draft14,
tasks,
);
let advert = subscriber.route(None, &cluster::Peer::default()).expect("route");
let _producer = subscriber
.start_announce(crate::Path::new("room/host").to_owned(), advert)
.unwrap();
tokio::time::sleep(std::time::Duration::from_millis(1)).await;
let broadcast = consumer.get_broadcast("room/host").unwrap();
let hops: Vec<_> = broadcast.routes()[0].hops.iter().copied().collect();
assert_eq!(hops, vec![assigned]);
}
fn cluster_subscriber(
self_origin: crate::Origin,
) -> (
Subscriber<crate::lite::test_transport::SinkSession>,
crate::origin::Producer,
) {
cluster_subscriber_with_linger(self_origin, std::time::Duration::ZERO)
}
fn cluster_subscriber_with_linger(
self_origin: crate::Origin,
linger: std::time::Duration,
) -> (
Subscriber<crate::lite::test_transport::SinkSession>,
crate::origin::Producer,
) {
let session = crate::lite::test_transport::SinkSession::new(Default::default());
let origin = crate::origin::Info::new(self_origin).with_linger(linger).produce();
let (tasks, task_set) = crate::util::TaskSet::new();
std::mem::forget(task_set);
let subscriber = Subscriber::new(
session,
origin.clone(),
Control::new(None, false),
None,
cluster::PeerSetup::default(),
self_origin,
None,
Version::Draft19,
tasks,
);
(subscriber, origin)
}
fn hop_path(ids: &[u64]) -> cluster::HopPath {
let hops = ids
.iter()
.map(|&id| crate::Origin::new(id).unwrap())
.collect::<Vec<_>>();
cluster::HopPath::new(crate::OriginList::try_from(hops).unwrap())
}
#[tokio::test]
async fn cluster_advert_becomes_a_route_with_the_link_charged() {
let (subscriber, origin) = cluster_subscriber(crate::Origin::new(1).unwrap());
let consumer = origin.consume();
let peer = cluster::Peer {
origin: Some(crate::Origin::new(9).unwrap()),
cost: Some(3),
};
let advert = cluster::Advert {
hops: hop_path(&[7, 9]),
cost: 4,
};
let advertised = subscriber.route(Some(&advert), &peer).expect("route");
assert_eq!(
advertised.route.cost, 7,
"the link's price is added to the advertised cost"
);
assert_eq!(advertised.route.hops, hop_path(&[7, 9]).hops().clone());
assert!(advertised.route.announce);
assert_eq!(advertised.publisher, Some(crate::Origin::new(7).unwrap()));
let mut subscriber = subscriber;
subscriber
.start_announce(crate::Path::new("room/host").to_owned(), advertised)
.unwrap();
tokio::time::sleep(std::time::Duration::from_millis(1)).await;
let broadcast = consumer.get_broadcast("room/host").unwrap();
let hops: Vec<_> = broadcast.routes()[0].hops.iter().map(|h| h.id()).collect();
assert_eq!(hops, vec![7, 9]);
assert_eq!(broadcast.routes()[0].cost, 7);
}
#[test]
fn cluster_advert_loop_is_discarded() {
let (subscriber, _origin) = cluster_subscriber(crate::Origin::new(5).unwrap());
let peer = cluster::Peer {
origin: Some(crate::Origin::new(9).unwrap()),
cost: None,
};
let looped = cluster::Advert {
hops: hop_path(&[7, 5, 9]),
cost: 0,
};
assert!(subscriber.route(Some(&looped), &peer).is_none());
let clean = cluster::Advert {
hops: hop_path(&[7, 9]),
cost: 0,
};
assert!(subscriber.route(Some(&clean), &peer).is_some());
}
#[test]
fn unpriced_link_costs_one() {
let (subscriber, _origin) = cluster_subscriber(crate::Origin::new(1).unwrap());
let peer = cluster::Peer {
origin: Some(crate::Origin::new(9).unwrap()),
cost: None,
};
let advert = cluster::Advert {
hops: hop_path(&[7, 9]),
cost: 2,
};
assert_eq!(subscriber.route(Some(&advert), &peer).unwrap().route.cost, 3);
let free = cluster::Peer {
origin: Some(crate::Origin::new(9).unwrap()),
cost: Some(0),
};
assert_eq!(subscriber.route(Some(&advert), &free).unwrap().route.cost, 2);
}
#[tokio::test(start_paused = true)]
async fn a_lost_namespace_stream_leaves_the_linger_window_open() {
const VERSION: Version = Version::Draft18;
let origin = crate::origin::Info::new(crate::Origin::new(1).unwrap())
.with_linger(std::time::Duration::from_secs(30))
.produce();
let consumer = origin.consume();
let session = crate::lite::test_transport::ScriptedSession::eof(namespace_response(VERSION, "x.hang").await);
let (tasks, task_set) = crate::util::TaskSet::new();
std::mem::forget(task_set);
let peer_setup = cluster::PeerSetup::default();
peer_setup.set(cluster::Peer::default());
let mut subscriber = Subscriber::new(
session.clone(),
origin,
Control::new(None, false),
None,
peer_setup,
crate::Origin::new(1).unwrap(),
None,
VERSION,
tasks,
);
let stream = Stream::open(&session, VERSION).await.unwrap();
subscriber
.run_subscribe_namespace(stream, crate::Path::new("").to_owned())
.await
.expect("a clean FIN is not an error");
settle().await;
assert!(
consumer.get_broadcast("x.hang").is_some(),
"an ended stream closed the broadcast instead of lingering for a reconnect",
);
}
#[tokio::test(start_paused = true)]
async fn an_explicit_namespace_done_closes_despite_the_linger_window() {
let (mut subscriber, origin) =
cluster_subscriber_with_linger(crate::Origin::new(1).unwrap(), std::time::Duration::from_secs(30));
let consumer = origin.consume();
let path = crate::Path::new("room/host").to_owned();
let advert = subscriber.route(None, &cluster::Peer::default()).expect("route");
subscriber.start_announce(path.clone(), advert).unwrap();
settle().await;
subscriber.stop_announce(path, Detach::Graceful).unwrap();
settle().await;
assert!(
consumer.get_broadcast("room/host").is_none(),
"an explicit NAMESPACE_DONE lingered instead of closing",
);
}
#[tokio::test(start_paused = true)]
async fn a_publish_namespace_done_retracts_without_faulting_the_session() {
const VERSION: Version = Version::Draft14;
let path = crate::Path::new("room/host").to_owned();
let log = crate::lite::test_transport::Log::default();
let mut writer = crate::coding::Writer::new(crate::lite::test_transport::SinkSend::new(log.clone()), VERSION);
writer
.encode_message(&ietf::PublishNamespaceDone {
track_namespace: path.borrow(),
request_id: RequestId(0),
})
.await
.unwrap();
let script = log.writes.lock().unwrap().clone();
let origin = crate::origin::Info::new(crate::Origin::new(1).unwrap())
.with_linger(std::time::Duration::from_secs(30))
.produce();
let consumer = origin.consume();
let session = crate::lite::test_transport::ScriptedSession::eof(script);
let (tasks, task_set) = crate::util::TaskSet::new();
std::mem::forget(task_set);
let mut subscriber = Subscriber::new(
session.clone(),
origin,
Control::new(None, false),
None,
cluster::PeerSetup::default(),
crate::Origin::new(1).unwrap(),
None,
VERSION,
tasks,
);
let stream = Stream::open(&session, VERSION).await.unwrap();
let msg = ietf::PublishNamespace {
request_id: RequestId(0),
track_namespace: path.borrow(),
cluster: None,
};
subscriber
.run_publish_namespace_stream(stream, msg, cluster::Peer::default())
.await
.expect("a withdrawal is not a protocol violation");
settle().await;
assert!(
consumer.get_broadcast("room/host").is_none(),
"an explicit withdrawal lingered instead of closing",
);
}
#[tokio::test(start_paused = true)]
async fn a_broken_publish_namespace_stream_leaves_the_linger_window_open() {
const VERSION: Version = Version::Draft19;
let origin = crate::origin::Info::new(crate::Origin::new(1).unwrap())
.with_linger(std::time::Duration::from_secs(30))
.produce();
let consumer = origin.consume();
let log = crate::lite::test_transport::Log::default();
let mut writer = crate::coding::Writer::new(crate::lite::test_transport::SinkSend::new(log.clone()), VERSION);
writer.encode(&ietf::NamespaceDone::ID).await.unwrap();
let script = log.writes.lock().unwrap().clone();
let session = crate::lite::test_transport::ScriptedSession::eof(script);
let (tasks, task_set) = crate::util::TaskSet::new();
std::mem::forget(task_set);
let mut subscriber = Subscriber::new(
session.clone(),
origin,
Control::new(None, false),
None,
cluster::PeerSetup::default(),
crate::Origin::new(1).unwrap(),
None,
VERSION,
tasks,
);
let path = crate::Path::new("room/host").to_owned();
let stream = Stream::open(&session, VERSION).await.unwrap();
let msg = ietf::PublishNamespace {
request_id: RequestId(0),
track_namespace: path.borrow(),
cluster: None,
};
subscriber
.run_publish_namespace_stream(stream, msg, cluster::Peer::default())
.await
.expect_err("an unexpected message ends the stream");
settle().await;
assert!(
consumer.get_broadcast("room/host").is_some(),
"a broken stream closed the broadcast instead of lingering for a reconnect",
);
}
#[tokio::test(start_paused = true)]
async fn the_last_owner_out_decides_the_detach() {
let (mut subscriber, origin) =
cluster_subscriber_with_linger(crate::Origin::new(1).unwrap(), std::time::Duration::from_secs(30));
let consumer = origin.consume();
let path = crate::Path::new("room/host").to_owned();
let peer = cluster::Peer::default();
for _ in 0..2 {
let advert = subscriber.route(None, &peer).expect("route");
subscriber.start_announce(path.clone(), advert).unwrap();
}
settle().await;
subscriber.stop_announce(path.clone(), Detach::Abrupt).unwrap();
subscriber.stop_announce(path.clone(), Detach::Graceful).unwrap();
settle().await;
assert!(
consumer.get_broadcast("room/host").is_none(),
"an explicit retraction lingered because an earlier owner was lost abruptly",
);
for _ in 0..2 {
let advert = subscriber.route(None, &peer).expect("route");
subscriber.start_announce(path.clone(), advert).unwrap();
}
settle().await;
subscriber.stop_announce(path.clone(), Detach::Graceful).unwrap();
subscriber.stop_announce(path, Detach::Abrupt).unwrap();
settle().await;
assert!(
consumer.get_broadcast("room/host").is_some(),
"an abrupt last detach closed the broadcast instead of lingering for a reconnect",
);
}
#[test]
fn a_pathless_advert_still_pays_for_its_link() {
let (unpriced, _origin) = cluster_subscriber(crate::Origin::new(1).unwrap());
let peer = cluster::Peer::default();
assert_eq!(unpriced.route(None, &peer).unwrap().route.cost, cluster::DEFAULT_COST);
let priced_peer = cluster::Peer {
origin: None,
cost: Some(4),
};
assert_eq!(unpriced.route(None, &priced_peer).unwrap().route.cost, 4);
let (mut priced, _origin) = cluster_subscriber(crate::Origin::new(1).unwrap());
priced.cost = Some(6);
assert_eq!(priced.route(None, &priced_peer).unwrap().route.cost, 6);
}
#[tokio::test]
async fn cluster_update_replaces_in_place() {
let (mut subscriber, origin) = cluster_subscriber(crate::Origin::new(1).unwrap());
let consumer = origin.consume();
let path = crate::Path::new("room/host").to_owned();
let first = crate::broadcast::Route::new()
.with_hops(hop_path(&[7, 9]).hops().clone())
.with_cost(4)
.with_announce(true);
let first = Advertised {
route: first,
publisher: Some(crate::Origin::new(7).unwrap()),
};
subscriber.start_announce(path.clone(), first).unwrap();
tokio::time::sleep(std::time::Duration::from_millis(1)).await;
let original = consumer.get_broadcast("room/host").unwrap();
let rerouted = crate::broadcast::Route::new()
.with_hops(hop_path(&[7, 11]).hops().clone())
.with_cost(2)
.with_announce(true);
let rerouted = Advertised {
route: rerouted,
publisher: Some(crate::Origin::new(7).unwrap()),
};
subscriber.update_announce(path.clone(), rerouted).unwrap();
tokio::time::sleep(std::time::Duration::from_millis(1)).await;
let broadcast = consumer.get_broadcast("room/host").unwrap();
let hops: Vec<_> = broadcast.routes()[0].hops.iter().map(|h| h.id()).collect();
assert_eq!(hops, vec![7, 11]);
assert_eq!(broadcast.routes()[0].cost, 2);
assert!(!original.is_closed(), "the source survived the update");
subscriber.stop_announce(path, Detach::Graceful).unwrap();
tokio::time::sleep(std::time::Duration::from_millis(1)).await;
assert!(consumer.get_broadcast("room/host").is_none());
}
#[tokio::test]
async fn cluster_publisher_change_replaces_the_source() {
let (mut subscriber, origin) = cluster_subscriber(crate::Origin::new(1).unwrap());
let consumer = origin.consume();
let path = crate::Path::new("room/host").to_owned();
let first = crate::broadcast::Route::new()
.with_hops(hop_path(&[7, 9]).hops().clone())
.with_announce(true);
let first = Advertised {
route: first,
publisher: Some(crate::Origin::new(7).unwrap()),
};
subscriber.start_announce(path.clone(), first).unwrap();
tokio::time::sleep(std::time::Duration::from_millis(1)).await;
let original = consumer.get_broadcast("room/host").unwrap();
let taken_over = crate::broadcast::Route::new()
.with_hops(hop_path(&[8, 9]).hops().clone())
.with_announce(true);
let taken_over = Advertised {
route: taken_over,
publisher: Some(crate::Origin::new(8).unwrap()),
};
subscriber.update_announce(path.clone(), taken_over).unwrap();
tokio::time::sleep(std::time::Duration::from_millis(1)).await;
assert!(original.is_closed(), "the old source was detached");
let broadcast = consumer.get_broadcast("room/host").unwrap();
let hops: Vec<_> = broadcast.routes()[0].hops.iter().map(|h| h.id()).collect();
assert_eq!(hops, vec![8, 9]);
subscriber.stop_announce(path, Detach::Graceful).unwrap();
tokio::time::sleep(std::time::Duration::from_millis(1)).await;
assert!(consumer.get_broadcast("room/host").is_none());
}
#[tokio::test]
async fn pathless_adverts_never_replace_the_source() {
let (mut subscriber, origin) = cluster_subscriber(crate::Origin::new(1).unwrap());
let consumer = origin.consume();
let path = crate::Path::new("room/host").to_owned();
let peer = cluster::Peer::default();
let first = subscriber.route(None, &peer).expect("route");
assert_eq!(first.publisher, None, "no path means no identity to compare");
subscriber.start_announce(path.clone(), first).unwrap();
tokio::time::sleep(std::time::Duration::from_millis(1)).await;
let original = consumer.get_broadcast("room/host").unwrap();
let second = subscriber.route(None, &peer).expect("route");
subscriber.start_announce(path.clone(), second).unwrap();
tokio::time::sleep(std::time::Duration::from_millis(1)).await;
assert!(!original.is_closed(), "the source survived the second advertisement");
assert!(consumer.get_broadcast("room/host").is_some());
subscriber.stop_announce(path.clone(), Detach::Graceful).unwrap();
tokio::time::sleep(std::time::Duration::from_millis(1)).await;
assert!(consumer.get_broadcast("room/host").is_some());
subscriber.stop_announce(path, Detach::Graceful).unwrap();
tokio::time::sleep(std::time::Duration::from_millis(1)).await;
assert!(consumer.get_broadcast("room/host").is_none());
}
#[tokio::test]
async fn reflected_replacement_retracts_the_route() {
let self_origin = crate::Origin::new(5).unwrap();
let (mut subscriber, origin) = cluster_subscriber(self_origin);
let consumer = origin.consume();
let path = crate::Path::new("room/host").to_owned();
let peer = cluster::Peer {
origin: Some(crate::Origin::new(9).unwrap()),
cost: None,
};
let clean = cluster::Advert {
hops: hop_path(&[7, 9]),
cost: 0,
};
let advert = subscriber.route(Some(&clean), &peer).expect("route");
subscriber.start_announce(path.clone(), advert).unwrap();
tokio::time::sleep(std::time::Duration::from_millis(1)).await;
assert!(consumer.get_broadcast("room/host").is_some());
let looped = cluster::Advert {
hops: hop_path(&[7, 5, 9]),
cost: 0,
};
assert!(
subscriber.route(Some(&looped), &peer).is_none(),
"a path containing our own Hop ID is a loop"
);
subscriber.stop_announce(path, Detach::Graceful).unwrap();
tokio::time::sleep(std::time::Duration::from_millis(1)).await;
assert!(
consumer.get_broadcast("room/host").is_none(),
"the superseded route must not stay attached"
);
}
async fn publish_namespace_updates(
request_id: RequestId,
path: &str,
updates: &[Option<cluster::Advert>],
) -> Vec<u8> {
const VERSION: Version = Version::Draft19;
let log = crate::lite::test_transport::Log::default();
let mut writer = crate::coding::Writer::new(crate::lite::test_transport::SinkSend::new(log.clone()), VERSION);
for cluster in updates {
writer.encode(&ietf::PublishNamespace::ID).await.unwrap();
writer
.encode(&ietf::PublishNamespace {
request_id,
track_namespace: crate::Path::new(path),
cluster: cluster.clone(),
})
.await
.unwrap();
}
let writes = log.writes.lock().unwrap();
writes.clone()
}
async fn reflected_harness(
self_origin: crate::Origin,
request_id: RequestId,
peer: &cluster::Peer,
attached: &cluster::Advert,
updates: &[Option<cluster::Advert>],
) -> (
Subscriber<crate::lite::test_transport::ScriptedSession>,
crate::origin::Consumer,
Stream<crate::lite::test_transport::ScriptedSession, Version>,
) {
const VERSION: Version = Version::Draft19;
let script = publish_namespace_updates(request_id, "room/host", updates).await;
let session = crate::lite::test_transport::ScriptedSession::new(script);
let origin = crate::origin::Info::new(self_origin).produce();
let consumer = origin.consume();
let (tasks, task_set) = crate::util::TaskSet::new();
std::mem::forget(task_set);
let mut subscriber = Subscriber::new(
session.clone(),
origin,
Control::new(None, false),
None,
cluster::PeerSetup::default(),
self_origin,
None,
VERSION,
tasks,
);
let path = crate::Path::new("room/host").to_owned();
let advert = subscriber.route(Some(attached), peer).expect("route");
subscriber.start_announce(path, advert).unwrap();
settle().await;
assert!(consumer.get_broadcast("room/host").is_some(), "attached to start with");
let stream = Stream::open(&session, VERSION).await.unwrap();
(subscriber, consumer, stream)
}
fn peer_9() -> cluster::Peer {
cluster::Peer {
origin: Some(crate::Origin::new(9).unwrap()),
cost: None,
}
}
fn clean_and_looped() -> (cluster::Advert, cluster::Advert) {
(
cluster::Advert {
hops: hop_path(&[7, 9]),
cost: 0,
},
cluster::Advert {
hops: hop_path(&[7, 5, 9]),
cost: 0,
},
)
}
#[tokio::test]
async fn a_reflected_update_detaches_but_keeps_the_stream() {
let self_origin = crate::Origin::new(5).unwrap();
let request_id = RequestId(1);
let peer = peer_9();
let (clean, looped) = clean_and_looped();
let (mut subscriber, consumer, mut stream) =
reflected_harness(self_origin, request_id, &peer, &clean, &[Some(looped)]).await;
let path = crate::Path::new("room/host").to_owned();
let mut attached = true;
{
let mut run = std::pin::pin!(subscriber.run_publish_namespace_updates(
&mut stream,
&path,
request_id,
peer,
&mut attached,
));
for _ in 0..100 {
assert!(
futures::poll!(run.as_mut()).is_pending(),
"the stream must stay open after a reflected update"
);
if consumer.get_broadcast("room/host").is_none() {
break;
}
settle().await;
}
}
assert!(
consumer.get_broadcast("room/host").is_none(),
"an unusable path must not stay attached"
);
assert!(!attached, "the caller must not release it a second time");
}
#[tokio::test]
async fn a_clean_update_after_a_reflection_reattaches() {
let self_origin = crate::Origin::new(5).unwrap();
let request_id = RequestId(1);
let peer = peer_9();
let (clean, looped) = clean_and_looped();
let (mut subscriber, consumer, mut stream) = reflected_harness(
self_origin,
request_id,
&peer,
&clean,
&[Some(looped), Some(clean.clone())],
)
.await;
let path = crate::Path::new("room/host").to_owned();
let mut attached = true;
{
let mut run = std::pin::pin!(subscriber.run_publish_namespace_updates(
&mut stream,
&path,
request_id,
peer,
&mut attached,
));
for _ in 0..20 {
assert!(futures::poll!(run.as_mut()).is_pending());
settle().await;
}
}
assert!(attached, "the clean path must re-attach");
assert!(
consumer.get_broadcast("room/host").is_some(),
"the namespace is routable again",
);
}
#[tokio::test]
async fn namespace_stream_close_releases_live_paths() {
let (mut subscriber, origin) = cluster_subscriber(crate::Origin::new(1).unwrap());
let consumer = origin.consume();
let peer = cluster::Peer::default();
let mut live = std::collections::HashSet::new();
for path in ["room/a", "room/b"] {
let path = crate::Path::new(path).to_owned();
let advert = subscriber.route(None, &peer).expect("route");
subscriber.start_announce(path.clone(), advert).unwrap();
live.insert(path);
}
tokio::time::sleep(std::time::Duration::from_millis(1)).await;
assert!(consumer.get_broadcast("room/a").is_some());
assert!(consumer.get_broadcast("room/b").is_some());
for path in live {
subscriber.stop_announce(path, Detach::Graceful).unwrap();
}
tokio::time::sleep(std::time::Duration::from_millis(1)).await;
assert!(consumer.get_broadcast("room/a").is_none(), "room/a leaked a refcount");
assert!(consumer.get_broadcast("room/b").is_none(), "room/b leaked a refcount");
}
#[tokio::test]
async fn publish_is_rejected_without_announcing() {
let gate = kio::Producer::new(true);
let session = crate::lite::test_transport::SinkSession::gated_bi(gate.consume());
let origin = crate::origin::Info::new(crate::Origin::new(1).unwrap()).produce();
let consumer = origin.consume();
let (tasks, task_set) = crate::util::TaskSet::new();
std::mem::forget(task_set);
let mut subscriber = Subscriber::new(
session.clone(),
origin,
Control::new(None, false),
None,
cluster::PeerSetup::default(),
crate::Origin::new(1).unwrap(),
None,
Version::Draft19,
tasks,
);
let stream = Stream::open(&session, Version::Draft19).await.unwrap();
let msg = ietf::Publish {
request_id: RequestId(1),
track_namespace: crate::Path::new("room/host"),
track_name: "video".into(),
track_alias: 7,
group_order: ietf::GroupOrder::Ascending,
largest_location: None,
forward: true,
timescale: None,
};
subscriber.run_publish_stream(stream, msg).await.unwrap();
tokio::time::sleep(std::time::Duration::from_millis(1)).await;
assert!(
consumer.get_broadcast("room/host").is_none(),
"a rejected PUBLISH must not announce a broadcast"
);
assert_eq!(
occurrences(&session.log, b"PUBLISH is not supported"),
1,
"the decline reaches the peer"
);
}
}