diff --git a/src/protocol/initialize.rs b/src/protocol/initialize.rs index 40b4d79..2c6dff4 100644 --- a/src/protocol/initialize.rs +++ b/src/protocol/initialize.rs @@ -4,25 +4,29 @@ use std::{ }; use axum::extract::ws::WebSocket; -use ed25519_dalek::Signer; +use ed25519_dalek::{Signer, VerifyingKey}; use crate::server::Server; use super::*; +use crate::server::UserConnections; -impl super::Client { - pub async fn initialize(server: &Arc, mut socket: WebSocket) -> anyhow::Result { +impl UserConnections { + pub async fn initialize( + server: &Arc, + mut socket: WebSocket, + ) -> anyhow::Result<(WebSocket, VerifyingKey, ClientMeta)> { let Some(ServerMethod::Initialize { public_key: public_key_string, signature, timestamp, hostname, - }) = Client::read_socket(&mut socket).await? + }) = read_socket(&mut socket).await? else { - Client::send_socket( + send_socket( &mut socket, - ClientMethod::Error { + &ClientMethod::Error { error: Cow::Borrowed("Initialization required"), }, ) @@ -33,24 +37,42 @@ impl super::Client { )); }; - let public_key = crate::signature::from_string(&public_key_string)?; + let Ok(public_key) = crate::signature::from_string(&public_key_string) else { + send_socket( + &mut socket, + &ClientMethod::Error { + error: Cow::Borrowed("Invalid public key"), + }, + ) + .await?; + + return Err(anyhow::anyhow!("Invalid public key")); + }; if public_key .verify_strict( format!("{timestamp}@{hostname}").as_bytes(), &crate::signature::from_string_sig(&signature)?, ) - .is_ok() + .is_err() { + send_socket( + &mut socket, + &ClientMethod::Error { + error: Cow::Borrowed("Invalid signature"), + }, + ) + .await?; + return Err(anyhow::anyhow!("Invalid signature")); } { let timestamp = SystemTime::now().duration_since(UNIX_EPOCH)?.as_secs(); - Client::send_socket( + send_socket( &mut socket, - ClientMethod::Initialized { + &ClientMethod::Initialized { public_key: crate::signature::to_string(&server.key.verifying_key()), signature: server .key @@ -64,10 +86,10 @@ impl super::Client { .await?; } - let Some(ServerMethod::Meta(meta)) = Client::read_socket(&mut socket).await? else { - Client::send_socket( + let Some(ServerMethod::Meta(meta)) = read_socket(&mut socket).await? else { + send_socket( &mut socket, - ClientMethod::Error { + &ClientMethod::Error { error: Cow::Borrowed("Expected meta"), }, ) @@ -78,10 +100,6 @@ impl super::Client { )); }; - Ok(Self { - socket, - meta, - public_key, - }) + Ok((socket, public_key, meta)) } } diff --git a/src/protocol/mod.rs b/src/protocol/mod.rs index 460530f..0503573 100644 --- a/src/protocol/mod.rs +++ b/src/protocol/mod.rs @@ -1,17 +1,11 @@ -use std::borrow::Cow; +use std::{borrow::Cow, sync::Arc}; use axum::extract::ws::{Message, Utf8Bytes, WebSocket}; -use ed25519_dalek::VerifyingKey; use serde::{Deserialize, Serialize}; +use tokio::sync::Mutex; pub mod initialize; -pub struct Client { - pub socket: WebSocket, - pub meta: ClientMeta, - pub public_key: VerifyingKey, -} - #[derive(Debug, Clone, Serialize, Deserialize)] pub struct ClientMeta {} @@ -55,73 +49,67 @@ pub enum ServerMethod { }, } -impl Client { - pub async fn read_loop(&mut self) -> anyhow::Result<()> { - while let Some(message) = self.read().await? { - match message { - ServerMethod::Initialize { .. } => { - self.send(ClientMethod::Error { +pub async fn read_loop(socket: &Arc>) -> anyhow::Result<()> { + while let Some(message) = read_socket(&mut *socket.lock().await).await? { + match message { + ServerMethod::Initialize { .. } => { + send_socket( + &mut *socket.lock().await, + &ClientMethod::Error { error: Cow::Borrowed("Already initialized"), - }) - .await?; - } + }, + ) + .await?; + } - ServerMethod::Meta(meta) => {} + #[allow(unused_variables)] + ServerMethod::Meta(meta) => {} - ServerMethod::Error { error } => { - eprintln!("Client error: {error}"); - } + ServerMethod::Error { error } => { + eprintln!("Client error: {error}"); } } - - Ok(()) } - pub async fn read_socket(socket: &mut WebSocket) -> anyhow::Result> { - match socket.recv().await.transpose()? { - Some(Message::Text(text)) => { - if let Ok(msg) = serde_json::from_str(&text.to_string()) { - Ok(Some(msg)) - } else { - Client::send_socket( - socket, - ClientMethod::Error { - error: Cow::Borrowed("Unable to parse message: {text}"), - }, - ) - .await?; + Ok(()) +} - Ok(None) - } - } +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)), - Some(Message::Ping(v)) => { - socket.send(Message::Pong(v)).await?; + Err(e) => { + send_socket( + socket, + &ClientMethod::Error { + error: Cow::Owned(format!("Unable to parse message: {e}")), + }, + ) + .await?; Ok(None) } + }, - Some(_) => Ok(None), + Some(Message::Ping(v)) => { + socket.send(Message::Pong(v)).await?; - None => Ok(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?; + Some(_) => Ok(None), - Ok(()) - } - - pub async fn read(&mut self) -> anyhow::Result> { - Self::read_socket(&mut self.socket).await - } - - pub async fn send(&mut self, message: ClientMethod) -> anyhow::Result<()> { - Self::send_socket(&mut self.socket, message).await + 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/server.rs b/src/server.rs index 0eb6d92..f76a9e4 100644 --- a/src/server.rs +++ b/src/server.rs @@ -1,16 +1,31 @@ -use std::sync::Arc; +use std::{ + collections::HashMap, + sync::{Arc, atomic::AtomicU16}, +}; use axum::{ extract::{WebSocketUpgrade, ws::WebSocket}, response::Response, }; -use ed25519_dalek::SigningKey; +use ed25519_dalek::{SigningKey, VerifyingKey}; +use tokio::sync::Mutex; -use crate::{config::Config, protocol::Client}; +use crate::{ + config::Config, + protocol::{ClientMeta, ClientMethod, read_loop, send_socket}, +}; + +pub struct UserConnections { + pub meta: ClientMeta, + pub counter: AtomicU16, + pub public_key: VerifyingKey, + pub connections: HashMap>>, +} pub struct Server { pub key: SigningKey, pub config: Config, + pub clients: Mutex>, } impl Server { @@ -18,6 +33,7 @@ impl Server { Ok(Arc::new(Self { key: crate::signature::get().await?, config: Config::get().await?, + clients: Mutex::new(HashMap::new()), })) } } @@ -27,9 +43,30 @@ impl Server { let s = self.clone(); ws.on_upgrade(move |socket: WebSocket| async move { - match Client::initialize(&s, socket).await { - Ok(mut client) => { - if let Err(e) = client.read_loop().await { + match UserConnections::initialize(&s, socket).await { + Ok((client, public_key, meta)) => { + let client = Arc::new(Mutex::new(client)); + + let mut clients_meta = s.clients.lock().await; + + let client_meta = + clients_meta + .entry(public_key) + .or_insert_with(|| UserConnections { + meta, + public_key: public_key, + counter: AtomicU16::new(0), + connections: HashMap::new(), + }); + + client_meta.connections.insert( + client_meta + .counter + .fetch_add(1, std::sync::atomic::Ordering::Relaxed), + client.clone(), + ); + + if let Err(e) = read_loop(&client).await { eprintln!("Failed to handle client: {e}"); } else { println!("Client connection closed") @@ -43,3 +80,23 @@ impl Server { }) } } + +impl UserConnections { + pub async fn send(&self, message: &ClientMethod) -> anyhow::Result<()> { + for (_, conn) in &self.connections { + send_socket(&mut *conn.lock().await, message).await?; + } + + Ok(()) + } + + pub async fn send_to(&self, id: u16, message: &ClientMethod) -> anyhow::Result { + if let Some(conn) = self.connections.get(&id) { + send_socket(&mut *conn.lock().await, message).await?; + + Ok(true) + } else { + Ok(false) + } + } +}