Commonize message serialization logic

This commit is contained in:
Josiah Glosson
2025-01-23 12:46:42 -06:00
parent 0ede08415f
commit 4d24c17c06
10 changed files with 141 additions and 148 deletions

1
Cargo.lock generated
View File

@@ -6771,7 +6771,6 @@ dependencies = [
name = "rust-common"
version = "0.1.0"
dependencies = [
"actix-ws",
"chrono",
"either",
"rand 0.8.5",

View File

@@ -122,66 +122,71 @@ pub async fn ws_init(
actix_web::rt::spawn(async move {
// receive messages from websocket
while let Some(msg) = stream.next().await {
if msg.is_err() {
continue;
}
match ClientToServerMessage::deserialize(msg.unwrap()) {
Ok(Either::Left(message)) => {
match message {
ClientToServerMessage::StatusUpdate {
profile_name,
} => {
if let Some(mut pair) =
db.auth_sockets.get_mut(&user.id)
{
let (status, _) = pair.value_mut();
if status
.profile_name
.as_ref()
.map(|x| x.len() > 64)
.unwrap_or(false)
{
return;
}
status.profile_name = profile_name;
status.last_update = Utc::now();
let user_status = status.clone();
// We drop the pair to avoid holding the lock for too long
drop(pair);
let _ = broadcast_friends(
user.id,
ServerToClientMessage::StatusUpdate {
status: user_status,
},
&pool,
&db,
None,
)
.await;
}
}
ClientToServerMessage::SocketOpen
| ClientToServerMessage::SocketClose { .. }
| ClientToServerMessage::SocketSend { .. } => todo!(),
}
let message = match msg {
Ok(Message::Text(text)) => {
ClientToServerMessage::deserialize(Either::Left(&text))
}
Ok(Either::Right(Message::Close(_))) => {
Ok(Message::Binary(bytes)) => {
ClientToServerMessage::deserialize(Either::Right(&bytes))
}
Ok(Message::Close(_)) => {
let _ = close_socket(user.id, &pool, &db).await;
continue;
}
Ok(Either::Right(Message::Ping(msg))) => {
Ok(Message::Ping(msg)) => {
if let Some(socket) = db.auth_sockets.get(&user.id) {
let (_, socket) = socket.value();
let _ = socket.clone().pong(&msg).await;
}
continue;
}
_ => {}
_ => continue,
};
if message.is_err() {
continue;
}
match message.unwrap() {
ClientToServerMessage::StatusUpdate { profile_name } => {
if let Some(mut pair) = db.auth_sockets.get_mut(&user.id) {
let (status, _) = pair.value_mut();
if status
.profile_name
.as_ref()
.map(|x| x.len() > 64)
.unwrap_or(false)
{
return;
}
status.profile_name = profile_name;
status.last_update = Utc::now();
let user_status = status.clone();
// We drop the pair to avoid holding the lock for too long
drop(pair);
let _ = broadcast_friends(
user.id,
ServerToClientMessage::StatusUpdate {
status: user_status,
},
&pool,
&db,
None,
)
.await;
}
}
ClientToServerMessage::SocketOpen
| ClientToServerMessage::SocketClose { .. }
| ClientToServerMessage::SocketSend { .. } => todo!(),
}
}
@@ -216,7 +221,14 @@ pub async fn broadcast_friends(
if let Some(socket) = sockets.auth_sockets.get(&friend_id.into()) {
let (_, socket) = socket.value();
let _ = message.send(&mut socket.clone()).await; // FIXME Probably shouldn't swallow this error
// FIXME Probably shouldn't swallow sending errors
let _ = match message.serialize() {
Ok(Either::Left(text)) => socket.clone().text(text).await,
Ok(Either::Right(bytes)) => {
socket.clone().binary(bytes).await
}
Err(_) => Ok(()), // TODO: Maybe should log these? Though it is the backend
};
}
}
}

View File

@@ -13,6 +13,11 @@ pub enum ErrorKind {
#[error("Serialization error (JSON): {0}")]
JSONError(#[from] serde_json::Error),
#[error("Serialization error (websocket): {0}")]
WebsocketSerializationError(
#[from] rust_common::networking::serialization::SerializationError,
),
#[error("Error parsing UUID: {0}")]
UUIDError(#[from] uuid::Error),

View File

@@ -10,6 +10,7 @@ use async_tungstenite::tungstenite::Message;
use async_tungstenite::WebSocketStream;
use chrono::{DateTime, Utc};
use dashmap::DashMap;
use either::Either;
use futures::stream::SplitSink;
use futures::{SinkExt, StreamExt};
use reqwest::header::HeaderValue;
@@ -108,21 +109,16 @@ impl FriendsSocket {
while let Some(msg_result) = read_stream.next().await {
match msg_result {
Ok(msg) => {
// TODO: Make wire package work with this other library
let server_message = match msg {
Message::Text(text) => {
serde_json::from_str::<
ServerToClientMessage,
>(
&text
ServerToClientMessage::deserialize(
Either::Left(&text),
)
.ok()
}
Message::Binary(bytes) => {
serde_json::from_slice::<
ServerToClientMessage,
>(
&bytes
ServerToClientMessage::deserialize(
Either::Right(&bytes),
)
.ok()
}
@@ -257,16 +253,8 @@ impl FriendsSocket {
&self,
profile_name: Option<String>,
) -> crate::Result<()> {
let mut write_lock = self.write.write().await;
if let Some(ref mut write_half) = *write_lock {
write_half
.send(Message::Text(serde_json::to_string(
&ClientToServerMessage::StatusUpdate { profile_name },
)?))
.await?;
}
Ok(())
self.send_message(ClientToServerMessage::StatusUpdate { profile_name })
.await
}
#[tracing::instrument(skip_all)]
@@ -334,4 +322,22 @@ impl FriendsSocket {
Ok(())
}
#[tracing::instrument(skip(self))]
async fn send_message(
&self,
message: ClientToServerMessage,
) -> crate::Result<()> {
let serialized = match message.serialize()? {
Either::Left(text) => Message::text(text),
Either::Right(bytes) => Message::binary(bytes),
};
let mut write_lock = self.write.write().await;
if let Some(ref mut write_half) = *write_lock {
write_half.send(serialized).await?;
}
Ok(())
}
}

View File

@@ -12,8 +12,6 @@ serde_cbor = "0.11"
uuid = { version = "1.12", features = ["serde"] }
chrono = { version = "0.4", features = ["serde"] }
actix-ws = "0.3"
thiserror = "2.0"
rand = "0.8"
either = "1.13"

View File

@@ -3,7 +3,7 @@ use crate::users::UserStatus;
use serde::{Deserialize, Serialize};
use uuid::Uuid;
#[derive(Serialize, Deserialize)]
#[derive(Debug, Serialize, Deserialize)]
#[serde(tag = "type", rename_all = "snake_case")]
pub enum ClientToServerMessage {
StatusUpdate {
@@ -21,7 +21,7 @@ pub enum ClientToServerMessage {
},
}
#[derive(Serialize, Deserialize)]
#[derive(Debug, Serialize, Deserialize)]
#[serde(tag = "type", rename_all = "snake_case")]
pub enum ServerToClientMessage {
StatusUpdate {

View File

@@ -1,2 +1,2 @@
pub mod message;
pub mod wire;
pub mod serialization;

View File

@@ -0,0 +1,45 @@
use super::message::{ClientToServerMessage, ServerToClientMessage};
use either::Either;
use thiserror::Error;
#[derive(Debug, Error)]
pub enum SerializationError {
#[error("Failed to (de)serialize message: {0}")]
SerializationFailed(#[from] serde_json::Error),
#[error("Failed to (de)serialize binary message: {0}")]
BinarySerializationFailed(#[from] serde_cbor::Error),
}
macro_rules! message_serialization {
($message_enum:ty, $binary_pattern:pat) => {
impl $message_enum {
pub fn serialize(
&self,
) -> Result<Either<String, Vec<u8>>, SerializationError> {
Ok(match self {
$binary_pattern => Either::Right(serde_cbor::to_vec(self)?),
_ => Either::Left(serde_json::to_string(self)?),
})
}
pub fn deserialize(
msg: Either<&str, &[u8]>,
) -> Result<Self, SerializationError> {
Ok(match msg {
Either::Left(text) => serde_json::from_str(&text)?,
Either::Right(bytes) => serde_cbor::from_slice(&bytes)?,
})
}
}
};
}
message_serialization!(
ClientToServerMessage,
ClientToServerMessage::SocketSend { .. }
);
message_serialization!(
ServerToClientMessage,
ServerToClientMessage::SocketData { .. }
);

View File

@@ -1,72 +0,0 @@
use super::message::{ClientToServerMessage, ServerToClientMessage};
use actix_ws::{Closed, Message, Session};
use either::Either;
use thiserror::Error;
#[derive(Debug, Error)]
pub enum DeserializationError {
#[error("Failed to deserialize message: {0}")]
DeserializationFailed(#[from] serde_json::Error),
#[error("Failed to deserialize binary message: {0}")]
BinaryDeserializationFailed(#[from] serde_cbor::Error),
}
#[derive(Debug, Error)]
pub enum SendError {
#[error("Failed to serialize message: {0}")]
SerializationFailed(#[from] serde_json::Error),
#[error("Failed to serialize binary message: {0}")]
BinarySerializationFailed(#[from] serde_cbor::Error),
#[error("Websocket closed")]
Closed,
}
impl From<Closed> for SendError {
fn from(_: Closed) -> Self {
SendError::Closed
}
}
macro_rules! message_wire {
($message_enum:ty, $binary_pattern:pat) => {
impl $message_enum {
pub fn deserialize(
msg: Message,
) -> Result<Either<Self, Message>, DeserializationError> {
Ok(match msg {
Message::Text(text) => {
Either::Left(serde_json::from_str(&text)?)
}
Message::Binary(bytes) => {
Either::Left(serde_cbor::from_slice(&bytes)?)
}
other => Either::Right(other),
})
}
pub async fn send(
&self,
session: &mut Session,
) -> Result<(), SendError> {
Ok(match self {
$binary_pattern => {
session.binary(serde_cbor::to_vec(self)?).await?
}
_ => session.text(serde_json::to_string(self)?).await?,
})
}
}
};
}
message_wire!(
ClientToServerMessage,
ClientToServerMessage::SocketSend { .. }
);
message_wire!(
ServerToClientMessage,
ServerToClientMessage::SocketData { .. }
);

View File

@@ -7,7 +7,7 @@ use serde::{Deserialize, Serialize};
#[serde(into = "Base62Id")]
pub struct UserId(pub u64);
#[derive(Serialize, Deserialize, Clone)]
#[derive(Debug, Serialize, Deserialize, Clone)]
pub struct UserStatus {
pub user_id: UserId,
pub profile_name: Option<String>,