This commit is contained in:
parent
590810ce59
commit
b6483b7f6d
6 changed files with 198 additions and 73 deletions
|
|
@ -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")]
|
||||||
{
|
{
|
||||||
|
|
|
||||||
|
|
@ -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)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -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);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
|
||||||
|
|
@ -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")]
|
||||||
|
|
|
||||||
|
|
@ -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))
|
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) {
|
||||||
|
|
@ -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);
|
self.state_callback.as_ref(),
|
||||||
|
new_state,
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
let global = js_sys::global();
|
fn set_state_if_current(&self, generation: u32, new_state: ConnectionState) {
|
||||||
let qmt = js_sys::Reflect::get(&global, &JsValue::from_str("queueMicrotask"))
|
if self.connection_generation.get() == generation {
|
||||||
.and_then(|f| f.dyn_into::<js_sys::Function>());
|
self.set_state(new_state);
|
||||||
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(),
|
|
||||||
&JsValue::from_f64(0.0),
|
|
||||||
)
|
|
||||||
})
|
|
||||||
.is_err()
|
|
||||||
{
|
|
||||||
self.pending_state_callbacks.borrow_mut().pop_back();
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
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);
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -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;
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue