use std::collections::HashMap;
#[cfg(feature = "json")]
use std::collections::HashSet;
use std::fmt;
use std::sync::{Arc, Mutex, MutexGuard, PoisonError};
use std::time::Duration;
use bevy_app::{App, First, Last, PostUpdate};
use bevy_ecs::message::Message;
use bevy_ecs::resource::Resource;
use bevy_ecs::schedule::common_conditions::on_message;
use bevy_ecs::schedule::IntoScheduleConfigs;
use http::header::{HeaderMap, HeaderName, HeaderValue};
use crate::request::RequestId;
use crate::response::BackendError;
use crate::BackendSystems;
mod link;
mod protocol;
mod proxy;
mod systems;
mod transport;
pub use link::TungsteniteTransport;
#[cfg(feature = "json")]
pub use protocol::JsonEnvelope;
pub use protocol::{WsIncoming, WsProtocol};
#[cfg(any(feature = "ssh", feature = "oauth"))]
pub(crate) use systems::{ws_exit, ws_receive, ws_send};
pub use transport::{FakeWsTransport, WsHandshake, WsLinkEvent, WsLinkId, WsTransport, WsTransportRes};
#[derive(Clone, PartialEq, Eq, Hash, PartialOrd, Ord)]
pub struct WsName(Arc<str>);
impl WsName {
pub fn as_str(&self) -> &str {
&self.0
}
}
impl fmt::Debug for WsName {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "WsName({:?})", &*self.0)
}
}
impl fmt::Display for WsName {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str(&self.0)
}
}
impl From<&str> for WsName {
fn from(name: &str) -> Self {
Self(Arc::from(name))
}
}
impl From<String> for WsName {
fn from(name: String) -> Self {
Self(Arc::from(name))
}
}
impl From<&String> for WsName {
fn from(name: &String) -> Self {
Self(Arc::from(name.as_str()))
}
}
impl From<&WsName> for WsName {
fn from(name: &WsName) -> Self {
name.clone()
}
}
impl PartialEq<str> for WsName {
fn eq(&self, other: &str) -> bool {
&*self.0 == other
}
}
impl PartialEq<&str> for WsName {
fn eq(&self, other: &&str) -> bool {
&*self.0 == *other
}
}
#[derive(Clone, PartialEq, Eq)]
#[non_exhaustive]
pub enum WsFrame {
Text(String),
Binary(Vec<u8>),
}
impl WsFrame {
pub fn len(&self) -> usize {
match self {
WsFrame::Text(text) => text.len(),
WsFrame::Binary(bytes) => bytes.len(),
}
}
pub fn is_empty(&self) -> bool {
self.len() == 0
}
pub fn as_text(&self) -> Option<&str> {
match self {
WsFrame::Text(text) => Some(text),
WsFrame::Binary(_) => None,
}
}
pub fn as_bytes(&self) -> &[u8] {
match self {
WsFrame::Text(text) => text.as_bytes(),
WsFrame::Binary(bytes) => bytes,
}
}
}
impl fmt::Debug for WsFrame {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
WsFrame::Text(text) => write!(f, "Text({} bytes)", text.len()),
WsFrame::Binary(bytes) => write!(f, "Binary({} bytes)", bytes.len()),
}
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
#[non_exhaustive]
pub enum WsState {
Connecting,
Connected,
Reconnecting {
attempt: u32,
retry_in: Duration,
},
Disconnected,
WaitingForCredentials,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct WsCredentialsRefresh {
pub(crate) timeout: Duration,
pub(crate) close_codes: Vec<u16>,
}
impl Default for WsCredentialsRefresh {
fn default() -> Self {
Self { timeout: Duration::from_secs(30), close_codes: Vec::new() }
}
}
impl WsCredentialsRefresh {
pub fn new() -> Self {
Self::default()
}
pub fn with_timeout(mut self, timeout: Duration) -> Self {
self.timeout = timeout.clamp(Duration::from_millis(1), crate::config::MAX_TIMEOUT);
self
}
pub fn with_close_code(mut self, code: u16) -> Self {
if !self.close_codes.contains(&code) {
self.close_codes.push(code);
}
self
}
pub fn timeout(&self) -> Duration {
self.timeout
}
pub fn close_codes(&self) -> &[u16] {
&self.close_codes
}
}
#[derive(Message, Clone, Debug)]
#[non_exhaustive]
pub struct WsCredentialsRefused {
pub name: WsName,
pub error: BackendError,
}
#[derive(Clone, Debug)]
pub struct WsReconnect {
base: Duration,
cap: Duration,
max_attempts: Option<u32>,
stable_after: Duration,
jitter: bool,
pub(crate) retry_tls: bool,
}
impl Default for WsReconnect {
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,
retry_tls: false,
}
}
}
impl WsReconnect {
pub fn never() -> Self {
Self { max_attempts: Some(0), ..Self::default() }
}
pub fn with_base(mut self, base: Duration) -> Self {
self.base = base.max(Duration::from_millis(1));
self
}
pub fn with_cap(mut self, cap: Duration) -> Self {
self.cap = cap;
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;
self
}
pub fn with_tls_retry(mut self, retry: bool) -> Self {
self.retry_tls = retry;
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)]
pub struct WsSettings {
pub(crate) url: String,
pub(crate) headers: HeaderMap,
header_error: Option<String>,
pub(crate) read_timeout: Duration,
pub(crate) connect_timeout: Duration,
pub(crate) ping_interval: Duration,
pub(crate) dead_after: Duration,
pub(crate) request_timeout: Duration,
pub(crate) max_message_bytes: usize,
pub(crate) reconnect: WsReconnect,
pub(crate) allow_insecure: bool,
pub(crate) credentials: bool,
pub(crate) protocol: Option<Arc<dyn WsProtocol>>,
pub(crate) outbox_limit: usize,
pub(crate) resend_limit: usize,
pub(crate) waiting_limit: usize,
pub(crate) auth_ack: Option<Duration>,
pub(crate) refresh: Option<WsCredentialsRefresh>,
}
impl fmt::Debug for WsSettings {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
let url = self.url.parse::<http::Uri>().map(|u| crate::request::redacted_url(&u)).unwrap_or_else(|_| "<invalid>".into());
let headers: Vec<&str> = self.headers.keys().map(HeaderName::as_str).collect();
f.debug_struct("WsSettings")
.field("url", &url)
.field("header_names", &headers)
.field("read_timeout", &self.read_timeout)
.field("connect_timeout", &self.connect_timeout)
.field("ping_interval", &self.ping_interval)
.field("dead_after", &self.dead_after)
.field("request_timeout", &self.request_timeout)
.field("max_message_bytes", &self.max_message_bytes)
.field("reconnect", &self.reconnect)
.field("allow_insecure", &self.allow_insecure)
.field("credentials", &self.credentials)
.field("protocol", &self.protocol.is_some())
.field("credentials_refresh", &self.refresh)
.finish_non_exhaustive()
}
}
pub const DEFAULT_WS_READ_TIMEOUT: Duration = Duration::from_millis(20);
pub const DEFAULT_WS_MAX_MESSAGE_BYTES: usize = 1024 * 1024;
impl WsSettings {
pub fn new(url: impl Into<String>) -> Self {
Self {
url: url.into(),
headers: HeaderMap::new(),
header_error: None,
read_timeout: DEFAULT_WS_READ_TIMEOUT,
connect_timeout: Duration::from_secs(10),
ping_interval: Duration::from_secs(15),
dead_after: Duration::from_secs(45),
request_timeout: Duration::from_secs(10),
max_message_bytes: DEFAULT_WS_MAX_MESSAGE_BYTES,
reconnect: WsReconnect::default(),
allow_insecure: false,
credentials: true,
#[cfg(feature = "json")]
protocol: Some(Arc::new(JsonEnvelope)),
#[cfg(not(feature = "json"))]
protocol: None,
outbox_limit: 64,
resend_limit: 32,
waiting_limit: 64,
auth_ack: None,
refresh: None,
}
}
pub fn with_read_timeout(mut self, timeout: Duration) -> Self {
self.read_timeout = timeout.clamp(Duration::from_millis(5), Duration::from_millis(250));
self
}
pub fn with_connect_timeout(mut self, timeout: Duration) -> Self {
self.connect_timeout = timeout.clamp(Duration::from_millis(100), crate::config::MAX_TIMEOUT);
self
}
pub fn with_heartbeat(mut self, interval: Duration, dead_after: Duration) -> Self {
self.ping_interval = interval.clamp(Duration::from_millis(10), crate::config::MAX_TIMEOUT);
self.dead_after = dead_after.clamp(self.ping_interval, crate::config::MAX_TIMEOUT);
self
}
pub fn with_request_timeout(mut self, timeout: Duration) -> Self {
self.request_timeout = timeout.clamp(Duration::from_millis(1), crate::config::MAX_TIMEOUT);
self
}
pub fn with_max_message_bytes(mut self, bytes: usize) -> Self {
self.max_message_bytes = bytes.max(1024);
self
}
pub fn with_reconnect(mut self, reconnect: WsReconnect) -> Self {
self.reconnect = reconnect;
self
}
pub fn allow_insecure_ws(mut self, allow: bool) -> Self {
self.allow_insecure = allow;
self
}
pub fn with_header(mut self, name: &str, value: &str) -> Self {
match (HeaderName::try_from(name), HeaderValue::try_from(value)) {
(Ok(name), Ok(value)) => {
self.headers.append(name, value);
}
_ => self.header_error = Some(format!("the handshake header `{name}` is not valid")),
}
self
}
pub fn without_credentials(mut self) -> Self {
self.credentials = false;
self
}
pub fn with_protocol(mut self, protocol: impl WsProtocol) -> Self {
self.protocol = Some(Arc::new(protocol));
self
}
pub fn without_protocol(mut self) -> Self {
self.protocol = None;
self
}
pub fn with_outbox_limit(mut self, frames: usize) -> Self {
self.outbox_limit = frames;
self
}
pub fn with_resend_limit(mut self, requests: usize) -> Self {
self.resend_limit = requests;
self
}
pub fn with_waiting_limit(mut self, requests: usize) -> Self {
self.waiting_limit = requests;
self
}
pub fn with_auth_ack(mut self, timeout: Duration) -> Self {
self.auth_ack = Some(timeout.clamp(Duration::from_millis(1), crate::config::MAX_TIMEOUT));
self
}
pub fn with_credentials_refresh(mut self, refresh: WsCredentialsRefresh) -> Self {
self.refresh = Some(refresh);
self
}
pub fn credentials_refresh(&self) -> Option<&WsCredentialsRefresh> {
self.refresh.as_ref()
}
pub fn connect_timeout(&self) -> Duration {
self.connect_timeout
}
pub fn heartbeat(&self) -> (Duration, Duration) {
(self.ping_interval, self.dead_after)
}
pub fn url(&self) -> &str {
&self.url
}
pub fn read_timeout(&self) -> Duration {
self.read_timeout
}
pub fn reconnect(&self) -> &WsReconnect {
&self.reconnect
}
pub fn validate(&self) -> Result<(), BackendError> {
if let Some(why) = &self.header_error {
return Err(BackendError::InvalidRequest(why.clone()));
}
systems::parse_ws_url(&self.url).map(|_| ())
}
}
#[derive(Clone)]
pub struct WsOutgoing {
pub(crate) kind: String,
pub(crate) payload: Vec<u8>,
pub(crate) resend: bool,
pub(crate) timeout: Option<Duration>,
}
impl fmt::Debug for WsOutgoing {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("WsOutgoing")
.field("kind", &self.kind)
.field("payload_bytes", &self.payload.len())
.field("resend", &self.resend)
.field("timeout", &self.timeout)
.finish()
}
}
impl WsOutgoing {
pub fn new(kind: impl Into<String>, payload: impl Into<Vec<u8>>) -> Self {
Self { kind: kind.into(), payload: payload.into(), resend: false, timeout: None }
}
pub fn resend_on_reconnect(mut self, resend: bool) -> Self {
self.resend = resend;
self
}
pub fn with_timeout(mut self, timeout: Duration) -> Self {
self.timeout = Some(timeout.clamp(Duration::from_millis(1), crate::config::MAX_TIMEOUT));
self
}
}
#[derive(Message, Clone, Debug)]
#[non_exhaustive]
pub struct WsStateChanged {
pub name: WsName,
pub state: WsState,
pub error: Option<BackendError>,
}
#[derive(Message, Clone, Debug)]
#[non_exhaustive]
pub struct WsMessage {
pub name: WsName,
pub frame: WsFrame,
}
#[derive(Message, Clone)]
#[non_exhaustive]
pub struct WsRawResponse {
pub id: RequestId,
pub name: WsName,
pub result: Result<Vec<u8>, BackendError>,
}
impl fmt::Debug for WsRawResponse {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
let result = self.result.as_ref().map(Vec::len);
f.debug_struct("WsRawResponse").field("id", &self.id).field("name", &self.name).field("result_bytes", &result).finish()
}
}
#[cfg(feature = "json")]
pub trait WsRequest: serde::Serialize + Send + Sync + 'static {
type Response: serde::de::DeserializeOwned + Send + Sync + 'static;
const KIND: &'static str;
fn resend_on_reconnect(&self) -> bool {
false
}
}
#[cfg(feature = "json")]
pub trait WsPushMessage: serde::de::DeserializeOwned + Send + Sync + 'static {
const KIND: &'static str;
}
#[cfg(feature = "json")]
#[derive(Message, Clone, Debug)]
#[non_exhaustive]
pub struct WsResponse<T: Send + Sync + 'static> {
pub id: RequestId,
pub name: WsName,
pub result: Result<T, BackendError>,
}
#[cfg(feature = "json")]
#[derive(Message, Clone, Debug)]
#[non_exhaustive]
pub struct WsPush<P: Send + Sync + 'static> {
pub name: WsName,
pub data: P,
}
#[derive(Clone, Debug, PartialEq, Eq)]
#[non_exhaustive]
pub struct WsConnectionInfo {
pub state: WsState,
pub attempt: u32,
pub last_error: Option<BackendError>,
pub pending_requests: usize,
pub queued_frames: usize,
}
#[derive(Resource, Default, Debug)]
pub struct WsConnections {
pub(crate) map: HashMap<WsName, WsConnectionInfo>,
}
impl WsConnections {
pub fn get(&self, name: &str) -> Option<&WsConnectionInfo> {
self.map.get(&WsName::from(name))
}
pub fn state(&self, name: &str) -> Option<WsState> {
self.get(name).map(|c| c.state)
}
pub fn is_connected(&self, name: &str) -> bool {
self.state(name) == Some(WsState::Connected)
}
pub fn iter(&self) -> impl Iterator<Item = (&WsName, &WsConnectionInfo)> {
self.map.iter()
}
}
pub(crate) enum WsRoute {
Raw,
#[cfg(feature = "json")]
Typed(Box<dyn TypedRoute>),
}
#[cfg(feature = "json")]
pub(crate) trait TypedRoute: Send + Sync + 'static {
fn deliver(&self, id: RequestId, name: WsName, result: Result<Vec<u8>, BackendError>, commands: &mut bevy_ecs::system::Commands);
}
#[cfg(feature = "json")]
struct TypedRouteFor<T>(std::marker::PhantomData<fn() -> T>);
#[cfg(feature = "json")]
impl<T: serde::de::DeserializeOwned + Send + Sync + 'static> TypedRoute for TypedRouteFor<T> {
fn deliver(&self, id: RequestId, name: WsName, result: Result<Vec<u8>, BackendError>, commands: &mut bevy_ecs::system::Commands) {
let result = result.and_then(|bytes| {
let bytes = if bytes.iter().all(u8::is_ascii_whitespace) { b"null".to_vec() } else { bytes };
serde_json::from_slice::<T>(&bytes)
.map_err(|e| BackendError::Decode { message: e.to_string(), response: Box::new(crate::RawResponse::new(http::StatusCode::OK, bytes.clone())) })
});
commands.queue(move |world: &mut bevy_ecs::world::World| {
if world.write_message(WsResponse::<T> { id, name, result }).is_none() {
tracing::error!(">>> NET-BACKEND: {id}: `WsResponse<{}>` is not registered; the answer is lost", std::any::type_name::<T>());
}
});
}
}
#[cfg(feature = "json")]
pub(crate) trait PushRoute: Send + Sync + 'static {
fn deliver(&self, name: WsName, data: &[u8], commands: &mut bevy_ecs::system::Commands);
}
#[cfg(feature = "json")]
struct PushRouteFor<P>(std::marker::PhantomData<fn() -> P>);
#[cfg(feature = "json")]
impl<P: WsPushMessage> PushRoute for PushRouteFor<P> {
fn deliver(&self, name: WsName, data: &[u8], commands: &mut bevy_ecs::system::Commands) {
match serde_json::from_slice::<P>(data) {
Ok(data) => commands.queue(move |world: &mut bevy_ecs::world::World| {
world.write_message(WsPush::<P> { name, data });
}),
Err(_) => tracing::debug!(">>> NET-BACKEND: ws `{name}`: a `{}` push did not decode as `{}`", P::KIND, std::any::type_name::<P>()),
}
}
}
pub(crate) enum WsQueued {
Connect(WsName, Box<WsSettings>),
Disconnect(WsName),
Send(WsName, WsFrame),
Request {
name: WsName,
id: RequestId,
out: WsOutgoing,
route: WsRoute,
},
#[cfg_attr(not(feature = "json"), allow(dead_code))]
Fail {
name: WsName,
id: RequestId,
route: WsRoute,
error: BackendError,
},
}
#[derive(Resource, Default)]
pub struct WsClient {
queue: Mutex<Vec<WsQueued>>,
cancels: crate::inflight::CancelList,
#[cfg(feature = "json")]
response_types: HashSet<std::any::TypeId>,
#[cfg(feature = "json")]
pub(crate) push_routes: HashMap<&'static str, Vec<Box<dyn PushRoute>>>,
}
impl fmt::Debug for WsClient {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("WsClient").field("queued", &self.lock().len()).finish_non_exhaustive()
}
}
impl WsClient {
fn lock(&self) -> MutexGuard<'_, Vec<WsQueued>> {
self.queue.lock().unwrap_or_else(PoisonError::into_inner)
}
pub(crate) fn drain(&self) -> Vec<WsQueued> {
std::mem::take(&mut *self.lock())
}
pub fn connect(&self, name: impl Into<WsName>, settings: WsSettings) {
self.lock().push(WsQueued::Connect(name.into(), Box::new(settings)));
}
pub fn disconnect(&self, name: impl Into<WsName>) {
self.lock().push(WsQueued::Disconnect(name.into()));
}
pub fn send(&self, name: impl Into<WsName>, frame: WsFrame) {
self.lock().push(WsQueued::Send(name.into(), frame));
}
pub fn send_text(&self, name: impl Into<WsName>, text: impl Into<String>) {
self.send(name, WsFrame::Text(text.into()));
}
pub fn send_binary(&self, name: impl Into<WsName>, bytes: impl Into<Vec<u8>>) {
self.send(name, WsFrame::Binary(bytes.into()));
}
pub fn request_raw(&self, name: impl Into<WsName>, request: WsOutgoing) -> RequestId {
let id = RequestId::next();
self.lock().push(WsQueued::Request { name: name.into(), id, out: request, route: WsRoute::Raw });
id
}
#[cfg(feature = "json")]
pub fn request<R: WsRequest>(&self, name: impl Into<WsName>, request: &R) -> RequestId {
let id = RequestId::next();
let name = name.into();
if !self.response_types.contains(&std::any::TypeId::of::<R::Response>()) {
let type_name = std::any::type_name::<R>();
tracing::error!(">>> NET-BACKEND: `{type_name}` is not registered: call `app.add_ws_request::<{type_name}>()`; answered on WsRawResponse");
let error = BackendError::InvalidRequest(format!("request type `{type_name}` is not registered; call `app.add_ws_request::<{type_name}>()`"));
self.lock().push(WsQueued::Fail { name, id, route: WsRoute::Raw, error });
return id;
}
let route = WsRoute::Typed(Box::new(TypedRouteFor::<R::Response>(std::marker::PhantomData)));
match serde_json::to_vec(request) {
Ok(payload) => {
let out = WsOutgoing::new(R::KIND, payload).resend_on_reconnect(request.resend_on_reconnect());
self.lock().push(WsQueued::Request { name, id, out, route });
}
Err(e) => self.lock().push(WsQueued::Fail { name, id, route, error: BackendError::Encode(e.to_string()) }),
}
id
}
pub fn cancel(&self, id: RequestId) {
self.cancels.push(id);
}
pub(crate) fn share_cancels(&mut self, cancels: crate::inflight::CancelList) {
self.cancels = cancels;
}
}
#[cfg(feature = "json")]
pub(crate) fn register_request<R: WsRequest>(app: &mut App) {
app.add_message::<WsResponse<R::Response>>();
app.init_resource::<WsClient>();
if let Some(mut client) = app.world_mut().get_resource_mut::<WsClient>() {
client.response_types.insert(std::any::TypeId::of::<R::Response>());
}
}
#[cfg(feature = "json")]
pub(crate) fn register_push<P: WsPushMessage>(app: &mut App) {
app.add_message::<WsPush<P>>();
app.init_resource::<WsClient>();
if let Some(mut client) = app.world_mut().get_resource_mut::<WsClient>() {
client.push_routes.entry(P::KIND).or_default().push(Box::new(PushRouteFor::<P>(std::marker::PhantomData)));
}
}
pub(crate) fn build(app: &mut App, tls: &crate::TlsSettings, proxy: &crate::ProxySettings) {
app.init_resource::<WsClient>();
let cancels = app.world().resource::<crate::InFlight>().cancel_list();
if let Some(mut client) = app.world_mut().get_resource_mut::<WsClient>() {
client.share_cancels(cancels);
}
app.init_resource::<WsConnections>()
.init_resource::<systems::WsRuntime>()
.add_message::<WsStateChanged>()
.add_message::<WsCredentialsRefused>()
.add_message::<WsMessage>()
.add_message::<WsRawResponse>()
.add_systems(First, systems::ws_receive.in_set(BackendSystems::Receive).after(crate::inflight::receive_answers))
.add_systems(PostUpdate, systems::ws_send.in_set(BackendSystems::Send).after(crate::inflight::send_requests))
.add_systems(Last, systems::ws_exit.in_set(BackendSystems::Exit).after(crate::inflight::shutdown_on_exit).run_if(on_message::<bevy_app::AppExit>));
if !app.world().contains_resource::<WsTransportRes>() {
app.insert_resource(WsTransportRes::new(TungsteniteTransport::with_settings(tls, proxy)));
}
}