use crate::origin;
use crate::{
Error, Origin, bandwidth,
coding::{Reader, Stream, Writer},
lite::SessionInfo,
util::{MaybeBoxedExt, MaybeSendBox, TaskSet, err_only},
};
use std::task::Poll;
use super::{
Connecting, DataType, PeerSetup, Publisher, PublisherConfig, Setup, Subscriber, SubscriberConfig, Version,
};
pub(crate) struct SessionStart {
pub recv_bandwidth: Option<bandwidth::Consumer>,
pub connecting: Connecting,
pub driver: MaybeSendBox<'static, Result<(), Error>>,
}
pub async fn accept_setup<S: web_transport_trait::Session>(session: &S, version: Version) -> Result<Setup, Error> {
loop {
let stream = session.accept_uni().await.map_err(Error::from_transport)?;
let mut reader = Reader::new(stream, version);
match reader.decode::<DataType>().await? {
DataType::Setup => return reader.decode::<Setup>().await,
_ => reader.abort(&Error::UnexpectedStream),
}
}
}
#[allow(clippy::too_many_arguments)]
pub fn start<S: web_transport_trait::Session>(
session: S,
setup_stream: Option<Stream<S, Version>>,
publish: Option<origin::Consumer>,
subscribe: Option<origin::Producer>,
version: Version,
our_setup: Setup,
peer_setup: Option<Setup>,
) -> Result<SessionStart, Error> {
let recv_bw = bandwidth::Producer::new();
let recv_bw_consumer = match version {
Version::Lite01 | Version::Lite02 => None,
_ => Some(recv_bw.consume()),
};
let recv_bw_for_sub = match version {
Version::Lite01 | Version::Lite02 => None,
_ => Some(recv_bw),
};
let (connecting_producer, connecting) = Connecting::new();
let sub_connecting = if matches!(version, Version::Lite01 | Version::Lite02) || version.has_announce_ok() {
Some(connecting_producer)
} else {
None
};
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_slot = PeerSetup::default();
if let Some(setup) = peer_setup {
peer_setup_slot.set(setup);
}
let peer_setup = peer_setup_slot;
let (tasks, task_set) = TaskSet::new();
let our_cost = our_setup.cost;
if version.has_setup_stream() {
let session = session.clone();
tasks.push(async move {
if let Err(err) = send_setup(&session, our_setup, version).await {
tracing::debug!(%err, "failed to send setup");
}
});
}
let publisher = Publisher::new(PublisherConfig {
session: session.clone(),
origin: publish,
version,
});
let subscriber = Subscriber::new(SubscriberConfig {
session: session.clone(),
origin: subscribe,
recv_bandwidth: recv_bw_for_sub,
version,
peer_setup,
cost: our_cost,
tasks,
});
let driver = async move {
let res = {
let mut session = std::pin::pin!(err_only(run_session(setup_stream)));
let mut publisher = std::pin::pin!(publisher.run());
let mut subscriber = std::pin::pin!(subscriber.run(sub_connecting, task_set));
kio::wait(|waiter| {
if let Poll::Ready(err) = waiter.poll_future(session.as_mut()) {
return Poll::Ready(Err(err));
}
if let Poll::Ready(res) = waiter.poll_future(publisher.as_mut()) {
return Poll::Ready(res);
}
if let Poll::Ready(res) = waiter.poll_future(subscriber.as_mut()) {
return Poll::Ready(res);
}
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(SessionStart {
recv_bandwidth: recv_bw_consumer,
connecting,
driver,
})
}
async fn send_setup<S: web_transport_trait::Session>(session: &S, setup: Setup, version: Version) -> Result<(), Error> {
let stream = session.open_uni().await.map_err(Error::from_transport)?;
let mut writer = Writer::new(stream, version);
writer.encode(&super::DataType::Setup).await?;
writer.encode(&setup).await?;
writer.finish()?;
writer.closed().await
}
async fn run_session<S: web_transport_trait::Session>(stream: Option<Stream<S, Version>>) -> Result<(), Error> {
if let Some(mut stream) = stream {
while let Some(_info) = stream.reader.decode_maybe::<SessionInfo>().await? {}
return Err(Error::Cancel);
}
Ok(())
}