diff --git a/client/src/connection.rs b/client/src/connection.rs index 6a3c843..bf44b43 100644 --- a/client/src/connection.rs +++ b/client/src/connection.rs @@ -159,7 +159,7 @@ impl MTPConnection { } } -pub(crate) fn connection_from_parts( +pub(crate) async fn connection_from_parts( config: ClientConfig, sender: mtp_transport::Sender, receiver: mtp_transport::Receiver, @@ -168,7 +168,7 @@ pub(crate) fn connection_from_parts( #[cfg(feature = "crypto")] client_id: u64, ) -> MTPConnection { let remote_addr = sender.handle().remote_addr(); - let ping = start_ping_session(&config, sender.clone(), &receiver); + let ping = start_ping_session(&config, sender.clone(), &receiver).await; #[cfg(feature = "pipes")] { diff --git a/client/src/lib.rs b/client/src/lib.rs index 7003686..f594200 100644 --- a/client/src/lib.rs +++ b/client/src/lib.rs @@ -142,9 +142,10 @@ impl MTPClient { negotiated, error::AuthState::Unauthenticated, client_id, - )); + ) + .await); #[cfg(not(feature = "crypto"))] - Ok(connection_from_parts(config, sender, receiver, negotiated)) + Ok(connection_from_parts(config, sender, receiver, negotiated).await) } } @@ -283,7 +284,8 @@ impl MTPClient { crypto::negotiated_version(&response)?, error::AuthState::Authenticated, client_id, - )) + ) + .await) } pub async fn auth_register( @@ -440,7 +442,8 @@ impl MTPClient { crypto::negotiated_version(&response)?, error::AuthState::Authenticated, assigned_id, - )) + ) + .await) } } diff --git a/client/src/ping.rs b/client/src/ping.rs index 9309cbc..b84145f 100644 --- a/client/src/ping.rs +++ b/client/src/ping.rs @@ -1,5 +1,4 @@ use rand::RngExt; -use std::collections::HashMap; use std::sync::Arc; use tokio::sync::{Mutex, mpsc}; use tokio::time::{Duration, Instant}; @@ -12,6 +11,38 @@ pub(crate) struct PingSession { pub(crate) task: tokio::task::JoinHandle<()>, } +#[derive(Default)] +struct PingTracker { + pending: Option<(u32, Instant)>, + missed_pings: usize, +} + +impl PingTracker { + fn begin_round(&mut self) -> usize { + if self.pending.take().is_some() { + self.missed_pings += 1; + } + self.missed_pings + } + + fn sent(&mut self, id: u32) { + self.pending = Some((id, Instant::now())); + } + + fn received(&mut self, id: u32) -> Option { + if !self + .pending + .as_ref() + .is_some_and(|(pending, _)| *pending == id) + { + return None; + } + let (_, sent_at) = self.pending.take()?; + self.missed_pings = 0; + Some(sent_at.elapsed()) + } +} + impl PingSession { pub(crate) fn get_ping(&self) -> Option { self.last_ping.try_lock().ok().and_then(|ping| *ping) @@ -24,7 +55,7 @@ impl Drop for PingSession { } } -pub(crate) fn start_ping_session( +pub(crate) async fn start_ping_session( config: &crate::config::ClientConfig, sender: Sender, receiver: &Receiver, @@ -34,7 +65,7 @@ pub(crate) fn start_ping_session( } let (pong_tx, mut pong_rx) = mpsc::unbounded_channel(); - receiver.observe_pongs(pong_tx); + receiver.observe_pongs(pong_tx).await; let last_ping = Arc::new(Mutex::new(None)); let ping_state = last_ping.clone(); let interval = config.ping_interval; @@ -46,7 +77,7 @@ pub(crate) fn start_ping_session( let task = tokio::spawn(async move { let mut ticker = tokio::time::interval(interval); ticker.tick().await; - let mut pending = HashMap::new(); + let mut tracker = PingTracker::default(); loop { tokio::select! { @@ -56,7 +87,8 @@ pub(crate) fn start_ping_session( } } _ = ticker.tick() => { - if max_missed_pings > 0 && !pending.is_empty() && pending.len() >= max_missed_pings { + let missed_pings = tracker.begin_round(); + if max_missed_pings > 0 && missed_pings >= max_missed_pings { sender.close().await; break; } @@ -83,13 +115,13 @@ pub(crate) fn start_ping_session( sender.close().await; break; } - pending.insert(id, Instant::now()); + tracker.sent(id); } pong = pong_rx.recv() => match pong { Some(pong) => { - if let Some(sent_at) = pending.remove(&pong.get_id()) { + if let Some(ping) = tracker.received(pong.get_id()) { let mut last_ping = ping_state.lock().await; - *last_ping = Some(sent_at.elapsed()); + *last_ping = Some(ping); } } None => break, @@ -100,3 +132,32 @@ pub(crate) fn start_ping_session( Some(PingSession { last_ping, task }) } + +#[cfg(test)] +mod tests { + use super::PingTracker; + + #[test] + fn successful_pong_resets_consecutive_misses() { + let mut tracker = PingTracker::default(); + tracker.sent(1); + assert_eq!(tracker.begin_round(), 1); + + tracker.sent(2); + assert!(tracker.received(2).is_some()); + + tracker.sent(3); + assert_eq!(tracker.begin_round(), 1); + } + + #[test] + fn stale_pong_does_not_acknowledge_current_round() { + let mut tracker = PingTracker::default(); + tracker.sent(1); + assert_eq!(tracker.begin_round(), 1); + tracker.sent(2); + + assert!(tracker.received(1).is_none()); + assert_eq!(tracker.begin_round(), 2); + } +} diff --git a/transport/src/connection.rs b/transport/src/connection.rs index be04588..c543dec 100644 --- a/transport/src/connection.rs +++ b/transport/src/connection.rs @@ -938,12 +938,8 @@ impl Receiver { } /* Route reserved Pong frames to a connection-level observer. */ - pub fn observe_pongs(&self, observer: mpsc::UnboundedSender) { - if let Ok(mut control) = self.inner.ping_control.try_write() { - control.pong_observer = Some(observer); - } else { - warn!("[Receiver] could not register Pong observer: control lock busy"); - } + pub async fn observe_pongs(&self, observer: mpsc::UnboundedSender) { + self.inner.ping_control.write().await.pong_observer = Some(observer); } #[instrument(skip(stream, policy), level = "trace")] diff --git a/wasm/src/client.rs b/wasm/src/client.rs index 91dc57a..3b9525b 100644 --- a/wasm/src/client.rs +++ b/wasm/src/client.rs @@ -53,8 +53,8 @@ fn route_incoming_frame( if let Some(sent_at) = frame_id(frame).and_then(|id| pending_pings.borrow_mut().remove(&id)) { ping_ms.set(Some(js_sys::Date::now() - sent_at)); + return; } - return; } if let Some(request_id) = frame_id(frame) { @@ -138,6 +138,33 @@ async fn wait_for_timeout(timeout_ms: u32) -> Result<(), JsValue> { Ok(()) } +fn set_shared_state( + state: &Rc>, + pending_state_callbacks: &Rc>>, + state_callback: &JsValue, + new_state: ConnectionState, +) { + state.set(new_state); + pending_state_callbacks.borrow_mut().push_back(new_state); + + let global = js_sys::global(); + let qmt = js_sys::Reflect::get(&global, &JsValue::from_str("queueMicrotask")) + .and_then(|f| f.dyn_into::()); + let scheduled = qmt + .and_then(|qmt| qmt.call1(&global, state_callback)) + .is_ok(); + if !scheduled + && js_sys::Reflect::get(&global, &JsValue::from_str("setTimeout")) + .and_then(|f| f.dyn_into::()) + .and_then(|set_timeout| { + set_timeout.call2(&global, state_callback, &JsValue::from_f64(0.0)) + }) + .is_err() + { + pending_state_callbacks.borrow_mut().pop_back(); + } +} + #[wasm_bindgen] #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub enum ConnectionState { @@ -150,6 +177,7 @@ pub enum ConnectionState { #[wasm_bindgen] pub struct WasmClient { transport: Rc>>, + connection_generation: Rc>, state: Rc>, pending_state_callbacks: Rc>>, state_callback: Closure, @@ -187,6 +215,7 @@ impl WasmClient { }) as Box); Self { transport: Rc::new(RefCell::new(None)), + connection_generation: Rc::new(Cell::new(0)), state: Rc::new(Cell::new(ConnectionState::Disconnected)), pending_state_callbacks, state_callback, @@ -221,7 +250,7 @@ impl WasmClient { #[wasm_bindgen] pub async fn connect(&self, config: &ConnectionConfig) -> Result<(), JsValue> { - self.set_state(ConnectionState::Connecting); + let generation = self.begin_connection(); let transport = WasmTransport::connect( &config.url, config.server_certificate_hashes.clone(), @@ -249,7 +278,7 @@ impl WasmClient { .map_err(|e| js_error(format!("parse handshake outcome: {e}")))?; let tm = mtp_codec::TypeMap::latest(); if Some(outcome.get_type()) == CommunicationType::ErrorBadVersion.try_to_id(&tm) { - self.set_state(ConnectionState::Disconnected); + self.set_state_if_current(generation, ConnectionState::Disconnected); return Err(js_error( outcome .get_str(DataType::ErrorMessage) @@ -262,7 +291,7 @@ impl WasmClient { if outcome.get_type() != expected || outcome.get_data(DataType::Connected) != &DataValue::BoolTrue { - self.set_state(ConnectionState::Disconnected); + self.set_state_if_current(generation, ConnectionState::Disconnected); return Err(js_error( outcome .get_str(DataType::ErrorMessage) @@ -272,12 +301,14 @@ impl WasmClient { match outcome.get_data(DataType::Version) { DataValue::Str(version) if mtp_codec::Version::parse(version).is_some() => {} _ => { - self.set_state(ConnectionState::Disconnected); + self.set_state_if_current(generation, ConnectionState::Disconnected); return Err(js_error("host omitted a valid negotiated protocol version")); } } - self.start_receive_loop(transport); + if !self.start_receive_loop(transport, generation) { + return Err(js_error("connection attempt superseded")); + } Ok(()) } @@ -289,7 +320,7 @@ impl WasmClient { keyring_bytes: &[u8], client_id: u64, ) -> Result { - self.set_state(ConnectionState::Connecting); + let generation = self.begin_connection(); let host_pk = mtp_crypto::PublicKeyBundle::from_bytes(host_public_key_bytes) .map_err(|e| js_error(format!("invalid host public key: {}", e)))?; @@ -326,6 +357,7 @@ impl WasmClient { "auth_connect challenge", config.require_pq, !keyring.sig_pq_secret_key.as_bytes().is_empty(), + generation, ) .await?; @@ -348,7 +380,7 @@ impl WasmClient { .try_to_id(&tm) .ok_or_else(|| js_error("IdentificationResponse is absent from the type map"))?; if resp_type != expected_type { - self.set_state(ConnectionState::Disconnected); + self.set_state_if_current(generation, ConnectionState::Disconnected); return Err(auth::unexpected_response_type_error( "auth_connect", expected_type, @@ -359,7 +391,7 @@ impl WasmClient { } if resp_comm.get_data(DataType::Connected) != &DataValue::BoolTrue { - self.set_state(ConnectionState::Disconnected); + self.set_state_if_current(generation, ConnectionState::Disconnected); return Err(js_error( resp_comm .get_str(DataType::ErrorMessage) @@ -376,19 +408,21 @@ impl WasmClient { server_challenge, config.require_pq, ) { - self.set_state(ConnectionState::Disconnected); + self.set_state_if_current(generation, ConnectionState::Disconnected); return Err(e); } let assigned_id = match resp_comm.get_data(DataType::Id) { DataValue::UnsignedNumber(n) => *n as u64, _ => { - self.set_state(ConnectionState::Disconnected); + self.set_state_if_current(generation, ConnectionState::Disconnected); return Err(js_error("missing assigned ID")); } }; - self.start_receive_loop(transport); + if !self.start_receive_loop(transport, generation) { + return Err(js_error("connection attempt superseded")); + } Ok(assigned_id) } @@ -400,7 +434,7 @@ impl WasmClient { host_public_key_bytes: &[u8], keyring_bytes: &[u8], ) -> Result { - self.set_state(ConnectionState::Connecting); + let generation = self.begin_connection(); let host_pk = mtp_crypto::PublicKeyBundle::from_bytes(host_public_key_bytes) .map_err(|e| js_error(format!("invalid host public key: {}", e)))?; @@ -438,6 +472,7 @@ impl WasmClient { "auth_register challenge", config.require_pq, !keyring.sig_pq_secret_key.as_bytes().is_empty(), + generation, ) .await?; @@ -460,7 +495,7 @@ impl WasmClient { .try_to_id(&tm) .ok_or_else(|| js_error("RegisterResponse is absent from the type map"))?; if resp_type != expected_type { - self.set_state(ConnectionState::Disconnected); + self.set_state_if_current(generation, ConnectionState::Disconnected); return Err(auth::unexpected_response_type_error( "auth_register", expected_type, @@ -471,7 +506,7 @@ impl WasmClient { } if resp_comm.get_data(DataType::Connected) != &DataValue::BoolTrue { - self.set_state(ConnectionState::Disconnected); + self.set_state_if_current(generation, ConnectionState::Disconnected); return Err(js_error( resp_comm .get_str(DataType::ErrorMessage) @@ -482,7 +517,7 @@ impl WasmClient { let assigned_id = match resp_comm.get_data(DataType::Id) { DataValue::UnsignedNumber(n) => *n as u64, _ => { - self.set_state(ConnectionState::Disconnected); + self.set_state_if_current(generation, ConnectionState::Disconnected); return Err(js_error("missing assigned ID")); } }; @@ -496,11 +531,13 @@ impl WasmClient { server_challenge, config.require_pq, ) { - self.set_state(ConnectionState::Disconnected); + self.set_state_if_current(generation, ConnectionState::Disconnected); return Err(e); } - self.start_receive_loop(transport); + if !self.start_receive_loop(transport, generation) { + return Err(js_error("connection attempt superseded")); + } Ok(assigned_id) } @@ -674,6 +711,8 @@ impl WasmClient { #[wasm_bindgen] pub fn disconnect(&self) { + self.connection_generation + .set(self.connection_generation.get().wrapping_add(1)); self.stop_protocol_pings(); if let Some(t) = self.transport.borrow_mut().take() { t.close(); @@ -733,39 +772,50 @@ impl WasmClient { } fn set_state(&self, new_state: ConnectionState) { - self.state.set(new_state); - self.pending_state_callbacks - .borrow_mut() - .push_back(new_state); + set_shared_state( + &self.state, + &self.pending_state_callbacks, + self.state_callback.as_ref(), + new_state, + ); + } - let global = js_sys::global(); - let qmt = js_sys::Reflect::get(&global, &JsValue::from_str("queueMicrotask")) - .and_then(|f| f.dyn_into::()); - let scheduled = qmt - .and_then(|qmt| qmt.call1(&global, self.state_callback.as_ref())) - .is_ok(); - if !scheduled - && js_sys::Reflect::get(&global, &JsValue::from_str("setTimeout")) - .and_then(|f| f.dyn_into::()) - .and_then(|set_timeout| { - set_timeout.call2( - &global, - self.state_callback.as_ref(), - &JsValue::from_f64(0.0), - ) - }) - .is_err() - { - self.pending_state_callbacks.borrow_mut().pop_back(); + fn set_state_if_current(&self, generation: u32, new_state: ConnectionState) { + if self.connection_generation.get() == generation { + self.set_state(new_state); } } - fn start_receive_loop(&self, transport: WasmTransport) { + fn begin_connection(&self) -> u32 { + let generation = self.connection_generation.get().wrapping_add(1); + self.connection_generation.set(generation); + self.stop_protocol_pings(); + if let Some(transport) = self.transport.borrow_mut().take() { + transport.close(); + } + self.reject_pending_requests("connection replaced"); + client_pipe::reject_pending_pipe_creations( + &self.pending_pipe_creations, + "connection replaced", + ); + self.set_state(ConnectionState::Connecting); + generation + } + + fn start_receive_loop(&self, transport: WasmTransport, generation: u32) -> bool { + if self.connection_generation.get() != generation { + transport.close(); + return false; + } let loop_transport = transport.clone(); *self.transport.borrow_mut() = Some(transport); self.set_state(ConnectionState::Connected); + let connection_generation = self.connection_generation.clone(); + let error_generation = connection_generation.clone(); let state = self.state.clone(); + let pending_state_callbacks = self.pending_state_callbacks.clone(); + let state_callback = self.state_callback.as_ref().clone(); let on_msg = self.on_message.clone(); let on_err = self.on_error.clone(); let subscriptions = self.subscriptions.clone(); @@ -847,7 +897,11 @@ impl WasmClient { &loop_ping_ms, ); }, - on_err.clone(), + move |error| { + if error_generation.get() == generation { + let _ = on_err.call1(&JsValue::NULL, &error); + } + }, move |pipe_reader: PipeReader| { let pipe_id = pipe_reader.pipe_id(); let mut pending = pending_pipes.borrow_mut(); @@ -857,13 +911,22 @@ impl WasmClient { }, ) .await; - state.set(ConnectionState::Disconnected); + if connection_generation.get() != generation { + return; + } + set_shared_state( + &state, + &pending_state_callbacks, + &state_callback, + ConnectionState::Disconnected, + ); stop_ping_timer(&ping_timer); pending_pings.borrow_mut().clear(); ping_ms.set(None); reject_pending_requests(&pending_requests, "disconnected"); client_pipe::reject_pending_pipe_creations(&pending_pipe_creations, "disconnected"); }); + true } fn reject_pending_requests(&self, message: &str) { @@ -879,6 +942,7 @@ impl WasmClient { context: &str, require_pq: bool, client_has_pq_key: bool, + generation: u32, ) -> Result { let challenge_bytes = transport.read_one_frame().await?; let challenge = CommunicationValue::from_bytes(&challenge_bytes) @@ -887,7 +951,7 @@ impl WasmClient { .try_to_id(tm) .ok_or_else(|| js_error("Challenge is absent from the type map"))?; if challenge.get_type() != expected { - self.set_state(ConnectionState::Disconnected); + self.set_state_if_current(generation, ConnectionState::Disconnected); return Err(auth::unexpected_response_type_error( context, expected, @@ -900,13 +964,13 @@ impl WasmClient { let server_challenge = match challenge.get_data(DataType::ServerNonce) { DataValue::UnsignedNumber(n) => *n, _ => { - self.set_state(ConnectionState::Disconnected); + self.set_state_if_current(generation, ConnectionState::Disconnected); return Err(js_error("missing server challenge")); } }; if challenge.get_data(DataType::RequirePq) == &DataValue::BoolTrue && !client_has_pq_key { - self.set_state(ConnectionState::Disconnected); + self.set_state_if_current(generation, ConnectionState::Disconnected); return Err(js_error( "host requires post-quantum authentication but the client PQ key is absent", )); @@ -920,7 +984,7 @@ impl WasmClient { server_challenge, require_pq, ) { - self.set_state(ConnectionState::Disconnected); + self.set_state_if_current(generation, ConnectionState::Disconnected); return Err(e); } diff --git a/wasm/src/transport.rs b/wasm/src/transport.rs index 1fb8ca8..8e151a8 100644 --- a/wasm/src/transport.rs +++ b/wasm/src/transport.rs @@ -477,14 +477,15 @@ impl WasmTransport { /// Pipe-aware receive loop. Identical to `receive_loop` but detects /// `PipeRequest` as the first frame on a new incoming stream and routes /// the stream to `on_pipe` instead of `on_message`. - pub async fn receive_loop_with_pipes( + pub async fn receive_loop_with_pipes( &self, mut on_message: F, - on_error: js_sys::Function, + mut on_error: H, mut on_pipe: G, ) where F: FnMut(JsValue), G: FnMut(crate::pipe::PipeReader), + H: FnMut(JsValue), { let pipe_request_type = mtp_codec::CommunicationType::PipeRequest.try_to_id(&mtp_codec::TypeMap::latest()); @@ -528,13 +529,13 @@ impl WasmTransport { } Err(e) => { let message = e.as_string().unwrap_or_else(|| format!("{:?}", e)); - let _ = on_error.call1(&JsValue::NULL, &JsValue::from_str(&message)); + on_error(JsValue::from_str(&message)); } } } Ok(FrameOutcome::Closed) | Ok(FrameOutcome::Ended) => break, Err(e) => { - let _ = on_error.call1(&JsValue::NULL, &e); + on_error(e); break; } }