diff --git a/Cargo.lock b/Cargo.lock index 4c7578f..7819992 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -2016,6 +2016,7 @@ dependencies = [ "base64 0.23.1", "dashmap", "dotenv", + "http", "livekit-api", "livekit-protocol", "log", diff --git a/Cargo.toml b/Cargo.toml index f1f7139..1328e6b 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -17,6 +17,7 @@ ansi_term = "0.12.1" uuid = { version = "1.24.0", features = ["v4"] } base64 = "0.23.0" dashmap = "6.2.1" +http = "1" once_cell = "1.21.4" rand = "0.10.2" rustls = { version = "0.23.42", default-features = false, features = [ diff --git a/src/omega/omega_connection.rs b/src/omega/omega_connection.rs index f785aed..5f28fc9 100644 --- a/src/omega/omega_connection.rs +++ b/src/omega/omega_connection.rs @@ -75,6 +75,33 @@ fn client_changed_target(value: &CommunicationValue) -> Option<(i64, i64)> { Some((receiver, session_id)) } +fn iota_user_snapshot(value: &CommunicationValue) -> Option<(i64, Vec)> { + if !value.is_type(CommunicationType::IotaUserData) { + return None; + } + let iota_id = value + .get_data(DataType::IotaId) + .as_number() + .and_then(|id| i64::try_from(id).ok()) + .filter(|id| *id > 0)?; + let Some(DataValue::Array(values)) = value.get_data(DataType::UserIds) else { + return None; + }; + let mut user_ids = Vec::with_capacity(values.len()); + for value in values { + let DataValue::SignedNumber(user_id) = value else { + return None; + }; + let user_id = i64::try_from(*user_id) + .ok() + .filter(|user_id| *user_id > 0)?; + user_ids.push(user_id); + } + user_ids.sort_unstable(); + user_ids.dedup(); + Some((iota_id, user_ids)) +} + fn parse_omega_capability_response(value: &CommunicationValue) -> Result { if !value.is_type(CommunicationType::IdentificationResponse) { return Err("not an identification response".to_string()); @@ -203,6 +230,10 @@ pub struct SessionPresenceState { } impl OmegaConnection { + pub fn omikron_id(&self) -> u64 { + self.omikron_id + } + pub fn from_config( config: &Config, keyring: mtp::crypto::Keyring, @@ -622,6 +653,11 @@ impl OmegaConnection { continue; } + if let Some((iota_id, user_ids)) = iota_user_snapshot(&cv) { + self.rho.replace_users_for_iota(iota_id, user_ids).await; + continue; + } + if cv.is_type(CommunicationType::Relay) { let request_id = match cv.require_id() { Ok(request_id) => request_id, @@ -701,26 +737,6 @@ impl OmegaConnection { } } - if cv.is_type(CommunicationType::IotaUserData) { - if let Some(DataValue::Array(users)) = cv.get_data(DataType::UserIds) { - let mut user_ids = Vec::new(); - for value in users { - if let DataValue::SignedNumber(user_id) = value { - if let Ok(user_id) = i64::try_from(*user_id) { - user_ids.push(user_id); - } - } - } - if let Some(iota_id) = cv.get_data(DataType::IotaId).as_number() { - self.rho.replace_users_for_iota(iota_id as i64, user_ids).await; - } else { - let iota_ids = self.rho.iota_ids().await; - for iota_id in iota_ids { - self.rho.replace_users_for_iota(iota_id, user_ids.clone()).await; - } - } - } - } if cv.is_type(CommunicationType::EraseHostedUserData) { let Some(iota_id) = cv .get_data(DataType::IotaId) @@ -1244,8 +1260,8 @@ impl OmegaConnection { mod tests { use super::{ ConnectionOutcome, MAX_RECONNECT_DELAY, PeerCapabilities, RECONNECT_DELAY, - client_changed_target, parse_omega_capability_response, reconnect_base_after_outcome, - reconnect_delay_with_jitter, + client_changed_target, iota_user_snapshot, parse_omega_capability_response, + reconnect_base_after_outcome, reconnect_delay_with_jitter, }; use mtp::codec::{CommunicationType, CommunicationValue, DataType, DataValue}; @@ -1283,6 +1299,33 @@ mod tests { assert_eq!(client_changed_target(&invalid), None); } + #[test] + fn iota_user_snapshot_requires_its_iota_and_valid_user_ids() { + let snapshot = CommunicationValue::new(CommunicationType::IotaUserData) + .add_typed_default(DataType::IotaId, DataValue::SignedNumber(8)) + .add_typed_default( + DataType::UserIds, + DataValue::Array(vec![ + DataValue::SignedNumber(20), + DataValue::SignedNumber(19), + DataValue::SignedNumber(20), + ]), + ); + assert_eq!(iota_user_snapshot(&snapshot), Some((8, vec![19, 20]))); + + let missing_iota = CommunicationValue::new(CommunicationType::IotaUserData) + .add_typed_default(DataType::UserIds, DataValue::Array(Vec::new())); + assert_eq!(iota_user_snapshot(&missing_iota), None); + + let malformed_user = CommunicationValue::new(CommunicationType::IotaUserData) + .add_typed_default(DataType::IotaId, DataValue::SignedNumber(8)) + .add_typed_default( + DataType::UserIds, + DataValue::Array(vec![DataValue::Str("not-a-user".into())]), + ); + assert_eq!(iota_user_snapshot(&malformed_user), None); + } + #[test] fn reconnect_jitter_is_bounded_without_changing_the_base_delay() { for _ in 0..32 { diff --git a/src/rho/client_connection.rs b/src/rho/client_connection.rs index 2899c2a..5625793 100644 --- a/src/rho/client_connection.rs +++ b/src/rho/client_connection.rs @@ -17,10 +17,7 @@ use tokio::sync::RwLock; use trust_dns_resolver::TokioAsyncResolver; use uuid::Uuid; -fn authenticated_peer_control_request( - cv: &CommunicationValue, - user_id: u64, -) -> CommunicationValue { +fn authenticated_peer_control_request(cv: &CommunicationValue, user_id: u64) -> CommunicationValue { cv.clone().with_sender(user_id) } @@ -173,9 +170,7 @@ impl ClientConnection { None => Err(relay_router::RelayRouteError::DestinationIotaNotLocal), }; let response = match result { - Ok(response) => { - response.with_id(message_id) - } + Ok(response) => response.with_id(message_id), Err(error) => { log_err!( self.user_id as i64, @@ -324,6 +319,28 @@ impl ClientConnection { } } + if cv.is_type(CommunicationType::ClientStateAck) { + let Some(rho) = self.get_rho_connection().await else { + self.send_error_response(message_id, CommunicationType::ErrorNoIota) + .await; + return; + }; + let request = cv.with_sender(self.user_id); + match rho + .get_iota_connection() + .clone() + .await_response(&request, Some(Duration::from_secs(20))) + .await + { + Ok(response) => self.send_message(&response).await, + Err(_) => { + self.send_error_response(message_id, CommunicationType::ErrorInternal) + .await; + } + } + return; + } + if cv.is_type(CommunicationType::ChangeUserData) || cv.is_type(CommunicationType::ReadNotification) || cv.is_type(CommunicationType::GetNotifications) diff --git a/src/rho/iota_connection.rs b/src/rho/iota_connection.rs index 813131b..5b9e08d 100755 --- a/src/rho/iota_connection.rs +++ b/src/rho/iota_connection.rs @@ -155,18 +155,24 @@ impl IotaConnection { } pub async fn add_user_id(&self, user_id: u64) { - let mut should_sync = false; + let Ok(user_id) = i64::try_from(user_id) else { + return; + }; + let manager_rho = self + .state + .rho + .bind_user_to_iota(user_id, self.iota_id as i64) + .await; { let mut guard = self.user_ids.write().await; - if !guard.contains(&user_id) { - guard.push(user_id); - should_sync = true; + if !guard.contains(&(user_id as u64)) { + guard.push(user_id as u64); } } - if should_sync { + if manager_rho.is_none() { if let Some(rho_conn) = self.get_rho_connection().await { - rho_conn.add_user_id(user_id as i64).await; + rho_conn.add_user_id(user_id).await; } } } @@ -340,6 +346,11 @@ impl IotaConnection { } } + if cv.is_type(CommunicationType::ClientStateSync) { + self.forward_to_client(cv).await; + return; + } + if cv.is_type(CommunicationType::StateSubscribe) { self.send_error_response( message_id, @@ -412,12 +423,14 @@ impl IotaConnection { self.send_message(&response_cv).await; } Err(error) => { - // Omega may have committed the insert even when its - // Success response was lost in transit. Verify the exact - // generated user ID before reporting failure; GetUserData - // uses the proven request/response path and keeps this - // recovery idempotent. + /* + * A lost Success can follow a committed insert. Verify + * the requested identity and this authenticated Iota + * before adding a local binding. + */ let user_id = cv.get_data(DataType::UserId).as_number(); + let username = cv.get_data(DataType::Username).as_str().map(str::to_owned); + let public_key = cv.get_data(DataType::PublicKey).as_str().map(str::to_owned); if let Some(user_id) = user_id { let verification = CommunicationValue::new(CommunicationType::GetUserData) .add_typed_default(DataType::UserId, DataValue::SignedNumber(user_id)); @@ -430,7 +443,13 @@ impl IotaConnection { { Ok(verified) if verified.get_data(DataType::UserId).as_number() - == Some(user_id) => + == Some(user_id) + && verified.get_data(DataType::IotaId).as_number() + == Some(self.iota_id.into()) + && verified.get_data(DataType::Username).as_str() + == username.as_deref() + && verified.get_data(DataType::PublicKey).as_str() + == public_key.as_deref() => { log_in!( self.iota_id as i64, @@ -479,8 +498,7 @@ impl IotaConnection { return; } - if cv.is_type(CommunicationType::ChangeIotaData) - || cv.is_type(CommunicationType::PushNotification) + if cv.is_type(CommunicationType::PushNotification) || cv.is_type(CommunicationType::GetUserData) || cv.is_type(CommunicationType::GetIotaData) || cv.is_type(CommunicationType::DeleteIota) diff --git a/src/rho/relay_router.rs b/src/rho/relay_router.rs index 19e46f8..a903256 100644 --- a/src/rho/relay_router.rs +++ b/src/rho/relay_router.rs @@ -210,13 +210,17 @@ async fn route_from_iota( if !target.has_local_client(user_id).await { return Err(RelayRouteError::ClientOffline); } - target.send_relay_to_client(&frame).await.map(|_| CommunicationValue::new(CommunicationType::Success)).map_err(|error| { - if error == "client offline" { - RelayRouteError::ClientOffline - } else { - RelayRouteError::Send(error) - } - }) + target + .send_relay_to_client(&frame) + .await + .map(|_| CommunicationValue::new(CommunicationType::Success)) + .map_err(|error| { + if error == "client offline" { + RelayRouteError::ClientOffline + } else { + RelayRouteError::Send(error) + } + }) } RouteTarget::Iota(iota_id) => { if iota_id == source_iota_id { @@ -326,6 +330,19 @@ mod tests { )); } + #[test] + fn iota_target_round_trips() { + let original = RouteTarget::Iota(42); + let wire = wire(original); + + assert_eq!(RouteTarget::from_wire_id(wire), Some(original)); + } + + #[test] + fn raw_iota_id_is_not_a_valid_typed_route() { + assert_ne!(RouteTarget::from_wire_id(42), Some(RouteTarget::Iota(42))); + } + #[test] fn relay_rejects_an_outer_sender() { let frame = relay(Some(wire(RouteTarget::Iota(7)))).with_sender(9); diff --git a/src/rho/server.rs b/src/rho/server.rs index b94885b..7290190 100644 --- a/src/rho/server.rs +++ b/src/rho/server.rs @@ -11,15 +11,38 @@ use crate::{ rho::connection::{ConnectionKind, GeneralConnection, OptionalDataValueCompat}, util::{file_util::load_file_vec, logger::PrintType}, }; +use http::{Method, StatusCode}; use mtp::codec::{CommunicationType, CommunicationValue, DataType, DataValue}; use mtp::crypto::PublicKeyBundle; use mtp::host::{AuthenticationPolicy, HostConfig, Policy, SendMode}; -use mtp::webserver::{MTPWebServer, WebServerConfig}; +use mtp::webserver::{HttpResponse, MTPWebServer, WebServerConfig}; + +fn with_cors(response: HttpResponse) -> HttpResponse { + response + .header("access-control-allow-origin", "*") + .header("access-control-allow-methods", "GET, POST, OPTIONS") + .header("access-control-allow-headers", "*") +} fn web_config(max_connections: usize) -> Result { WebServerConfig::new() .max_connections(max_connections) - .route("/", |_request, response| async move { response.body("OK") }) + .route("/", |request, response| async move { + let response = if request.method == Method::OPTIONS { + response.status(StatusCode::NO_CONTENT) + } else { + response.body("OK") + }; + with_cors(response) + })? + .fallback(|request, response| async move { + let status = if request.method == Method::OPTIONS { + StatusCode::NO_CONTENT + } else { + StatusCode::NOT_FOUND + }; + with_cors(response.status(status)) + }) } fn rho_policy() -> Policy { @@ -49,16 +72,7 @@ pub async fn get_by_connector_id( client_id: u64, description: Option, ) -> Option { - let request = match description.as_deref() { - Some("iota") => { - println!("Iota connection request for client_id: {}", client_id); - CommunicationValue::new(CommunicationType::GetIotaData) - .add_typed_default(DataType::IotaId, DataValue::SignedNumber(client_id as i128)) - } - Some("client") => CommunicationValue::new(CommunicationType::GetUserData) - .add_typed_default(DataType::UserId, DataValue::SignedNumber(client_id as i128)), - _ => return None, - }; + let request = public_key_lookup_request(omega.omikron_id(), client_id, description.as_deref())?; let response = match omega .await_response(&request, Some(Duration::from_secs(20))) @@ -83,6 +97,25 @@ pub async fn get_by_connector_id( PublicKeyBundle::from_bytes(&bytes).ok() } +fn public_key_lookup_request( + omikron_id: u64, + client_id: u64, + description: Option<&str>, +) -> Option { + Some(match description { + Some("iota") => { + println!("Iota connection request for client_id: {}", client_id); + CommunicationValue::new(CommunicationType::GetIotaData) + .with_sender(omikron_id) + .add_typed_default(DataType::IotaId, DataValue::SignedNumber(client_id as i128)) + } + Some("client") => CommunicationValue::new(CommunicationType::GetUserData) + .with_sender(omikron_id) + .add_typed_default(DataType::UserId, DataValue::SignedNumber(client_id as i128)), + _ => return None, + }) +} + /* Only Iota registration goes through mtp's Register flow; users are registered out of band. */ pub async fn complete_register( omega: Arc, @@ -279,7 +312,8 @@ pub async fn start(state: Arc) -> Result<(), Box