use crate::{stats, track};
use std::{
collections::{HashMap, VecDeque},
sync::Arc,
task::{Poll, ready},
};
use crate::Error;
use super::{OriginList, Requests, WeakCache};
#[derive(Clone, Debug, Default)]
#[non_exhaustive]
pub struct Info {
pub origin: super::origin::Info,
}
impl Info {
pub fn new() -> Self {
Self::default()
}
pub fn produce(self) -> Producer {
Producer::new(self)
}
}
#[derive(Clone, Debug, Default, PartialEq, Eq)]
#[non_exhaustive]
pub struct Route {
pub hops: OriginList,
pub cost: u64,
pub(crate) advertised: u64,
pub announce: bool,
}
impl Route {
pub fn new() -> Self {
Self::default()
}
pub fn announced() -> Self {
Self {
announce: true,
..Self::default()
}
}
pub fn with_hop(mut self, origin: super::Origin) -> Result<Self, super::TooManyOrigins> {
self.hops.push(origin)?;
Ok(self)
}
pub fn with_hops(mut self, hops: OriginList) -> Self {
self.hops = hops;
self
}
pub fn with_cost(mut self, cost: u64) -> Self {
self.cost = cost;
self
}
pub fn with_announce(mut self, announce: bool) -> Self {
self.announce = announce;
self
}
}
#[derive(Default)]
struct BroadcastState {
tracks: WeakCache<Arc<str>, track::TrackWeak>,
requests: Requests<Arc<str>, track::Request>,
spliced: Option<SplicedState>,
route: Route,
route_epoch: u64,
closing: bool,
finished: bool,
abort: Option<Error>,
}
#[derive(Default)]
struct SplicedState {
tracks: HashMap<Arc<str>, super::resume::Producer>,
pending: VecDeque<Arc<str>>,
}
impl BroadcastState {
fn insert_track(&mut self, weak: track::TrackWeak) -> Result<(), Error> {
match self.tracks.insert(weak.name().clone(), weak) {
Some(_) => Err(Error::Duplicate),
None => Ok(()),
}
}
fn is_used(&self) -> bool {
if let Some(spliced) = &self.spliced {
return spliced.tracks.values().any(|track| track.is_used());
}
!self.requests.is_empty() || self.tracks.iter().any(|track| track.is_used())
}
fn register_demand(&self, waiter: &kio::Waiter, want: bool) {
if let Some(spliced) = &self.spliced {
for track in spliced.tracks.values() {
match want {
true => track.poll_used(waiter),
false => track.poll_unused(waiter),
}
}
return;
}
for track in self.tracks.iter() {
match want {
true => track.poll_used(waiter),
false => track.poll_unused(waiter),
}
}
}
}
#[derive(Clone)]
pub struct Producer {
info: Arc<Info>,
alive: kio::Producer<()>,
state: kio::Shared<BroadcastState>,
stats: stats::Scope,
}
impl Producer {
pub fn new(info: Info) -> Self {
Self {
info: Arc::new(info),
alive: Default::default(),
state: Default::default(),
stats: stats::Scope::default(),
}
}
pub(crate) fn with_stats(mut self, scope: stats::Scope) -> Self {
self.stats = scope;
self
}
pub(crate) fn new_spliced(info: Info) -> Self {
Self {
info: Arc::new(info),
alive: Default::default(),
state: kio::Shared::new(BroadcastState {
spliced: Some(SplicedState::default()),
..Default::default()
}),
stats: stats::Scope::default(),
}
}
pub fn info(&self) -> &Info {
&self.info
}
pub fn demand(&self) -> Demand {
Demand {
alive: self.alive.consume().weak(),
state: self.state.clone(),
}
}
pub fn remove_track(&mut self, name: &str) -> Result<(), Error> {
self.state.lock().tracks.remove(name).ok_or(Error::NotFound)?;
Ok(())
}
pub fn create_track(
&mut self,
name: impl Into<Arc<str>>,
info: impl Into<Option<track::Info>>,
) -> Result<track::Producer, Error> {
let info = info.into().unwrap_or_default();
let track = track::Producer::new(self.info.clone(), name, info).with_stats(self.stats.clone());
self.state.lock().insert_track(track.weak())?;
Ok(track)
}
pub fn reserve_track(&mut self, name: impl Into<Arc<str>>) -> Result<track::Request, Error> {
let request = track::Request::new(self.info.clone(), name).with_stats(self.stats.clone());
self.state.lock().insert_track(request.weak())?;
Ok(request)
}
pub fn unique_track(
&mut self,
suffix: &str,
info: impl Into<Option<track::Info>>,
) -> Result<track::Producer, Error> {
let name = self.unique_name(suffix);
self.create_track(name, info)
}
pub fn unique_name(&self, suffix: &str) -> String {
let state = self.state.read();
(0u16..)
.map(|i| format!("{i}{suffix}"))
.find(|name| !state.tracks.contains_key(name.as_str()))
.expect("u16 namespace exhausted; wow")
}
pub fn dynamic(&self) -> Dynamic {
Dynamic::new(
self.info.clone(),
self.alive.clone(),
self.state.clone(),
self.stats.clone(),
)
}
pub fn set_route(&mut self, route: Route) -> Result<(), Error> {
let mut state = self.state.lock();
if state.route == route {
return Ok(());
}
state.route = route;
state.route_epoch += 1;
Ok(())
}
pub(crate) fn poll_spliced_assigned(&self, waiter: &kio::Waiter) -> Poll<(Arc<str>, super::resume::Producer)> {
let mut state = ready!(self.state.poll(waiter, |state| {
match &state.spliced {
Some(spliced) if !spliced.pending.is_empty() => Poll::Ready(()),
_ => Poll::Pending,
}
}));
let spliced = state.spliced.as_mut().expect("predicate guaranteed spliced");
let name = spliced.pending.pop_front().expect("predicate guaranteed a request");
let producer = spliced.tracks.get(&name).expect("pending name without a track").clone();
Poll::Ready((name, producer))
}
pub(crate) fn abort_spliced(&self, err: Error) {
let mut state = self.state.lock();
if let Some(spliced) = state.spliced.as_mut() {
spliced.pending.clear();
for producer in spliced.tracks.values_mut() {
let _ = producer.abort(err.clone());
}
}
}
pub fn consume(&self) -> Consumer {
Consumer {
info: self.info.clone(),
alive: self.alive.consume(),
state: self.state.clone(),
route_seen: None,
stats: stats::Scope::default(),
}
}
pub fn finish(&mut self) {
{
let mut state = self.state.lock();
state.closing = true;
state.finished = true;
}
let _ = self.alive.close();
}
pub fn abort(self, err: Error) -> Result<(), Error> {
{
let mut state = self.state.lock();
if state.closing {
return Err(Error::Closed);
}
state.closing = true;
state.abort = Some(err);
}
let _ = self.alive.close();
Ok(())
}
pub fn is_clone(&self, other: &Self) -> bool {
self.state.same_channel(&other.state)
}
}
impl Drop for Producer {
fn drop(&mut self) {
if !self.alive.is_last() {
return;
}
if !self.state.read().closing {
tracing::warn!(
"broadcast::Producer dropped without finish(). Keep the producer alive while publishing, then call finish()."
);
}
}
}
#[cfg(test)]
#[allow(missing_docs)] impl Producer {
pub fn assert_create_track(
&mut self,
name: impl Into<Arc<str>>,
info: impl Into<Option<track::Info>>,
) -> track::Producer {
self.create_track(name, info).expect("should not have errored")
}
}
pub(crate) struct SourceGuard {
producer: Option<Producer>,
}
impl SourceGuard {
pub fn new(producer: Producer) -> Self {
Self {
producer: Some(producer),
}
}
pub fn producer(&self) -> Producer {
self.producer.clone().expect("guard holds a producer until finished")
}
pub fn finish(mut self) {
if let Some(mut producer) = self.producer.take() {
producer.finish();
}
}
pub fn set_route(&mut self, route: Route) {
if let Some(producer) = &mut self.producer {
let _ = producer.set_route(route);
}
}
}
impl Drop for SourceGuard {
fn drop(&mut self) {
if let Some(producer) = self.producer.take() {
let _ = producer.abort(Error::Dropped);
}
}
}
pub struct Dynamic {
info: Arc<Info>,
alive: kio::Producer<()>,
state: kio::Shared<BroadcastState>,
stats: stats::Scope,
}
impl Clone for Dynamic {
fn clone(&self) -> Self {
self.state.lock().requests.add_handler();
Self {
info: self.info.clone(),
alive: self.alive.clone(),
state: self.state.clone(),
stats: self.stats.clone(),
}
}
}
impl Dynamic {
fn new(info: Arc<Info>, alive: kio::Producer<()>, state: kio::Shared<BroadcastState>, stats: stats::Scope) -> Self {
state.lock().requests.add_handler();
Self {
info,
alive,
state,
stats,
}
}
pub fn info(&self) -> &Info {
&self.info
}
pub fn poll_requested_track(&mut self, waiter: &kio::Waiter) -> Poll<Result<track::Request, Error>> {
let mut state = ready!(self.state.poll(waiter, |state| {
if state.requests.has_queued() || state.closing {
Poll::Ready(())
} else {
Poll::Pending
}
}));
if state.closing && !state.requests.has_queued() {
return Poll::Ready(Err(Error::Closed));
}
let name = state.requests.pop().expect("predicate guaranteed a request");
let pending = state.requests.remove(&name).expect("popped key must be pending");
let _ = state.tracks.insert(name, pending.weak());
Poll::Ready(Ok(pending.with_stats(self.stats.clone())))
}
pub async fn requested_track(&mut self) -> Result<track::Request, Error> {
kio::wait(|waiter| self.poll_requested_track(waiter)).await
}
pub fn consume(&self) -> Consumer {
Consumer {
info: self.info.clone(),
alive: self.alive.consume(),
state: self.state.clone(),
route_seen: None,
stats: stats::Scope::default(),
}
}
pub async fn closed(&self) -> Error {
kio::wait(|waiter| self.poll_closed(waiter)).await
}
pub fn poll_closed(&self, waiter: &kio::Waiter) -> Poll<Error> {
ready!(self.alive.poll_closed(waiter));
Poll::Ready(self.state.read().abort.clone().unwrap_or(Error::Dropped))
}
pub fn is_clone(&self, other: &Self) -> bool {
self.state.same_channel(&other.state)
}
}
impl Drop for Dynamic {
fn drop(&mut self) {
let mut state = self.state.lock();
if state.requests.remove_handler() {
for request in state.requests.drain_queued() {
request.reject(Error::Dropped);
}
}
}
}
#[cfg(test)]
use futures::FutureExt;
#[cfg(test)]
#[allow(missing_docs)] impl Dynamic {
pub fn assert_request(&mut self) -> track::Request {
self.requested_track()
.now_or_never()
.expect("should not have blocked")
.expect("should not have errored")
}
pub fn assert_no_request(&mut self) {
assert!(self.requested_track().now_or_never().is_none(), "should have blocked");
}
}
pub struct Consumer {
info: Arc<Info>,
alive: kio::Consumer<()>,
state: kio::Shared<BroadcastState>,
route_seen: Option<u64>,
stats: stats::Scope,
}
impl Clone for Consumer {
fn clone(&self) -> Self {
Self {
info: self.info.clone(),
alive: self.alive.clone(),
state: self.state.clone(),
route_seen: None,
stats: self.stats.clone(),
}
}
}
impl Consumer {
pub(crate) fn with_stats(mut self, scope: stats::Scope) -> Self {
self.stats = scope;
self
}
pub fn info(&self) -> &Info {
&self.info
}
pub fn route(&self) -> Route {
self.state.read().route.clone()
}
pub fn poll_route_changed(&mut self, waiter: &kio::Waiter) -> Poll<Result<Route, Error>> {
let seen = self.route_seen;
if let Poll::Ready(state) = self.state.poll(waiter, |state| {
if seen != Some(state.route_epoch) {
Poll::Ready(())
} else {
Poll::Pending
}
}) {
self.route_seen = Some(state.route_epoch);
return Poll::Ready(Ok(state.route.clone()));
}
ready!(self.alive.poll_closed(waiter));
Poll::Ready(Err(Error::Dropped))
}
pub async fn route_changed(&mut self) -> Result<Route, Error> {
kio::wait(|waiter| self.poll_route_changed(waiter)).await
}
pub fn track(&self, name: &str) -> Result<track::Consumer, Error> {
self.track_inner(name).map(|track| track.with_stats(self.stats.clone()))
}
fn track_inner(&self, name: &str) -> Result<track::Consumer, Error> {
if self.is_closed() {
return Err(Error::Dropped);
}
let mut state = self.state.lock();
if let Some(spliced) = state.spliced.as_mut() {
if let Some(producer) = spliced.tracks.get(name) {
return Ok(track::Consumer::spliced(name.into(), producer.consume()));
}
let name: Arc<str> = name.into();
let producer = super::resume::Producer::new();
let consumer = producer.consume();
spliced.tracks.insert(name.clone(), producer);
spliced.pending.push_back(name.clone());
return Ok(track::Consumer::spliced(name, consumer));
}
if let Some(weak) = state.tracks.get(name) {
return Ok(weak.consume());
}
if let Some(pending) = state.requests.join(name) {
return Ok(pending.consume());
}
if state.closing {
return Err(Error::NotFound);
}
let name: Arc<str> = name.into();
let request = track::Request::new(self.info.clone(), name.clone());
let consumer = request.consume();
if state.requests.insert(name, request).is_err() {
return Err(Error::NotFound);
}
Ok(consumer)
}
pub(crate) fn demand(&self) -> Demand {
Demand {
alive: self.alive.weak(),
state: self.state.clone(),
}
}
pub async fn closed(&self) -> Error {
self.alive.closed().await;
self.state.read().abort.clone().unwrap_or(Error::Dropped)
}
pub fn is_closed(&self) -> bool {
self.alive.is_closed()
}
pub(crate) fn is_closing(&self) -> bool {
self.is_closed() || self.state.read().closing
}
pub fn is_finished(&self) -> bool {
self.state.read().finished
}
pub fn poll_closed(&self, waiter: &kio::Waiter) -> Poll<()> {
self.alive.poll_closed(waiter)
}
pub fn is_clone(&self, other: &Self) -> bool {
self.state.same_channel(&other.state)
}
pub(crate) fn weak(&self) -> WeakConsumer {
WeakConsumer {
info: self.info.clone(),
alive: self.alive.weak(),
state: self.state.clone(),
}
}
}
#[derive(Clone)]
pub(crate) struct WeakConsumer {
info: Arc<Info>,
alive: kio::ConsumerWeak<()>,
state: kio::Shared<BroadcastState>,
}
impl WeakConsumer {
pub fn consume(&self) -> Consumer {
Consumer {
info: self.info.clone(),
alive: self.alive.consume(),
state: self.state.clone(),
route_seen: None,
stats: stats::Scope::default(),
}
}
}
impl super::WeakEntry for WeakConsumer {
fn is_closed(&self) -> bool {
self.alive.is_closed()
}
fn same_channel(&self, other: &Self) -> bool {
self.state.same_channel(&other.state)
}
}
#[derive(Clone)]
pub struct Demand {
alive: kio::ConsumerWeak<()>,
state: kio::Shared<BroadcastState>,
}
impl Demand {
pub fn is_used(&self) -> bool {
self.state.read().is_used()
}
pub async fn used(&self) -> Result<(), Error> {
kio::wait(|waiter| self.poll_used(waiter)).await
}
pub async fn unused(&self) -> Result<(), Error> {
kio::wait(|waiter| self.poll_unused(waiter)).await
}
pub fn poll_used(&self, waiter: &kio::Waiter) -> Poll<Result<(), Error>> {
self.poll_demand(waiter, true)
}
pub fn poll_unused(&self, waiter: &kio::Waiter) -> Poll<Result<(), Error>> {
self.poll_demand(waiter, false)
}
fn poll_demand(&self, waiter: &kio::Waiter, want: bool) -> Poll<Result<(), Error>> {
if self.alive.poll_closed(waiter).is_ready() {
return Poll::Ready(Err(Error::Dropped));
}
let ready = self.state.poll(waiter, |state| {
state.register_demand(waiter, want);
match state.is_used() == want {
true => Poll::Ready(()),
false => Poll::Pending,
}
});
match ready {
Poll::Ready(_) => Poll::Ready(Ok(())),
Poll::Pending => Poll::Pending,
}
}
}
#[cfg(test)]
#[allow(missing_docs)] impl Consumer {
pub fn assert_not_closed(&self) {
assert!(self.closed().now_or_never().is_none(), "should not be closed");
}
pub fn assert_closed(&self) {
assert!(self.closed().now_or_never().is_some(), "should be closed");
}
}
#[cfg(test)]
mod test {
use super::*;
async fn expect<T>(fut: impl Future<Output = T>) -> T {
tokio::time::timeout(std::time::Duration::from_secs(1), fut)
.await
.expect("timed out waiting for a demand edge")
}
#[tokio::test]
async fn demand_ordinary() {
tokio::time::pause();
let mut producer = Info::new().produce();
let consumer = producer.consume();
let demand = producer.demand();
assert!(!demand.is_used());
demand.unused().await.unwrap();
let _track = producer.create_track("a", None).unwrap();
assert!(!demand.is_used());
let (used, handle) = tokio::join!(expect(demand.used()), async { consumer.track("a").unwrap() });
used.unwrap();
assert!(demand.is_used());
let (unused, ()) = tokio::join!(expect(demand.unused()), async { drop(handle) });
unused.unwrap();
assert!(!demand.is_used());
producer.finish();
assert!(matches!(demand.used().await, Err(Error::Dropped)));
assert!(matches!(demand.unused().await, Err(Error::Dropped)));
}
#[tokio::test]
async fn demand_spliced() {
tokio::time::pause();
let producer = Producer::new_spliced(Info::new());
let consumer = producer.consume();
let demand = producer.demand();
assert!(!demand.is_used());
let track = consumer.track("video").unwrap();
assert!(demand.is_used());
let (unused, ()) = tokio::join!(expect(demand.unused()), async { drop(track) });
unused.unwrap();
assert!(!demand.is_used());
let _track = consumer.track("video").unwrap();
assert!(demand.is_used());
}
macro_rules! subscribe_pending {
($consumer:expr, $name:expr) => {{
let pending = $consumer.track($name).unwrap().subscribe(None);
assert!(
pending.poll_ok(&kio::Waiter::noop()).is_pending(),
"subscribe should stay pending until the request is accepted"
);
pending
}};
}
#[tokio::test]
async fn insert() {
let mut producer = Info::new().produce();
let mut track1 = producer.assert_create_track("track1", None);
track1.append_group().unwrap();
let consumer = producer.consume();
let mut track1_sub = consumer.track("track1").unwrap().subscribe(None).await.unwrap();
track1_sub.assert_group();
let mut track2 = producer.assert_create_track("track2", None);
let consumer2 = producer.consume();
let mut track2_consumer = consumer2.track("track2").unwrap().subscribe(None).await.unwrap();
track2_consumer.assert_no_group();
track2.append_group().unwrap();
track2_consumer.assert_group();
}
#[tokio::test]
async fn closed() {
let mut producer = Info::new().produce();
let dynamic = producer.dynamic();
let consumer = producer.consume();
consumer.assert_not_closed();
let track1 = producer.assert_create_track("track1", None);
let mut track1c = consumer.track("track1").unwrap().subscribe(None).await.unwrap();
let track2_fut = subscribe_pending!(consumer, "track2");
drop(dynamic);
assert!(track2_fut.await.is_err());
assert!(!track1.is_closed());
track1c.assert_not_closed();
}
#[tokio::test]
async fn closed_cause() {
let producer = Info::new().produce();
let consumer = producer.consume();
producer.abort(Error::Timeout).unwrap();
assert!(matches!(consumer.closed().await, Error::Timeout));
assert!(!consumer.is_finished());
let mut producer = Info::new().produce();
let consumer = producer.consume();
producer.finish();
assert!(matches!(consumer.closed().await, Error::Dropped));
assert!(consumer.is_finished());
let producer = Info::new().produce();
let consumer = producer.consume();
drop(producer);
assert!(matches!(consumer.closed().await, Error::Dropped));
assert!(!consumer.is_finished());
}
#[tokio::test]
async fn requests() {
let mut producer = Info::new().produce().dynamic();
let consumer = producer.consume();
let consumer2 = consumer.clone();
let track1_fut = subscribe_pending!(consumer, "track1");
let track2_fut = subscribe_pending!(consumer2, "track1");
let request = producer.assert_request();
producer.assert_no_request();
assert_eq!(request.name(), "track1");
let mut track3 = request.accept(None);
let mut track1 = track1_fut.await.unwrap();
let mut track2 = track2_fut.await.unwrap();
track1.assert_not_closed();
track1.assert_is_clone(&track2);
track3.subscribe(None).assert_is_clone(&track1);
track3.append_group().unwrap();
track1.assert_group();
track2.assert_group();
let track4_fut = subscribe_pending!(consumer, "track2");
drop(producer);
assert!(track4_fut.await.is_err());
let track5 = consumer2.track("track3");
assert!(track5.is_err(), "should have errored");
}
#[tokio::test]
async fn stale_producer() {
let mut broadcast = Info::new().produce().dynamic();
let consumer = broadcast.consume();
let track1_fut = subscribe_pending!(consumer, "track1");
let mut producer1 = broadcast.assert_request().accept(None);
let mut track1 = track1_fut.await.unwrap();
producer1.append_group().unwrap();
producer1.finish().unwrap();
drop(producer1);
track1.assert_closed();
let track2_fut = subscribe_pending!(consumer, "track1");
let mut producer2 = broadcast.assert_request().accept(None);
let mut track2 = track2_fut.await.unwrap();
track2.assert_not_closed();
track2.assert_not_clone(&track1);
producer2.append_group().unwrap();
track2.assert_group();
}
#[tokio::test(start_paused = true)]
async fn requested_unused() {
let mut broadcast = Info::new().produce().dynamic();
let bc = broadcast.consume();
let c1_fut = subscribe_pending!(bc, "unknown_track");
let producer1 = broadcast.assert_request().accept(None);
let consumer1 = c1_fut.await.unwrap();
assert!(
producer1.unused().now_or_never().is_none(),
"track producer should be used"
);
let consumer2 = bc.track("unknown_track").unwrap().subscribe(None).await.unwrap();
consumer2.assert_is_clone(&consumer1);
drop(consumer1);
assert!(
producer1.unused().now_or_never().is_none(),
"track producer should be used"
);
drop(consumer2);
assert!(
producer1.unused().now_or_never().is_some(),
"track producer should be unused after all consumers are dropped"
);
let consumer3 = bc.track("unknown_track").unwrap().subscribe(None).await.unwrap();
consumer3.assert_is_clone(&producer1.subscribe(None));
broadcast.assert_no_request();
drop(consumer3);
producer1.abort(Error::Cancel).unwrap();
let c4_fut = subscribe_pending!(bc, "unknown_track");
let producer2 = broadcast.assert_request().accept(None);
let consumer4 = c4_fut.await.unwrap();
drop(consumer4);
assert!(
producer2.unused().now_or_never().is_some(),
"new track producer should be unused after its consumer is dropped"
);
}
#[tokio::test]
async fn route_clone_observes_current_route() {
let mut producer = Info::new().produce();
let mut consumer = producer.consume();
consumer.route_changed().await.unwrap();
let route = Route::new().with_cost(7);
producer.set_route(route.clone()).unwrap();
assert_eq!(consumer.route_changed().await.unwrap(), route);
assert!(consumer.route_changed().now_or_never().is_none());
let mut clone = consumer.clone();
let seen = clone
.route_changed()
.now_or_never()
.expect("clone should observe the current route immediately")
.unwrap();
assert_eq!(seen, route);
}
#[tokio::test]
async fn dynamic_clone_keeps_alive() {
let broadcast = Info::new().produce().dynamic();
let consumer = broadcast.consume();
let clone = broadcast.clone();
drop(clone);
let _fut = subscribe_pending!(consumer, "track1");
}
}