diff --git a/src/protocol/message.rs b/src/protocol/message.rs index 0f3b76f..b19542c 100644 --- a/src/protocol/message.rs +++ b/src/protocol/message.rs @@ -45,6 +45,12 @@ pub async fn send_message( server.message_store.insert_message(&channel_id, &stored)?; + server + .broadcast(&ClientMethod::Messages { + messages: HashMap::from([(channel_id, vec![stored])]), + }) + .await?; + Ok(()) } diff --git a/src/server.rs b/src/server.rs index 6497b57..a48b07c 100644 --- a/src/server.rs +++ b/src/server.rs @@ -9,7 +9,7 @@ use axum::{ response::Response, }; use ed25519_dalek::{SigningKey, VerifyingKey}; -use tokio::sync::Mutex; +use tokio::{sync::Mutex, task::JoinSet}; use crate::{ data::{config::Config, messages::MessageStore}, @@ -21,13 +21,13 @@ pub struct UserConnections { pub meta: ClientMeta, pub counter: AtomicU16, pub public_key: VerifyingKey, - pub connections: HashMap>>, + pub connections: Mutex>>>, } pub struct Server { pub key: SigningKey, pub config: Config, - pub clients: Mutex>, + pub clients: Mutex>>, pub message_store: MessageStore, } @@ -53,21 +53,27 @@ impl Server { let mut clients_meta = s.clients.lock().await; - let clients = - clients_meta - .entry(public_key) - .or_insert_with(|| UserConnections { + let clients = clients_meta + .entry(public_key) + .or_insert_with(|| { + Arc::new(UserConnections { meta, public_key: public_key, counter: AtomicU16::new(0), - connections: HashMap::new(), - }); + connections: Mutex::new(HashMap::new()), + }) + }) + .clone(); let conid = clients .counter .fetch_add(1, std::sync::atomic::Ordering::Relaxed); - clients.connections.insert(conid, client.clone()); + clients + .connections + .lock() + .await + .insert(conid, client.clone()); if let Err(e) = read_loop(&s, public_key, &client).await { eprintln!("Failed to handle client: {e}"); @@ -75,9 +81,11 @@ impl Server { println!("Client connection closed") } - clients.connections.remove(&conid); + let mut connections = clients.connections.lock().await; - if clients.connections.len() == 0 { + connections.remove(&conid); + + if connections.len() == 0 { clients_meta.remove(&public_key); } } @@ -88,11 +96,32 @@ impl Server { } }) } + + pub async fn broadcast(self: &Arc, message: &ClientMethod) -> anyhow::Result<()> { + let mut set = JoinSet::new(); + + for (_, client) in self.clients.lock().await.iter() { + let msg = message.clone(); + let client = client.clone(); + + set.spawn(async move { client.send(&msg).await }); + } + + // Await all spawned tasks to finish + while let Some(res) = set.join_next().await { + // handle task panic or errors if necessary + if let Ok(Err(e)) = res { + eprintln!("Failed to send to a client: {:?}", e); + } + } + + Ok(()) + } } impl UserConnections { pub async fn send(&self, message: &ClientMethod) -> anyhow::Result<()> { - for (_, conn) in &self.connections { + for (_, conn) in self.connections.lock().await.iter() { send_socket(&mut *conn.lock().await, message).await?; } @@ -100,7 +129,7 @@ impl UserConnections { } pub async fn send_to(&self, id: u16, message: &ClientMethod) -> anyhow::Result { - if let Some(conn) = self.connections.get(&id) { + if let Some(conn) = self.connections.lock().await.get(&id) { send_socket(&mut *conn.lock().await, message).await?; Ok(true)