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};
#[allow(clippy::too_many_arguments)]
pub fn start<S: web_transport_trait::Session>(
session: S,
setup: Option<Stream<S, Version>>,
request_id_max: Option<RequestId>,
client: bool,
publish: Option<origin::Consumer>,
subscribe: Option<origin::Producer>,
version: Version,
path: Option<String>,
peer_setup: Option<Reader<S::RecvStream, crate::Version>>,
) -> Result<MaybeSendBox<'static, Result<(), Error>>, Error> {
let driver = async move {
let publish = publish.unwrap_or_else(|| origin::Producer::empty(Origin::random()).consume());
let subscribe = subscribe.unwrap_or_else(|| origin::Producer::empty(Origin::random()));
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(), version);
let (tasks, mut task_set) = TaskSet::new();
let subscriber = Subscriber::new(adapter.clone(), subscribe, control, version, tasks);
let dispatch_session = adapter.clone();
let mut 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(), version)));
let mut dispatch = std::pin::pin!(err_only(run_dispatch(
dispatch_session,
publisher.clone(),
subscriber.clone(),
version
)));
let mut publisher_run = std::pin::pin!(err_only(publisher.run()));
let mut sub_ns_run = std::pin::pin!(err_only(async {
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).await {
tracing::warn!(%err, "subscribe_namespace failed, continuing without");
}
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 let Poll::Ready(err) = waiter.poll_future(publisher_run.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).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(), version);
let (tasks, mut task_set) = TaskSet::new();
let subscriber = Subscriber::new(session.clone(), subscribe, control, version, tasks);
let sub_ns_session = session.clone();
let mut sub_ns = subscriber.clone();
let goaway = async move {
match peer_setup {
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(), version)));
let mut dispatch = std::pin::pin!(err_only(run_dispatch(
session.clone(),
publisher.clone(),
subscriber.clone(),
version
)));
let mut publisher_run = std::pin::pin!(err_only(publisher.run()));
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 stream = Stream::open(&sub_ns_session, version).await?;
if let Err(err) = sub_ns.run_subscribe_namespace(stream).await {
tracing::warn!(%err, "subscribe_namespace failed, continuing without");
}
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(publisher_run.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 async fn accept_setup<S: web_transport_trait::Session>(
session: &S,
version: Version,
) -> Result<(Reader<S::RecvStream, crate::Version>, Option<String>), 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;
let path = match ietf::Parameters::decode(&mut bytes, version)?.get_bytes(ietf::ParameterBytes::Path) {
Some(bytes) => Some(
std::str::from_utf8(bytes)
.map_err(|_| Error::Decode(crate::DecodeError::InvalidValue))?
.to_owned(),
),
None => None,
};
return Ok((reader, path));
}
}
async fn run_setup<S: web_transport_trait::Session>(
session: S,
version: Version,
path: Option<String>,
) -> 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());
}
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>,
version: Version,
) -> Result<(), Error> {
let outer_version = crate::Version::Ietf(version);
let mut tasks = TaskSet::owned();
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 {
tasks.push(async move {
if let Err(err) = reader.decode::<setup::Setup>().await {
tracing::warn!(%err, "setup decode error");
return;
}
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 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)?);
}
_ => {
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)
}
}