use crate::origin;
use crate::{
Error, Hop, SessionError, StreamError,
coding::{Decode, DecodeError, 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, hidden, peer, solicit,
subscriber::is_protocol_violation,
};
pub struct Config<S: crate::transport::poll::Session> {
pub runtime: crate::time::Clock,
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_hop: Option<Hop>,
pub cost: Option<u64>,
pub version: Version,
pub path: Option<String>,
pub peer_setup_stream: Option<Reader<S::RecvStream, crate::Version>>,
pub peer_declared: Option<peer::Peer>,
}
pub fn start<S>(config: Config<S>) -> Result<(MaybeSendBox<'static, Result<(), Error>>, crate::goaway::Handle), Error>
where
S: crate::transport::poll::Boxable,
{
let Config {
runtime,
mut session,
setup,
request_id_max,
client,
publish,
subscribe,
peer_hop,
cost,
version,
path,
peer_setup_stream,
peer_declared,
} = config;
let (goaway_handle, goaway) = crate::goaway::Handle::new(!client);
let driver = async move {
let self_origin = self_origin(publish.as_ref(), subscribe.as_ref());
let publish = publish.unwrap_or_else(|| origin::Producer::empty(Hop::random()).consume());
let subscribe = subscribe.unwrap_or_else(|| origin::Producer::empty(Hop::random()));
let peer_setup = peer::PeerSetup::default();
let setup_read = peer_declared.is_some();
match peer_declared {
Some(declared) => peer_setup.set(declared),
None if !cluster::supported(version) => peer_setup.set(peer::Peer::default()),
None => {}
}
let res = match version {
Version::Draft14 | Version::Draft15 | Version::Draft16 => {
let Some(setup) = setup else {
let err = Error::ProtocolViolation;
session.close(SessionError::from(&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(
runtime.clone(),
adapter.clone(),
publish,
control.clone(),
peer_hop,
peer_setup.clone(),
version,
);
let (tasks, mut task_set) = TaskSet::new();
let subscriber = Subscriber::new(
runtime.clone(),
adapter.clone(),
subscribe,
control,
peer_hop,
peer_setup.clone(),
self_origin,
cost,
version,
tasks.clone(),
goaway.going_away.clone(),
);
{
let mut session = session.clone();
let adapter = adapter.clone();
let goaway = goaway.clone();
let runtime = runtime.clone();
tasks.push(async move {
let payload = kio::wait(|waiter| {
let mut cx = std::task::Context::from_waker(waiter.waker());
if session.poll_closed(&mut cx).is_ready() {
return std::task::Poll::Ready(None);
}
goaway.poll_triggered(waiter)
})
.await;
let Some(payload) = payload else {
return;
};
let timeout_ms = payload.timeout.map(|d| d.as_millis() as u64).unwrap_or(0);
adapter.send_goaway(&payload.uri, timeout_ms, version);
crate::goaway::enforce(&runtime, &mut session, payload.timeout).await;
});
}
drop(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, goaway.clone())));
let mut unis = std::pin::pin!(err_only(run_unis(
adapter.clone(),
subscriber.clone(),
None,
false,
version,
goaway,
)));
let mut dispatch = std::pin::pin!(err_only(run_dispatch(
dispatch_session,
publisher.clone(),
subscriber.clone(),
version
)));
let mut pub_ns_run = std::pin::pin!(err_only(publisher.clone().run_publish_namespaces()));
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 mut sub_ns_adapter = sub_ns_adapter;
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(&mut sub_ns_adapter.clone(), version).await?,
};
if let Err(err) = sub_ns.run_subscribe_namespace(stream, prefix).await {
if is_protocol_violation(&err) {
return Err(err);
}
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));
}
if let Poll::Ready(err) = waiter.poll_future(pub_ns_run.as_mut()) {
return Poll::Ready(Err(err));
}
Poll::Pending
})
.await
}
_ => {
let setup = {
let runtime = runtime.clone();
let session = session.clone();
let goaway = goaway.clone();
async move {
if let Err(err) = run_setup(runtime, session, version, path, self_origin, cost, goaway).await {
tracing::warn!(%err, "setup send error");
}
std::future::pending::<()>().await;
}
};
let control = Control::new(None, client);
let publisher = Publisher::new(
runtime.clone(),
session.clone(),
publish,
control.clone(),
peer_hop,
peer_setup.clone(),
version,
);
let (tasks, mut task_set) = TaskSet::new();
let subscriber = Subscriber::new(
runtime.clone(),
session.clone(),
subscribe,
control,
peer_hop,
peer_setup.clone(),
self_origin,
cost,
version,
tasks,
goaway.going_away.clone(),
);
let sub_ns_session = session.clone();
let sub_ns = subscriber.clone();
let goaway_recv = {
let goaway = goaway.clone();
async move {
match peer_setup_stream {
Some(reader) => run_goaway(reader.with_version(version), version, goaway).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,
goaway,
)));
let mut dispatch = std::pin::pin!(err_only(run_dispatch(
session.clone(),
publisher.clone(),
subscriber.clone(),
version
)));
let mut goaway_recv = std::pin::pin!(err_only(goaway_recv));
let mut setup = std::pin::pin!(setup);
let mut pub_ns_run = std::pin::pin!(err_only(publisher.clone().run_publish_namespaces()));
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 mut sub_ns_session = sub_ns_session;
let stream = Stream::open(&mut sub_ns_session, version).await?;
if let Err(err) = sub_ns.run_subscribe_namespace(stream, prefix).await {
if is_protocol_violation(&err) {
return Err(err);
}
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_recv.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));
}
if let Poll::Ready(err) = waiter.poll_future(pub_ns_run.as_mut()) {
return Poll::Ready(Err(err));
}
Poll::Pending
})
.await
}
};
match &res {
Err(Error::Transport(_)) => {
tracing::info!("session terminated");
session.close(SessionError::Internal.to_code(), "");
}
Err(err) => {
tracing::warn!(%err, "session error");
session.close(SessionError::from(err).to_code(), err.to_string().as_ref());
}
_ => {
tracing::info!("session closed");
session.close(SessionError::Cancel.to_code(), "");
}
}
res
}
.maybe_boxed();
Ok((driver, goaway_handle))
}
pub struct PeerSetup<S: crate::transport::poll::Session> {
pub stream: Reader<S::RecvStream, crate::Version>,
pub path: Option<String>,
pub declared: peer::Peer,
}
fn self_origin(publish: Option<&origin::Consumer>, subscribe: Option<&origin::Producer>) -> Hop {
publish
.map(|origin| origin.hop())
.or_else(|| subscribe.map(|origin| origin.hop()))
.unwrap_or_else(Hop::random)
}
pub async fn accept_setup<S: crate::transport::poll::Session>(
session: &mut 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 declared = peer_from_params(¶ms, version)?;
return Ok(PeerSetup {
stream: reader,
path,
declared,
});
}
}
fn decode_peer_setup(parameters: bytes::Bytes, version: Version) -> Result<peer::Peer, crate::DecodeError> {
let mut bytes = parameters;
let params = ietf::Parameters::decode(&mut bytes, version)?;
peer_from_params(¶ms, version)
}
fn peer_from_params(params: &ietf::Parameters, version: Version) -> Result<peer::Peer, crate::DecodeError> {
Ok(peer::Peer {
cluster: cluster::peer_from_setup(params, version)?,
solicit: solicit::from_setup(params, version)?,
hidden: hidden::from_setup(params, version),
})
}
async fn run_setup<S: crate::transport::poll::Session>(
runtime: crate::time::Clock,
mut session: S,
version: Version,
path: Option<String>,
self_origin: Hop,
cost: Option<u64>,
goaway: crate::goaway::Protocol,
) -> 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);
solicit::into_setup(&mut parameters, version);
hidden::into_setup(&mut parameters, version);
let parameters = parameters.encode_bytes(version)?;
writer.encode(&setup::Setup { parameters }).await?;
let payload = kio::wait(|waiter| {
let mut cx = std::task::Context::from_waker(waiter.waker());
if session.poll_closed(&mut cx).is_ready() {
return std::task::Poll::Ready(None);
}
goaway.poll_triggered(waiter)
})
.await;
if let Some(payload) = payload {
let timeout_ms = payload.timeout.map(|d| d.as_millis() as u64).unwrap_or(0);
let msg = ietf::GoAway {
new_session_uri: std::borrow::Cow::Borrowed(payload.uri.as_ref()),
timeout: timeout_ms,
};
let mut body = bytes::BytesMut::new();
msg.encode_msg(&mut body, version)?;
let size: u16 = body
.len()
.try_into()
.map_err(|_| Error::BoundsExceeded(crate::coding::BoundsExceeded))?;
let mut writer = writer.with_version(version);
writer.encode(&ietf::GoAway::ID).await?;
writer.encode(&size).await?;
writer.write_all(&mut std::io::Cursor::new(body)).await?;
crate::goaway::enforce(&runtime, &mut session, payload.timeout).await;
session.closed().await;
writer.finish().ok();
} else {
writer.finish().ok();
}
Ok(())
}
async fn run_unis<S>(
mut session: S,
subscriber: Subscriber<S>,
peer_setup: Option<peer::PeerSetup>,
setup_read: bool,
version: Version,
goaway: crate::goaway::Protocol,
) -> Result<(), Error>
where
S: crate::transport::poll::Boxable,
{
let outer_version = crate::Version::Ietf(version);
let mut tasks = TaskSet::owned();
let mut seen_setup = setup_read;
loop {
let recv = tasks
.drive(|waiter| {
let mut cx = std::task::Context::from_waker(waiter.waker());
session.poll_accept_uni(&mut cx)
})
.await
.map_err(Error::from_transport)?;
let mut reader: Reader<S::RecvStream, crate::Version> = Reader::new(recv, outer_version);
let kind: u64 = match tasks
.drive(|waiter| {
let mut cx = std::task::Context::from_waker(waiter.waker());
reader.poll_decode_peek(&mut cx)
})
.await
{
Ok(kind) => kind,
Err(err @ (Error::Cancel | Error::Stream(_) | Error::Remote(_) | Error::Decode(DecodeError::Short))) => {
tracing::debug!(%err, "dropping uni stream that died before its type");
continue;
}
Err(err) => return Err(err),
};
if kind == setup::SETUP_V17 {
if std::mem::replace(&mut seen_setup, true) {
return Err(Error::ProtocolViolation);
}
let peer_setup = peer_setup.clone();
let mut session = session.clone();
let goaway = goaway.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(SessionError::ProtocolViolation.to_code(), "invalid setup");
return;
}
};
if let Some(peer_setup) = peer_setup {
let peer = match decode_peer_setup(msg.parameters, version) {
Ok(peer) => peer,
Err(err) => {
tracing::warn!(%err, "setup parameter decode error");
session.close(SessionError::ProtocolViolation.to_code(), "invalid setup parameters");
return;
}
};
peer_setup.set(peer);
}
if let Err(err) = run_goaway(reader.with_version(version), version, goaway).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");
let reset = match StreamError::from(&err) {
StreamError::Session(_) => StreamError::Internal,
reset => reset,
};
reader.abort(reset);
}
});
}
}
async fn run_uni_group<S>(
subscriber: &mut Subscriber<S>,
stream: &mut Reader<S::RecvStream, Version>,
) -> Result<(), Error>
where
S: crate::transport::poll::Boxable,
{
let kind: u64 = stream.decode_peek().await?;
if kind <= 0xff && (kind & 0x90) == 0x10 {
return subscriber.recv_group(stream).await;
}
match kind {
FetchHeader::TYPE => subscriber.recv_fill(stream).await,
_ => Err(Error::UnexpectedStream),
}
}
async fn run_dispatch<S>(
session: S,
publisher: Publisher<S>,
mut subscriber: Subscriber<S>,
version: Version,
) -> Result<(), Error>
where
S: crate::transport::poll::Boxable,
{
let peer = subscriber.peer().await;
let declared = subscriber.solicit().await;
let mut tasks = TaskSet::owned();
let mut accept = session.clone();
loop {
let mut stream = tasks
.drive(|waiter| {
let mut cx = std::task::Context::from_waker(waiter.waker());
Stream::poll_accept(&mut accept, version, &mut cx)
})
.await?;
let mut hdr_id: Option<u64> = None;
let mut hdr_size: Option<u16> = None;
let header = tasks
.drive(|waiter| {
let mut cx = std::task::Context::from_waker(waiter.waker());
let id = match hdr_id {
Some(id) => id,
None => *hdr_id.insert(std::task::ready!(stream.reader.poll_decode(&mut cx))?),
};
let size = match hdr_size {
Some(size) => size,
None => *hdr_size.insert(std::task::ready!(stream.reader.poll_decode(&mut cx))?),
};
let data = std::task::ready!(stream.reader.poll_read_exact(&mut cx, size as usize))?;
std::task::Poll::Ready(Ok::<_, Error>((id, data)))
})
.await;
let (id, data) = match header {
Ok(header) => header,
Err(err @ (Error::Cancel | Error::Stream(_) | Error::Remote(_) | Error::Decode(DecodeError::Short))) => {
tracing::debug!(%err, "dropping bidi stream that died before its header");
continue;
}
Err(err) => return Err(err),
};
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, declared)?);
}
_ => {
tracing::warn!(id, "unexpected bidi stream type");
return Err(Error::UnexpectedStream);
}
}
}
}
async fn run_goaway<R: crate::transport::poll::RecvStream>(
mut reader: Reader<R, Version>,
version: Version,
goaway: crate::goaway::Protocol,
) -> 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 {
return Err(Error::UnexpectedMessage);
}
let msg = ietf::GoAway::decode_msg(&mut data, version)?;
tracing::info!(message = ?msg, "received GOAWAY");
let timeout = (msg.timeout > 0).then(|| std::time::Duration::from_millis(msg.timeout));
goaway.record(crate::goaway::Goaway {
uri: msg.new_session_uri.into_owned(),
timeout,
})?;
loop {
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)?;
let timeout = (msg.timeout > 0).then(|| std::time::Duration::from_millis(msg.timeout));
goaway.record(crate::goaway::Goaway {
uri: msg.new_session_uri.into_owned(),
timeout,
})?;
continue;
}
tracing::warn!(id, "unexpected message after GOAWAY on the SETUP stream; ignoring");
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::model::ProduceTest;
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()
}
async fn namespace_without_hop_path(version: Version) -> 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("cam"),
cluster: None,
})
.await
.unwrap();
let writes = log.writes.lock().unwrap();
writes.clone()
}
#[tokio::test]
async fn a_namespace_without_a_hop_path_closes_the_session() {
const VERSION: Version = Version::Draft19;
tokio::time::pause();
let origin = crate::origin::Config::new(crate::Hop::new(1).unwrap()).produce();
let session = crate::lite::test_transport::ScriptedSession::new(namespace_without_hop_path(VERSION).await);
let log = session.log.clone();
let (driver, _goaway) = start(Config {
runtime: crate::time::Clock::tokio(),
session,
setup: None,
request_id_max: None,
client: true,
publish: None,
subscribe: Some(origin),
peer_hop: None,
cost: None,
version: VERSION,
path: None,
peer_setup_stream: None,
peer_declared: Some(peer::Peer {
cluster: cluster::Peer {
hop: Some(crate::Hop::new(2).unwrap()),
cost: None,
},
..Default::default()
}),
})
.expect("start the session");
let err = tokio::time::timeout(std::time::Duration::from_secs(10), driver)
.await
.expect("the session ended rather than carrying on")
.expect_err("a malformed NAMESPACE fails the session");
assert!(is_protocol_violation(&err), "not treated as the peer's fault: {err}");
let (code, reason) = log.closes().first().cloned().expect("the session was closed");
assert_eq!(code, SessionError::ProtocolViolation.to_code());
assert_eq!(reason, err.to_string());
}
#[tokio::test]
async fn every_permitted_prefix_gets_its_own_subscribe_namespace() {
let origin = crate::origin::Config::new(crate::Hop::new(1).unwrap()).produce();
let scope: crate::Patterns = ["cam", "mic"]
.into_iter()
.map(|prefix| crate::Pattern::subtree(prefix).unwrap())
.collect();
let scoped = origin
.scope("rootns", &scope)
.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, _goaway) = start(Config {
runtime: crate::time::Clock::tokio(),
session,
setup: None,
request_id_max: None,
client: true,
publish: None,
subscribe: Some(scoped),
peer_hop: None,
cost: None,
version: Version::Draft18,
path: None,
peer_setup_stream: None,
peer_declared: Some(peer::Peer::default()),
})
.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");
}
const ANNOUNCE_TURNS: usize = 100;
async fn announce_occurrences(peer_declared: Option<peer::Peer>) -> usize {
let origin = crate::origin::Config::new(crate::Hop::new(1).unwrap()).produce();
let _cam = origin.announce("solo-cam", crate::origin::Route::default()).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, _goaway) = start(Config {
runtime: crate::time::Clock::tokio(),
session,
setup: None,
request_id_max: None,
client: true,
publish: Some(origin.consume()),
subscribe: None,
peer_hop: None,
cost: None,
version: Version::Draft18,
path: None,
peer_setup_stream: None,
peer_declared,
})
.expect("start the session");
let _driver = tokio::spawn(driver);
for _ in 0..ANNOUNCE_TURNS {
if occurrences(&log, b"solo-cam") > 0 {
break;
}
tokio::time::sleep(std::time::Duration::from_millis(1)).await;
}
occurrences(&log, b"solo-cam")
}
#[tokio::test(start_paused = true)]
async fn no_announce_before_the_peer_setup() {
assert_eq!(
announce_occurrences(None).await,
0,
"advertised before knowing what the peer wants"
);
}
#[tokio::test(start_paused = true)]
async fn a_peer_requiring_solicitation_is_not_told_unasked() {
let declared = peer::Peer {
solicit: Some(true),
..Default::default()
};
assert_eq!(
announce_occurrences(Some(declared)).await,
0,
"PUBLISH_NAMESPACE despite the peer asking to be told on request"
);
}
#[tokio::test(start_paused = true)]
async fn a_peer_declaring_nothing_is_told_unasked() {
assert_eq!(
announce_occurrences(Some(peer::Peer::default())).await,
1,
"no unsolicited PUBLISH_NAMESPACE"
);
}
#[test]
fn the_hop_id_comes_from_whichever_origin_the_caller_set() {
let ours = crate::Hop::new(42).unwrap();
let publish = crate::origin::Config::new(ours).produce();
assert_eq!(self_origin(Some(&publish.consume()), None), ours, "the publish half");
let subscribe = crate::origin::Config::new(ours).produce();
assert_eq!(self_origin(None, Some(&subscribe)), ours, "the subscribe half alone");
let publish = crate::origin::Config::new(ours).produce();
assert_eq!(self_origin(Some(&publish.consume()), Some(&subscribe)), ours);
}
async fn a_dead_incoming_stream_is_not_fatal(session: crate::lite::test_transport::DeadStreamSession) {
const VERSION: Version = Version::Draft19;
tokio::time::pause();
let origin = crate::origin::Config::new(crate::Hop::new(1).unwrap()).produce();
let log = session.log.clone();
let (driver, _goaway) = start(Config {
runtime: crate::time::Clock::tokio(),
session,
setup: None,
request_id_max: None,
client: true,
publish: None,
subscribe: Some(origin),
peer_hop: None,
cost: None,
version: VERSION,
path: None,
peer_setup_stream: None,
peer_declared: Some(peer::Peer::default()),
})
.expect("start the session");
tokio::time::timeout(std::time::Duration::from_secs(10), driver)
.await
.expect_err("the session ended over one dead stream");
assert_eq!(log.closes(), vec![], "nothing may close the transport");
}
#[tokio::test]
async fn a_uni_stream_dead_before_its_type_does_not_end_the_session() {
a_dead_incoming_stream_is_not_fatal(crate::lite::test_transport::DeadStreamSession::unis(1)).await;
}
#[tokio::test]
async fn a_bidi_stream_dead_before_its_header_does_not_end_the_session() {
a_dead_incoming_stream_is_not_fatal(crate::lite::test_transport::DeadStreamSession::bis(1)).await;
}
async fn subgroup_header(version: Version, track_alias: u64) -> 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::GroupHeader {
track_alias,
group_id: 0,
sub_group_id: 0,
publisher_priority: 128,
flags: ietf::GroupFlags::default(),
})
.await
.unwrap();
let writes = log.writes.lock().unwrap();
writes.clone()
}
async fn dispatch_uni(payload: Vec<u8>, retired_alias: Option<u64>) -> crate::lite::test_transport::Log {
const VERSION: Version = Version::Draft19;
let origin = crate::origin::Config::new(crate::Hop::new(1).unwrap()).produce();
let session = crate::lite::test_transport::ScriptedSession::new(Vec::new()).with_incoming_unis(vec![payload]);
let log = session.log.clone();
let (tasks, _task_set) = TaskSet::new();
let peer_setup = peer::PeerSetup::default();
let subscriber = Subscriber::new(
crate::time::Clock::tokio(),
session.clone(),
origin,
Control::new(None, true),
None,
peer_setup.clone(),
crate::Hop::new(1).unwrap(),
None,
VERSION,
tasks,
Default::default(),
);
if let Some(alias) = retired_alias {
subscriber.retire_alias(alias);
}
let (_goaway, goaway) = crate::goaway::Handle::new(false);
let mut unis = std::pin::pin!(run_unis(session, subscriber, Some(peer_setup), false, VERSION, goaway));
for _ in 0..100 {
if let std::task::Poll::Ready(result) = futures::poll!(unis.as_mut()) {
panic!("the dispatch loop ended over one rejected stream: {result:?}");
}
if !log.stops().is_empty() {
break;
}
tokio::time::sleep(std::time::Duration::from_millis(1)).await;
}
log
}
#[tokio::test(start_paused = true)]
async fn a_group_for_a_retired_alias_is_stopped_with_cancelled() {
let log = dispatch_uni(subgroup_header(Version::Draft19, 7).await, Some(7)).await;
assert_eq!(
log.stops(),
vec![crate::ietf::error::CANCELLED],
"the group stream must be stopped with the cancelled code",
);
assert_eq!(log.closes(), vec![], "one dropped group may not close the session");
}
#[tokio::test(start_paused = true)]
async fn unknown_uni_type_does_not_claim_the_session_closed() {
let log = dispatch_uni(vec![0], None).await;
assert_eq!(log.stops(), vec![crate::ietf::error::INTERNAL_ERROR]);
assert!(log.closes().is_empty());
}
async fn publish_namespace_then_two_withdrawals(version: Version) -> 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::PublishNamespace::ID).await.unwrap();
writer
.encode(&ietf::PublishNamespace {
request_id: RequestId(1),
track_namespace: crate::Path::new("room/host"),
cluster: None,
})
.await
.unwrap();
for _ in 0..2 {
writer.encode(&ietf::PublishNamespaceDone::ID).await.unwrap();
writer
.encode(&ietf::PublishNamespaceDone {
track_namespace: crate::Path::new("room/host"),
request_id: RequestId(0),
})
.await
.unwrap();
}
let writes = log.writes.lock().unwrap();
writes.clone()
}
async fn publish_namespace_ok(request_id: RequestId) -> 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::Draft14,
);
writer.encode(&ietf::PublishNamespaceOk::ID).await.unwrap();
writer.encode(&ietf::PublishNamespaceOk { request_id }).await.unwrap();
let writes = log.writes.lock().unwrap();
writes.clone()
}
#[tokio::test(start_paused = true)]
async fn a_repeated_publish_namespace_done_does_not_end_the_session() {
const VERSION: Version = Version::Draft14;
let origin = crate::origin::Config::new(crate::Hop::new(1).unwrap()).produce();
let consumer = origin.consume();
let mut session =
crate::lite::test_transport::ScriptedSession::new(publish_namespace_then_two_withdrawals(VERSION).await);
let log = session.log.clone();
let setup = Stream::open(&mut session, VERSION)
.await
.expect("open the control stream");
let (driver, _goaway) = start(Config {
runtime: crate::time::Clock::tokio(),
session,
setup: Some(setup),
request_id_max: None,
client: true,
publish: None,
subscribe: Some(origin),
peer_hop: None,
cost: None,
version: VERSION,
path: None,
peer_setup_stream: None,
peer_declared: None,
})
.expect("start the session");
let driver = tokio::spawn(driver);
let accepted = publish_namespace_ok(RequestId(1)).await;
for _ in 0..ANNOUNCE_TURNS {
if occurrences(&log, &accepted) > 0 && consumer.get_broadcast("room/host").is_none() {
break;
}
tokio::time::sleep(std::time::Duration::from_millis(1)).await;
}
assert_eq!(occurrences(&log, &accepted), 1, "the advertisement was not accepted");
assert!(
consumer.get_broadcast("room/host").is_none(),
"the withdrawal did not reach the request that advertised it"
);
assert!(!driver.is_finished(), "the repeated withdrawal ended the session");
assert!(log.closes().is_empty(), "closed the session: {:?}", log.closes());
}
}