use crate::origin;
use crate::{
Error, Hop, SessionError, bandwidth,
coding::{Reader, Stream, Writer},
lite::SessionInfo,
};
use std::task::{Context, Poll, ready};
use super::{
DataType, PeerSetup, Publisher, PublisherConfig, Setup, Subscriber, SubscriberConfig, SubscriberDriver, Version,
};
pub(crate) struct SessionStart<S: crate::transport::poll::Session> {
pub recv_bandwidth: Option<bandwidth::Consumer>,
pub driver: Driver<S>,
pub goaway: crate::goaway::Handle,
}
pub async fn accept_setup<S: crate::transport::poll::Session>(
session: &mut 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),
}
}
}
pub struct Config<S: crate::transport::poll::Session> {
pub runtime: crate::time::Clock,
pub session: S,
pub setup_stream: Option<Stream<S, Version>>,
pub publish: Option<origin::Consumer>,
pub subscribe: Option<origin::Producer>,
pub peer_hop: Option<Hop>,
pub version: Version,
pub our_setup: Setup,
pub peer_setup: Option<Setup>,
}
pub fn start<S>(config: Config<S>) -> Result<SessionStart<S>, Error>
where
S: crate::transport::poll::Session,
{
let Config {
runtime,
session,
setup_stream,
publish,
subscribe,
peer_hop,
version,
mut our_setup,
peer_setup,
} = config;
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),
};
if our_setup.hop.is_none() {
our_setup.hop = publish
.as_ref()
.map(|origin| origin.hop())
.or_else(|| subscribe.as_ref().map(|origin| origin.hop()))
.filter(|hop| hop.id() != 0);
}
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_slot = PeerSetup::default();
if let Some(setup) = peer_setup {
peer_setup_slot.set(setup);
}
let peer_setup = peer_setup_slot;
let (goaway_handle, goaway) = crate::goaway::Handle::new(true);
let our_cost = our_setup.cost;
let publisher = Publisher::new(PublisherConfig {
runtime: runtime.clone(),
session: session.clone(),
origin: publish,
version,
peer_setup: peer_setup.clone(),
goaway: goaway.clone(),
peer_hop,
});
let subscriber = Subscriber::new(SubscriberConfig {
runtime: runtime.clone(),
session: session.clone(),
origin: subscribe,
recv_bandwidth: recv_bw_for_sub,
version,
peer_setup,
peer_hop,
cost: our_cost,
going_away: goaway.going_away.clone(),
});
let driver = Driver {
setup: version
.has_setup_stream()
.then(|| SendSetup::new(session.clone(), our_setup, version)),
goaway: Some(SendGoaway::new(runtime, session.clone(), goaway, version)),
session_stream: setup_stream,
publisher,
subscriber: SubscriberDriver::new(subscriber),
session,
};
Ok(SessionStart {
recv_bandwidth: recv_bw_consumer,
driver,
goaway: goaway_handle,
})
}
pub(crate) struct Driver<S: crate::transport::poll::Session> {
setup: Option<SendSetup<S>>,
goaway: Option<SendGoaway<S>>,
session_stream: Option<Stream<S, Version>>,
publisher: Publisher<S>,
subscriber: SubscriberDriver<S>,
session: S,
}
impl<S> Driver<S>
where
S: crate::transport::poll::Session,
{
pub(crate) fn poll(&mut self, waiter: &kio::Waiter) -> Poll<Result<(), Error>> {
let res = std::task::ready!(self.poll_protocol(waiter));
match &res {
Err(Error::Transport(_)) => {
tracing::info!("session terminated");
self.session.close(SessionError::Internal.to_code(), "");
}
Err(err) => {
tracing::warn!(%err, "session error");
self.session
.close(SessionError::from(err).to_code(), err.to_string().as_ref());
}
_ => {
tracing::info!("session closed");
self.session.close(SessionError::Cancel.to_code(), "");
}
}
Poll::Ready(res)
}
fn poll_protocol(&mut self, waiter: &kio::Waiter) -> Poll<Result<(), Error>> {
let mut cx = Context::from_waker(waiter.waker());
if let Some(setup) = &mut self.setup
&& setup.poll(&mut cx).is_ready()
{
self.setup = None;
}
if let Some(goaway) = &mut self.goaway
&& goaway.poll(waiter).is_ready()
{
self.goaway = None;
}
if let Some(stream) = &mut self.session_stream
&& let Poll::Ready(err) = poll_session_stream(stream, &mut cx)
{
return Poll::Ready(Err(err));
}
if let Poll::Ready(res) = self.publisher.poll(waiter) {
return Poll::Ready(res);
}
if let Poll::Ready(res) = self.subscriber.poll(waiter) {
return Poll::Ready(res);
}
Poll::Pending
}
}
fn poll_session_stream<S: crate::transport::poll::Session>(
stream: &mut Stream<S, Version>,
cx: &mut Context<'_>,
) -> Poll<Error> {
loop {
match ready!(stream.reader.poll_decode_maybe::<SessionInfo>(cx)) {
Ok(Some(_info)) => {}
Ok(None) => return Poll::Ready(Error::Cancel),
Err(err) => return Poll::Ready(err),
}
}
}
struct SendSetup<S: crate::transport::poll::Session> {
version: Version,
state: SendSetupState<S>,
}
enum SendSetupState<S: crate::transport::poll::Session> {
Open {
session: S,
setup: Box<Setup>,
},
Send {
writer: Writer<S::SendStream, Version>,
finished: bool,
},
Done,
}
impl<S: crate::transport::poll::Session> SendSetup<S> {
fn new(session: S, setup: Setup, version: Version) -> Self {
Self {
version,
state: SendSetupState::Open {
session,
setup: Box::new(setup),
},
}
}
fn poll(&mut self, cx: &mut Context<'_>) -> Poll<()> {
match ready!(self.poll_send(cx)) {
Ok(()) => {}
Err(err) => tracing::debug!(%err, "failed to send setup"),
}
self.state = SendSetupState::Done;
Poll::Ready(())
}
fn poll_send(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), Error>> {
loop {
match &mut self.state {
SendSetupState::Open { session, setup } => {
let stream = ready!(session.poll_open_uni(cx)).map_err(Error::from_transport)?;
let mut writer = Writer::new(stream, self.version);
writer.buffer(&super::DataType::Setup)?;
writer.buffer(&**setup)?;
self.state = SendSetupState::Send {
writer,
finished: false,
};
}
SendSetupState::Send { writer, finished } => {
if !*finished {
ready!(writer.poll_flush(cx))?;
writer.finish()?;
*finished = true;
}
return writer.poll_closed(cx);
}
SendSetupState::Done => return Poll::Ready(Ok(())),
}
}
}
}
struct SendGoaway<S: crate::transport::poll::Session> {
version: Version,
runtime: crate::time::Clock,
goaway: crate::goaway::Protocol,
closed: S,
session: S,
state: SendGoawayState<S>,
}
enum SendGoawayState<S: crate::transport::poll::Session> {
Waiting,
Open { payload: crate::goaway::Goaway },
Send {
stream: Stream<S, Version>,
timeout: Option<std::time::Duration>,
finished: bool,
},
Enforce(crate::goaway::Enforce<S>),
}
impl<S: crate::transport::poll::Session> SendGoaway<S> {
fn new(runtime: crate::time::Clock, session: S, goaway: crate::goaway::Protocol, version: Version) -> Self {
Self {
version,
runtime,
goaway,
closed: session.clone(),
session,
state: SendGoawayState::Waiting,
}
}
fn enforce_after(&mut self, err: Error, timeout: Option<std::time::Duration>) {
tracing::warn!(%err, "failed to send goaway");
self.state = SendGoawayState::Enforce(crate::goaway::Enforce::new(
&self.runtime,
self.session.clone(),
timeout,
));
}
fn poll(&mut self, waiter: &kio::Waiter) -> Poll<()> {
let mut cx = Context::from_waker(waiter.waker());
loop {
match &mut self.state {
SendGoawayState::Waiting => {
if self.closed.poll_closed(&mut cx).is_ready() {
return Poll::Ready(());
}
let Some(payload) = ready!(self.goaway.poll_triggered(waiter)) else {
return Poll::Ready(());
};
self.state = if self.version.has_goaway() {
SendGoawayState::Open { payload }
} else {
SendGoawayState::Enforce(crate::goaway::Enforce::new(
&self.runtime,
self.session.clone(),
payload.timeout,
))
};
}
SendGoawayState::Open { payload } => {
let timeout = payload.timeout;
let mut stream = match ready!(Stream::poll_open(&mut self.session, self.version, &mut cx)) {
Ok(stream) => stream,
Err(err) => {
self.enforce_after(err, timeout);
continue;
}
};
let msg = super::Goaway {
uri: std::borrow::Cow::Borrowed(payload.uri.as_str()),
};
if let Err(err) = stream
.writer
.buffer(&super::ControlType::Goaway)
.and_then(|()| stream.writer.buffer(&msg))
{
self.enforce_after(err, timeout);
continue;
}
self.state = SendGoawayState::Send {
stream,
timeout,
finished: false,
};
}
SendGoawayState::Send {
stream,
timeout,
finished,
} => {
let timeout = *timeout;
let res = if !*finished {
match ready!(stream.writer.poll_flush(&mut cx)) {
Ok(()) => {
*finished = true;
stream.writer.finish()
}
Err(err) => Err(err),
}
} else {
Ok(())
};
let res = match res {
Ok(()) => ready!(stream.writer.poll_closed(&mut cx)),
Err(err) => Err(err),
};
match res {
Ok(()) => {
self.state = SendGoawayState::Enforce(crate::goaway::Enforce::new(
&self.runtime,
self.session.clone(),
timeout,
));
}
Err(err) => self.enforce_after(err, timeout),
}
}
SendGoawayState::Enforce(enforce) => return enforce.poll(waiter),
}
}
}
}