use crate::origin;
use crate::{
Error, Origin,
coding::{Decode, Encode, Reader, Stream, Writer},
ietf::{self, FetchHeader, RequestId},
setup,
util::{MaybeBoxedExt, MaybeSendBox, TaskSet, err_only},
};
use super::{Control, Message, Publisher, Subscriber, Version, adapter::ControlStreamAdapter, cluster};
pub struct Config<S: web_transport_trait::Session> {
pub session: S,
pub setup: Option<Stream<S, Version>>,
pub request_id_max: Option<RequestId>,
pub client: bool,
pub publish: Option<origin::Consumer>,
pub subscribe: Option<origin::Producer>,
pub peer_origin: Option<Origin>,
pub cost: Option<u64>,
pub version: Version,
pub path: Option<String>,
pub peer_setup_stream: Option<Reader<S::RecvStream, crate::Version>>,
pub peer_cluster: Option<cluster::Peer>,
}
pub fn start<S: web_transport_trait::Session>(
config: Config<S>,
) -> Result<MaybeSendBox<'static, Result<(), Error>>, Error> {
let Config {
session,
setup,
request_id_max,
client,
publish,
subscribe,
peer_origin,
cost,
version,
path,
peer_setup_stream,
peer_cluster,
} = config;
let driver = async move {
let self_origin = self_origin(publish.as_ref(), subscribe.as_ref());
let publish = publish.unwrap_or_else(|| origin::Producer::empty(Origin::random()).consume());
let subscribe = subscribe.unwrap_or_else(|| origin::Producer::empty(Origin::random()));
let peer_setup = cluster::PeerSetup::default();
let setup_read = peer_cluster.is_some();
match peer_cluster {
Some(peer) => peer_setup.set(peer),
None if !cluster::supported(version) => peer_setup.set(cluster::Peer::default()),
None => {}
}
let res = match version {
Version::Draft14 | Version::Draft15 | Version::Draft16 => {
let Some(setup) = setup else {
let err = Error::ProtocolViolation;
session.close(err.to_code(), "setup stream required");
return Err(err);
};
let control = Control::new(request_id_max, client);
let adapter = ControlStreamAdapter::new(session.clone(), control.clone(), version);
let publisher = Publisher::new(
adapter.clone(),
publish,
control.clone(),
peer_origin,
peer_setup.clone(),
version,
);
let (tasks, mut task_set) = TaskSet::new();
let subscriber = Subscriber::new(
adapter.clone(),
subscribe,
control,
peer_origin,
peer_setup.clone(),
self_origin,
cost,
version,
tasks,
);
let dispatch_session = adapter.clone();
let sub_ns = subscriber.clone();
let sub_ns_adapter = adapter.clone();
let mut adapter_run = std::pin::pin!(err_only(adapter.run(setup.reader, setup.writer)));
let mut unis = std::pin::pin!(err_only(run_unis(
adapter.clone(),
subscriber.clone(),
None,
false,
version
)));
let mut dispatch = std::pin::pin!(err_only(run_dispatch(
dispatch_session,
publisher.clone(),
subscriber.clone(),
version
)));
let mut sub_ns_run = std::pin::pin!(err_only(async {
let mut prefixes = futures::stream::FuturesUnordered::new();
for prefix in sub_ns.subscribe_prefixes() {
let mut sub_ns = sub_ns.clone();
let sub_ns_adapter = sub_ns_adapter.clone();
prefixes.push(async move {
let stream = match version {
Version::Draft16 => {
let (send, recv) = sub_ns_adapter.open_native_bi().await?;
Stream {
writer: crate::coding::Writer::new(send, version),
reader: crate::coding::Reader::new(recv, version),
}
}
_ => Stream::open(&sub_ns_adapter, version).await?,
};
if let Err(err) = sub_ns.run_subscribe_namespace(stream, prefix).await {
tracing::warn!(%err, "subscribe_namespace failed, continuing without");
}
Ok::<(), Error>(())
});
}
while let Some(result) = futures::StreamExt::next(&mut prefixes).await {
result?;
}
Ok(())
}));
kio::wait(|waiter| {
use std::task::Poll;
if let Poll::Ready(err) = waiter.poll_future(adapter_run.as_mut()) {
return Poll::Ready(Err::<(), Error>(err));
}
if let Poll::Ready(err) = waiter.poll_future(unis.as_mut()) {
return Poll::Ready(Err(err));
}
if let Poll::Ready(err) = waiter.poll_future(dispatch.as_mut()) {
return Poll::Ready(Err(err));
}
if task_set.poll(waiter).is_ready() {
return Poll::Ready(Ok(()));
}
if let Poll::Ready(err) = waiter.poll_future(sub_ns_run.as_mut()) {
return Poll::Ready(Err(err));
}
Poll::Pending
})
.await
}
_ => {
let setup = {
let session = session.clone();
async move {
if let Err(err) = run_setup(session, version, path, self_origin, cost).await {
tracing::warn!(%err, "setup send error");
}
std::future::pending::<()>().await;
}
};
let control = Control::new(None, client);
let publisher = Publisher::new(
session.clone(),
publish,
control.clone(),
peer_origin,
peer_setup.clone(),
version,
);
let (tasks, mut task_set) = TaskSet::new();
let subscriber = Subscriber::new(
session.clone(),
subscribe,
control,
peer_origin,
peer_setup.clone(),
self_origin,
cost,
version,
tasks,
);
let sub_ns_session = session.clone();
let sub_ns = subscriber.clone();
let goaway = async move {
match peer_setup_stream {
Some(reader) => run_goaway(reader.with_version(version), version).await,
None => std::future::pending().await,
}
};
let mut unis = std::pin::pin!(err_only(run_unis(
session.clone(),
subscriber.clone(),
Some(peer_setup.clone()),
setup_read,
version
)));
let mut dispatch = std::pin::pin!(err_only(run_dispatch(
session.clone(),
publisher.clone(),
subscriber.clone(),
version
)));
let mut goaway = std::pin::pin!(err_only(goaway));
let mut setup = std::pin::pin!(setup);
let mut sub_ns_run = std::pin::pin!(err_only(async {
let mut prefixes = futures::stream::FuturesUnordered::new();
for prefix in sub_ns.subscribe_prefixes() {
let mut sub_ns = sub_ns.clone();
let sub_ns_session = sub_ns_session.clone();
prefixes.push(async move {
let stream = Stream::open(&sub_ns_session, version).await?;
if let Err(err) = sub_ns.run_subscribe_namespace(stream, prefix).await {
tracing::warn!(%err, "subscribe_namespace failed, continuing without");
}
Ok::<(), Error>(())
});
}
while let Some(result) = futures::StreamExt::next(&mut prefixes).await {
result?;
}
Ok(())
}));
kio::wait(|waiter| {
use std::task::Poll;
if let Poll::Ready(err) = waiter.poll_future(unis.as_mut()) {
return Poll::Ready(Err::<(), Error>(err));
}
if let Poll::Ready(err) = waiter.poll_future(dispatch.as_mut()) {
return Poll::Ready(Err(err));
}
if let Poll::Ready(err) = waiter.poll_future(goaway.as_mut()) {
return Poll::Ready(Err(err));
}
if waiter.poll_future(setup.as_mut()).is_ready() {
return Poll::Ready(Ok(()));
}
if task_set.poll(waiter).is_ready() {
return Poll::Ready(Ok(()));
}
if let Poll::Ready(err) = waiter.poll_future(sub_ns_run.as_mut()) {
return Poll::Ready(Err(err));
}
Poll::Pending
})
.await
}
};
match &res {
Err(Error::Transport(_)) => {
tracing::info!("session terminated");
session.close(1, "");
}
Err(err) => {
tracing::warn!(%err, "session error");
session.close(err.to_code(), err.to_string().as_ref());
}
_ => {
tracing::info!("session closed");
session.close(0, "");
}
}
res
}
.maybe_boxed();
Ok(driver)
}
pub struct PeerSetup<S: web_transport_trait::Session> {
pub stream: Reader<S::RecvStream, crate::Version>,
pub path: Option<String>,
pub cluster: cluster::Peer,
}
fn self_origin(publish: Option<&origin::Consumer>, subscribe: Option<&origin::Producer>) -> Origin {
publish
.map(|origin| **origin)
.or_else(|| subscribe.map(|origin| **origin))
.unwrap_or_else(Origin::random)
}
pub async fn accept_setup<S: web_transport_trait::Session>(
session: &S,
version: Version,
) -> Result<PeerSetup<S>, Error> {
let outer_version = crate::Version::Ietf(version);
loop {
let recv = session.accept_uni().await.map_err(Error::from_transport)?;
let mut reader: Reader<S::RecvStream, crate::Version> = Reader::new(recv, outer_version);
if reader.decode_peek::<u64>().await? != setup::SETUP_V17 {
reader.abort(&Error::UnexpectedStream);
continue;
}
let setup: setup::Setup = reader.decode().await?;
let mut bytes = setup.parameters.clone();
let params = ietf::Parameters::decode(&mut bytes, version)?;
let path = match params.get_bytes(ietf::ParameterBytes::Path) {
Some(bytes) => Some(
std::str::from_utf8(bytes)
.map_err(|_| Error::Decode(crate::DecodeError::InvalidValue))?
.to_owned(),
),
None => None,
};
let cluster = cluster::peer_from_setup(¶ms, version)?;
return Ok(PeerSetup {
stream: reader,
path,
cluster,
});
}
}
fn decode_peer_cluster(parameters: bytes::Bytes, version: Version) -> Result<cluster::Peer, crate::DecodeError> {
let mut bytes = parameters;
let params = ietf::Parameters::decode(&mut bytes, version)?;
cluster::peer_from_setup(¶ms, version)
}
async fn run_setup<S: web_transport_trait::Session>(
session: S,
version: Version,
path: Option<String>,
self_origin: Origin,
cost: Option<u64>,
) -> Result<(), Error> {
let outer_version = crate::Version::Ietf(version);
let send = session.open_uni().await.map_err(Error::from_transport)?;
let mut writer: Writer<S::SendStream, crate::Version> = Writer::new(send, outer_version);
let mut parameters = ietf::Parameters::default();
parameters.set_bytes(ietf::ParameterBytes::Implementation, b"moq-lite-rs".to_vec());
if let Some(path) = path {
parameters.set_bytes(ietf::ParameterBytes::Path, path.into_bytes());
}
cluster::peer_into_setup(&mut parameters, self_origin, cost, version);
let parameters = parameters.encode_bytes(version)?;
writer.encode(&setup::Setup { parameters }).await?;
session.closed().await;
writer.finish().ok();
Ok(())
}
async fn run_unis<S: web_transport_trait::Session>(
session: S,
subscriber: Subscriber<S>,
peer_setup: Option<cluster::PeerSetup>,
setup_read: bool,
version: Version,
) -> Result<(), Error> {
let outer_version = crate::Version::Ietf(version);
let mut tasks = TaskSet::owned();
let mut seen_setup = setup_read;
loop {
let recv = tasks.drive(session.accept_uni()).await.map_err(Error::from_transport)?;
let mut reader: Reader<S::RecvStream, crate::Version> = Reader::new(recv, outer_version);
let kind: u64 = tasks.drive(reader.decode_peek()).await?;
if kind == setup::SETUP_V17 {
if std::mem::replace(&mut seen_setup, true) {
return Err(Error::ProtocolViolation);
}
let peer_setup = peer_setup.clone();
let session = session.clone();
tasks.push(async move {
let msg = match reader.decode::<setup::Setup>().await {
Ok(msg) => msg,
Err(err) => {
tracing::warn!(%err, "setup decode error");
session.close(Error::ProtocolViolation.to_code(), "invalid setup");
return;
}
};
if let Some(peer_setup) = peer_setup {
let peer = match decode_peer_cluster(msg.parameters, version) {
Ok(peer) => peer,
Err(err) => {
tracing::warn!(%err, "setup parameter decode error");
session.close(Error::ProtocolViolation.to_code(), "invalid setup parameters");
return;
}
};
peer_setup.set(peer);
}
if let Err(err) = run_goaway(reader.with_version(version), version).await {
tracing::warn!(%err, "goaway error");
}
});
continue;
}
let mut sub = subscriber.clone();
tasks.push(async move {
let mut reader = reader.with_version(version);
if let Err(err) = run_uni_group(&mut sub, &mut reader).await {
tracing::debug!(%err, "uni stream error");
reader.abort(&err);
}
});
}
}
async fn run_uni_group<S: web_transport_trait::Session>(
subscriber: &mut Subscriber<S>,
stream: &mut Reader<S::RecvStream, Version>,
) -> Result<(), Error> {
let kind: u64 = stream.decode_peek().await?;
if kind <= 0xff && (kind & 0x90) == 0x10 {
return subscriber.recv_group(stream).await;
}
match kind {
FetchHeader::TYPE => Err(Error::Unsupported),
_ => Err(Error::UnexpectedStream),
}
}
async fn run_dispatch<S: web_transport_trait::Session>(
session: S,
publisher: Publisher<S>,
mut subscriber: Subscriber<S>,
version: Version,
) -> Result<(), Error> {
let peer = subscriber.peer().await;
let mut tasks = TaskSet::owned();
loop {
let mut stream = tasks.drive(Stream::accept(&session, version)).await?;
let header = tasks
.drive(async {
let id: u64 = stream.reader.decode().await?;
let size: u16 = stream.reader.decode().await?;
let data = stream.reader.read_exact(size as usize).await?;
Ok::<_, Error>((id, data))
})
.await;
let (id, data) = header?;
match id {
ietf::Subscribe::ID
| ietf::Fetch::ID
| ietf::SubscribeNamespace::ID
| ietf::SubscribeNamespaceLegacy::ID
| ietf::TrackStatus::ID => {
tasks.push(publisher.handle_stream(id, data, stream)?);
}
ietf::Publish::ID | ietf::PublishNamespace::ID => {
tasks.push(subscriber.handle_stream(id, data, stream, peer)?);
}
_ => {
tracing::warn!(id, "unexpected bidi stream type");
return Err(Error::UnexpectedStream);
}
}
}
}
async fn run_goaway<R: web_transport_trait::RecvStream>(
mut reader: Reader<R, Version>,
version: Version,
) -> Result<(), Error> {
let id: u64 = match reader.decode_maybe().await? {
Some(id) => id,
None => return Ok(()),
};
let size: u16 = reader.decode::<u16>().await?;
let mut data = reader.read_exact(size as usize).await?;
if id == ietf::GoAway::ID {
let msg = ietf::GoAway::decode_msg(&mut data, version)?;
tracing::debug!(message = ?msg, "received GOAWAY");
Err(Error::Unsupported)
} else {
Err(Error::UnexpectedMessage)
}
}
#[cfg(test)]
mod tests {
use super::*;
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 every_permitted_prefix_gets_its_own_subscribe_namespace() {
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"), crate::Path::new("mic")]))
.expect("scope the origin to two prefixes");
let gate = kio::Producer::new(true);
let session = crate::lite::test_transport::SinkSession::gated_bi(gate.consume());
let log = session.log.clone();
let driver = start(Config {
session,
setup: None,
request_id_max: None,
client: true,
publish: None,
subscribe: Some(scoped),
peer_origin: None,
cost: None,
version: Version::Draft18,
path: None,
peer_setup_stream: None,
peer_cluster: None,
})
.expect("start the session");
let _driver = tokio::spawn(driver);
for _ in 0..100 {
if occurrences(&log, b"cam") > 0 && occurrences(&log, b"mic") > 0 {
break;
}
tokio::time::sleep(std::time::Duration::from_millis(5)).await;
}
assert_eq!(occurrences(&log, b"cam"), 1, "one SUBSCRIBE_NAMESPACE for cam");
assert_eq!(occurrences(&log, b"mic"), 1, "one SUBSCRIBE_NAMESPACE for mic");
assert_eq!(occurrences(&log, b"rootns"), 0, "asked the peer for our local root");
}
#[tokio::test]
async fn announces_wait_for_a_subscribe_namespace() {
let origin = crate::origin::Info::new(crate::Origin::new(1).unwrap()).produce();
let _cam = origin
.create_broadcast("solo-cam", crate::broadcast::Route::new().with_announce(true))
.unwrap();
let gate = kio::Producer::new(true);
let session = crate::lite::test_transport::SinkSession::gated_bi(gate.consume());
let log = session.log.clone();
let driver = start(Config {
session,
setup: None,
request_id_max: None,
client: true,
publish: Some(origin.consume()),
subscribe: None,
peer_origin: None,
cost: None,
version: Version::Draft18,
path: None,
peer_setup_stream: None,
peer_cluster: None,
})
.expect("start the session");
let _driver = tokio::spawn(driver);
tokio::time::sleep(std::time::Duration::from_millis(50)).await;
assert_eq!(
occurrences(&log, b"solo-cam"),
0,
"PUBLISH_NAMESPACE must wait for a SUBSCRIBE_NAMESPACE"
);
}
#[test]
fn the_hop_id_comes_from_whichever_origin_the_caller_set() {
let ours = crate::Origin::new(42).unwrap();
let publish = crate::origin::Info::new(ours).produce();
assert_eq!(self_origin(Some(&publish.consume()), None), ours, "the publish half");
let subscribe = crate::origin::Info::new(ours).produce();
assert_eq!(self_origin(None, Some(&subscribe)), ours, "the subscribe half alone");
let publish = crate::origin::Info::new(ours).produce();
assert_eq!(self_origin(Some(&publish.consume()), Some(&subscribe)), ours);
}
}