Fix native ping connection lifecycle
Some checks failed
CI / checks (push) Failing after 2m49s

This commit is contained in:
Alois 2026-07-28 22:26:45 +02:00
commit b6483b7f6d
Signed by: alois
SSH key fingerprint: SHA256:GBzT2DXvAuGV9XIV5W3WrzVpjU54FThmxHXdbz95J24
6 changed files with 198 additions and 73 deletions

View file

@ -159,7 +159,7 @@ impl MTPConnection {
} }
} }
pub(crate) fn connection_from_parts( pub(crate) async fn connection_from_parts(
config: ClientConfig, config: ClientConfig,
sender: mtp_transport::Sender, sender: mtp_transport::Sender,
receiver: mtp_transport::Receiver, receiver: mtp_transport::Receiver,
@ -168,7 +168,7 @@ pub(crate) fn connection_from_parts(
#[cfg(feature = "crypto")] client_id: u64, #[cfg(feature = "crypto")] client_id: u64,
) -> MTPConnection { ) -> MTPConnection {
let remote_addr = sender.handle().remote_addr(); 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")] #[cfg(feature = "pipes")]
{ {

View file

@ -142,9 +142,10 @@ impl MTPClient {
negotiated, negotiated,
error::AuthState::Unauthenticated, error::AuthState::Unauthenticated,
client_id, client_id,
)); )
.await);
#[cfg(not(feature = "crypto"))] #[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)?, crypto::negotiated_version(&response)?,
error::AuthState::Authenticated, error::AuthState::Authenticated,
client_id, client_id,
)) )
.await)
} }
pub async fn auth_register( pub async fn auth_register(
@ -440,7 +442,8 @@ impl MTPClient {
crypto::negotiated_version(&response)?, crypto::negotiated_version(&response)?,
error::AuthState::Authenticated, error::AuthState::Authenticated,
assigned_id, assigned_id,
)) )
.await)
} }
} }

View file

