use std::collections::{HashMap, VecDeque};
use std::future::Future;
use std::sync::atomic::AtomicU64;
use std::sync::Arc;
use std::time::Duration;
use futures_util::{SinkExt, StreamExt};
use net_backend_protocol::{CloseCode, WsServerFrame};
use serde_json::Value;
use tokio::sync::{broadcast, mpsc, watch};
use tokio::time::Instant;
use tokio_tungstenite::tungstenite::protocol::frame::coding::CloseCode as WireCloseCode;
use tokio_tungstenite::tungstenite::protocol::CloseFrame;
use tokio_tungstenite::tungstenite::Message;
use super::link::{self, Link};
use super::{Shared, WsConnection, WsEvent, WsPush, WsSettings, WsState};
use crate::runtime::RuntimeThread;
use crate::{Client, Error};
const GOODBYE: Duration = Duration::from_secs(1);
pub(crate) struct Pending {
pub(crate) id: u64,
pub(crate) text: String,
pub(crate) deadline: Instant,
pub(crate) answer: Box<dyn FnOnce(Result<Value, Error>) + Send>,
}
pub(crate) enum Command {
Request(Pending),
Close,
}
enum End {
App,
Closed(CloseCode, String),
Lost(Error),
}
struct Task {
client: Client,
settings: WsSettings,
commands: mpsc::UnboundedReceiver<Command>,
state: watch::Sender<WsState>,
events: broadcast::Sender<WsEvent>,
pushes: broadcast::Sender<Arc<WsPush>>,
waiting: VecDeque<Pending>,
in_flight: HashMap<u64, Pending>,
}
pub(crate) async fn connect(client: Client, settings: WsSettings, runtime: Option<Arc<RuntimeThread>>) -> Result<WsConnection, Error> {
let link = open_with_refresh(&client, &settings).await?;
let (commands, receiver) = mpsc::unbounded_channel();
let (state_sender, state) = watch::channel(WsState::Connected);
let (events, events_template) = broadcast::channel(settings.event_buffer);
let (pushes, pushes_template) = broadcast::channel(settings.push_buffer);
let task = Task {
client,
settings: settings.clone(),
commands: receiver,
state: state_sender,
events,
pushes,
waiting: VecDeque::new(),
in_flight: HashMap::new(),
};
let handle = crate::runtime::current()?;
handle.spawn(task.run(link));
let (events, pushes) = (std::sync::Mutex::new(events_template), std::sync::Mutex::new(pushes_template));
Ok(WsConnection::new(Shared { commands, next_id: AtomicU64::new(1), state, events, pushes, settings, _runtime: runtime }))
}
async fn open_with_refresh(client: &Client, settings: &WsSettings) -> Result<Link, Error> {
let deadline = deadline_after(settings.connect_timeout);
let (token, _) = client.access_token(deadline).await?;
match link::open(client, settings, token.expose(), deadline).await {
Err(error) if link::wants_refresh(&error) => {
tracing::debug!("net_backend_client: the WebSocket handshake refused the token ({error}); one refresh");
let deadline = deadline_after(settings.connect_timeout);
client.refresh_shared(deadline).await?;
let (token, _) = client.access_token(deadline).await?;
link::open(client, settings, token.expose(), deadline).await
}
other => other,
}
}
fn deadline_after(timeout: Duration) -> Instant {
let now = Instant::now();
now.checked_add(timeout).unwrap_or(now)
}
impl Task {
async fn run(mut self, first: Link) {
let mut link = Some(first);
let mut reconnected = false;
let mut attempt: u32 = 0;
let mut refreshed_after_4001 = false;
loop {
if let Some(current) = link.take() {
self.set_state(WsState::Connected);
self.event(WsEvent::Connected { reconnected });
let up_since = Instant::now();
let end = self.serve(current).await;
self.fail_in_flight();
if up_since.elapsed() >= self.settings.reconnect.as_ref().map_or(Duration::MAX, |r| r.stable_after) {
attempt = 0;
refreshed_after_4001 = false;
}
let error = match end {
End::App => return self.finish(None),
End::Closed(code, reason) if code == CloseCode::UNAUTHORIZED && !refreshed_after_4001 => {
refreshed_after_4001 = true;
let closed = Error::Closed { code, reason };
match self.client.refresh_shared(deadline_after(self.settings.connect_timeout)).await {
Ok(_) => {}
Err(error @ Error::SessionEnded { .. }) => return self.finish(Some(error)),
Err(_) => return self.finish(Some(closed)),
}
match self.attempt_now().await {
Some(Ok(new)) => {
link = Some(new);
reconnected = true;
continue;
}
Some(Err(error)) if link::is_permanent(&error) => return self.finish(Some(error)),
Some(Err(error)) => error,
None => return self.finish(None),
}
}
End::Closed(code, reason) if code.is_permanent() => return self.finish(Some(Error::Closed { code, reason })),
End::Closed(code, reason) => Error::Closed { code, reason },
End::Lost(error) => error,
};
tracing::debug!("net_backend_client: the WebSocket link ended ({error})");
let Some(reconnect) = self.settings.reconnect.clone() else { return self.finish(Some(error)) };
let mut last = error;
loop {
attempt = attempt.saturating_add(1);
if !reconnect.may_retry(attempt) {
return self.finish(Some(last));
}
let mut delay = reconnect.delay(attempt, crate::tls::random_u64());
if let Some(wait) = last.retry_after() {
delay = delay.max(wait.min(crate::MAX_TIMEOUT));
}
self.set_state(WsState::Reconnecting { attempt, retry_in: delay });
self.event(WsEvent::Reconnecting { attempt, retry_in: delay, error: last.clone() });
if self.wait_queueing(tokio::time::sleep(delay)).await.is_none() {
return self.finish(None);
}
match self.attempt_now().await {
Some(Ok(new)) => {
link = Some(new);
reconnected = true;
break;
}
Some(Err(error)) if link::is_permanent(&error) => return self.finish(Some(error)),
Some(Err(error)) => last = error,
None => return self.finish(None),
}
}
}
}
}
async fn attempt_now(&mut self) -> Option<Result<Link, Error>> {
self.set_state(WsState::Connecting);
let client = self.client.clone();
let settings = self.settings.clone();
self.wait_queueing(async move { open_with_refresh(&client, &settings).await }).await
}
async fn wait_queueing<T>(&mut self, future: impl Future<Output = T>) -> Option<T> {
tokio::pin!(future);
loop {
let next = self.next_deadline();
tokio::select! {
biased;
command = self.commands.recv() => match command {
None | Some(Command::Close) => return None,
Some(Command::Request(pending)) => self.queue(pending),
},
output = &mut future => return Some(output),
() = sleep_until(next) => self.expire(),
}
}
}
fn queue(&mut self, pending: Pending) {
if self.waiting.len().saturating_add(self.in_flight.len()) >= self.settings.max_pending {
let limit = self.settings.max_pending;
(pending.answer)(Err(Error::invalid(format!("more than {limit} WebSocket requests are waiting or running"))));
return;
}
if Instant::now() >= pending.deadline {
(pending.answer)(Err(Error::timeout("not sent: the request timed out before the connection was up", Some(false))));
return;
}
self.waiting.push_back(pending);
}
fn next_deadline(&self) -> Option<Instant> {
self.waiting.iter().chain(self.in_flight.values()).map(|p| p.deadline).min()
}
fn expire(&mut self) {
let now = Instant::now();
let mut kept = VecDeque::with_capacity(self.waiting.len());
for pending in self.waiting.drain(..) {
if now >= pending.deadline {
(pending.answer)(Err(Error::timeout("not sent: the WebSocket was not connected before the request timed out", Some(false))));
} else {
kept.push_back(pending);
}
}
self.waiting = kept;
let overdue: Vec<u64> = self.in_flight.iter().filter(|(_, p)| now >= p.deadline).map(|(id, _)| *id).collect();
for id in overdue {
if let Some(pending) = self.in_flight.remove(&id) {
(pending.answer)(Err(Error::timeout("the server did not answer in time", Some(true))));
}
}
}
async fn serve(&mut self, mut link: Link) -> End {
while let Some(pending) = self.waiting.pop_front() {
if let Err(end) = self.write(&mut link, pending).await {
return end;
}
}
let mut ping = tokio::time::interval_at(Instant::now() + self.settings.ping_interval, self.settings.ping_interval);
ping.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Delay);
let mut last_seen = Instant::now();
loop {
let next = self.next_deadline();
tokio::select! {
biased;
command = self.commands.recv() => match command {
None | Some(Command::Close) => {
let frame = CloseFrame { code: WireCloseCode::Normal, reason: "".into() };
let _ = tokio::time::timeout(GOODBYE, link.close(Some(frame))).await;
return End::App;
}
Some(Command::Request(pending)) => {
if self.waiting.len().saturating_add(self.in_flight.len()) >= self.settings.max_pending || Instant::now() >= pending.deadline {
self.queue(pending);
} else if let Err(end) = self.write(&mut link, pending).await {
return end;
}
}
},
message = link.next() => {
last_seen = Instant::now();
match message {
Some(Ok(Message::Text(text))) => self.frame(text.as_str()),
Some(Ok(Message::Close(frame))) => {
let (code, reason) = frame.map_or((CloseCode(1005), String::new()), |f| (CloseCode(u16::from(f.code)), f.reason.as_str().to_string()));
let _ = tokio::time::timeout(GOODBYE, async { while link.next().await.is_some() {} }).await;
return End::Closed(code, reason);
}
Some(Ok(_)) => {}
Some(Err(error)) => return End::Lost(link::map_ws_error(error)),
None => return End::Lost(Error::network("the connection closed without a close frame", None)),
}
}
_ = ping.tick() => {
if last_seen.elapsed() >= self.settings.dead_after {
return End::Lost(Error::timeout(format!("no frame from the server for {:?} (heartbeat)", self.settings.dead_after), None));
}
match tokio::time::timeout(self.settings.dead_after, link.send(Message::Ping(Default::default()))).await {
Ok(Ok(())) => {}
Ok(Err(error)) => return End::Lost(link::map_ws_error(error)),
Err(_) => return End::Lost(Error::timeout("the heartbeat could not be written", None)),
}
}
() = sleep_until(next) => self.expire(),
}
}
}
async fn write(&mut self, link: &mut Link, pending: Pending) -> Result<(), End> {
let message = Message::text(pending.text.clone());
match tokio::time::timeout(self.settings.dead_after, link.send(message)).await {
Ok(Ok(())) => {
self.in_flight.insert(pending.id, pending);
Ok(())
}
Ok(Err(error)) => {
let error = link::map_ws_error(error);
(pending.answer)(Err(Error::disconnected(format!("the request could not be written: {error}"), None)));
Err(End::Lost(error))
}
Err(_) => {
(pending.answer)(Err(Error::disconnected("the request could not be written in time", None)));
Err(End::Lost(Error::timeout("a request could not be written (the server stopped reading)", None)))
}
}
}
fn frame(&mut self, text: &str) {
match WsServerFrame::parse(text) {
Ok(WsServerFrame::Response(response)) => match self.in_flight.remove(&response.id) {
Some(pending) => (pending.answer)(response.result.map_err(|error| Error::api(None, error, None))),
None => tracing::debug!("net_backend_client: an answer to an unknown request id {} (timed out already?)", response.id),
},
Ok(WsServerFrame::Push(push)) => {
let _ = self.pushes.send(Arc::new(WsPush { kind: push.kind, data: push.data }));
}
Ok(WsServerFrame::AuthOk(_)) => {}
Ok(WsServerFrame::AuthFailed(error)) => tracing::debug!("net_backend_client: auth.failed on an open socket ({error})"),
Ok(_) | Err(_) => tracing::debug!("net_backend_client: a frame that is not part of the protocol was ignored"),
}
}
fn fail_in_flight(&mut self) {
for (_, pending) in self.in_flight.drain() {
(pending.answer)(Err(Error::disconnected("the connection was lost after the request was sent", Some(true))));
}
}
fn set_state(&self, state: WsState) {
self.state.send_replace(state);
}
fn event(&self, event: WsEvent) {
let _ = self.events.send(event);
}
fn finish(mut self, error: Option<Error>) {
self.fail_in_flight();
let reason = match &error {
None => "the connection was closed by the app".to_string(),
Some(error) => format!("the connection closed: {error}"),
};
for pending in self.waiting.drain(..) {
(pending.answer)(Err(Error::disconnected(reason.clone(), Some(false))));
}
self.commands.close();
while let Ok(command) = self.commands.try_recv() {
if let Command::Request(pending) = command {
(pending.answer)(Err(Error::disconnected(reason.clone(), Some(false))));
}
}
if let Some(error) = &error {
tracing::info!("net_backend_client: the WebSocket closed for good ({error})");
}
self.set_state(WsState::Closed);
self.event(WsEvent::Closed { error });
}
}
async fn sleep_until(deadline: Option<Instant>) {
match deadline {
Some(deadline) => tokio::time::sleep_until(deadline).await,
None => std::future::pending().await,
}
}