diff --git a/Cargo.lock b/Cargo.lock index 25589f5..28940fa 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -256,6 +256,7 @@ dependencies = [ "axum", "bs58", "ed25519-dalek", + "futures-util", "rand 0.8.7", "rusqlite", "serde", @@ -313,6 +314,17 @@ version = "0.3.34" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "92d699e522242e69e3003b94ecc1f960f3a5e015aa7c5d7486e65ad01dd94f5e" +[[package]] +name = "futures-macro" +version = "0.3.34" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9fb9654ba8355388abeb8dcb4fc62f511300867002afc858860463bdd9fe0c44" +dependencies = [ + "proc-macro2", + "quote", + "syn 3.0.3", +] + [[package]] name = "futures-sink" version = "0.3.34" @@ -332,6 +344,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "0d50a92467f8ba5dd6e3ee5d4bd04d73ab2e4e1c44474a0674821dfce14b79bc" dependencies = [ "futures-core", + "futures-macro", "futures-sink", "futures-task", "pin-project-lite", diff --git a/Cargo.toml b/Cargo.toml index 219512b..2d0a27a 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -15,3 +15,4 @@ bs58 = "0.5.1" tower-http = { version = "0.7.0", features = ["fs", "cors"] } rusqlite = { version = "0.31", features = ["bundled"] } uuid = { version = "1.24.1", features = ["v4"] } +futures-util = "0.3.34" diff --git a/src/main.rs b/src/main.rs index 68a08b9..a4524bf 100644 --- a/src/main.rs +++ b/src/main.rs @@ -4,6 +4,7 @@ pub mod protocol; pub mod server; pub mod types; pub mod vc_server; +pub mod ws; use std::{ net::{IpAddr, Ipv4Addr, SocketAddr}, diff --git a/src/protocol/initialize.rs b/src/protocol/initialize.rs index 7b0140e..f5e4778 100644 --- a/src/protocol/initialize.rs +++ b/src/protocol/initialize.rs @@ -3,10 +3,9 @@ use std::{ time::{SystemTime, UNIX_EPOCH}, }; -use axum::extract::ws::WebSocket; use ed25519_dalek::{Signer, VerifyingKey}; -use crate::server::Server; +use crate::{server::Server, ws::EnclaveWebSocket}; use super::*; use crate::server::UserConnections; @@ -14,23 +13,21 @@ use crate::server::UserConnections; impl UserConnections { pub async fn initialize( server: &Arc, - mut socket: WebSocket, - ) -> anyhow::Result<(WebSocket, VerifyingKey, ClientMeta)> { + socket: Arc, + ) -> anyhow::Result<(Arc, VerifyingKey, ClientMeta)> { let Some(ServerMethod::Initialize { public_key: public_key_string, signature, timestamp, hostname, - }) = read_socket(&mut socket).await? + }) = socket.read().await? else { - send_socket( - &mut socket, - &ClientMethod::Error { + socket + .send(&ClientMethod::Error { error: Cow::Borrowed("Initialization required"), - }, - ) - .await?; + }) + .await?; return Err(anyhow::anyhow!( "Failed to initialize: Client sent the wrong method" @@ -40,23 +37,20 @@ impl UserConnections { let server_timestamp = SystemTime::now().duration_since(UNIX_EPOCH)?.as_millis() as u64; if server_timestamp.saturating_sub(timestamp) > 2000 { - send_socket( - &mut socket, - &ClientMethod::Error { + socket + .send(&ClientMethod::Error { error: Cow::Borrowed( "Timestamp doesn't match, make sure it's in secs and is (<= 2secs)", ), - }, - ) - .await?; + }) + .await?; return Err(anyhow::anyhow!("Client tampstamp wasn't correct")); } if hostname != server.config.public_hostname || !server.config.hostnames.contains(&hostname) { - send_socket( - &mut socket, + socket.send( &ClientMethod::Error { error: Cow::Owned(format!("Invalid Hostname, to avoid man-in-the-middle attacks, please use the correct hostname: {}", server.config.public_hostname)), }, @@ -67,13 +61,11 @@ impl UserConnections { } let Ok(public_key) = crate::crypto::from_string(&public_key_string) else { - send_socket( - &mut socket, - &ClientMethod::Error { + socket + .send(&ClientMethod::Error { error: Cow::Borrowed("Invalid public key"), - }, - ) - .await?; + }) + .await?; return Err(anyhow::anyhow!("Invalid public key")); }; @@ -85,21 +77,18 @@ impl UserConnections { ) .is_err() { - send_socket( - &mut socket, - &ClientMethod::Error { + socket + .send(&ClientMethod::Error { error: Cow::Borrowed("Invalid signature"), - }, - ) - .await?; + }) + .await?; return Err(anyhow::anyhow!("Invalid signature")); } { - send_socket( - &mut socket, - &ClientMethod::Initialized { + socket + .send(&ClientMethod::Initialized { public_key: crate::crypto::to_string(&server.key.verifying_key()), signature: crate::crypto::to_string_sig(&server.key.sign( format!("{server_timestamp}@{hostname}@{public_key_string}").as_bytes(), @@ -107,19 +96,16 @@ impl UserConnections { timestamp: server_timestamp, hostname, - }, - ) - .await?; + }) + .await?; } - let Some(ServerMethod::Meta(meta)) = read_socket(&mut socket).await? else { - send_socket( - &mut socket, - &ClientMethod::Error { + let Some(ServerMethod::Meta(meta)) = socket.read().await? else { + socket + .send(&ClientMethod::Error { error: Cow::Borrowed("Expected meta"), - }, - ) - .await?; + }) + .await?; return Err(anyhow::anyhow!( "Expected meta, client called another method" diff --git a/src/protocol/message.rs b/src/protocol/message.rs index ec96e43..91d7a1e 100644 --- a/src/protocol/message.rs +++ b/src/protocol/message.rs @@ -1,17 +1,15 @@ use crate::data::messages::{MessageData, StoredMessage}; -use crate::protocol::{ClientMethod, send_socket}; +use crate::protocol::ClientMethod; use crate::server::Server; -use axum::extract::ws::WebSocket; use ed25519_dalek::{Verifier, VerifyingKey}; use std::collections::HashMap; use std::sync::Arc; use std::time::{SystemTime, UNIX_EPOCH}; -use tokio::sync::Mutex; pub async fn send_message( server: &Arc, verifying_key: VerifyingKey, - _socket: &Arc>, + _socket: &Arc, message: MessageData, channel_id: String, ) -> anyhow::Result<()> { @@ -58,7 +56,7 @@ pub async fn send_message( pub async fn get_messages( server: &Arc, _verifying_key: VerifyingKey, - socket: &Arc>, + socket: &Arc, channel_id: String, chunk: u32, ) -> anyhow::Result<()> { @@ -68,13 +66,11 @@ pub async fn get_messages( .message_store .get_recent_messages(&channel_id, CHUNK_SIZE, chunk)?; - send_socket( - &mut *socket.lock().await, - &ClientMethod::Messages { + socket + .send(&ClientMethod::Messages { messages: HashMap::from([(channel_id, messages)]), - }, - ) - .await?; + }) + .await?; Ok(()) } diff --git a/src/protocol/mod.rs b/src/protocol/mod.rs index 7fbd942..f8e1814 100644 --- a/src/protocol/mod.rs +++ b/src/protocol/mod.rs @@ -1,9 +1,7 @@ use std::{borrow::Cow, collections::HashMap, sync::Arc}; -use axum::extract::ws::{Message, Utf8Bytes, WebSocket}; use ed25519_dalek::VerifyingKey; use serde::{Deserialize, Serialize}; -use tokio::sync::Mutex; use crate::{data::messages::StoredMessage, server::Server, types::ClientMeta}; @@ -110,22 +108,16 @@ pub enum ServerMethod { pub async fn read_loop( server: &Arc, verifying_key: VerifyingKey, - socket: &Arc>, + socket: &Arc, ) -> anyhow::Result<()> { - let mut socket_lock = socket.lock().await; - - while let Some(message) = read_socket(&mut *socket_lock).await? { - drop(socket_lock); - + while let Some(message) = socket.read().await? { match message { ServerMethod::Initialize { .. } => { - send_socket( - &mut *socket.lock().await, - &ClientMethod::Error { + socket + .send(&ClientMethod::Error { error: Cow::Borrowed("Already initialized"), - }, - ) - .await?; + }) + .await?; } #[allow(unused_variables)] @@ -181,14 +173,12 @@ pub async fn read_loop( .await .insert(pin, (verifying_key, channel_id.clone())); - send_socket( - &mut *socket.lock().await, - &ClientMethod::JoinVoice { + socket + .send(&ClientMethod::JoinVoice { channel_id: channel_id.clone(), pin, - }, - ) - .await?; + }) + .await?; } server @@ -199,49 +189,7 @@ pub async fn read_loop( .await?; } } - - socket_lock = socket.lock().await; } Ok(()) } - -pub async fn read_socket(socket: &mut WebSocket) -> anyhow::Result> { - match socket.recv().await.transpose()? { - Some(Message::Text(text)) => match serde_json::from_str(&text.to_string()) { - Ok(msg) => Ok(Some(msg)), - - Err(e) => { - send_socket( - socket, - &ClientMethod::Error { - error: Cow::Owned(format!("Unable to parse message: {e}")), - }, - ) - .await?; - - Ok(None) - } - }, - - Some(Message::Ping(v)) => { - socket.send(Message::Pong(v)).await?; - - Ok(None) - } - - Some(_) => Ok(None), - - None => Ok(None), - } -} - -pub async fn send_socket(socket: &mut WebSocket, message: &ClientMethod) -> anyhow::Result<()> { - socket - .send(Message::Text(Utf8Bytes::from(serde_json::to_string( - message, - )?))) - .await?; - - Ok(()) -} diff --git a/src/protocol/user.rs b/src/protocol/user.rs index 8bc0ac1..c56e76f 100644 --- a/src/protocol/user.rs +++ b/src/protocol/user.rs @@ -1,23 +1,18 @@ use std::sync::Arc; -use axum::extract::ws::WebSocket; use ed25519_dalek::VerifyingKey; -use tokio::sync::Mutex; -use crate::{ - protocol::{ClientMethod, send_socket}, - server::Server, -}; +use crate::{protocol::ClientMethod, server::Server}; pub async fn get_users( server: &Arc, _verifying_key: VerifyingKey, - socket: &Arc>, + socket: &Arc, pubkeys: Vec, ) -> anyhow::Result<()> { let users = server.user_store.get_users(&pubkeys).await?; - send_socket(&mut *socket.lock().await, &ClientMethod::Users { users }).await?; + socket.send(&ClientMethod::Users { users }).await?; Ok(()) } diff --git a/src/server.rs b/src/server.rs index c5bbc74..39895f4 100644 --- a/src/server.rs +++ b/src/server.rs @@ -19,8 +19,9 @@ use tokio::{ use crate::{ data::{config::Config, messages::MessageStore, users::UserMetaStore}, - protocol::{ClientMethod, read_loop, send_socket}, + protocol::{ClientMethod, read_loop}, types::ClientMeta, + ws::EnclaveWebSocket, }; pub struct VoiceConnection { @@ -33,7 +34,7 @@ pub struct UserConnections { pub meta: ClientMeta, pub counter: AtomicU16, pub public_key: VerifyingKey, - pub connections: Mutex>>>, + pub connections: Mutex>>, pub voice: Mutex>, } @@ -66,7 +67,7 @@ impl Server { let s = self.clone(); ws.on_upgrade(move |socket: WebSocket| async move { - match UserConnections::initialize(&s, socket).await { + match UserConnections::initialize(&s, Arc::new(EnclaveWebSocket::new(socket))).await { Ok((client, public_key, meta)) => { if let Err(e) = s .user_store @@ -76,8 +77,6 @@ impl Server { eprintln!("Failed to upsert client: {e}"); } - let client = Arc::new(Mutex::new(client)); - let mut clients_meta = s.clients.lock().await; let clients = clients_meta @@ -154,7 +153,7 @@ impl Server { impl UserConnections { pub async fn send(&self, message: &ClientMethod) -> anyhow::Result<()> { for (_, conn) in self.connections.lock().await.iter() { - send_socket(&mut *conn.lock().await, message).await?; + conn.send(message).await?; } Ok(()) @@ -162,7 +161,7 @@ impl UserConnections { pub async fn send_to(&self, id: u16, message: &ClientMethod) -> anyhow::Result { if let Some(conn) = self.connections.lock().await.get(&id) { - send_socket(&mut *conn.lock().await, message).await?; + conn.send(message).await?; Ok(true) } else { diff --git a/src/vc_server.rs b/src/vc_server.rs index 289c289..cabb20c 100644 --- a/src/vc_server.rs +++ b/src/vc_server.rs @@ -4,10 +4,7 @@ use anyhow::Context; use ed25519_dalek::VerifyingKey; use tokio::net::UdpSocket; -use crate::{ - protocol::{ClientMethod, send_socket}, - server::Server, -}; +use crate::{protocol::ClientMethod, server::Server}; use tokio::time::Instant; @@ -107,13 +104,11 @@ impl Server { if now.duration_since(voice.last_speaking_sent).as_millis() >= 600 { for conn in user.connections.lock().await.values() { - let _ = send_socket( - &mut *conn.lock().await, - &ClientMethod::Speaking { + let _ = conn + .send(&ClientMethod::Speaking { pubkey: crate::crypto::to_string(sender), - }, - ) - .await; + }) + .await; } voice.last_speaking_sent = now; diff --git a/src/ws.rs b/src/ws.rs new file mode 100644 index 0000000..de5e692 --- /dev/null +++ b/src/ws.rs @@ -0,0 +1,65 @@ +use std::borrow::Cow; + +use axum::extract::ws::{Message, Utf8Bytes, WebSocket}; +use futures_util::{ + SinkExt, StreamExt, + stream::{SplitSink, SplitStream}, +}; +use tokio::sync::Mutex; + +use crate::protocol::{ClientMethod, ServerMethod}; + +pub struct EnclaveWebSocket { + tx: Mutex>, + rx: Mutex>, +} + +impl EnclaveWebSocket { + pub fn new(ws: WebSocket) -> Self { + let (tx, rx) = ws.split(); + + Self { + tx: Mutex::new(tx), + rx: Mutex::new(rx), + } + } + + pub async fn read(&self) -> anyhow::Result> { + match self.rx.lock().await.next().await.transpose()? { + Some(Message::Text(text)) => match serde_json::from_str(&text.to_string()) { + Ok(msg) => Ok(Some(msg)), + + Err(e) => { + self.send(&ClientMethod::Error { + error: Cow::Owned(format!("Unable to parse message: {e}")), + }) + .await?; + + Ok(None) + } + }, + + Some(Message::Ping(v)) => { + self.tx.lock().await.send(Message::Pong(v)).await?; + + Ok(None) + } + + Some(_) => Ok(None), + + None => Ok(None), + } + } + + pub async fn send(&self, message: &ClientMethod) -> anyhow::Result<()> { + self.tx + .lock() + .await + .send(Message::Text(Utf8Bytes::from(serde_json::to_string( + message, + )?))) + .await?; + + Ok(()) + } +}