use std::{
collections::VecDeque,
fmt::Debug,
sync::{
Arc,
atomic::{AtomicBool, AtomicU8, Ordering},
},
time::Duration,
};
use crate::{config::EngineIoConfig, errors::Error};
use bytes::Bytes;
use engineioxide_core::{Packet, PacketBuf, ProtocolVersion, Str, TransportType};
use futures_util::FutureExt;
use http::request::Parts;
use smallvec::{SmallVec, smallvec};
use tokio::sync::{
Mutex,
mpsc::{self, Receiver, error::TrySendError},
watch,
};
pub use engineioxide_core::Sid;
use tokio_util::sync::CancellationToken;
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum DisconnectReason {
TransportClose,
MultipleHttpPollingError,
PacketParsingError,
TransportError,
HeartbeatTimeout,
ClosingServer,
}
impl From<&Error> for Option<DisconnectReason> {
fn from(err: &Error) -> Self {
use Error::*;
match err {
WsTransport(_) => Some(DisconnectReason::TransportError),
BadPacket(_) | PacketParse(_) => Some(DisconnectReason::PacketParsingError),
HeartbeatTimeout => Some(DisconnectReason::HeartbeatTimeout),
_ => None,
}
}
}
pub struct Permit<'a> {
inner: mpsc::Permit<'a, PacketBuf>,
}
impl Permit<'_> {
#[inline]
pub fn emit(self, msg: Str) {
self.inner.send(smallvec![Packet::Message(msg)]);
}
#[inline]
pub fn emit_binary(self, data: Bytes) {
self.inner.send(smallvec![Packet::Binary(data)]);
}
pub fn emit_many(self, msg: Str, data: VecDeque<Bytes>) {
let mut packets = SmallVec::with_capacity(data.len() + 1);
packets.push(Packet::Message(msg));
for d in data {
packets.push(Packet::Binary(d));
}
self.inner.send(packets);
}
pub fn emit_many_binary(self, bin: Bytes, data: Vec<Bytes>) {
let mut packets = SmallVec::with_capacity(data.len() + 1);
packets.push(Packet::Binary(bin));
for d in data {
packets.push(Packet::Binary(d));
}
self.inner.send(packets);
}
}
#[derive(Debug)]
pub(crate) struct InternalRx {
pub buffered_rx: Receiver<PacketBuf>,
pub peeked_packet: Option<PacketBuf>,
pub volatile_rx: watch::Receiver<Option<PacketBuf>>,
}
impl InternalRx {
fn new(
buffered_rx: Receiver<PacketBuf>,
volatile_rx: watch::Receiver<Option<PacketBuf>>,
) -> Self {
Self {
buffered_rx,
peeked_packet: None,
volatile_rx,
}
}
}
pub struct Socket<D>
where
D: Default + Send + Sync + 'static,
{
pub id: Sid,
pub protocol: ProtocolVersion,
transport: AtomicU8,
upgrading: AtomicBool,
pub(crate) internal_rx: Mutex<InternalRx>,
pub(crate) internal_tx: mpsc::Sender<PacketBuf>,
volatile_tx: watch::Sender<Option<PacketBuf>>,
heartbeat_rx: Mutex<Receiver<()>>,
pub(crate) heartbeat_tx: mpsc::Sender<()>,
pub(crate) cancellation_token: CancellationToken,
close_fn: Box<dyn Fn(Sid, DisconnectReason) + Send + Sync>,
pub data: D,
pub req_parts: Parts,
pub(crate) supports_binary: bool,
}
impl<D> Socket<D>
where
D: Default + Send + Sync + 'static,
{
pub(crate) fn new(
protocol: ProtocolVersion,
transport: TransportType,
config: &EngineIoConfig,
req_parts: Parts,
close_fn: Box<dyn Fn(Sid, DisconnectReason) + Send + Sync>,
supports_binary: bool,
) -> Self {
let (internal_tx, internal_rx) = mpsc::channel(config.max_buffer_size);
let (heartbeat_tx, heartbeat_rx) = mpsc::channel(1);
let (volatile_tx, volatile_rx) = watch::channel(None);
Self {
id: Sid::new(),
protocol,
transport: AtomicU8::new(transport as u8),
upgrading: AtomicBool::new(false),
internal_rx: Mutex::new(InternalRx::new(internal_rx, volatile_rx)),
internal_tx,
volatile_tx,
heartbeat_rx: Mutex::new(heartbeat_rx),
heartbeat_tx,
cancellation_token: CancellationToken::new(),
close_fn,
data: D::default(),
req_parts,
supports_binary,
}
}
pub(crate) fn send(&self, packet: Packet) -> Result<(), TrySendError<Packet>> {
#[cfg(feature = "tracing")]
tracing::debug!(?packet, "sending packet");
self.internal_tx
.try_send(smallvec![packet])
.map_err(|p| match p {
TrySendError::Full(mut p) => TrySendError::Full(p.pop().unwrap()),
TrySendError::Closed(mut p) => TrySendError::Closed(p.pop().unwrap()),
})?;
Ok(())
}
pub(crate) fn is_upgrading(&self) -> bool {
self.upgrading.load(Ordering::Relaxed)
}
pub(crate) fn start_upgrade(&self) {
self.upgrading.store(true, Ordering::Relaxed);
}
pub(crate) fn spawn_heartbeat(self: Arc<Self>, interval: Duration, timeout: Duration) {
let cancellation_token = self.cancellation_token.clone();
tokio::spawn(
cancellation_token
.run_until_cancelled_owned(async move {
if let Err(_e) = self.heartbeat_job(interval, timeout).await {
self.close(DisconnectReason::HeartbeatTimeout);
#[cfg(feature = "tracing")]
tracing::debug!(id = ?self.id, "heartbeat error: {_e}");
}
})
.inspect(|_v| {
#[cfg(feature = "tracing")]
tracing::debug!(aborted = _v.is_none(), "heartbeat job completed");
}),
);
}
#[cfg(feature = "v3")]
async fn heartbeat_job(&self, interval: Duration, timeout: Duration) -> Result<(), Error> {
match self.protocol {
ProtocolVersion::V3 => self.heartbeat_job_v3(interval, timeout).await,
ProtocolVersion::V4 => self.heartbeat_job_v4(interval, timeout).await,
}
}
#[cfg(not(feature = "v3"))]
async fn heartbeat_job(&self, interval: Duration, timeout: Duration) -> Result<(), Error> {
self.heartbeat_job_v4(interval, timeout).await
}
async fn heartbeat_job_v4(&self, interval: Duration, timeout: Duration) -> Result<(), Error> {
let mut heartbeat_rx = self
.heartbeat_rx
.try_lock()
.expect("Pong rx should be locked only once");
#[cfg(feature = "tracing")]
tracing::debug!(sid = ?self.id, "heartbeat sender routine started");
let mut interval_tick = tokio::time::interval(interval);
interval_tick.tick().await;
heartbeat_rx.try_recv().ok();
loop {
if self.is_upgrading() {
interval_tick.tick().await;
continue;
}
#[cfg(feature = "tracing")]
tracing::trace!(sid = ?self.id, "emitting ping");
self.internal_tx
.try_send(smallvec![Packet::Ping])
.map_err(|_| Error::HeartbeatTimeout)?;
#[cfg(feature = "tracing")]
tracing::trace!(sid = ?self.id, "waiting for pong");
tokio::time::timeout(timeout, heartbeat_rx.recv())
.await
.map_err(|_| Error::HeartbeatTimeout)?
.ok_or(Error::HeartbeatTimeout)?;
#[cfg(feature = "tracing")]
tracing::trace!(sid = ?self.id, "pong received");
interval_tick.tick().await;
}
}
#[cfg(feature = "v3")]
async fn heartbeat_job_v3(&self, interval: Duration, timeout: Duration) -> Result<(), Error> {
let mut heartbeat_rx = self
.heartbeat_rx
.try_lock()
.expect("Pong rx should be locked only once");
#[cfg(feature = "tracing")]
tracing::debug!(sid = ?self.id, "heartbeat receiver routine started");
loop {
tokio::time::timeout(interval + timeout, heartbeat_rx.recv())
.await
.map_err(|_| Error::HeartbeatTimeout)?
.ok_or(Error::HeartbeatTimeout)?;
#[cfg(feature = "tracing")]
tracing::trace!(sid = ?self.id, "ping received, sending pong");
self.internal_tx
.try_send(smallvec![Packet::Pong])
.map_err(|_| Error::HeartbeatTimeout)?;
}
}
pub(crate) fn is_ws(&self) -> bool {
self.transport.load(Ordering::Relaxed) == TransportType::Websocket as u8
}
pub(crate) fn is_http(&self) -> bool {
self.transport.load(Ordering::Relaxed) == TransportType::Polling as u8
}
pub(crate) fn upgrade_to_websocket(&self) {
self.upgrading.store(false, Ordering::Relaxed);
self.transport
.store(TransportType::Websocket as u8, Ordering::Relaxed);
}
pub fn transport_type(&self) -> TransportType {
TransportType::from(self.transport.load(Ordering::Relaxed))
}
#[inline]
pub fn reserve(&self) -> Result<Permit<'_>, TrySendError<()>> {
let permit = self.internal_tx.try_reserve()?;
Ok(Permit { inner: permit })
}
pub fn emit(&self, msg: impl Into<Str>) -> Result<(), TrySendError<Str>> {
self.send(Packet::Message(msg.into())).map_err(|e| match e {
TrySendError::Full(p) => TrySendError::Full(p.into_message()),
TrySendError::Closed(p) => TrySendError::Closed(p.into_message()),
})
}
#[cfg_attr(feature = "tracing", tracing::instrument(skip(self)))]
pub fn close(&self, reason: DisconnectReason) {
self.send(Packet::Close).ok();
(self.close_fn)(self.id, reason);
}
pub fn is_closed(&self) -> bool {
self.internal_tx.is_closed()
}
pub async fn closed(&self) {
self.internal_tx.closed().await
}
pub fn emit_binary<B: Into<Bytes>>(&self, data: B) -> Result<(), TrySendError<Bytes>> {
if self.protocol == ProtocolVersion::V3 {
self.send(Packet::BinaryV3(data.into()))
} else {
self.send(Packet::Binary(data.into()))
}
.map_err(|e| match e {
TrySendError::Full(p) => TrySendError::Full(p.into_binary()),
TrySendError::Closed(p) => TrySendError::Closed(p.into_binary()),
})
}
#[inline]
pub fn emit_volatile(&self, msg: impl Into<Str>) -> bool {
self.send_volatile(smallvec![Packet::Message(msg.into())])
}
#[inline]
pub fn emit_binary_volatile<B: Into<Bytes>>(&self, data: B) -> bool {
if self.protocol == ProtocolVersion::V3 {
self.send_volatile(smallvec![Packet::BinaryV3(data.into())])
} else {
self.send_volatile(smallvec![Packet::Binary(data.into())])
}
}
#[inline]
pub fn emit_many_volatile(&self, msg: Str, data: VecDeque<Bytes>) -> bool {
let mut packets = SmallVec::with_capacity(1 + data.len());
packets.push(Packet::Message(msg));
for bin in data {
packets.push(Packet::Binary(bin));
}
self.send_volatile(packets)
}
pub(crate) fn send_volatile(&self, packets: PacketBuf) -> bool {
self.volatile_tx.send(Some(packets)).is_ok()
}
}
impl<D: Default + Send + Sync + 'static> std::fmt::Debug for Socket<D> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("Socket")
.field("sid", &self.id)
.field("protocol", &self.protocol)
.field("conn", &self.transport)
.field("internal_rx", &self.internal_rx)
.field("internal_tx", &self.internal_tx)
.field("heartbeat_rx", &self.heartbeat_rx)
.field("heartbeat_tx", &self.heartbeat_tx)
.field("cancellation_token", &self.cancellation_token)
.field("req_data", &self.req_parts)
.finish()
}
}
#[doc(hidden)]
#[cfg(feature = "__test_harness")]
impl<D> Drop for Socket<D>
where
D: Default + Send + Sync + 'static,
{
fn drop(&mut self) {
#[cfg(feature = "tracing")]
tracing::debug!("[sid={}] dropping socket", self.id);
}
}
#[doc(hidden)]
#[cfg(feature = "__test_harness")]
impl<D> Socket<D>
where
D: Default + Send + Sync + 'static,
{
pub fn new_dummy(
sid: Sid,
close_fn: Box<dyn Fn(Sid, DisconnectReason) + Send + Sync>,
) -> Arc<Socket<D>> {
let (s, mut rx) = Socket::new_dummy_piped(sid, close_fn, 1024);
tokio::spawn(async move {
while let Some(_el) = rx.recv().await {
#[cfg(feature = "tracing")]
tracing::debug!(?sid, ?_el, "emitting eio msg");
}
});
s
}
pub fn new_dummy_piped(
sid: Sid,
close_fn: Box<dyn Fn(Sid, DisconnectReason) + Send + Sync>,
buffer_size: usize,
) -> (Arc<Socket<D>>, tokio::sync::mpsc::Receiver<Packet>) {
let (internal_tx, internal_rx) = mpsc::channel(buffer_size);
let (heartbeat_tx, heartbeat_rx) = mpsc::channel(1);
let (volatile_tx, volatile_rx) = watch::channel(None);
let sock = Self {
id: sid,
protocol: ProtocolVersion::V4,
transport: AtomicU8::new(TransportType::Websocket as u8),
upgrading: AtomicBool::new(false),
internal_rx: Mutex::new(InternalRx::new(internal_rx, volatile_rx)),
internal_tx,
volatile_tx,
heartbeat_rx: Mutex::new(heartbeat_rx),
heartbeat_tx,
cancellation_token: CancellationToken::new(),
close_fn,
data: D::default(),
req_parts: http::Request::<()>::default().into_parts().0,
supports_binary: true,
};
let sock = Arc::new(sock);
let (tx, rx) = mpsc::channel(buffer_size);
let sock_clone = sock.clone();
tokio::spawn(async move {
let mut internal_rx = sock_clone.internal_rx.try_lock().unwrap();
while let Some(packets) = internal_rx.buffered_rx.recv().await {
for packet in packets {
tx.send(packet).await.unwrap();
}
}
});
(sock, rx)
}
}