diff --git a/Cargo.lock b/Cargo.lock index b5d82de..d195ce0 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -76,6 +76,15 @@ dependencies = [ "memchr", ] +[[package]] +name = "android_system_properties" +version = "0.1.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "819e7219dbd41043ac279b19830f2efc897156490d7fd6ea916720117ee66311" +dependencies = [ + "libc", +] + [[package]] name = "ansi_term" version = "0.12.1" @@ -284,9 +293,9 @@ checksum = "c08606f8c3cbf4ce6ec8e28fb0014a2c086708fe954eaa885384a6165172e7e8" [[package]] name = "aws-lc-rs" -version = "1.16.1" +version = "1.16.2" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "94bffc006df10ac2a68c83692d734a465f8ee6c5b384d8545a636f81d858f4bf" +checksum = "a054912289d18629dc78375ba2c3726a3afe3ff71b4edba9dedfca0e3446d1fc" dependencies = [ "aws-lc-sys", "untrusted 0.7.1", @@ -295,9 +304,9 @@ dependencies = [ [[package]] name = "aws-lc-sys" -version = "0.38.0" +version = "0.39.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "4321e568ed89bb5a7d291a7f37997c2c0df89809d7b6d12062c81ddb54aa782e" +checksum = "1fa7e52a4c5c547c741610a2c6f123f3881e409b714cd27e6798ef020c514f0a" dependencies = [ "cc", "cmake", @@ -344,6 +353,15 @@ dependencies = [ "generic-array", ] +[[package]] +name = "block2" +version = "0.6.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "cdeb9d870516001442e364c5220d3574d2da8dc765554b4a617230d33fa58ef5" +dependencies = [ + "objc2", +] + [[package]] name = "blocking" version = "1.6.2" @@ -598,6 +616,16 @@ dependencies = [ "subtle", ] +[[package]] +name = "dispatch2" +version = "0.3.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1e0e367e4e7da84520dedcac1901e4da967309406d1e51017ae1abfb97adbd38" +dependencies = [ + "bitflags", + "objc2", +] + [[package]] name = "displaydoc" version = "0.2.5" @@ -1323,9 +1351,9 @@ checksum = "d98f6fed1fde3f8c21bc40a1abb88dd75e67924f9cffc3ef95607bad8017f8e2" [[package]] name = "iri-string" -version = "0.7.10" +version = "0.7.11" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "c91338f0783edbd6195decb37bae672fd3b165faffb89bf7b9e6942f8b1a731a" +checksum = "d8e7418f59cc01c88316161279a7f665217ae316b388e58a0d10e29f54f1e5eb" dependencies = [ "memchr", "serde", @@ -1351,9 +1379,9 @@ dependencies = [ [[package]] name = "itoa" -version = "1.0.17" +version = "1.0.18" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "92ecc6618181def0457392ccd0ee51198e065e016d1d527a7ac1b6dc7c1f09d2" +checksum = "8f42a60cbdf9a97f5d2305f08a87dc4e09308d1276d28c869c684d7777685682" [[package]] name = "jobserver" @@ -1446,9 +1474,9 @@ checksum = "6373607a59f0be73a39b6fe456b8192fcc3585f602af20751600e974dd455e77" [[package]] name = "livekit-api" -version = "0.4.15" +version = "0.4.16" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "1afaefea8ba98d6d51a9b7c50e5eca9ee2cbaf74ca6682e2ed62aa823f698aa8" +checksum = "020c7eaab7cec16b69a058659f9acff32067e0e2d4f81e80639f3ffa29464b09" dependencies = [ "async-tungstenite", "base64 0.21.7", @@ -1456,6 +1484,7 @@ dependencies = [ "jsonwebtoken", "livekit-protocol", "log", + "os_info", "parking_lot", "pbjson-types", "prost", @@ -1472,9 +1501,9 @@ dependencies = [ [[package]] name = "livekit-protocol" -version = "0.7.1" +version = "0.7.2" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "01875112fb961da484dedf0c2d5705d8875c5e99da6d5aa1d567be1b7d03532a" +checksum = "c21d56b4dc2598b1296ca891abb225bf1c05ea007d6f43bc4fd48c31be64bf59" dependencies = [ "futures-util", "livekit-runtime", @@ -1567,6 +1596,18 @@ dependencies = [ "tempfile", ] +[[package]] +name = "nix" +version = "0.30.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "74523f3a35e05aba87a1d978330aef40f67b0304ac79c1c00b294c9830543db6" +dependencies = [ + "bitflags", + "cfg-if", + "cfg_aliases", + "libc", +] + [[package]] name = "nom" version = "7.1.3" @@ -1605,9 +1646,9 @@ dependencies = [ [[package]] name = "num-conv" -version = "0.2.0" +version = "0.2.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "cf97ec579c3c42f953ef76dbf8d55ac91fb219dde70e49aa4a6b7d74e9919050" +checksum = "c6673768db2d862beb9b39a78fdcb1a69439615d5794a1be50caa9bc92c81967" [[package]] name = "num-integer" @@ -1639,6 +1680,165 @@ dependencies = [ "libm", ] +[[package]] +name = "objc2" +version = "0.6.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "3a12a8ed07aefc768292f076dc3ac8c48f3781c8f2d5851dd3d98950e8c5a89f" +dependencies = [ + "objc2-encode", +] + +[[package]] +name = "objc2-cloud-kit" +version = "0.3.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "73ad74d880bb43877038da939b7427bba67e9dd42004a18b809ba7d87cee241c" +dependencies = [ + "bitflags", + "objc2", + "objc2-foundation", +] + +[[package]] +name = "objc2-core-data" +version = "0.3.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0b402a653efbb5e82ce4df10683b6b28027616a2715e90009947d50b8dd298fa" +dependencies = [ + "objc2", + "objc2-foundation", +] + +[[package]] +name = "objc2-core-foundation" +version = "0.3.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2a180dd8642fa45cdb7dd721cd4c11b1cadd4929ce112ebd8b9f5803cc79d536" +dependencies = [ + "bitflags", + "dispatch2", + "objc2", +] + +[[package]] +name = "objc2-core-graphics" +version = "0.3.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e022c9d066895efa1345f8e33e584b9f958da2fd4cd116792e15e07e4720a807" +dependencies = [ + "bitflags", + "dispatch2", + "objc2", + "objc2-core-foundation", + "objc2-io-surface", +] + +[[package]] +name = "objc2-core-image" +version = "0.3.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e5d563b38d2b97209f8e861173de434bd0214cf020e3423a52624cd1d989f006" +dependencies = [ + "objc2", + "objc2-foundation", +] + +[[package]] +name = "objc2-core-location" +version = "0.3.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ca347214e24bc973fc025fd0d36ebb179ff30536ed1f80252706db19ee452009" +dependencies = [ + "objc2", + "objc2-foundation", +] + +[[package]] +name = "objc2-core-text" +version = "0.3.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0cde0dfb48d25d2b4862161a4d5fcc0e3c24367869ad306b0c9ec0073bfed92d" +dependencies = [ + "bitflags", + "objc2", + "objc2-core-foundation", + "objc2-core-graphics", +] + +[[package]] +name = "objc2-encode" +version = "4.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ef25abbcd74fb2609453eb695bd2f860d389e457f67dc17cafc8b8cbc89d0c33" + +[[package]] +name = "objc2-foundation" +version = "0.3.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e3e0adef53c21f888deb4fa59fc59f7eb17404926ee8a6f59f5df0fd7f9f3272" +dependencies = [ + "bitflags", + "block2", + "libc", + "objc2", + "objc2-core-foundation", +] + +[[package]] +name = "objc2-io-surface" +version = "0.3.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "180788110936d59bab6bd83b6060ffdfffb3b922ba1396b312ae795e1de9d81d" +dependencies = [ + "bitflags", + "objc2", + "objc2-core-foundation", +] + +[[package]] +name = "objc2-quartz-core" +version = "0.3.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "96c1358452b371bf9f104e21ec536d37a650eb10f7ee379fff67d2e08d537f1f" +dependencies = [ + "bitflags", + "objc2", + "objc2-core-foundation", + "objc2-foundation", +] + +[[package]] +name = "objc2-ui-kit" +version = "0.3.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d87d638e33c06f577498cbcc50491496a3ed4246998a7fbba7ccb98b1e7eab22" +dependencies = [ + "bitflags", + "block2", + "objc2", + "objc2-cloud-kit", + "objc2-core-data", + "objc2-core-foundation", + "objc2-core-graphics", + "objc2-core-image", + "objc2-core-location", + "objc2-core-text", + "objc2-foundation", + "objc2-quartz-core", + "objc2-user-notifications", +] + +[[package]] +name = "objc2-user-notifications" +version = "0.3.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9df9128cbbfef73cda168416ccf7f837b62737d748333bfe9ab71c245d76613e" +dependencies = [ + "objc2", + "objc2-foundation", +] + [[package]] name = "octets" version = "0.3.5" @@ -1710,6 +1910,22 @@ dependencies = [ "vcpkg", ] +[[package]] +name = "os_info" +version = "3.14.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e4022a17595a00d6a369236fdae483f0de7f0a339960a53118b818238e132224" +dependencies = [ + "android_system_properties", + "log", + "nix", + "objc2", + "objc2-foundation", + "objc2-ui-kit", + "serde", + "windows-sys 0.61.2", +] + [[package]] name = "p256" version = "0.13.2" @@ -2366,9 +2582,9 @@ dependencies = [ [[package]] name = "rustls-webpki" -version = "0.103.9" +version = "0.103.10" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "d7df23109aa6c1567d1c575b9952556388da57401e4ace1d15f79eedad0d8f53" +checksum = "df33b2b81ac578cabaf06b89b0631153a3f416b0a886e8a7a1707fb51abbd1ef" dependencies = [ "aws-lc-rs", "ring", @@ -2915,7 +3131,7 @@ checksum = "e421abadd41a4225275504ea4d6566923418b7f05506fbc9c0fe86ba7396114b" [[package]] name = "ttp-core" version = "0.1.0" -source = "git+https://github.com/t3kkm0tt/TTP.git#8a3cd5cea756d62e7f85824ebfa1e64c7594c7de" +source = "git+https://github.com/Tensamin/TTP.git#c127d196082401f37b51ca5417b341ecbdb74318" dependencies = [ "base64 0.22.1", "byteorder", @@ -2927,7 +3143,7 @@ dependencies = [ [[package]] name = "ttp-native" version = "0.1.0" -source = "git+https://github.com/t3kkm0tt/TTP.git#8a3cd5cea756d62e7f85824ebfa1e64c7594c7de" +source = "git+https://github.com/Tensamin/TTP.git#c127d196082401f37b51ca5417b341ecbdb74318" dependencies = [ "quinn", "rustls", @@ -3611,18 +3827,18 @@ dependencies = [ [[package]] name = "zerocopy" -version = "0.8.42" +version = "0.8.47" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "f2578b716f8a7a858b7f02d5bd870c14bf4ddbbcf3a4c05414ba6503640505e3" +checksum = "efbb2a062be311f2ba113ce66f697a4dc589f85e78a4aea276200804cea0ed87" dependencies = [ "zerocopy-derive", ] [[package]] name = "zerocopy-derive" -version = "0.8.42" +version = "0.8.47" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "7e6cc098ea4d3bd6246687de65af3f920c430e236bee1e3bf2e441463f08a02f" +checksum = "0e8bc7269b54418e7aeeef514aa68f8690b8c0489a06b0136e5f57c4c5ccab89" dependencies = [ "proc-macro2", "quote", diff --git a/Cargo.toml b/Cargo.toml index cd9d652..3127b12 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -23,12 +23,12 @@ rustls = { version = "0.23.37", default-features = false, features = [ "aws-lc-rs", "prefer-post-quantum", ] } -sha2 = "*" +aes-gcm = "0.10.3" +sha2 = "0.10.9" tokio = { version = "*", features = ["full"] } x448 = { version = "*" } log = "0.4" dotenv = "0.15.0" -aes-gcm = "0.10.3" hkdf = "0.12.4" strum = "0.28.0" strum_macros = "0.28.0" diff --git a/src/anonymous_clients/anonymous_client_connection.rs b/src/anonymous_clients/anonymous_client_connection.rs index ed0ef65..8c57490 100644 --- a/src/anonymous_clients/anonymous_client_connection.rs +++ b/src/anonymous_clients/anonymous_client_connection.rs @@ -49,6 +49,7 @@ impl AnonymousClientConnection { while let Ok(cv) = self_clone.receiver.receive().await { self_clone.clone().handle_message(cv).await; } + self_clone.handle_close().await; }); } @@ -457,11 +458,10 @@ impl AnonymousClientConnection { let error = CommunicationValue::new(error_type).with_id(*message_id); self.send_message(&error).await; } - #[allow(dead_code)] /// Close the connection pub async fn close(&self) { let mut is_open_guard = self.is_open.write().await; - if *is_open_guard { + if !*is_open_guard { return; } *is_open_guard = false; @@ -494,7 +494,6 @@ impl AnonymousClientConnection { } } - #[allow(dead_code)] /// Handle connection close pub async fn handle_close(&self) { diff --git a/src/rho/client_connection.rs b/src/rho/client_connection.rs index bb379bb..ad3fc63 100644 --- a/src/rho/client_connection.rs +++ b/src/rho/client_connection.rs @@ -5,7 +5,7 @@ use crate::rho::connection::GeneralConnection; use crate::rho::{rho_connection::RhoConnection, rho_manager}; use crate::util::logger::PrintType; use crate::{data::user::UserStatus, omega::omega_connection::OmegaConnection}; -use crate::{log_cv_in, log_cv_out, log_out}; +use crate::{log_cv_in, log_cv_out, log_err, log_in, log_out}; use std::str::FromStr; use std::sync::Arc; use std::time::{Duration, SystemTime, UNIX_EPOCH}; @@ -31,7 +31,7 @@ impl ClientConnection { Arc::new(Self { ping: Arc::new(RwLock::new(0)), pub_key: Arc::new(RwLock::new(None)), - rho_connection: Arc::new(RwLock::new(None)), + rho_connection: general.rho_connection.clone(), interested_users: Arc::new(RwLock::new(Vec::new())), is_open: Arc::new(RwLock::new(true)), sender: general.sender.clone(), @@ -45,6 +45,7 @@ impl ClientConnection { while let Ok(cv) = self_clone.receiver.receive().await { self_clone.clone().handle_message(cv).await; } + self_clone.handle_close().await; }); } @@ -393,6 +394,20 @@ impl ClientConnection { /// Forward message to Iota async fn forward_to_iota(self: Arc, cv: CommunicationValue) { + let sender_user_id = self.get_user_id().await; + let msg_id = cv.get_id(); + let msg_type = cv.get_type(); + + log_in!( + sender_user_id as i64, + PrintType::Client, + "Forwarding client->iota: sender={} type={:?} id={} receiver={}", + sender_user_id, + msg_type, + msg_id, + cv.get_receiver() + ); + if cv.is_type(CommunicationType::add_conversation) && cv .get_data(DataTypes::chat_partner_id) @@ -434,17 +449,56 @@ impl ClientConnection { }; if let Some(rho_conn) = self.get_rho_connection().await { + let iota_id = rho_conn.get_iota_id().await; + log_in!( + sender_user_id as i64, + PrintType::Client, + "Resolved rho for add_conversation: sender={} -> iota_id={} id={}", + sender_user_id, + iota_id, + msg_id + ); + let updated_cv = cv - .with_sender(self.get_user_id().await as u64) + .with_sender(sender_user_id as u64) .add_data(DataTypes::chat_partner_id, chat_partner_id); rho_conn.message_to_iota(updated_cv).await; + } else { + log_err!( + sender_user_id as i64, + PrintType::Client, + "No rho/iota mapping found for add_conversation sender={} type={:?} id={}", + sender_user_id, + msg_type, + msg_id + ); } return; } if let Some(rho_conn) = self.get_rho_connection().await { - let updated_cv = cv.with_sender(self.get_user_id().await as u64); + let iota_id = rho_conn.get_iota_id().await; + log_in!( + sender_user_id as i64, + PrintType::Client, + "Resolved rho for forward: sender={} -> iota_id={} type={:?} id={}", + sender_user_id, + iota_id, + msg_type, + msg_id + ); + + let updated_cv = cv.with_sender(sender_user_id as u64); rho_conn.message_to_iota(updated_cv).await; + } else { + log_err!( + sender_user_id as i64, + PrintType::Client, + "No rho/iota mapping found for sender={} type={:?} id={}", + sender_user_id, + msg_type, + msg_id + ); } } @@ -455,10 +509,9 @@ impl ClientConnection { } /// Close the connection - #[allow(dead_code)] pub async fn close(&self) { let mut is_open_guard = self.is_open.write().await; - if *is_open_guard { + if !*is_open_guard { return; } *is_open_guard = false; @@ -491,7 +544,6 @@ impl ClientConnection { } /// Handle connection close - #[allow(dead_code)] pub async fn handle_close(&self) { let user_id = self.get_user_id().await; if let Some(rho_conn) = rho_manager::get_rho_con_for_user(user_id as i64).await { diff --git a/src/rho/connection.rs b/src/rho/connection.rs index a187c32..a676321 100755 --- a/src/rho/connection.rs +++ b/src/rho/connection.rs @@ -37,6 +37,7 @@ pub struct GeneralConnection { challenge: Arc>, connection_kind: Arc>>, + pub rho_connection: Arc>>>, id: Arc>, pub_key: Arc>>>, @@ -50,6 +51,7 @@ impl GeneralConnection { challenged: Arc::new(RwLock::new(false)), challenge: Arc::new(RwLock::new(String::new())), connection_kind: Arc::new(RwLock::new(None)), + rho_connection: Arc::new(RwLock::new(None)), id: Arc::new(RwLock::new(0)), pub_key: Arc::new(RwLock::new(None)), }) @@ -76,13 +78,13 @@ impl GeneralConnection { if !*self.challenged.read().await { self.handle_challenge_response(cv).await; - if *self.challenged.read().await { - break; - } - continue; } - if self.migrate().await { + if *self.challenged.read().await { + let self_clone = self.clone(); + tokio::spawn(async move { + self_clone.migrate().await; + }); break; } } @@ -233,10 +235,6 @@ impl GeneralConnection { if let Err(_) = self.sender.send(&response).await { return; } - - if self.migrate().await { - log_out!(id, PrintType::Iota, "Immediate migration"); - } } else { log_err!( id, @@ -270,7 +268,48 @@ impl GeneralConnection { .add_data(DataTypes::user_id, DataValue::Number(id as i64)); get_omega_connection().send_message(¬ify).await; + let user_id = id as i64; + + let mut rho = rho_manager::get_rho_con_for_user(user_id).await; + + if rho.is_none() { + let get_user_msg = CommunicationValue::new(CommunicationType::get_user_data) + .add_data(DataTypes::user_id, DataValue::Number(user_id)); + + if let Ok(user_data_cv) = get_omega_connection() + .await_response(&get_user_msg, Some(Duration::from_secs(20))) + .await + { + if let DataValue::Number(iota_id) = + user_data_cv.get_data(DataTypes::iota_id) + { + if let Some(bound_rho) = + rho_manager::bind_user_to_iota(user_id, *iota_id).await + { + bound_rho.bind_user_id(user_id).await; + rho = Some(bound_rho); + } + } + } + } + + *self.rho_connection.write().await = rho.clone(); + let client = ClientConnection::from_general(self.clone(), id).await; + + if let Some(rho_conn) = rho { + // Make sure user is bound before the client starts forwarding + rho_conn.bind_user_id(user_id).await; + rho_conn.add_client_connection(client.clone()).await; + } else { + log_err!( + user_id, + PrintType::Client, + "No RhoConnection found for user {}, client not attached to iota", + id + ); + } + client.start(); } ConnectionKind::Iota => { @@ -282,7 +321,25 @@ impl GeneralConnection { let rho = Arc::new(RhoConnection::new(iota.clone(), Vec::new()).await); - iota.set_rho_connection(Arc::downgrade(&rho)).await; + iota.set_rho_connection(rho.clone()).await; + + let get_iota_msg = CommunicationValue::new(CommunicationType::get_iota_data) + .add_data(DataTypes::iota_id, DataValue::Number(id as i64)); + + if let Ok(iota_data_cv) = get_omega_connection() + .await_response(&get_iota_msg, Some(Duration::from_secs(20))) + .await + { + if let DataValue::Array(users) = iota_data_cv.get_data(DataTypes::user_ids) { + let mut user_ids: Vec = Vec::new(); + for value in users { + if let DataValue::Number(user_id) = value { + user_ids.push(*user_id as u64); + } + } + iota.set_user_ids(user_ids).await; + } + } rho_manager::add_rho(rho).await; diff --git a/src/rho/iota_connection.rs b/src/rho/iota_connection.rs index 05d3004..a52f8a2 100755 --- a/src/rho/iota_connection.rs +++ b/src/rho/iota_connection.rs @@ -3,16 +3,13 @@ use crate::calls::call_manager; use crate::log_cv_in; use crate::log_cv_out; use crate::log_err; +use crate::log_in; use crate::omega::omega_connection::get_omega_connection; use crate::rho::connection::GeneralConnection; use crate::util::logger::PrintType; use dashmap::DashMap; use std::collections::BTreeMap; -use std::{ - collections::HashMap, - sync::{Arc, Weak}, - time::Duration, -}; +use std::{collections::HashMap, sync::Arc, time::Duration}; use tokio::sync::RwLock; use tokio::sync::mpsc; use ttp_core::CommunicationType; @@ -36,7 +33,7 @@ pub struct IotaConnection { pub_key: Arc>>>, pub waiting_tasks: DashMap, CommunicationValue) -> bool + Send + Sync>>, - pub rho_connection: Arc>>>, + pub rho_connection: Arc>>>, } impl IotaConnection { @@ -44,7 +41,7 @@ impl IotaConnection { Arc::new(Self { ping: Arc::new(RwLock::new(0)), pub_key: Arc::new(RwLock::new(None)), - rho_connection: Arc::new(RwLock::new(None)), + rho_connection: general.rho_connection.clone(), user_ids: Arc::new(RwLock::new(Vec::new())), sender: general.sender.clone(), receiver: general.receiver.clone(), @@ -65,6 +62,7 @@ impl IotaConnection { } } } + self_clone.handle_close().await; }); } @@ -87,13 +85,43 @@ impl IotaConnection { self.user_ids.read().await.clone() } + /// Replace all users linked to this iota and synchronize the attached rho mapping. + pub async fn set_user_ids(&self, user_ids: Vec) { + { + let mut guard = self.user_ids.write().await; + *guard = user_ids.clone(); + } + + if let Some(rho_conn) = self.get_rho_connection().await { + let user_ids_i64: Vec = user_ids.into_iter().map(|u| u as i64).collect(); + rho_conn.set_user_ids(user_ids_i64); + } + } + + pub async fn add_user_id(&self, user_id: u64) { + let mut should_sync = false; + { + let mut guard = self.user_ids.write().await; + if !guard.contains(&user_id) { + guard.push(user_id); + should_sync = true; + } + } + + if should_sync { + if let Some(rho_conn) = self.get_rho_connection().await { + rho_conn.add_user_id(user_id as i64).await; + } + } + } + /// Get current ping pub async fn get_ping(&self) -> i64 { *self.ping.read().await } /// Set the RhoConnection reference - pub async fn set_rho_connection(&self, rho_connection: Weak) { + pub async fn set_rho_connection(&self, rho_connection: Arc) { let mut rho_ref = self.rho_connection.write().await; *rho_ref = Some(rho_connection); } @@ -102,7 +130,7 @@ impl IotaConnection { pub async fn get_rho_connection(&self) -> Option> { let rho_ref = self.rho_connection.read().await; if let Some(weak_ref) = rho_ref.as_ref() { - weak_ref.upgrade() + Some(weak_ref.clone()) } else { None } @@ -217,8 +245,20 @@ impl IotaConnection { async fn handle_forward_message(&self, cv: CommunicationValue) { let receiver_id = cv.get_receiver(); let sender_id = cv.get_sender(); + let my_user_ids = self.get_user_ids().await; - if self.get_user_ids().await.contains(&(sender_id as u64)) { + log_in!( + self.iota_id as i64, + PrintType::Iota, + "Authority check: sender_id={} receiver_id={} iota_user_ids={:?} msg_type={:?} msg_id={}", + sender_id, + receiver_id, + my_user_ids, + cv.get_type(), + cv.get_id() + ); + + if my_user_ids.contains(&(sender_id as u64)) { if let Some(target_rho) = rho_manager::get_rho_con_for_user(receiver_id as i64).await { target_rho.message_to_iota(cv).await; } else { @@ -228,6 +268,14 @@ impl IotaConnection { self.send_message(&error).await; } } else { + log_err!( + self.iota_id as i64, + PrintType::Iota, + "Rejected client->iota forward: sender_id={} is not authorized for this iota. Known users={:?}", + sender_id, + my_user_ids + ); + self.send_message( &CommunicationValue::new(CommunicationType::error_invalid_user_id).add_data( DataTypes::error_type, @@ -350,7 +398,6 @@ impl IotaConnection { } } - #[allow(dead_code)] pub async fn handle_close(&self) { if let Some(rho_conn) = self.get_rho_connection().await { rho_conn.close_iota_connection().await; diff --git a/src/rho/rho_connection.rs b/src/rho/rho_connection.rs index af37187..6a740ff 100644 --- a/src/rho/rho_connection.rs +++ b/src/rho/rho_connection.rs @@ -9,7 +9,7 @@ use ttp_core::{CommunicationType, CommunicationValue, DataTypes, DataValue}; pub struct RhoConnection { iota_connection: Arc, - user_ids: Vec, + user_ids: Arc>>, client_connections: Arc>>>, } @@ -18,7 +18,7 @@ impl RhoConnection { pub async fn new(iota_connection: Arc, user_ids: Vec) -> Self { let rho_connection = Self { iota_connection, - user_ids: user_ids.clone(), + user_ids: Arc::new(RwLock::new(user_ids.clone())), client_connections: Arc::new(RwLock::new(Vec::new())), }; @@ -29,8 +29,25 @@ impl RhoConnection { self.iota_connection.iota_id } - pub fn get_user_ids(&self) -> &Vec { - &self.user_ids + pub async fn get_user_ids(&self) -> Vec { + self.user_ids.read().await.clone() + } + + pub async fn set_user_ids(&self, user_ids: Vec) { + let mut guard = self.user_ids.write().await; + *guard = user_ids; + } + + pub async fn add_user_id(&self, user_id: i64) { + let mut guard = self.user_ids.write().await; + if !guard.contains(&user_id) { + guard.push(user_id); + } + } + + pub async fn bind_user_id(&self, user_id: i64) { + self.add_user_id(user_id).await; + self.iota_connection.add_user_id(user_id as u64).await; } pub fn get_iota_connection(&self) -> &Arc { @@ -169,8 +186,8 @@ impl RhoConnection { /// Check if this RhoConnection contains a specific user ID #[allow(dead_code)] - pub fn contains_user(&self, user_id: &i64) -> bool { - self.user_ids.contains(user_id) + pub async fn contains_user(&self, user_id: &i64) -> bool { + self.user_ids.read().await.contains(user_id) } /// Get count of active client connections diff --git a/src/rho/rho_manager.rs b/src/rho/rho_manager.rs index dd73497..e823a04 100644 --- a/src/rho/rho_manager.rs +++ b/src/rho/rho_manager.rs @@ -13,13 +13,14 @@ pub static RHO_CONNECTIONS: LazyLock> pub async fn get_rho_con_for_user(user_id: i64) -> Option> { let connections = RHO_CONNECTIONS.read().await; for rho_connection in connections.values() { + let rho_user_ids = rho_connection.get_user_ids().await; log_in!( user_id, PrintType::Client, "Comparing user IDs: {:?}", - rho_connection.get_user_ids().to_vec() + rho_user_ids ); - if rho_connection.get_user_ids().contains(&user_id) { + if rho_user_ids.contains(&user_id) { return Some(Arc::clone(rho_connection)); } } @@ -32,6 +33,29 @@ pub async fn contains_iota(iota_id: i64) -> bool { connections.contains_key(&iota_id) } +/// Bind a user ID to an already tracked iota/rho connection. +pub async fn bind_user_to_iota(user_id: i64, iota_id: i64) -> Option> { + let connections = RHO_CONNECTIONS.read().await; + if let Some(rho_connection) = connections.get(&iota_id) { + let rho = Arc::clone(rho_connection); + drop(connections); + + rho.add_user_id(user_id).await; + + log_in!( + user_id, + PrintType::Client, + "Bound user {} to iota {}", + user_id, + iota_id + ); + + Some(rho) + } else { + None + } +} + /// Remove a RhoConnection by Iota ID pub async fn remove_rho(iota_id: i64) -> Option> { let mut connections = RHO_CONNECTIONS.write().await; diff --git a/src/util/crypto_util.rs b/src/util/crypto_util.rs index 26448e0..bb784e8 100644 --- a/src/util/crypto_util.rs +++ b/src/util/crypto_util.rs @@ -4,12 +4,12 @@ use aes_gcm::{ }; use base64::{Engine as _, engine::general_purpose::STANDARD as BASE64_STD}; use hkdf::Hkdf; -use sha2::{Digest, Sha256}; +type HkdfSha256 = sha2::Sha256; +use sha2::{Digest, Sha256 as HashSha256}; use x448::{PublicKey, Secret}; -// --- Custom Errors --- -#[allow(dead_code)] #[derive(Debug)] +#[allow(dead_code)] pub enum SecurePayloadError { InvalidBase64, InvalidHex, @@ -18,9 +18,8 @@ pub enum SecurePayloadError { InvalidKeyLength, } -// --- Data Format Enum --- -#[allow(dead_code)] #[derive(Clone, Copy, Debug)] +#[allow(dead_code)] pub enum DataFormat { Raw, Base64, @@ -41,8 +40,8 @@ impl Clone for SecurePayload { } } +#[allow(dead_code)] impl SecurePayload { - /// Clear Constructor: Takes data in any format and the user's private key. pub fn new>( data: T, format: DataFormat, @@ -67,13 +66,10 @@ impl SecurePayload { }) } - /// Helper to get the public key associated with this instance's private key. - #[allow(dead_code)] pub fn get_public_key(&self) -> [u8; 56] { *PublicKey::from(&self.private_key).as_bytes() } - /// Exports the internal data to the requested format pub fn export(&self, format: DataFormat) -> String { match format.into() { DataFormat::Raw => String::from_utf8_lossy(&self.inner_data).to_string(), @@ -82,16 +78,12 @@ impl SecurePayload { } } - /// Access raw bytes directly - #[allow(dead_code)] pub fn get_bytes(&self) -> &[u8] { &self.inner_data } - /// Returns the SHA-256 Hash of the data in the requested format - #[allow(dead_code)] pub fn get_hash(&self, format: DataFormat) -> String { - let mut hasher = Sha256::new(); + let mut hasher = HashSha256::new(); hasher.update(&self.inner_data); let result = hasher.finalize(); @@ -102,8 +94,6 @@ impl SecurePayload { } } - /// Encrypts the held data for a specific recipient using AES-256-GCM. - /// The message will contain ONLY the ciphertext. pub fn encrypt_x448(&self, public_key: S) -> Result where S: Into, @@ -111,8 +101,9 @@ impl SecurePayload { let peer_pub = public_key.into(); let shared_secret = self.private_key.as_diffie_hellman(&peer_pub).unwrap(); - let hkdf = Hkdf::::new(None, shared_secret.as_bytes()); + let hkdf = Hkdf::::new(None, shared_secret.as_bytes()); let mut okm = [0u8; 44]; + hkdf.expand(b"x448-aes-gcm-no-overhead", &mut okm) .map_err(|_| SecurePayloadError::EncryptionError)?; @@ -138,26 +129,27 @@ impl SecurePayload { }) } - /// Decrypts the held data providing the sender's public key manually. - #[allow(dead_code)] pub fn decrypt_to_format( &self, peer_public_key_bytes: &[u8; 56], output_format: DataFormat, ) -> Result { - let decrypted_instance = self.decrypt_x448(peer_public_key_bytes)?; + let decrypted_instance = + self.decrypt_x448(PublicKey::from_bytes(peer_public_key_bytes).unwrap())?; Ok(decrypted_instance.export(output_format)) } - /// Decrypts the held data using the internal Private Key and the provided Peer Public Key. - pub fn decrypt_x448( + pub fn decrypt_x448( &self, - peer_public_key_bytes: &[u8; 56], - ) -> Result { - let peer_pub = PublicKey::from_bytes(peer_public_key_bytes).unwrap(); + peer_public_key_bytes: S, + ) -> Result + where + S: Into, + { + let peer_pub = peer_public_key_bytes.into(); let shared_secret = self.private_key.as_diffie_hellman(&peer_pub).unwrap(); - let hkdf = Hkdf::::new(None, shared_secret.as_bytes()); + let hkdf = Hkdf::::new(None, shared_secret.as_bytes()); let mut okm = [0u8; 44]; hkdf.expand(b"x448-aes-gcm-no-overhead", &mut okm) .map_err(|_| SecurePayloadError::DecryptionError)?;