diff --git a/Cargo.lock b/Cargo.lock index 741985a34..e5f189917 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -6771,7 +6771,6 @@ dependencies = [ name = "rust-common" version = "0.1.0" dependencies = [ - "actix-ws", "chrono", "either", "rand 0.8.5", diff --git a/apps/labrinth/src/routes/internal/statuses.rs b/apps/labrinth/src/routes/internal/statuses.rs index dff4af3f5..dd4f8dd18 100644 --- a/apps/labrinth/src/routes/internal/statuses.rs +++ b/apps/labrinth/src/routes/internal/statuses.rs @@ -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 + }; } } } diff --git a/packages/app-lib/src/error.rs b/packages/app-lib/src/error.rs index fc068e305..2a16e4bde 100644 --- a/packages/app-lib/src/error.rs +++ b/packages/app-lib/src/error.rs @@ -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), diff --git a/packages/app-lib/src/state/friends.rs b/packages/app-lib/src/state/friends.rs index 5cc23d48b..acdbcbd63 100644 --- a/packages/app-lib/src/state/friends.rs +++ b/packages/app-lib/src/state/friends.rs @@ -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, ) -> 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(()) + } } diff --git a/packages/rust-common/Cargo.toml b/packages/rust-common/Cargo.toml index d68110ec7..2b3271f4f 100644 --- a/packages/rust-common/Cargo.toml +++ b/packages/rust-common/Cargo.toml @@ -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" diff --git a/packages/rust-common/src/networking/message.rs b/packages/rust-common/src/networking/message.rs index 983d269a8..f59301a6b 100644 --- a/packages/rust-common/src/networking/message.rs +++ b/packages/rust-common/src/networking/message.rs @@ -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 { diff --git a/packages/rust-common/src/networking/mod.rs b/packages/rust-common/src/networking/mod.rs index b4d18c775..b7855d592 100644 --- a/packages/rust-common/src/networking/mod.rs +++ b/packages/rust-common/src/networking/mod.rs @@ -1,2 +1,2 @@ pub mod message; -pub mod wire; +pub mod serialization; diff --git a/packages/rust-common/src/networking/serialization.rs b/packages/rust-common/src/networking/serialization.rs new file mode 100644 index 000000000..f94fec3a9 --- /dev/null +++ b/packages/rust-common/src/networking/serialization.rs @@ -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>, 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 { + 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 { .. } +); diff --git a/packages/rust-common/src/networking/wire.rs b/packages/rust-common/src/networking/wire.rs deleted file mode 100644 index d194e6646..000000000 --- a/packages/rust-common/src/networking/wire.rs +++ /dev/null @@ -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 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, 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 { .. } -); diff --git a/packages/rust-common/src/users.rs b/packages/rust-common/src/users.rs index 5ca7d55f5..2fe087283 100644 --- a/packages/rust-common/src/users.rs +++ b/packages/rust-common/src/users.rs @@ -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,