use std::collections::HashMap;
use tokio::sync::{mpsc, oneshot, watch};
use crate::client::ClientId;
use crate::client::events::{
ClosedReason, Connected, Resubscribed, ServerInfo, SessionEvent, SubscriptionId,
};
use crate::client::message::MessageOutcome;
use crate::client::updates::SubscriptionEvent;
use crate::error::{Error, ServerError};
use crate::session::{
ControlOutcome, ControlTarget, SessionEvent as WireEvent, SessionHandle, SubscriptionError,
SubscriptionEvent as WireSubscriptionEvent, SubscriptionKey, SubscriptionOperation,
};
#[derive(Debug)]
pub(crate) enum RouterCommand {
Register {
id: SubscriptionId,
events: mpsc::Sender<SubscriptionEvent>,
},
Unregister {
id: SubscriptionId,
},
StreamDropped {
id: SubscriptionId,
},
}
fn control_failure(outcome: ControlOutcome) -> Option<SessionEvent> {
match outcome {
ControlOutcome::Rejected { cause } => {
Some(SessionEvent::RequestRejected(ServerError::from(cause)))
}
ControlOutcome::NotSent { reason } => {
tracing::warn!(reason, "a control request never reached the server");
Some(SessionEvent::RequestNotSent { reason })
}
ControlOutcome::Accepted => None,
}
}
#[derive(Debug)]
struct Registered {
events: mpsc::Sender<SubscriptionEvent>,
}
#[derive(Debug)]
pub(crate) struct Router {
client: ClientId,
session_events: mpsc::Receiver<WireEvent>,
commands: mpsc::UnboundedReceiver<RouterCommand>,
session_out: Option<mpsc::Sender<SessionEvent>>,
unsubscribe: mpsc::UnboundedSender<SubscriptionKey>,
ready: Option<oneshot::Sender<Result<Box<Connected>, Error>>>,
stop: watch::Receiver<bool>,
subscriptions: HashMap<SubscriptionKey, Registered>,
}
impl Router {
pub(crate) fn new(
client: ClientId,
session_events: mpsc::Receiver<WireEvent>,
commands: mpsc::UnboundedReceiver<RouterCommand>,
session_out: mpsc::Sender<SessionEvent>,
unsubscribe: mpsc::UnboundedSender<SubscriptionKey>,
ready: oneshot::Sender<Result<Box<Connected>, Error>>,
stop: watch::Receiver<bool>,
) -> Self {
Self {
client,
session_events,
commands,
session_out: Some(session_out),
unsubscribe,
ready: Some(ready),
stop,
subscriptions: HashMap::new(),
}
}
pub(crate) async fn run(mut self) {
let mut taking_commands = true;
loop {
let done = tokio::select! {
biased;
command = self.commands.recv(), if taking_commands => match command {
Some(command) => {
self.on_command(command);
false
}
None => {
taking_commands = false;
false
}
},
() = stopped(&mut self.stop) => true,
event = self.session_events.recv() => match event {
Some(event) => self.on_session_event(event).await,
None => true,
},
};
if done {
break;
}
}
self.fail_ready(Error::Disconnected);
tracing::debug!("event router stopped");
}
fn on_command(&mut self, command: RouterCommand) {
match command {
RouterCommand::Register { id, events } => {
self.subscriptions.insert(id.key(), Registered { events });
}
RouterCommand::Unregister { id } => {
self.subscriptions.remove(&id.key());
tracing::debug!(id = id.get(), "registration rolled back");
}
RouterCommand::StreamDropped { id } => {
self.subscriptions.remove(&id.key());
if self.unsubscribe.send(id.key()).is_err() {
tracing::debug!(id = id.get(), "cannot unsubscribe: the client has stopped");
}
}
}
}
async fn on_session_event(&mut self, event: WireEvent) -> bool {
match event {
WireEvent::Bound(info) => {
let connected = Box::new(Connected::from(*info));
self.signal_ready(connected.clone());
self.emit(SessionEvent::Connected(connected)).await;
false
}
WireEvent::Recovered(outcome) => {
self.emit(SessionEvent::Recovered(outcome.into())).await;
false
}
WireEvent::Resubscribed(entries) => {
let client = self.client;
let entries: Vec<Resubscribed> = entries
.into_iter()
.map(|entry| Resubscribed::from_entry(client, entry))
.collect();
self.emit(SessionEvent::Resubscribed(entries)).await;
false
}
WireEvent::Unbound { reason, retry_in } => {
self.emit(SessionEvent::Disconnected {
reason: reason.into(),
retry_in,
})
.await;
false
}
WireEvent::Message {
sequence,
prog,
result,
..
} => {
self.emit(SessionEvent::Message(Box::new(MessageOutcome::new(
sequence, prog, result,
))))
.await;
false
}
WireEvent::Subscription { key, outcome, .. } => {
self.route(key, *outcome).await;
false
}
WireEvent::ControlResponse {
target, outcome, ..
} => {
self.on_control_response(target, outcome).await;
false
}
WireEvent::ServerInfo(announcement) => {
self.emit(SessionEvent::ServerInfo(ServerInfo::from(announcement)))
.await;
false
}
WireEvent::Unparsed { line, error } => {
tracing::debug!(%error, "surfacing a line this client cannot parse");
self.emit(SessionEvent::Unrecognized { line }).await;
false
}
WireEvent::Closed(closed) => {
let reason = ClosedReason::from(closed);
self.fail_ready(reason.clone().into_error());
self.emit(SessionEvent::Closed(reason)).await;
true
}
}
}
async fn on_control_response(&mut self, target: ControlTarget, outcome: ControlOutcome) {
let ControlTarget::Subscription { key, operation } = target else {
if let Some(event) = control_failure(outcome) {
self.emit(event).await;
}
return;
};
match (operation, outcome) {
(SubscriptionOperation::Subscribe, ControlOutcome::Rejected { cause }) => {
let cause = ServerError::from(cause);
tracing::warn!(key = key.get(), code = cause.code(), "subscription refused");
match self.subscriptions.remove(&key) {
Some(entry) => {
deliver(
&entry.events,
SubscriptionEvent::Rejected(cause),
&mut self.stop,
)
.await;
}
None => self.emit(SessionEvent::RequestRejected(cause)).await,
}
}
(SubscriptionOperation::Subscribe, ControlOutcome::NotSent { reason }) => {
tracing::warn!(
key = key.get(),
reason,
"subscription request could not be sent; it will be retried at the next bind"
);
self.notify(key, SubscriptionEvent::Deferred { reason })
.await;
}
(SubscriptionOperation::Unsubscribe | SubscriptionOperation::Reconfigure, outcome) => {
match outcome {
ControlOutcome::Accepted => {}
ControlOutcome::Rejected { cause } => {
let cause = ServerError::from(cause);
tracing::warn!(
key = key.get(),
?operation,
code = cause.code(),
"a subscription request was refused; the subscription is unchanged"
);
self.emit(SessionEvent::RequestRejected(cause)).await;
}
ControlOutcome::NotSent { reason } => {
tracing::warn!(
key = key.get(),
?operation,
reason,
"a subscription request never reached the server; \
the subscription is unchanged"
);
self.emit(SessionEvent::RequestNotSent { reason }).await;
}
}
}
(SubscriptionOperation::Subscribe, ControlOutcome::Accepted) => {}
}
}
async fn notify(&mut self, key: SubscriptionKey, event: SubscriptionEvent) {
let Some(entry) = self.subscriptions.get(&key) else {
return;
};
let sender = entry.events.clone();
deliver(&sender, event, &mut self.stop).await;
}
async fn route(
&mut self,
key: SubscriptionKey,
outcome: Result<WireSubscriptionEvent, SubscriptionError>,
) {
if !self.subscriptions.contains_key(&key) {
tracing::debug!(
key = key.get(),
"a notification arrived for a subscription with no stream"
);
return;
}
let event = match outcome {
Ok(event) => SubscriptionEvent::from_wire(event),
Err(error) => {
tracing::warn!(key = key.get(), %error, "a notification did not decode");
SubscriptionEvent::Undecodable {
detail: error.to_string(),
}
}
};
let terminal = event.is_terminal();
let Some(entry) = self.subscriptions.get(&key) else {
return;
};
let sender = entry.events.clone();
deliver(&sender, event, &mut self.stop).await;
if terminal {
self.subscriptions.remove(&key);
}
}
async fn emit(&mut self, event: SessionEvent) {
let Some(sender) = self.session_out.clone() else {
return;
};
if deliver(&sender, event, &mut self.stop).await == Delivery::ReceiverGone {
tracing::debug!("session event stream dropped; no longer forwarding session events");
self.session_out = None;
}
}
fn signal_ready(&mut self, connected: Box<Connected>) {
if let Some(ready) = self.ready.take() {
let _ = ready.send(Ok(connected));
}
}
fn fail_ready(&mut self, error: Error) {
if let Some(ready) = self.ready.take() {
let _ = ready.send(Err(error));
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum Delivery {
Delivered,
ReceiverGone,
Stopped,
}
async fn deliver<T>(
sender: &mpsc::Sender<T>,
event: T,
stop: &mut watch::Receiver<bool>,
) -> Delivery {
tokio::select! {
biased;
permit = sender.reserve() => match permit {
Ok(permit) => {
permit.send(event);
Delivery::Delivered
}
Err(_) => Delivery::ReceiverGone,
},
() = stopped(stop) => {
tracing::debug!("a stop was ordered while a caller's stream was full");
Delivery::Stopped
}
}
}
async fn stopped(stop: &mut watch::Receiver<bool>) {
loop {
if *stop.borrow_and_update() {
return;
}
if stop.changed().await.is_err() {
return;
}
}
}
pub(crate) async fn issue_unsubscriptions(
handle: SessionHandle,
mut keys: mpsc::UnboundedReceiver<SubscriptionKey>,
mut shutdown: mpsc::Receiver<()>,
mut stop: watch::Receiver<bool>,
) {
loop {
let key = tokio::select! {
key = keys.recv() => match key {
Some(key) => key,
None => break,
},
_ = shutdown.recv() => break,
() = stopped(&mut stop) => break,
};
let sent = tokio::select! {
biased;
result = handle.unsubscribe(key) => result.is_ok(),
() = stopped(&mut stop) => break,
};
if !sent {
tracing::debug!(key = key.get(), "the session stopped before unsubscribing");
break;
}
}
drop(handle);
tracing::debug!("unsubscription task stopped");
}