use std::collections::{HashSet, VecDeque};
use std::fmt;
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::{Arc, Mutex, MutexGuard, PoisonError};
use std::time::Duration;
use bevy_ecs::resource::Resource;
use http::header::{HeaderMap, HeaderName};
use http::Uri;
use super::WsFrame;
use crate::response::BackendError;
#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash, PartialOrd, Ord)]
pub struct WsLinkId(u64);
static NEXT_LINK: AtomicU64 = AtomicU64::new(1);
impl WsLinkId {
pub(crate) fn next() -> Self {
Self(NEXT_LINK.fetch_add(1, Ordering::Relaxed))
}
}
impl fmt::Display for WsLinkId {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "link#{}", self.0)
}
}
#[derive(Clone)]
#[non_exhaustive]
pub struct WsHandshake {
pub uri: Uri,
pub headers: HeaderMap,
pub connect_timeout: Duration,
pub read_timeout: Duration,
pub ping_interval: Duration,
pub dead_after: Duration,
pub max_message_bytes: usize,
}
impl fmt::Debug for WsHandshake {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
let headers: Vec<&str> = self.headers.keys().map(HeaderName::as_str).collect();
f.debug_struct("WsHandshake")
.field("url", &crate::request::redacted_url(&self.uri))
.field("header_names", &headers)
.field("connect_timeout", &self.connect_timeout)
.field("read_timeout", &self.read_timeout)
.field("ping_interval", &self.ping_interval)
.field("dead_after", &self.dead_after)
.field("max_message_bytes", &self.max_message_bytes)
.finish()
}
}
impl WsHandshake {
pub fn is_secure(&self) -> bool {
self.uri.scheme_str().is_some_and(|s| s.eq_ignore_ascii_case("wss"))
}
}
#[derive(Clone, Debug, PartialEq, Eq)]
#[non_exhaustive]
pub enum WsLinkEvent {
Opened,
Frame(WsFrame),
Closed {
code: Option<u16>,
reason: String,
},
Failed(BackendError),
}
pub trait WsTransport: Send + Sync + 'static {
fn open(&mut self, link: WsLinkId, handshake: WsHandshake);
fn send(&mut self, link: WsLinkId, frame: WsFrame);
fn close(&mut self, link: WsLinkId, code: u16);
fn poll(&mut self) -> Vec<(WsLinkId, WsLinkEvent)>;
fn shutdown(&mut self) {}
}
static NEXT_GENERATION: AtomicU64 = AtomicU64::new(1);
#[derive(Resource)]
pub struct WsTransportRes {
inner: Box<dyn WsTransport>,
generation: u64,
}
impl fmt::Debug for WsTransportRes {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("WsTransportRes").field("generation", &self.generation).finish_non_exhaustive()
}
}
impl WsTransportRes {
pub fn new(transport: impl WsTransport) -> Self {
Self { inner: Box::new(transport), generation: NEXT_GENERATION.fetch_add(1, Ordering::Relaxed) }
}
pub(crate) fn generation(&self) -> u64 {
self.generation
}
pub(crate) fn get_mut(&mut self) -> &mut dyn WsTransport {
self.inner.as_mut()
}
}
#[derive(Default)]
struct FakeState {
opened: Vec<(WsLinkId, WsHandshake)>,
live: HashSet<WsLinkId>,
sent: Vec<(WsLinkId, WsFrame)>,
closed: Vec<(WsLinkId, u16)>,
events: VecDeque<(WsLinkId, WsLinkEvent)>,
manual_accept: bool,
reject_next: VecDeque<BackendError>,
echo_envelope: bool,
shutdowns: usize,
}
#[derive(Clone, Default)]
pub struct FakeWsTransport {
state: Arc<Mutex<FakeState>>,
}
impl fmt::Debug for FakeWsTransport {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
let state = self.lock();
f.debug_struct("FakeWsTransport").field("opened", &state.opened.len()).field("live", &state.live.len()).finish_non_exhaustive()
}
}
impl FakeWsTransport {
pub fn new() -> Self {
Self::default()
}
fn lock(&self) -> MutexGuard<'_, FakeState> {
self.state.lock().unwrap_or_else(PoisonError::into_inner)
}
fn event(&self, link: WsLinkId, event: WsLinkEvent) {
let mut state = self.lock();
if matches!(event, WsLinkEvent::Closed { .. } | WsLinkEvent::Failed(_)) {
state.live.remove(&link);
}
state.events.push_back((link, event));
}
pub fn manual_accept(&self, manual: bool) -> &Self {
self.lock().manual_accept = manual;
self
}
pub fn reject_next(&self, error: BackendError) -> &Self {
self.lock().reject_next.push_back(error);
self
}
pub fn echo_envelope(&self, echo: bool) -> &Self {
self.lock().echo_envelope = echo;
self
}
pub fn accept(&self, link: WsLinkId) {
self.event(link, WsLinkEvent::Opened);
}
pub fn push(&self, link: WsLinkId, frame: WsFrame) {
self.event(link, WsLinkEvent::Frame(frame));
}
pub fn drop_link(&self, link: WsLinkId, code: u16) {
self.event(link, WsLinkEvent::Closed { code: Some(code), reason: String::new() });
}
pub fn fail_link(&self, link: WsLinkId, error: BackendError) {
self.event(link, WsLinkEvent::Failed(error));
}
pub fn opened(&self) -> Vec<(WsLinkId, WsHandshake)> {
self.lock().opened.clone()
}
pub fn last_link(&self) -> Option<WsLinkId> {
self.lock().opened.last().map(|(link, _)| *link)
}
pub fn live_links(&self) -> Vec<WsLinkId> {
let mut links: Vec<WsLinkId> = self.lock().live.iter().copied().collect();
links.sort_unstable();
links
}
pub fn sent(&self, link: WsLinkId) -> Vec<WsFrame> {
self.lock().sent.iter().filter(|(l, _)| *l == link).map(|(_, f)| f.clone()).collect()
}
pub fn all_sent(&self) -> Vec<(WsLinkId, WsFrame)> {
self.lock().sent.clone()
}
pub fn closed(&self) -> Vec<(WsLinkId, u16)> {
self.lock().closed.clone()
}
pub fn shutdown_count(&self) -> usize {
self.lock().shutdowns
}
}
impl WsTransport for FakeWsTransport {
fn open(&mut self, link: WsLinkId, handshake: WsHandshake) {
let mut state = self.lock();
state.opened.push((link, handshake));
if let Some(error) = state.reject_next.pop_front() {
state.events.push_back((link, WsLinkEvent::Failed(error)));
return;
}
state.live.insert(link);
if !state.manual_accept {
state.events.push_back((link, WsLinkEvent::Opened));
}
}
fn send(&mut self, link: WsLinkId, frame: WsFrame) {
let mut state = self.lock();
if !state.live.contains(&link) {
return;
}
#[cfg(feature = "json")]
if state.echo_envelope {
if let Some(answer) = frame.as_text().and_then(|t| serde_json::from_str::<serde_json::Value>(t).ok()).and_then(|v| {
let id = v.get("id")?.as_u64()?;
Some(serde_json::json!({ "id": id, "ok": true, "data": v.get("data").cloned().unwrap_or(serde_json::Value::Null) }).to_string())
}) {
state.events.push_back((link, WsLinkEvent::Frame(WsFrame::Text(answer))));
}
}
state.sent.push((link, frame));
}
fn close(&mut self, link: WsLinkId, code: u16) {
let mut state = self.lock();
state.live.remove(&link);
state.closed.push((link, code));
}
fn poll(&mut self) -> Vec<(WsLinkId, WsLinkEvent)> {
self.lock().events.drain(..).collect()
}
fn shutdown(&mut self) {
let mut state = self.lock();
state.shutdowns += 1;
state.live.clear();
}
}