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/signature.rs b/src/crypto.rs similarity index 100% rename from src/signature.rs rename to src/crypto.rs diff --git a/src/data/config.rs b/src/data/config.rs index ea4141b..5a8693e 100644 --- a/src/data/config.rs +++ b/src/data/config.rs @@ -19,18 +19,32 @@ impl Config { name: "New Server".to_string(), description: String::new(), - channels: vec![Channel { - id: "text-channels".to_string(), - name: "Text Channels".to_string(), + channels: vec![ + Channel { + id: "text-channels".to_string(), + name: "Text Channels".to_string(), - data: ChannelKind::Category { - channels: vec![Channel { - id: "general".to_string(), - name: "General".to_string(), - data: ChannelKind::Text, - }], + data: ChannelKind::Category { + channels: vec![Channel { + id: "general".to_string(), + name: "General".to_string(), + data: ChannelKind::Text, + }], + }, }, - }], + Channel { + id: "voice-channels".to_string(), + name: "Voice Channels".to_string(), + + data: ChannelKind::Category { + channels: vec![Channel { + id: "vc-1".to_string(), + name: "VC 1".to_string(), + data: ChannelKind::Voice { max_users: 255 }, + }], + }, + }, + ], }, port: 3415, diff --git a/src/main.rs b/src/main.rs index bacef3d..a4524bf 100644 --- a/src/main.rs +++ b/src/main.rs @@ -1,8 +1,10 @@ +pub mod crypto; pub mod data; pub mod protocol; pub mod server; -pub mod signature; pub mod types; +pub mod vc_server; +pub mod ws; use std::{ net::{IpAddr, Ipv4Addr, SocketAddr}, @@ -26,6 +28,8 @@ use crate::server::Server; async fn main() -> anyhow::Result<()> { let server = Server::new().await?; + let udp_server = tokio::spawn(server.clone().start_udp_server()); + let cors = CorsLayer::new() .allow_origin(Any) .allow_methods(Any) @@ -46,6 +50,8 @@ async fn main() -> anyhow::Result<()> { axum::serve(listener, app).await?; + udp_server.abort(); + Ok(()) } diff --git a/src/protocol/initialize.rs b/src/protocol/initialize.rs index 650cc48..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)), }, @@ -66,14 +60,12 @@ impl UserConnections { return Err(anyhow::anyhow!("Client's hostname wasn't correct")); } - let Ok(public_key) = crate::signature::from_string(&public_key_string) else { - send_socket( - &mut socket, - &ClientMethod::Error { + let Ok(public_key) = crate::crypto::from_string(&public_key_string) else { + socket + .send(&ClientMethod::Error { error: Cow::Borrowed("Invalid public key"), - }, - ) - .await?; + }) + .await?; return Err(anyhow::anyhow!("Invalid public key")); }; @@ -81,45 +73,39 @@ impl UserConnections { if public_key .verify_strict( format!("{timestamp}@{hostname}").as_bytes(), - &crate::signature::from_string_sig(&signature)?, + &crate::crypto::from_string_sig(&signature)?, ) .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 { - public_key: crate::signature::to_string(&server.key.verifying_key()), - signature: crate::signature::to_string_sig(&server.key.sign( + 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(), )), 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 3c4b2c8..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<()> { @@ -24,13 +22,13 @@ pub async fn send_message( ); } - let server_pubkey_string = crate::signature::to_string(&server.key.verifying_key()); + let server_pubkey_string = crate::crypto::to_string(&server.key.verifying_key()); let signed_string = format!( "{}@{}@{}", message.timestamp, server_pubkey_string, message.content ); - let signature = crate::signature::from_string_sig(&message.signature) + let signature = crate::crypto::from_string_sig(&message.signature) .map_err(|_| anyhow::anyhow!("Invalid signature encoding"))?; verifying_key @@ -39,7 +37,7 @@ pub async fn send_message( let stored = StoredMessage { id: uuid::Uuid::new_v4().to_string(), - author: crate::signature::to_string(&verifying_key), + author: crate::crypto::to_string(&verifying_key), is_edited: false, data: message, }; @@ -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(()) } @@ -92,18 +88,18 @@ pub async fn edit_message( .get_message(&channel_id, &message_id)? .ok_or_else(|| anyhow::anyhow!("Message not found"))?; - let author_pubkey = crate::signature::to_string(&verifying_key); + let author_pubkey = crate::crypto::to_string(&verifying_key); if existing.author != author_pubkey { anyhow::bail!("Not authorized to edit this message"); } - let server_pubkey_string = crate::signature::to_string(&server.key.verifying_key()); + let server_pubkey_string = crate::crypto::to_string(&server.key.verifying_key()); let signed_string = format!( "{}@{}@{}", existing.data.timestamp, server_pubkey_string, new_content ); - let signature = crate::signature::from_string_sig(&new_signature) + let signature = crate::crypto::from_string_sig(&new_signature) .map_err(|_| anyhow::anyhow!("Invalid signature encoding"))?; verifying_key @@ -146,7 +142,7 @@ pub async fn delete_message( .get_message(&channel_id, &message_id)? .ok_or_else(|| anyhow::anyhow!("Message not found"))?; - let author_pubkey = crate::signature::to_string(&verifying_key); + let author_pubkey = crate::crypto::to_string(&verifying_key); if existing.author != author_pubkey { anyhow::bail!("Not authorized to delete this message"); } diff --git a/src/protocol/mod.rs b/src/protocol/mod.rs index d526c29..9cb3a1b 100644 --- a/src/protocol/mod.rs +++ b/src/protocol/mod.rs @@ -1,15 +1,14 @@ 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}; pub mod initialize; pub mod message; pub mod user; +pub mod voice; #[derive(Debug, Clone, Serialize, Deserialize)] #[serde(tag = "method")] @@ -43,6 +42,25 @@ pub enum ClientMethod { channel_id: String, message_id: String, }, + + JoinVoice { + channel_id: String, + pin: u64, + }, + + UserJoinedVoice { + channel_id: String, + pubkey: String, + }, + + UserLeftVoice { + channel_id: String, + pubkey: String, + }, + + Speaking { + pubkey: String, + }, } #[derive(Debug, Clone, Serialize, Deserialize)] @@ -87,27 +105,27 @@ pub enum ServerMethod { message_id: String, channel_id: String, }, + + JoinVoice { + channel_id: String, + }, + + LeaveVoice, } 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)] @@ -152,50 +170,16 @@ pub async fn read_loop( ServerMethod::GetUsers { pubkeys } => { user::get_users(server, verifying_key, socket, pubkeys).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) + ServerMethod::JoinVoice { channel_id } => { + voice::join(server, verifying_key, socket, channel_id).await?; } - }, - Some(Message::Ping(v)) => { - socket.send(Message::Pong(v)).await?; - - Ok(None) + ServerMethod::LeaveVoice => { + voice::leave(server, verifying_key).await?; + } } - - 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/protocol/voice.rs b/src/protocol/voice.rs new file mode 100644 index 0000000..563225b --- /dev/null +++ b/src/protocol/voice.rs @@ -0,0 +1,67 @@ +use std::sync::Arc; + +use ed25519_dalek::VerifyingKey; + +use crate::{protocol::ClientMethod, server::Server}; + +pub async fn join( + server: &Arc, + verifying_key: VerifyingKey, + socket: &Arc, + channel_id: String, +) -> anyhow::Result<()> { + { + let pin = rand::random::() % (1 << 53); + + server + .voice_pins + .lock() + .await + .insert(pin, (verifying_key, channel_id.clone())); + + socket + .send(&ClientMethod::JoinVoice { + channel_id: channel_id.clone(), + pin, + }) + .await?; + } + + server + .broadcast(&ClientMethod::UserJoinedVoice { + channel_id, + pubkey: crate::crypto::to_string(&verifying_key), + }) + .await?; + + Ok(()) +} + +pub async fn leave(server: &Arc, verifying_key: VerifyingKey) -> anyhow::Result<()> { + let Some(channel_id) = ({ + let clients = server.clients.lock().await; + + let Some(client) = clients.get(&verifying_key) else { + return Ok(()); + }; + + client.voice.lock().await.take().map(|v| v.channel_id) + }) else { + return Ok(()); + }; + + server + .voice_pins + .lock() + .await + .retain(|_, v| v.0 != verifying_key); + + server + .broadcast(&ClientMethod::UserLeftVoice { + channel_id, + pubkey: crate::crypto::to_string(&verifying_key), + }) + .await?; + + Ok(()) +} diff --git a/src/server.rs b/src/server.rs index c6ca1f3..39895f4 100644 --- a/src/server.rs +++ b/src/server.rs @@ -1,5 +1,6 @@ use std::{ collections::HashMap, + net::SocketAddr, path::PathBuf, sync::{Arc, atomic::AtomicU16}, }; @@ -9,37 +10,54 @@ use axum::{ response::Response, }; use ed25519_dalek::{SigningKey, VerifyingKey}; -use tokio::{sync::Mutex, task::JoinSet}; +use tokio::{ + net::UdpSocket, + sync::{Mutex, OnceCell}, + task::JoinSet, + time::Instant, +}; 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 { + pub addr: SocketAddr, + pub channel_id: String, + pub last_speaking_sent: Instant, +} + pub struct UserConnections { pub meta: ClientMeta, pub counter: AtomicU16, pub public_key: VerifyingKey, - pub connections: Mutex>>>, + pub connections: Mutex>>, + pub voice: Mutex>, } pub struct Server { pub key: SigningKey, pub config: Config, pub clients: Mutex>>, + pub voice_pins: Mutex>, pub message_store: MessageStore, pub user_store: UserMetaStore, + pub voice_socket: OnceCell, } impl Server { pub async fn new() -> anyhow::Result> { Ok(Arc::new(Self { - key: crate::signature::get().await?, + key: crate::crypto::get().await?, config: Config::get().await?, clients: Mutex::new(HashMap::new()), + voice_pins: Mutex::new(HashMap::new()), message_store: MessageStore::new(PathBuf::from("messages"))?, user_store: UserMetaStore::new(PathBuf::from("users.db"))?, + voice_socket: OnceCell::new(), })) } } @@ -49,18 +67,16 @@ 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 - .upsert_user(&crate::signature::to_string(&public_key), &meta) + .upsert_user(&crate::crypto::to_string(&public_key), &meta) .await { 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 @@ -71,6 +87,7 @@ impl Server { public_key: public_key, counter: AtomicU16::new(0), connections: Mutex::new(HashMap::new()), + voice: Mutex::new(None), }) }) .clone(); @@ -136,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(()) @@ -144,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/types.rs b/src/types.rs index 9a0ef5a..f35f431 100644 --- a/src/types.rs +++ b/src/types.rs @@ -19,8 +19,9 @@ pub struct ServerMeta { #[serde(tag = "kind")] #[serde(rename_all = "camelCase")] pub enum ChannelKind { - Text, Category { channels: Vec }, + Voice { max_users: u8 }, + Text, } #[derive(Debug, Clone, Serialize, Deserialize)] diff --git a/src/vc_server.rs b/src/vc_server.rs new file mode 100644 index 0000000..5a0a946 --- /dev/null +++ b/src/vc_server.rs @@ -0,0 +1,143 @@ +use std::{net::SocketAddr, sync::Arc}; + +use anyhow::Context; +use ed25519_dalek::VerifyingKey; +use tokio::net::UdpSocket; + +use crate::{protocol::ClientMethod, server::Server}; + +use tokio::time::Instant; + +impl Server { + pub async fn start_udp_server(self: Arc) -> anyhow::Result<()> { + let socket = UdpSocket::bind(("0.0.0.0", self.config.port)).await?; + + self.voice_socket + .set(socket) + .map_err(|_| anyhow::anyhow!("UDP server already started"))?; + + let mut buf = [0u8; 4096]; + + eprintln!("[vc] UDP server listening on port {}", self.config.port); + + loop { + let (len, addr) = self.get_voice_socket()?.recv_from(&mut buf).await?; + + if len < 8 { + eprintln!("[vc] dropping packet too short for pin"); + continue; + } + + let pin_bytes: [u8; 8] = buf[..8].try_into().unwrap(); + let pin = u64::from_be_bytes(pin_bytes); + + let (sender_pubkey, channel_id, payload) = { + let mut pins = self.voice_pins.lock().await; + + if let Some((pubkey, channel_id)) = pins.remove(&pin) { + (pubkey, channel_id, &buf[8..len]) + } else { + drop(pins); + + match self.find_voice_sender(&addr).await { + Some((pubkey, channel_id)) => (pubkey, channel_id, &buf[..]), + None => { + continue; + } + } + } + }; + + let clients = self.clients.lock().await; + let Some(user) = clients.get(&sender_pubkey).cloned() else { + eprintln!("[vc] sender pubkey not found in clients, dropping"); + continue; + }; + drop(clients); + + let mut voice = user.voice.lock().await; + + if let Some(voice) = &mut *voice { + voice.addr = addr; + } else { + *voice = Some(crate::server::VoiceConnection { + addr, + channel_id: channel_id.clone(), + last_speaking_sent: Instant::now(), + }); + } + + drop(voice); + + self.relay_voice(&sender_pubkey, &channel_id, payload) + .await?; + } + } + + /// Looks up which known voice participant a UDP address belongs to, + /// for packets arriving after the initial pin-bearing packet. + async fn find_voice_sender(&self, addr: &SocketAddr) -> Option<(VerifyingKey, String)> { + let clients = self.clients.lock().await; + + for (pubkey, user) in clients.iter() { + if let Some(voice) = &*user.voice.lock().await { + if *addr == voice.addr { + return Some((*pubkey, voice.channel_id.clone())); + } + } + } + + None + } + + /// Sends `payload` to every voice participant currently in `channel_id`. + async fn relay_voice( + &self, + sender: &VerifyingKey, + channel_id: &str, + payload: &[u8], + ) -> anyhow::Result<()> { + let clients = self.clients.lock().await; + + for (_pubkey, user) in clients.iter() { + let Some(voice) = &mut *user.voice.lock().await else { + continue; + }; + + if channel_id != voice.channel_id { + continue; + } + + let now = Instant::now(); + + if now.duration_since(voice.last_speaking_sent).as_millis() >= 600 { + for conn in user.connections.lock().await.values() { + conn.send(&ClientMethod::Speaking { + pubkey: crate::crypto::to_string(sender), + }) + .await?; + } + + voice.last_speaking_sent = now; + } + + let _ = self.udp_send_to(&voice.addr, payload).await; + } + + Ok(()) + } + + pub fn get_voice_socket(&self) -> anyhow::Result<&UdpSocket> { + self.voice_socket + .get() + .context("Failed to get voice socket") + } + + pub async fn udp_send_to(&self, addr: &SocketAddr, payload: &[u8]) -> anyhow::Result<()> { + let socket = self.get_voice_socket()?; + + socket.send_to(payload, addr).await?; + + Ok(()) + } +} 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(()) + } +}