mod link;
mod task;
use std::fmt;
use std::marker::PhantomData;
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::{Arc, Mutex, PoisonError};
use std::time::Duration;
use net_backend_protocol::{ServerPush, WsCall, WsRequestFrame};
use serde_json::Value;
use tokio::sync::{broadcast, mpsc, watch};
use tokio::time::Instant;
use crate::runtime::RuntimeThread;
use crate::{Client, Error, Reply, MAX_TIMEOUT};
pub(crate) use task::{Command, Pending};
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq, Hash)]
#[non_exhaustive]
pub enum WsAuthMode {
#[default]
Header,
FirstMessage,
Both,
}
#[derive(Clone, Debug)]
pub struct Reconnect {
base: Duration,
cap: Duration,
max_attempts: Option<u32>,
stable_after: Duration,
jitter: bool,
}
impl Default for Reconnect {
fn default() -> Self {
Self { base: Duration::from_millis(500), cap: Duration::from_secs(30), max_attempts: None, stable_after: Duration::from_secs(10), jitter: true }
}
}
impl Reconnect {
pub fn with_base(mut self, base: Duration) -> Self {
self.base = base.clamp(Duration::from_millis(1), MAX_TIMEOUT);
self
}
pub fn with_cap(mut self, cap: Duration) -> Self {
self.cap = cap.min(MAX_TIMEOUT);
self
}
pub fn with_max_attempts(mut self, max: Option<u32>) -> Self {
self.max_attempts = max;
self
}
pub fn with_stable_after(mut self, stable_after: Duration) -> Self {
self.stable_after = stable_after.min(MAX_TIMEOUT);
self
}
pub fn with_jitter(mut self, jitter: bool) -> Self {
self.jitter = jitter;
self
}
pub fn delay_bound(&self, attempt: u32) -> Duration {
let factor = 2u32.checked_pow(attempt.saturating_sub(1).min(30)).unwrap_or(u32::MAX);
self.base.saturating_mul(factor).min(self.cap.max(self.base))
}
pub(crate) fn delay(&self, attempt: u32, random: u64) -> Duration {
let bound = self.delay_bound(attempt);
if !self.jitter {
return bound;
}
let nanos = u64::try_from(bound.as_nanos()).unwrap_or(u64::MAX);
Duration::from_nanos(random % nanos.saturating_add(1))
}
pub(crate) fn may_retry(&self, attempt: u32) -> bool {
self.max_attempts.is_none_or(|max| attempt <= max)
}
}
#[derive(Clone, Debug)]
pub struct WsSettings {
pub(crate) auth: WsAuthMode,
pub(crate) connect_timeout: Duration,
pub(crate) request_timeout: Duration,
pub(crate) ping_interval: Duration,
pub(crate) dead_after: Duration,
pub(crate) reconnect: Option<Reconnect>,
pub(crate) push_buffer: usize,
pub(crate) event_buffer: usize,
pub(crate) max_message_bytes: usize,
pub(crate) max_pending: usize,
}
impl Default for WsSettings {
fn default() -> Self {
Self {
auth: WsAuthMode::Header,
connect_timeout: Duration::from_secs(10),
request_timeout: Duration::from_secs(10),
ping_interval: Duration::from_secs(15),
dead_after: Duration::from_secs(45),
reconnect: Some(Reconnect::default()),
push_buffer: 256,
event_buffer: 64,
max_message_bytes: net_backend_protocol::envelope::MAX_MESSAGE_BYTES,
max_pending: 256,
}
}
}
impl WsSettings {
pub fn with_auth(mut self, auth: WsAuthMode) -> Self {
self.auth = auth;
self
}
pub fn with_connect_timeout(mut self, timeout: Duration) -> Self {
self.connect_timeout = timeout.clamp(Duration::from_millis(100), MAX_TIMEOUT);
self
}
pub fn with_request_timeout(mut self, timeout: Duration) -> Self {
self.request_timeout = timeout.clamp(Duration::from_millis(1), MAX_TIMEOUT);
self
}
pub fn with_heartbeat(mut self, interval: Duration, dead_after: Duration) -> Self {
self.ping_interval = interval.clamp(Duration::from_millis(10), MAX_TIMEOUT);
self.dead_after = dead_after.clamp(self.ping_interval, MAX_TIMEOUT);
self
}
pub fn with_reconnect(mut self, reconnect: Reconnect) -> Self {
self.reconnect = Some(reconnect);
self
}
pub fn without_reconnect(mut self) -> Self {
self.reconnect = None;
self
}
pub fn with_push_buffer(mut self, pushes: usize) -> Self {
self.push_buffer = pushes.max(1);
self
}
pub fn with_event_buffer(mut self, events: usize) -> Self {
self.event_buffer = events.max(1);
self
}
pub fn with_max_message_bytes(mut self, bytes: usize) -> Self {
self.max_message_bytes = bytes.max(1024);
self
}
pub fn with_max_pending(mut self, requests: usize) -> Self {
self.max_pending = requests.max(1);
self
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
#[non_exhaustive]
pub enum WsState {
Connected,
Reconnecting {
attempt: u32,
retry_in: Duration,
},
Connecting,
Closed,
}
#[derive(Clone, Debug)]
#[non_exhaustive]
pub enum WsEvent {
Connected {
reconnected: bool,
},
Reconnecting {
attempt: u32,
retry_in: Duration,
error: Error,
},
Closed {
error: Option<Error>,
},
}
#[derive(Clone, Debug, PartialEq)]
#[non_exhaustive]
pub struct WsPush {
pub kind: String,
pub data: Value,
}
impl WsPush {
pub fn decode<P: ServerPush>(&self) -> Result<P, Error> {
P::deserialize(&self.data).map_err(|e| Error::Decode { status: None, message: e.to_string() })
}
}
pub(crate) struct Shared {
pub(crate) commands: mpsc::UnboundedSender<Command>,
pub(crate) next_id: AtomicU64,
pub(crate) state: watch::Receiver<WsState>,
pub(crate) events: Mutex<broadcast::Receiver<WsEvent>>,
pub(crate) pushes: Mutex<broadcast::Receiver<Arc<WsPush>>>,
pub(crate) settings: WsSettings,
pub(crate) _runtime: Option<Arc<RuntimeThread>>,
}
#[derive(Clone)]
pub struct WsConnection {
shared: Arc<Shared>,
}
impl fmt::Debug for WsConnection {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("WsConnection").field("state", &self.state()).finish()
}
}
impl Client {
pub async fn connect_ws(&self, settings: WsSettings) -> Result<WsConnection, Error> {
crate::runtime::current()?;
task::connect(self.clone(), settings, None).await
}
}
impl crate::blocking::Client {
pub fn connect_ws(&self, settings: WsSettings) -> Result<WsConnection, Error> {
let client = self.inner.clone();
let runtime = Arc::clone(&self.runtime);
self.runtime.block(async move { task::connect(client, settings, Some(runtime)).await })
}
}
impl WsConnection {
pub(crate) fn new(shared: Shared) -> Self {
Self { shared: Arc::new(shared) }
}
pub fn request<C: WsCall>(&self, call: &C) -> Reply<C::Response>
where
C::Response: Send + 'static,
{
self.request_with_timeout(call, self.shared.settings.request_timeout)
}
pub fn request_with_timeout<C: WsCall>(&self, call: &C, timeout: Duration) -> Reply<C::Response>
where
C::Response: Send + 'static,
{
let id = self.shared.next_id.fetch_add(1, Ordering::Relaxed);
let text = match serde_json::to_string(&WsRequestFrame::new(id, C::KIND, call)) {
Ok(text) => text,
Err(e) => return Reply::ready(Err(Error::invalid(format!("the request cannot be encoded: {e}")))),
};
let (sender, reply) = Reply::channel();
let answer = Box::new(move |result: Result<Value, Error>| {
let typed =
result.and_then(|value| serde_json::from_value::<C::Response>(value).map_err(|e| Error::Decode { status: None, message: e.to_string() }));
let _ = sender.send(typed);
});
self.enqueue(id, text, timeout, answer);
reply
}
pub fn request_raw(&self, kind: &str, data: Value) -> Reply<Value> {
let id = self.shared.next_id.fetch_add(1, Ordering::Relaxed);
let text = match serde_json::to_string(&WsRequestFrame::new(id, kind, data)) {
Ok(text) => text,
Err(e) => return Reply::ready(Err(Error::invalid(format!("the request cannot be encoded: {e}")))),
};
let (sender, reply) = Reply::channel();
let answer = Box::new(move |result: Result<Value, Error>| {
let _ = sender.send(result);
});
self.enqueue(id, text, self.shared.settings.request_timeout, answer);
reply
}
fn enqueue(&self, id: u64, text: String, timeout: Duration, answer: Box<dyn FnOnce(Result<Value, Error>) + Send>) {
let limit = self.shared.settings.max_message_bytes;
if text.len() > limit {
answer(Err(Error::RequestTooLarge { limit: limit as u64, size: text.len() as u64 }));
return;
}
let now = Instant::now();
let deadline = now.checked_add(timeout.clamp(Duration::from_millis(1), MAX_TIMEOUT)).unwrap_or(now);
let pending = Pending { id, text, deadline, answer };
if let Err(mpsc::error::SendError(Command::Request(pending))) = self.shared.commands.send(Command::Request(pending)) {
(pending.answer)(Err(Error::disconnected("the connection is closed", Some(false))));
}
}
pub fn subscribe<P: ServerPush>(&self) -> PushStream<P> {
PushStream { receiver: self.shared.pushes.lock().unwrap_or_else(PoisonError::into_inner).resubscribe(), kind: Some(P::KIND), _type: PhantomData }
}
pub fn pushes(&self) -> PushStream<WsPush> {
PushStream { receiver: self.shared.pushes.lock().unwrap_or_else(PoisonError::into_inner).resubscribe(), kind: None, _type: PhantomData }
}
pub fn events(&self) -> WsEvents {
WsEvents { receiver: self.shared.events.lock().unwrap_or_else(PoisonError::into_inner).resubscribe() }
}
pub fn state(&self) -> WsState {
*self.shared.state.borrow()
}
pub fn is_closed(&self) -> bool {
self.state() == WsState::Closed
}
pub fn close(&self) {
let _ = self.shared.commands.send(Command::Close);
}
pub async fn closed(&self) {
let mut state = self.shared.state.clone();
let _ = state.wait_for(|s| *s == WsState::Closed).await;
}
}
pub struct PushStream<T> {
receiver: broadcast::Receiver<Arc<WsPush>>,
kind: Option<&'static str>,
_type: PhantomData<fn() -> T>,
}
impl<T> fmt::Debug for PushStream<T> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("PushStream").field("kind", &self.kind).finish()
}
}
pub trait PushItem: Sized {
#[doc(hidden)]
fn from_push(push: &WsPush, kind: Option<&str>) -> Option<Result<Self, Error>>;
}
impl<P: ServerPush> PushItem for P {
fn from_push(push: &WsPush, kind: Option<&str>) -> Option<Result<Self, Error>> {
(kind == Some(push.kind.as_str())).then(|| push.decode::<P>())
}
}
impl<T: PushItem> PushStream<T> {
pub async fn next(&mut self) -> Option<Result<T, Error>> {
loop {
match self.receiver.recv().await {
Ok(push) => {
if let Some(item) = T::from_push(&push, self.kind) {
return Some(item);
}
}
Err(broadcast::error::RecvError::Lagged(missed)) => return Some(Err(Error::Lagged { missed })),
Err(broadcast::error::RecvError::Closed) => return None,
}
}
}
pub fn try_next(&mut self) -> Option<Result<T, Error>> {
loop {
match self.receiver.try_recv() {
Ok(push) => {
if let Some(item) = T::from_push(&push, self.kind) {
return Some(item);
}
}
Err(broadcast::error::TryRecvError::Lagged(missed)) => return Some(Err(Error::Lagged { missed })),
Err(_) => return None,
}
}
}
}
impl PushItem for WsPush {
fn from_push(push: &WsPush, _kind: Option<&str>) -> Option<Result<Self, Error>> {
Some(Ok(push.clone()))
}
}
pub struct WsEvents {
receiver: broadcast::Receiver<WsEvent>,
}
impl fmt::Debug for WsEvents {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str("WsEvents")
}
}
impl WsEvents {
pub async fn next(&mut self) -> Option<WsEvent> {
loop {
match self.receiver.recv().await {
Ok(event) => return Some(event),
Err(broadcast::error::RecvError::Lagged(_)) => {}
Err(broadcast::error::RecvError::Closed) => return None,
}
}
}
pub fn try_next(&mut self) -> Option<WsEvent> {
loop {
match self.receiver.try_recv() {
Ok(event) => return Some(event),
Err(broadcast::error::TryRecvError::Lagged(_)) => {}
Err(_) => return None,
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn backoff_is_bounded() {
let policy = Reconnect::default().with_base(Duration::from_millis(100)).with_cap(Duration::from_secs(2)).with_jitter(false);
assert_eq!(policy.delay_bound(1), Duration::from_millis(100));
assert_eq!(policy.delay_bound(3), Duration::from_millis(400));
assert_eq!(policy.delay_bound(40), Duration::from_secs(2));
assert_eq!(policy.delay(u32::MAX, u64::MAX), Duration::from_secs(2));
let jitter = Reconnect::default().with_base(Duration::MAX).with_cap(Duration::MAX);
assert!(jitter.delay(5, u64::MAX) <= MAX_TIMEOUT);
assert!(Reconnect::default().with_max_attempts(Some(2)).may_retry(2));
assert!(!Reconnect::default().with_max_attempts(Some(2)).may_retry(3));
let settings = WsSettings::default().with_heartbeat(Duration::from_secs(20), Duration::from_secs(1));
assert_eq!(settings.dead_after, Duration::from_secs(20), "dead_after is at least the interval");
}
}