use base64::{Engine, engine::general_purpose};
use bytes::Bytes;
use engineioxide_core::{Sid, Str};
use serde::{Deserialize, Deserializer, Serialize, Serializer};
use smallvec::{SmallVec, smallvec};
use std::time::Duration;
use crate::TransportType;
use crate::config::EngineIoConfig;
use crate::errors::Error;
#[derive(Debug, Clone, PartialEq, PartialOrd)]
pub enum Packet {
Open(OpenPacket),
Close,
Ping,
Pong,
PingUpgrade,
PongUpgrade,
Message(Str),
Upgrade,
Noop,
Binary(Bytes),
BinaryV3(Bytes), }
impl Packet {
pub fn is_binary(&self) -> bool {
matches!(self, Packet::Binary(_) | Packet::BinaryV3(_))
}
pub(crate) fn into_message(self) -> Str {
match self {
Packet::Message(msg) => msg,
_ => panic!("Packet is not a message"),
}
}
pub(crate) fn into_binary(self) -> Bytes {
match self {
Packet::Binary(data) => data,
Packet::BinaryV3(data) => data,
_ => panic!("Packet is not a binary"),
}
}
pub(crate) fn get_size_hint(&self, b64: bool) -> usize {
match self {
Packet::Open(_) => 156, Packet::Close => 1,
Packet::Ping => 1,
Packet::Pong => 1,
Packet::PingUpgrade => 6,
Packet::PongUpgrade => 6,
Packet::Message(msg) => 1 + msg.len(),
Packet::Upgrade => 1,
Packet::Noop => 1,
Packet::Binary(data) => {
if b64 {
1 + base64::encoded_len(data.len(), true).unwrap_or(usize::MAX - 1)
} else {
1 + data.len()
}
}
Packet::BinaryV3(data) => {
if b64 {
2 + base64::encoded_len(data.len(), true).unwrap_or(usize::MAX - 2)
} else {
1 + data.len()
}
}
}
}
}
impl From<Packet> for String {
fn from(packet: Packet) -> String {
let len = packet.get_size_hint(true);
let mut buffer = String::with_capacity(len);
match packet {
Packet::Open(open) => {
buffer.push('0');
buffer.push_str(&serde_json::to_string(&open).unwrap());
}
Packet::Close => buffer.push('1'),
Packet::Ping => buffer.push('2'),
Packet::Pong => buffer.push('3'),
Packet::PingUpgrade => buffer.push_str("2probe"),
Packet::PongUpgrade => buffer.push_str("3probe"),
Packet::Message(msg) => {
buffer.push('4');
buffer.push_str(&msg);
}
Packet::Upgrade => buffer.push('5'),
Packet::Noop => buffer.push('6'),
Packet::Binary(data) => {
buffer.push('b');
general_purpose::STANDARD.encode_string(data, &mut buffer);
}
Packet::BinaryV3(data) => {
buffer.push_str("b4");
general_purpose::STANDARD.encode_string(data, &mut buffer);
}
};
buffer
}
}
impl From<Packet> for tokio_tungstenite::tungstenite::Utf8Bytes {
fn from(value: Packet) -> Self {
String::from(value).into()
}
}
impl From<Packet> for Bytes {
fn from(value: Packet) -> Self {
String::from(value).into()
}
}
impl TryFrom<Str> for Packet {
type Error = Error;
fn try_from(value: Str) -> Result<Self, Self::Error> {
let packet_type = value
.as_bytes()
.first()
.ok_or(Error::InvalidPacketType(None))?;
let is_upgrade = value.len() == 6 && &value[1..6] == "probe";
let res = match packet_type {
b'1' => Packet::Close,
b'2' if is_upgrade => Packet::PingUpgrade,
b'2' => Packet::Ping,
b'3' if is_upgrade => Packet::PongUpgrade,
b'3' => Packet::Pong,
b'4' => Packet::Message(value.slice(1..)),
b'5' => Packet::Upgrade,
b'6' => Packet::Noop,
b'b' if value.as_bytes().get(1) == Some(&b'4') => Packet::BinaryV3(
general_purpose::STANDARD
.decode(value.slice(2..).as_bytes())?
.into(),
),
b'b' => Packet::Binary(
general_purpose::STANDARD
.decode(value.slice(1..).as_bytes())?
.into(),
),
c => Err(Error::InvalidPacketType(Some(*c as char)))?,
};
Ok(res)
}
}
impl TryFrom<tokio_tungstenite::tungstenite::Utf8Bytes> for Packet {
type Error = Error;
fn try_from(value: tokio_tungstenite::tungstenite::Utf8Bytes) -> Result<Self, Self::Error> {
Packet::try_from(unsafe { Str::from_bytes_unchecked(value.into()) })
}
}
impl TryFrom<String> for Packet {
type Error = Error;
fn try_from(value: String) -> Result<Self, Self::Error> {
Packet::try_from(Str::from(value))
}
}
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, PartialOrd)]
#[serde(rename_all = "camelCase")]
pub struct OpenPacket {
sid: Sid,
upgrades: SmallVec<[TransportType; 1]>,
#[serde(
serialize_with = "serialize_duration_millis",
deserialize_with = "deserialize_duration_from_millis"
)]
ping_interval: Duration,
#[serde(
serialize_with = "serialize_duration_millis",
deserialize_with = "deserialize_duration_from_millis"
)]
ping_timeout: Duration,
max_payload: u64,
}
pub fn serialize_duration_millis<S>(duration: &Duration, serializer: S) -> Result<S::Ok, S::Error>
where
S: Serializer,
{
serializer.serialize_u64(duration.as_millis() as u64)
}
pub fn deserialize_duration_from_millis<'de, D>(deserializer: D) -> Result<Duration, D::Error>
where
D: Deserializer<'de>,
{
let millis = u64::deserialize(deserializer)?;
Ok(Duration::from_millis(millis))
}
impl OpenPacket {
pub fn new(transport: TransportType, sid: Sid, config: &EngineIoConfig) -> Self {
let upgrades = if transport == TransportType::Polling {
smallvec![TransportType::Websocket]
} else {
smallvec![]
};
OpenPacket {
sid,
upgrades,
ping_interval: config.ping_interval,
ping_timeout: config.ping_timeout,
max_payload: config.max_payload,
}
}
}
#[cfg(test)]
mod tests {
use crate::config::EngineIoConfig;
use super::*;
use std::{convert::TryInto, time::Duration};
#[test]
fn test_open_packet() {
let sid = Sid::new();
let packet = Packet::Open(OpenPacket::new(
TransportType::Polling,
sid,
&EngineIoConfig::default(),
));
let packet_str: String = packet.into();
assert_eq!(
packet_str,
format!(
"0{{\"sid\":\"{sid}\",\"upgrades\":[\"websocket\"],\"pingInterval\":25000,\"pingTimeout\":20000,\"maxPayload\":100000}}"
)
);
}
#[test]
fn test_message_packet() {
let packet = Packet::Message("hello".into());
let packet_str: String = packet.into();
assert_eq!(packet_str, "4hello");
}
#[test]
fn test_message_packet_deserialize() {
let packet_str = "4hello".to_string();
let packet: Packet = packet_str.try_into().unwrap();
assert_eq!(packet, Packet::Message("hello".into()));
}
#[test]
fn test_binary_packet() {
let packet = Packet::Binary(vec![1, 2, 3].into());
let packet_str: String = packet.into();
assert_eq!(packet_str, "bAQID");
}
#[test]
fn test_binary_packet_deserialize() {
let packet_str = "bAQID".to_string();
let packet: Packet = packet_str.try_into().unwrap();
assert_eq!(packet, Packet::Binary(vec![1, 2, 3].into()));
}
#[test]
fn test_binary_packet_v3() {
let packet = Packet::BinaryV3(vec![1, 2, 3].into());
let packet_str: String = packet.into();
assert_eq!(packet_str, "b4AQID");
}
#[test]
fn test_binary_packet_v3_deserialize() {
let packet_str = "b4AQID".to_string();
let packet: Packet = packet_str.try_into().unwrap();
assert_eq!(packet, Packet::BinaryV3(vec![1, 2, 3].into()));
}
#[test]
fn test_packet_get_size_hint() {
let open = OpenPacket::new(
TransportType::Polling,
Sid::new(),
&EngineIoConfig {
max_buffer_size: usize::MAX,
max_payload: u64::MAX,
ping_interval: Duration::MAX,
ping_timeout: Duration::MAX,
transports: TransportType::Polling as u8 | TransportType::Websocket as u8,
..Default::default()
},
);
let size = serde_json::to_string(&open).unwrap().len();
let packet = Packet::Open(open);
assert_eq!(packet.get_size_hint(false), size);
let packet = Packet::Close;
assert_eq!(packet.get_size_hint(false), 1);
let packet = Packet::Ping;
assert_eq!(packet.get_size_hint(false), 1);
let packet = Packet::Pong;
assert_eq!(packet.get_size_hint(false), 1);
let packet = Packet::PingUpgrade;
assert_eq!(packet.get_size_hint(false), 6);
let packet = Packet::PongUpgrade;
assert_eq!(packet.get_size_hint(false), 6);
let packet = Packet::Message("hello".into());
assert_eq!(packet.get_size_hint(false), 6);
let packet = Packet::Upgrade;
assert_eq!(packet.get_size_hint(false), 1);
let packet = Packet::Noop;
assert_eq!(packet.get_size_hint(false), 1);
let packet = Packet::Binary(vec![1, 2, 3].into());
assert_eq!(packet.get_size_hint(false), 4);
assert_eq!(packet.get_size_hint(true), 5);
let packet = Packet::BinaryV3(vec![1, 2, 3].into());
assert_eq!(packet.get_size_hint(false), 4);
assert_eq!(packet.get_size_hint(true), 6);
}
}