@ -1,5 +1,4 @@
use rand::RngExt; use rand::RngExt;
use std::collections::HashMap;
use std::sync::Arc; use std::sync::Arc;
use tokio::sync::{Mutex, mpsc}; use tokio::sync::{Mutex, mpsc};
use tokio::time::{Duration, Instant}; use tokio::time::{Duration, Instant};
@ -12,6 +11,38 @@ pub(crate) struct PingSession {
pub(crate) task: tokio::task::JoinHandle<()>, 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<Duration> {
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 { impl PingSession {
pub(crate) fn get_ping(&self) -> Option<Duration> { pub(crate) fn get_ping(&self) -> Option<Duration> {
self.last_ping.try_lock().ok().and_then(|ping| *ping) 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, config: &crate::config::ClientConfig,
sender: Sender, sender: Sender,
receiver: &Receiver, receiver: &Receiver,
@ -34,7 +65,7 @@ pub(crate) fn start_ping_session(
} }
let (pong_tx, mut pong_rx) = mpsc::unbounded_channel(); 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 last_ping = Arc::new(Mutex::new(None));
let ping_state = last_ping.clone(); let ping_state = last_ping.clone();
let interval = config.ping_interval; let interval = config.ping_interval;
@ -46,7 +77,7 @@ pub(crate) fn start_ping_session(
let task = tokio::spawn(async move { let task = tokio::spawn(async move {
let mut ticker = tokio::time::interval(interval); let mut ticker = tokio::time::interval(interval);
ticker.tick().await; ticker.tick().await;
let mut pending = HashMap::new(); let mut tracker = PingTracker::default();
loop { loop {
tokio::select! { tokio::select! {
@ -56,7 +87,8 @@ pub(crate) fn start_ping_session(
} }
} }
_ = ticker.tick() => { _ = 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; sender.close().await;
break; break;
} }
@ -83,13 +115,13 @@ pub(crate) fn start_ping_session(
sender.close().await; sender.close().await;
break; break;
} }
pending.insert(id, Instant::now()); tracker.sent(id);
} }
pong = pong_rx.recv() => match pong { pong = pong_rx.recv() => match pong {
Some(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; let mut last_ping = ping_state.lock().await;
*last_ping = Some(sent_at.elapsed()); *last_ping = Some(ping);
} }
} }
None => break, None => break,
@ -100,3 +132,32 @@ pub(crate) fn start_ping_session(
Some(PingSession { last_ping, task }) 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);
}
}

View file

@ -938,12 +938,8 @@ impl Receiver {
} }
/* Route reserved Pong frames to a connection-level observer. */ /* Route reserved Pong frames to a connection-level observer. */
pub fn observe_pongs(&self, observer: mpsc::UnboundedSender<CommunicationValue>) { pub async fn observe_pongs(&self, observer: mpsc::UnboundedSender<CommunicationValue>) {
if let Ok(mut control) = self.inner.ping_control.try_write() { self.inner.ping_control.write().await.pong_observer = Some(observer);
control.pong_observer = Some(observer);
} else {
warn!("[Receiver] could not register Pong observer: control lock busy");
}
} }
#[instrument(skip(stream, policy), level = "trace")] #[instrument(skip(stream, policy), level = "trace")]

View file

@ -53,9 +53,9 @@ fn route_incoming_frame(
if let Some(sent_at) = frame_id(frame).and_then(|id| pending_pings.borrow_mut().remove(&id)) 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)); ping_ms.set(Some(js_sys::Date::now() - sent_at));
}
return; return;
} }
}
if let Some(request_id) = frame_id(frame) { if let Some(request_id) = frame_id(frame) {
let pending = pending_requests.borrow_mut().remove(&request_id); let pending = pending_requests.borrow_mut().remove(&request_id);
@ -138,6 +138,33 @@ async fn wait_for_timeout(timeout_ms: u32) -> Result<(), JsValue> {
Ok(()) Ok(())
} }
fn set_shared_state(
state: &Rc<Cell<ConnectionState>>,
pending_state_callbacks: &Rc<RefCell<VecDeque<ConnectionState>>>,
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::<js_sys::Function>());
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::<js_sys::Function>())
.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] #[wasm_bindgen]
#[derive(Debug, Clone, Copy, PartialEq, Eq)] #[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ConnectionState { pub enum ConnectionState {
@ -150,6 +177,7 @@ pub enum ConnectionState {
#[wasm_bindgen] #[wasm_bindgen]
pub struct WasmClient { pub struct WasmClient {
transport: Rc<RefCell<Option<WasmTransport>>>, transport: Rc<RefCell<Option<WasmTransport>>>,
connection_generation: Rc<Cell<u32>>,
state: Rc<Cell<ConnectionState>>, state: Rc<Cell<ConnectionState>>,
pending_state_callbacks: Rc<RefCell<VecDeque<ConnectionState>>>, pending_state_callbacks: Rc<RefCell<VecDeque<ConnectionState>>>,
state_callback: Closure<dyn FnMut()>, state_callback: Closure<dyn FnMut()>,
@ -187,6 +215,7 @@ impl WasmClient {
}) as Box<dyn FnMut()>); }) as Box<dyn FnMut()>);
Self { Self {
transport: Rc::new(RefCell::new(None)), transport: Rc::new(RefCell::new(None)),
connection_generation: Rc::new(Cell::new(0)),
state: Rc::new(Cell::new(ConnectionState::Disconnected)), state: Rc::new(Cell::new(ConnectionState::Disconnected)),
pending_state_callbacks, pending_state_callbacks,
state_callback, state_callback,
@ -221,7 +250,7 @@ impl WasmClient {
#[wasm_bindgen] #[wasm_bindgen]
pub async fn connect(&self, config: &ConnectionConfig) -> Result<(), JsValue> { pub async fn connect(&self, config: &ConnectionConfig) -> Result<(), JsValue> {
self.set_state(ConnectionState::Connecting); let generation = self.begin_connection();
let transport = WasmTransport::connect( let transport = WasmTransport::connect(
&config.url, &config.url,
config.server_certificate_hashes.clone(), config.server_certificate_hashes.clone(),
@ -249,7 +278,7 @@ impl WasmClient {
.map_err(|e| js_error(format!("parse handshake outcome: {e}")))?; .map_err(|e| js_error(format!("parse handshake outcome: {e}")))?;
let tm = mtp_codec::TypeMap::latest(); let tm = mtp_codec::TypeMap::latest();
if Some(outcome.get_type()) == CommunicationType::ErrorBadVersion.try_to_id(&tm) { 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( return Err(js_error(
outcome outcome
.get_str(DataType::ErrorMessage) .get_str(DataType::ErrorMessage)
@ -262,7 +291,7 @@ impl WasmClient {
if outcome.get_type() != expected if outcome.get_type() != expected
|| outcome.get_data(DataType::Connected) != &DataValue::BoolTrue || outcome.get_data(DataType::Connected) != &DataValue::BoolTrue
{ {
self.set_state(ConnectionState::Disconnected); self.set_state_if_current(generation, ConnectionState::Disconnected);
return Err(js_error( return Err(js_error(
outcome outcome
.get_str(DataType::ErrorMessage) .get_str(DataType::ErrorMessage)
@ -272,12 +301,14 @@ impl WasmClient {
match outcome.get_data(DataType::Version) { match outcome.get_data(DataType::Version) {
DataValue::Str(version) if mtp_codec::Version::parse(version).is_some() => {} 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")); 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(()) Ok(())
} }
@ -289,7 +320,7 @@ impl WasmClient {
keyring_bytes: &[u8], keyring_bytes: &[u8],
client_id: u64, client_id: u64,
) -> Result<u64, JsValue> { ) -> Result<u64, JsValue> {
self.set_state(ConnectionState::Connecting); let generation = self.begin_connection();
let host_pk = mtp_crypto::PublicKeyBundle::from_bytes(host_public_key_bytes) let host_pk = mtp_crypto::PublicKeyBundle::from_bytes(host_public_key_bytes)
.map_err(|e| js_error(format!("invalid host public key: {}", e)))?; .map_err(|e| js_error(format!("invalid host public key: {}", e)))?;
@ -326,6 +357,7 @@ impl WasmClient {
"auth_connect challenge", "auth_connect challenge",
config.require_pq, config.require_pq,
!keyring.sig_pq_secret_key.as_bytes().is_empty(), !keyring.sig_pq_secret_key.as_bytes().is_empty(),
generation,
) )
.await?; .await?;
@ -348,7 +380,7 @@ impl WasmClient {
.try_to_id(&tm) .try_to_id(&tm)
.ok_or_else(|| js_error("IdentificationResponse is absent from the type map"))?; .ok_or_else(|| js_error("IdentificationResponse is absent from the type map"))?;
if resp_type != expected_type { 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( return Err(auth::unexpected_response_type_error(
"auth_connect", "auth_connect",
expected_type, expected_type,
@ -359,7 +391,7 @@ impl WasmClient {
} }
if resp_comm.get_data(DataType::Connected) != &DataValue::BoolTrue { 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( return Err(js_error(
resp_comm resp_comm
.get_str(DataType::ErrorMessage) .get_str(DataType::ErrorMessage)
@ -376,19 +408,21 @@ impl WasmClient {
server_challenge, server_challenge,
config.require_pq, config.require_pq,
) { ) {
self.set_state(ConnectionState::Disconnected); self.set_state_if_current(generation, ConnectionState::Disconnected);
return Err(e); return Err(e);
} }
let assigned_id = match resp_comm.get_data(DataType::Id) { let assigned_id = match resp_comm.get_data(DataType::Id) {
DataValue::UnsignedNumber(n) => *n as u64, 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")); 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) Ok(assigned_id)
} }
@ -400,7 +434,7 @@ impl WasmClient {
host_public_key_bytes: &[u8], host_public_key_bytes: &[u8],
keyring_bytes: &[u8], keyring_bytes: &[u8],
) -> Result<u64, JsValue> { ) -> Result<u64, JsValue> {
self.set_state(ConnectionState::Connecting); let generation = self.begin_connection();
let host_pk = mtp_crypto::PublicKeyBundle::from_bytes(host_public_key_bytes) let host_pk = mtp_crypto::PublicKeyBundle::from_bytes(host_public_key_bytes)
.map_err(|e| js_error(format!("invalid host public key: {}", e)))?; .map_err(|e| js_error(format!("invalid host public key: {}", e)))?;
@ -438,6 +472,7 @@ impl WasmClient {
"auth_register challenge", "auth_register challenge",
config.require_pq, config.require_pq,
!keyring.sig_pq_secret_key.as_bytes().is_empty(), !keyring.sig_pq_secret_key.as_bytes().is_empty(),
generation,
) )
.await?; .await?;
@ -460,7 +495,7 @@ impl WasmClient {
.try_to_id(&tm) .try_to_id(&tm)
.ok_or_else(|| js_error("RegisterResponse is absent from the type map"))?; .ok_or_else(|| js_error("RegisterResponse is absent from the type map"))?;
if resp_type != expected_type { 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( return Err(auth::unexpected_response_type_error(
"auth_register", "auth_register",
expected_type, expected_type,
@ -471,7 +506,7 @@ impl WasmClient {
} }
if resp_comm.get_data(DataType::Connected) != &DataValue::BoolTrue { 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( return Err(js_error(
resp_comm resp_comm
.get_str(DataType::ErrorMessage) .get_str(DataType::ErrorMessage)
@ -482,7 +517,7 @@ impl WasmClient {
let assigned_id = match resp_comm.get_data(DataType::Id) { let assigned_id = match resp_comm.get_data(DataType::Id) {
DataValue::UnsignedNumber(n) => *n as u64, 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")); return Err(js_error("missing assigned ID"));
} }
}; };
@ -496,11 +531,13 @@ impl WasmClient {
server_challenge, server_challenge,
config.require_pq, config.require_pq,
) { ) {
self.set_state(ConnectionState::Disconnected); self.set_state_if_current(generation, ConnectionState::Disconnected);
return Err(e); 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) Ok(assigned_id)
} }
@ -674,6 +711,8 @@ impl WasmClient {
#[wasm_bindgen] #[wasm_bindgen]
pub fn disconnect(&self) { pub fn disconnect(&self) {
self.connection_generation
.set(self.connection_generation.get().wrapping_add(1));
self.stop_protocol_pings(); self.stop_protocol_pings();
if let Some(t) = self.transport.borrow_mut().take() { if let Some(t) = self.transport.borrow_mut().take() {
t.close(); t.close();
@ -733,39 +772,50 @@ impl WasmClient {
} }
fn set_state(&self, new_state: ConnectionState) { fn set_state(&self, new_state: ConnectionState) {
self.state.set(new_state); set_shared_state(
self.pending_state_callbacks &self.state,
.borrow_mut() &self.pending_state_callbacks,
.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::<js_sys::Function>());
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::<js_sys::Function>())
.and_then(|set_timeout| {
set_timeout.call2(
&global,
self.state_callback.as_ref(), self.state_callback.as_ref(),
&JsValue::from_f64(0.0), new_state,
) );
}) }
.is_err()
{ fn set_state_if_current(&self, generation: u32, new_state: ConnectionState) {
self.pending_state_callbacks.borrow_mut().pop_back(); 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(); let loop_transport = transport.clone();
*self.transport.borrow_mut() = Some(transport); *self.transport.borrow_mut() = Some(transport);
self.set_state(ConnectionState::Connected); self.set_state(ConnectionState::Connected);
let connection_generation = self.connection_generation.clone();
let error_generation = connection_generation.clone();
let state = self.state.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_msg = self.on_message.clone();
let on_err = self.on_error.clone(); let on_err = self.on_error.clone();
let subscriptions = self.subscriptions.clone(); let subscriptions = self.subscriptions.clone();
@ -847,7 +897,11 @@ impl WasmClient {
&loop_ping_ms, &loop_ping_ms,
); );
}, },
on_err.clone(), move |error| {
if error_generation.get() == generation {
let _ = on_err.call1(&JsValue::NULL, &error);
}
},
move |pipe_reader: PipeReader| { move |pipe_reader: PipeReader| {
let pipe_id = pipe_reader.pipe_id(); let pipe_id = pipe_reader.pipe_id();
let mut pending = pending_pipes.borrow_mut(); let mut pending = pending_pipes.borrow_mut();
@ -857,13 +911,22 @@ impl WasmClient {
}, },
) )
.await; .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); stop_ping_timer(&ping_timer);
pending_pings.borrow_mut().clear(); pending_pings.borrow_mut().clear();
ping_ms.set(None); ping_ms.set(None);
reject_pending_requests(&pending_requests, "disconnected"); reject_pending_requests(&pending_requests, "disconnected");
client_pipe::reject_pending_pipe_creations(&pending_pipe_creations, "disconnected"); client_pipe::reject_pending_pipe_creations(&pending_pipe_creations, "disconnected");
}); });
true
} }
fn reject_pending_requests(&self, message: &str) { fn reject_pending_requests(&self, message: &str) {
@ -879,6 +942,7 @@ impl WasmClient {
context: &str, context: &str,
require_pq: bool, require_pq: bool,
client_has_pq_key: bool, client_has_pq_key: bool,
generation: u32,
) -> Result<u128, JsValue> { ) -> Result<u128, JsValue> {
let challenge_bytes = transport.read_one_frame().await?; let challenge_bytes = transport.read_one_frame().await?;
let challenge = CommunicationValue::from_bytes(&challenge_bytes) let challenge = CommunicationValue::from_bytes(&challenge_bytes)
@ -887,7 +951,7 @@ impl WasmClient {
.try_to_id(tm) .try_to_id(tm)
.ok_or_else(|| js_error("Challenge is absent from the type map"))?; .ok_or_else(|| js_error("Challenge is absent from the type map"))?;
if challenge.get_type() != expected { 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( return Err(auth::unexpected_response_type_error(
context, context,
expected, expected,
@ -900,13 +964,13 @@ impl WasmClient {
let server_challenge = match challenge.get_data(DataType::ServerNonce) { let server_challenge = match challenge.get_data(DataType::ServerNonce) {
DataValue::UnsignedNumber(n) => *n, DataValue::UnsignedNumber(n) => *n,
_ => { _ => {
self.set_state(ConnectionState::Disconnected); self.set_state_if_current(generation, ConnectionState::Disconnected);
return Err(js_error("missing server challenge")); return Err(js_error("missing server challenge"));
} }
}; };
if challenge.get_data(DataType::RequirePq) == &DataValue::BoolTrue && !client_has_pq_key { 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( return Err(js_error(
"host requires post-quantum authentication but the client PQ key is absent", "host requires post-quantum authentication but the client PQ key is absent",
)); ));
@ -920,7 +984,7 @@ impl WasmClient {
server_challenge, server_challenge,
require_pq, require_pq,
) { ) {
self.set_state(ConnectionState::Disconnected); self.set_state_if_current(generation, ConnectionState::Disconnected);
return Err(e); return Err(e);
} }

View file

@ -477,14 +477,15 @@ impl WasmTransport {
/// Pipe-aware receive loop. Identical to `receive_loop` but detects /// Pipe-aware receive loop. Identical to `receive_loop` but detects
/// `PipeRequest` as the first frame on a new incoming stream and routes /// `PipeRequest` as the first frame on a new incoming stream and routes
/// the stream to `on_pipe` instead of `on_message`. /// the stream to `on_pipe` instead of `on_message`.
pub async fn receive_loop_with_pipes<F, G>( pub async fn receive_loop_with_pipes<F, G, H>(
&self, &self,
mut on_message: F, mut on_message: F,
on_error: js_sys::Function, mut on_error: H,
mut on_pipe: G, mut on_pipe: G,
) where ) where
F: FnMut(JsValue), F: FnMut(JsValue),
G: FnMut(crate::pipe::PipeReader), G: FnMut(crate::pipe::PipeReader),
H: FnMut(JsValue),
{ {
let pipe_request_type = let pipe_request_type =
mtp_codec::CommunicationType::PipeRequest.try_to_id(&mtp_codec::TypeMap::latest()); mtp_codec::CommunicationType::PipeRequest.try_to_id(&mtp_codec::TypeMap::latest());
@ -528,13 +529,13 @@ impl WasmTransport {
} }
Err(e) => { Err(e) => {
let message = e.as_string().unwrap_or_else(|| format!("{:?}", 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, Ok(FrameOutcome::Closed) | Ok(FrameOutcome::Ended) => break,
Err(e) => { Err(e) => {
let _ = on_error.call1(&JsValue::NULL, &e); on_error(e);
break; break;
} }
} }