diff --git a/src/protocol/initialize.rs b/src/protocol/initialize.rs index aff2a34..75dfe9a 100644 --- a/src/protocol/initialize.rs +++ b/src/protocol/initialize.rs @@ -4,28 +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 { +impl UserConnections { pub async fn initialize( server: &Arc, mut socket: WebSocket, - ) -> anyhow::Result<(Self, ClientMeta)> { + ) -> 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"), }, ) @@ -51,9 +52,9 @@ impl super::Client { { 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 @@ -67,10 +68,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"), }, ) @@ -81,12 +82,6 @@ impl super::Client { )); }; - Ok(( - Self { - socket: Arc::new(Mutex::new(socket)), - public_key, - }, - meta, - )) + Ok((socket, public_key, meta)) } } diff --git a/src/protocol/mod.rs b/src/protocol/mod.rs index 9124a09..94bd18a 100644 --- a/src/protocol/mod.rs +++ b/src/protocol/mod.rs @@ -1,18 +1,11 @@ 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; -#[derive(Clone)] -pub struct Client { - pub socket: Arc>, - pub public_key: VerifyingKey, -} - #[derive(Debug, Clone, Serialize, Deserialize)] pub struct ClientMeta {} @@ -56,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) - } - } - - Some(Message::Ping(v)) => { - socket.send(Message::Pong(v)).await?; +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 { + send_socket( + socket, + &ClientMethod::Error { + error: Cow::Borrowed("Unable to parse message: {text}"), + }, + ) + .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?; + Some(Message::Ping(v)) => { + socket.send(Message::Pong(v)).await?; - Ok(()) - } + Ok(None) + } - pub async fn read(&mut self) -> anyhow::Result> { - Self::read_socket(&mut *self.socket.lock().await).await - } + Some(_) => Ok(None), - pub async fn send(&mut self, message: ClientMethod) -> anyhow::Result<()> { - Self::send_socket(&mut *self.socket.lock().await, 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 cbfd019..f76a9e4 100644 --- a/src/server.rs +++ b/src/server.rs @@ -12,19 +12,20 @@ use tokio::sync::Mutex; use crate::{ config::Config, - protocol::{Client, ClientMeta}, + protocol::{ClientMeta, ClientMethod, read_loop, send_socket}, }; -pub struct OnlineClientMeta { +pub struct UserConnections { pub meta: ClientMeta, pub counter: AtomicU16, - pub connections: HashMap, + pub public_key: VerifyingKey, + pub connections: HashMap>>, } pub struct Server { pub key: SigningKey, pub config: Config, - pub clients: Mutex>, + pub clients: Mutex>, } impl Server { @@ -42,15 +43,18 @@ impl Server { let s = self.clone(); ws.on_upgrade(move |socket: WebSocket| async move { - match Client::initialize(&s, socket).await { - Ok((mut client, meta)) => { + 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(client.public_key) - .or_insert_with(|| OnlineClientMeta { + .entry(public_key) + .or_insert_with(|| UserConnections { meta, + public_key: public_key, counter: AtomicU16::new(0), connections: HashMap::new(), }); @@ -62,7 +66,7 @@ impl Server { client.clone(), ); - if let Err(e) = client.read_loop().await { + if let Err(e) = read_loop(&client).await { eprintln!("Failed to handle client: {e}"); } else { println!("Client connection closed") @@ -76,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) + } + } +}