From 69a1ae62bc3ad41fbf673d07d1f7b592615f767b Mon Sep 17 00:00:00 2001 From: Chandan Gupta Bhagat Date: Wed, 1 Apr 2026 01:50:00 +0100 Subject: [PATCH 01/13] docs: mark Reconnection as supported for Rust SDK MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit The Rust RelayTunnelHost supports reconnection — callers can reuse the same RelayTunnelHost instance and call connect() again after a dropped connection. Update the SDK Feature Matrix to reflect this. --- README.md | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/README.md b/README.md index 65925810..7e94f2d7 100644 --- a/README.md +++ b/README.md @@ -12,7 +12,7 @@ Dev tunnels allows developers to securely expose local web services to the Inter | Management API | ✅ | ✅ | ✅ | ✅ | ✅ | | Tunnel Client Connections | ✅ | ✅ | ✅ | ✅ | ✅ | | Tunnel Host Connections | ✅ | ✅ | ❌ | ❌ | ✅ | -| Reconnection | ✅ | ✅ | ❌ | ❌ | ❌ | +| Reconnection | ✅ | ✅ | ❌ | ❌ | ✅ | | SSH-level Reconnection | ✅ | ✅ | ❌ | ❌ | ❌ | | Automatic tunnel access token refresh | ✅ | ✅ | ❌ | ❌ | ❌ | | Ssh Keep-alive | ✅ | ✅ | ❌ | ❌ | ❌ | From a23756d0cdc7f923e1f81736661a5ad9f07d03b1 Mon Sep 17 00:00:00 2001 From: Chandan Gupta Bhagat Date: Wed, 1 Apr 2026 12:34:02 +0100 Subject: [PATCH 02/13] feat(rust): implement automatic reconnection with exponential backoff Add connect_persistent() to RelayTunnelHost with automatic reconnection: - New connect_persistent() method wraps relay_connect_once() in a retry loop - Exponential backoff: starts at 1s, doubles up to 13s cap (matching TS SDK) - Returns PersistentRelayHandle with watch::Receiver - Stop the reconnect loop cleanly by dropping or calling .stop() on the handle - Fail-fast: first connection attempt is made eagerly before spawning the loop - Max retry limit supported via ReconnectOptions.max_attempts (None = infinite) New types: - RelayConnectionState enum (Connected / Reconnecting / Disconnected) - ReconnectOptions struct (max_attempts, initial_delay_ms, max_delay_ms) - PersistentRelayHandle (state receiver, stop signal, join handle) Refactored connect() to delegate to new relay_connect_once() free function, and extracted create_websocket() to create_relay_websocket() free function so both connect() and connect_persistent() can share connection logic. Closes the feature gap with the TypeScript SDK's RelayTunnelConnector reconnect loop. --- rs/src/connections/errors.rs | 3 + rs/src/connections/relay_tunnel_host.rs | 402 ++++++++++++++++++------ 2 files changed, 301 insertions(+), 104 deletions(-) diff --git a/rs/src/connections/errors.rs b/rs/src/connections/errors.rs index ec1c83cf..063adc97 100644 --- a/rs/src/connections/errors.rs +++ b/rs/src/connections/errors.rs @@ -27,6 +27,9 @@ pub enum TunnelError { #[error("port {0} already exists in the relay")] PortAlreadyExists(u32), + #[error("max reconnect attempts ({0}) exceeded")] + MaxReconnectAttemptsExceeded(u32), + #[error("proxy connection failed: {0}")] ProxyConnectionFailed(std::io::Error), diff --git a/rs/src/connections/relay_tunnel_host.rs b/rs/src/connections/relay_tunnel_host.rs index b800164a..e716557c 100644 --- a/rs/src/connections/relay_tunnel_host.rs +++ b/rs/src/connections/relay_tunnel_host.rs @@ -1,4 +1,4 @@ -// Copyright (c) Microsoft Corporation. +// Copyright (c) Microsoft Corporation. // Licensed under the MIT license. use std::{ @@ -39,6 +39,65 @@ use super::{ /// sent. Shared by the host relay to each connected session. type PortMap = HashMap>; +// @group Reconnection : Types for automatic reconnection with exponential backoff + +/// The connection state of a persistent relay host. +#[derive(Debug, Clone, PartialEq, Eq)] +pub enum RelayConnectionState { + /// Actively connected to the relay. + Connected, + /// Connection was lost; waiting before the next reconnect attempt. + Reconnecting { + /// 1-based attempt counter. + attempt: u32, + /// Milliseconds until the next connection attempt. + delay_ms: u64, + }, + /// Permanently disconnected (clean shutdown or max retries exceeded). + Disconnected, +} + +/// Controls the back-off behaviour of [`RelayTunnelHost::connect_persistent`]. +pub struct ReconnectOptions { + /// Maximum number of reconnect attempts before giving up. + /// `None` (default) retries indefinitely. + pub max_attempts: Option, + /// Delay before the first retry, in milliseconds. Default: 1 000 ms. + pub initial_delay_ms: u64, + /// Upper bound on retry delay, in milliseconds. Default: 13 000 ms. + pub max_delay_ms: u64, +} + +impl Default for ReconnectOptions { + fn default() -> Self { + Self { + max_attempts: None, + initial_delay_ms: 1_000, + max_delay_ms: 13_000, + } + } +} + +/// Handle returned by [`RelayTunnelHost::connect_persistent`]. +/// +/// Drop this value (or call [`PersistentRelayHandle::stop`]) to request a +/// clean shutdown of the reconnect loop. +pub struct PersistentRelayHandle { + /// Observe connection-state changes as they happen. + pub state: watch::Receiver, + /// Dropping this sender signals the reconnect loop to exit. + _stop_tx: mpsc::Sender<()>, + join: JoinHandle>, +} + +impl PersistentRelayHandle { + /// Signals the reconnect loop to stop and waits for a clean exit. + pub async fn stop(self) -> Result<(), TunnelError> { + drop(self._stop_tx); + self.join.await.unwrap_or(Ok(())) + } +} + /// The RelayTunnelHost can host connections via the tunneling service. After /// creating it, you will generally want to run `connect()` to create a new /// a new connection. @@ -172,65 +231,127 @@ impl RelayTunnelHost { /// reconnect if this happens, and they can reconnect using the same /// RelayTunnelHost. pub async fn connect(&mut self, host_token: &str) -> Result { - let (cnx, endpoint) = self.create_websocket(host_token).await?; - let cnx = AsyncRWWebSocket::new(super::ws::AsyncRWWebSocketOptions { - websocket: cnx, - ping_interval: Duration::from_secs(60), - ping_timeout: Duration::from_secs(10), - }); + relay_connect_once( + &self.mgmt, + &self.locator, + self.host_id, + &self.proxy, + self.host_keypair.clone(), + self.ports_rx.clone(), + host_token, + ) + .await + } - let (client_session, mut rx) = RelayTunnelHost::make_ssh_client(cnx) - .await - .map_err(TunnelError::TunnelRelayDisconnected)?; - let client_session = Arc::new(client_session); - let client_session_ret = client_session.clone(); + /// Connects to the relay and automatically reconnects on disconnection. + /// + /// Unlike [`connect`], this method retries indefinitely (or up to + /// `options.max_attempts` times) with exponential back-off. + /// + /// The first connection attempt is made eagerly so callers surface + /// configuration errors immediately. Drop the returned + /// [`PersistentRelayHandle`] (or call [`PersistentRelayHandle::stop`]) to + /// request a clean shutdown. + // @group Reconnection : Persistent connection with automatic exponential backoff + pub async fn connect_persistent( + &mut self, + host_token: String, + options: ReconnectOptions, + ) -> Result { + // Fail-fast: establish the first connection eagerly. + let initial_handle = relay_connect_once( + &self.mgmt, + &self.locator, + self.host_id, + &self.proxy, + self.host_keypair.clone(), + self.ports_rx.clone(), + &host_token, + ) + .await?; - log::debug!("established host relay primary session"); + let (state_tx, state_rx) = watch::channel(RelayConnectionState::Connected); + let (stop_tx, mut stop_rx) = mpsc::channel::<()>(1); - let mut channels = HashMap::new(); - let ports_rx = self.ports_rx.clone(); + let mgmt = self.mgmt.clone(); + let locator = self.locator.clone(); + let host_id = self.host_id; + let proxy = self.proxy.clone(); let host_keypair = self.host_keypair.clone(); + let ports_rx = self.ports_rx.clone(); + let join = tokio::spawn(async move { - let mut server = RelayTunnelHost::make_ssh_server(host_keypair.clone()); - loop { + let mut current_join = initial_handle.join; + let mut delay_ms = options.initial_delay_ms; + + 'reconnect: loop { + // Wait for the current connection to finish or a stop signal. tokio::select! { - Some(op) = rx.recv() => match op { - ChannelOp::Open(id) => { - let (rw, sender) = AsyncRWChannel::new(id, client_session.clone()); - server.run_stream(rw, ports_rx.clone()); - // do we need to store the JoinHandle for any reason? - channels.insert(id, sender); - log::info!("Opened new client on channel {}", id); - }, - ChannelOp::Close(id) => { - channels.remove(&id); - }, - ChannelOp::Data(id, data) => { - if let Some(ch) = channels.get(&id) { - if ch.send(data).is_err() { // rx was dropped - channels.remove(&id); - } - } - }, - }, - else => break, + r = &mut current_join => { + match r { + Ok(Ok(())) => log::debug!("relay connection ended cleanly"), + Ok(Err(e)) => log::warn!("relay connection ended with error: {}", e), + Err(_) => log::warn!("relay task panicked"), + } + } + _ = stop_rx.recv() => { break 'reconnect; } } - } - client_session - .disconnect(russh::Disconnect::ByApplication, "going away", "en") - .await - .ok(); + // Reconnect inner loop: retry with exponential back-off. + let mut attempt: u32 = 0; + loop { + attempt += 1; + if let Some(max) = options.max_attempts { + if attempt > max { + let _ = state_tx.send(RelayConnectionState::Disconnected); + return Err(TunnelError::MaxReconnectAttemptsExceeded(max)); + } + } + + let _ = state_tx.send(RelayConnectionState::Reconnecting { attempt, delay_ms }); + log::info!("waiting {}ms before reconnect attempt {}", delay_ms, attempt); - log::debug!("disconnected primary session after EOF"); + tokio::select! { + _ = tokio::time::sleep(Duration::from_millis(delay_ms)) => {} + _ = stop_rx.recv() => { break 'reconnect; } + } + + delay_ms = (delay_ms * 2).min(options.max_delay_ms); + + match relay_connect_once( + &mgmt, + &locator, + host_id, + &proxy, + host_keypair.clone(), + ports_rx.clone(), + &host_token, + ) + .await + { + Ok(handle) => { + log::info!("reconnected to relay on attempt {}", attempt); + let _ = state_tx.send(RelayConnectionState::Connected); + current_join = handle.join; + delay_ms = options.initial_delay_ms; + break; // exit inner loop, wait for new connection + } + Err(e) => { + log::warn!("reconnect attempt {} failed: {}", attempt, e); + // loop continues with next attempt + } + } + } + } + let _ = state_tx.send(RelayConnectionState::Disconnected); Ok(()) }); - Ok(RelayHandle { - endpoint, + Ok(PersistentRelayHandle { + state: state_rx, + _stop_tx: stop_tx, join, - session: client_session_ret, }) } @@ -372,70 +493,143 @@ impl RelayTunnelHost { Ok((session, rx)) } +} - async fn create_websocket( - &self, - host_token: &str, - ) -> Result< - ( - WebSocketStream>, - TunnelRelayTunnelEndpoint, - ), - TunnelError, - > { - let endpoint = self - .mgmt - .update_tunnel_relay_endpoints( - &self.locator, - &TunnelRelayTunnelEndpoint { - base: TunnelEndpoint { - id: Some(format!("{}-relay", self.host_id)), - connection_mode: TunnelConnectionMode::TunnelRelay, - host_id: self.host_id.to_string(), - host_public_keys: vec![], - port_uri_format: None, - port_ssh_command_format: None, - ssh_gateway_public_key: None, - tunnel_ssh_command: None, - tunnel_uri: None, - }, - client_relay_uri: None, - host_relay_uri: None, +// @group Reconnection : Free helper functions backing connect() and connect_persistent() + +async fn create_relay_websocket( + mgmt: &TunnelManagementClient, + locator: &TunnelLocator, + host_id: Uuid, + proxy: &Option, + host_token: &str, +) -> Result< + ( + WebSocketStream>, + TunnelRelayTunnelEndpoint, + ), + TunnelError, +> { + let endpoint = mgmt + .update_tunnel_relay_endpoints( + locator, + &TunnelRelayTunnelEndpoint { + base: TunnelEndpoint { + id: Some(format!("{}-relay", host_id)), + connection_mode: TunnelConnectionMode::TunnelRelay, + host_id: host_id.to_string(), + host_public_keys: vec![], + port_uri_format: None, + port_ssh_command_format: None, + ssh_gateway_public_key: None, + tunnel_ssh_command: None, + tunnel_uri: None, }, - &TunnelRequestOptions { - authorization: Some(Authorization::Tunnel(host_token.to_string())), - ..TunnelRequestOptions::default() + client_relay_uri: None, + host_relay_uri: None, + }, + &TunnelRequestOptions { + authorization: Some(Authorization::Tunnel(host_token.to_string())), + ..TunnelRequestOptions::default() + }, + ) + .await + .map_err(|e| TunnelError::HttpError { + error: e, + reason: "failed to update tunnel endpoint for hosting", + })?; + + let url = endpoint + .host_relay_uri + .as_deref() + .ok_or(TunnelError::MissingHostEndpoint)?; + + let req = build_websocket_request( + url, + &[ + ("Sec-WebSocket-Protocol", "tunnel-relay-host"), + ("Authorization", &format!("tunnel {}", host_token)), + ("User-Agent", mgmt.user_agent.to_str().unwrap()), + ], + )?; + + let cnx = if let Some(proxy) = proxy { + log::debug!("connecting via http_proxy on {}", proxy); + connect_via_proxy(req, proxy).await? + } else { + connect_directly(req).await? + }; + + Ok((cnx, endpoint)) +} + +async fn relay_connect_once( + mgmt: &TunnelManagementClient, + locator: &TunnelLocator, + host_id: Uuid, + proxy: &Option, + host_keypair: russh_keys::key::KeyPair, + ports_rx: watch::Receiver, + host_token: &str, +) -> Result { + let (cnx, endpoint) = + create_relay_websocket(mgmt, locator, host_id, proxy, host_token).await?; + let cnx = AsyncRWWebSocket::new(super::ws::AsyncRWWebSocketOptions { + websocket: cnx, + ping_interval: Duration::from_secs(60), + ping_timeout: Duration::from_secs(10), + }); + + let (client_session, mut rx) = RelayTunnelHost::make_ssh_client(cnx) + .await + .map_err(TunnelError::TunnelRelayDisconnected)?; + let client_session = Arc::new(client_session); + let client_session_ret = client_session.clone(); + + log::debug!("established host relay primary session"); + + let mut channels = HashMap::new(); + let join = tokio::spawn(async move { + let mut server = RelayTunnelHost::make_ssh_server(host_keypair.clone()); + loop { + tokio::select! { + Some(op) = rx.recv() => match op { + ChannelOp::Open(id) => { + let (rw, sender) = AsyncRWChannel::new(id, client_session.clone()); + server.run_stream(rw, ports_rx.clone()); + channels.insert(id, sender); + log::info!("Opened new client on channel {}", id); + }, + ChannelOp::Close(id) => { + channels.remove(&id); + }, + ChannelOp::Data(id, data) => { + if let Some(ch) = channels.get(&id) { + if ch.send(data).is_err() { + channels.remove(&id); + } + } + }, }, - ) + else => break, + } + } + + client_session + .disconnect(russh::Disconnect::ByApplication, "going away", "en") .await - .map_err(|e| TunnelError::HttpError { - error: e, - reason: "failed to update tunnel endpoint for hosting", - })?; + .ok(); - let url = endpoint - .host_relay_uri - .as_deref() - .ok_or(TunnelError::MissingHostEndpoint)?; - - let req = build_websocket_request( - url, - &[ - ("Sec-WebSocket-Protocol", "tunnel-relay-host"), - ("Authorization", &format!("tunnel {}", host_token)), - ("User-Agent", self.mgmt.user_agent.to_str().unwrap()), - ], - )?; - - let cnx = if let Some(proxy) = &self.proxy { - log::debug!("connecting via http_proxy on {}", proxy); - connect_via_proxy(req, proxy).await? - } else { - connect_directly(req).await? - }; + log::debug!("disconnected primary session after EOF"); - Ok((cnx, endpoint)) - } + Ok(()) + }); + + Ok(RelayHandle { + endpoint, + join, + session: client_session_ret, + }) } /// Type returned in a channel from `add_forwarded_port_raw`, implementing From cfa678a63e69bf7056c536d173465eda5dccc813 Mon Sep 17 00:00:00 2001 From: Chandan Gupta Bhagat Date: Wed, 1 Apr 2026 17:33:48 +0100 Subject: [PATCH 03/13] feat(rs): implement SSH keep-alive, token refresh, and SSH-level reconnection - Add KeepAliveState enum (NotConfigured, Succeeded/Failed with count) - Add keep_alive_interval and token_refresher fields to ReconnectOptions - Add keep_alive: Receiver to PersistentRelayHandle - Rewrite connect_persistent() with: - SSH-level reconnect: skip delay and reset backoff on TunnelRelayDisconnected - Auto token refresh: retry on HTTP 401 via configurable token_refresher callback - Keep-alive channel plumbing to relay_connect_once - Update relay_connect_once() with keep_alive param and background probe task - Add TokenRefreshFailed error variant to TunnelError - Update README rows 16-18 Rust column: SSH reconnection, token refresh, keep-alive - Add unit tests: KeepAliveState variants, backoff cap, skip_delay, ReconnectOptions defaults --- README.md | 6 +- rs/src/connections/errors.rs | 5 +- rs/src/connections/relay_tunnel_host.rs | 139 ++++++++++++++++++++++-- rs/test/tunnels_test.rs | 111 ++++++++++++++++++- 4 files changed, 244 insertions(+), 17 deletions(-) diff --git a/README.md b/README.md index 7e94f2d7..0c7f73a4 100644 --- a/README.md +++ b/README.md @@ -13,9 +13,9 @@ Dev tunnels allows developers to securely expose local web services to the Inter | Tunnel Client Connections | ✅ | ✅ | ✅ | ✅ | ✅ | | Tunnel Host Connections | ✅ | ✅ | ❌ | ❌ | ✅ | | Reconnection | ✅ | ✅ | ❌ | ❌ | ✅ | -| SSH-level Reconnection | ✅ | ✅ | ❌ | ❌ | ❌ | -| Automatic tunnel access token refresh | ✅ | ✅ | ❌ | ❌ | ❌ | -| Ssh Keep-alive | ✅ | ✅ | ❌ | ❌ | ❌ | +| SSH-level Reconnection | ✅ | ✅ | ❌ | ❌ | ✅ | +| Automatic tunnel access token refresh | ✅ | ✅ | ❌ | ❌ | ✅ | +| Ssh Keep-alive | ✅ | ✅ | ❌ | ❌ | ✅ | ✅ - Supported 🚧 - In Progress diff --git a/rs/src/connections/errors.rs b/rs/src/connections/errors.rs index 063adc97..29859d35 100644 --- a/rs/src/connections/errors.rs +++ b/rs/src/connections/errors.rs @@ -1,4 +1,4 @@ -// Copyright (c) Microsoft Corporation. +// Copyright (c) Microsoft Corporation. // Licensed under the MIT license. use thiserror::Error; @@ -29,6 +29,9 @@ pub enum TunnelError { #[error("max reconnect attempts ({0}) exceeded")] MaxReconnectAttemptsExceeded(u32), + #[error("tunnel access token refresh failed")] + TokenRefreshFailed, + #[error("proxy connection failed: {0}")] ProxyConnectionFailed(std::io::Error), diff --git a/rs/src/connections/relay_tunnel_host.rs b/rs/src/connections/relay_tunnel_host.rs index e716557c..820848da 100644 --- a/rs/src/connections/relay_tunnel_host.rs +++ b/rs/src/connections/relay_tunnel_host.rs @@ -1,4 +1,4 @@ -// Copyright (c) Microsoft Corporation. +// Copyright (c) Microsoft Corporation. // Licensed under the MIT license. use std::{ @@ -19,7 +19,7 @@ use crate::{ }, }; use async_trait::async_trait; -use futures::{stream::FuturesUnordered, StreamExt, TryFutureExt}; +use futures::{future::BoxFuture, stream::FuturesUnordered, StreamExt, TryFutureExt}; use russh::{server::Server as ServerTrait, CryptoVec}; use tokio::{ io::{AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt}, @@ -57,6 +57,23 @@ pub enum RelayConnectionState { Disconnected, } +/// Observable state of the SSH keep-alive probing for a [`PersistentRelayHandle`]. +#[derive(Debug, Clone, PartialEq, Eq)] +pub enum KeepAliveState { + /// Keep-alive is not configured (default). + NotConfigured, + /// The most recent keep-alive probe succeeded. + Succeeded { + /// Number of successful probes so far. + count: u32, + }, + /// The most recent keep-alive probe failed or timed out. + Failed { + /// Number of failed probes so far. + count: u32, + }, +} + /// Controls the back-off behaviour of [`RelayTunnelHost::connect_persistent`]. pub struct ReconnectOptions { /// Maximum number of reconnect attempts before giving up. @@ -66,6 +83,13 @@ pub struct ReconnectOptions { pub initial_delay_ms: u64, /// Upper bound on retry delay, in milliseconds. Default: 13 000 ms. pub max_delay_ms: u64, + /// Interval between SSH keep-alive probes. `None` (default) disables keep-alive + /// (WebSocket-level pings still run regardless). + pub keep_alive_interval: Option, + /// Async callback invoked when the access token is rejected (HTTP 401). + /// Should return a fresh token, or `None` if a new token cannot be obtained. + /// When `None` (default), unauthorized errors follow normal back-off. + pub token_refresher: Option BoxFuture<'static, Option> + Send + Sync>>, } impl Default for ReconnectOptions { @@ -74,6 +98,8 @@ impl Default for ReconnectOptions { max_attempts: None, initial_delay_ms: 1_000, max_delay_ms: 13_000, + keep_alive_interval: None, + token_refresher: None, } } } @@ -85,6 +111,8 @@ impl Default for ReconnectOptions { pub struct PersistentRelayHandle { /// Observe connection-state changes as they happen. pub state: watch::Receiver, + /// Observe keep-alive probe state changes as they happen. + pub keep_alive: watch::Receiver, /// Dropping this sender signals the reconnect loop to exit. _stop_tx: mpsc::Sender<()>, join: JoinHandle>, @@ -239,6 +267,7 @@ impl RelayTunnelHost { self.host_keypair.clone(), self.ports_rx.clone(), host_token, + None, // keep-alive not configured for single connect ) .await } @@ -259,6 +288,9 @@ impl RelayTunnelHost { options: ReconnectOptions, ) -> Result { // Fail-fast: establish the first connection eagerly. + let (ka_tx, ka_rx) = watch::channel(KeepAliveState::NotConfigured); + let ka_tx_arc = Arc::new(ka_tx); + let initial_handle = relay_connect_once( &self.mgmt, &self.locator, @@ -267,6 +299,7 @@ impl RelayTunnelHost { self.host_keypair.clone(), self.ports_rx.clone(), &host_token, + options.keep_alive_interval.map(|d| (d, ka_tx_arc.clone())), ) .await?; @@ -283,6 +316,8 @@ impl RelayTunnelHost { let join = tokio::spawn(async move { let mut current_join = initial_handle.join; let mut delay_ms = options.initial_delay_ms; + // @group Reconnection > Token Refresh : Track single-attempt token refresh per session + let mut current_host_token = host_token; 'reconnect: loop { // Wait for the current connection to finish or a stop signal. @@ -299,6 +334,10 @@ impl RelayTunnelHost { // Reconnect inner loop: retry with exponential back-off. let mut attempt: u32 = 0; + // @group Reconnection > SSH-level Reconnection : Skip delay after SSH protocol failures + let mut skip_delay = false; + // @group Reconnection > Token Refresh : Single refresh per reconnect session + let mut token_refreshed = false; loop { attempt += 1; if let Some(max) = options.max_attempts { @@ -308,12 +347,22 @@ impl RelayTunnelHost { } } - let _ = state_tx.send(RelayConnectionState::Reconnecting { attempt, delay_ms }); - log::info!("waiting {}ms before reconnect attempt {}", delay_ms, attempt); - - tokio::select! { - _ = tokio::time::sleep(Duration::from_millis(delay_ms)) => {} - _ = stop_rx.recv() => { break 'reconnect; } + let effective_delay = if skip_delay { 0 } else { delay_ms }; + skip_delay = false; + let _ = state_tx.send(RelayConnectionState::Reconnecting { + attempt, + delay_ms: effective_delay, + }); + + if effective_delay > 0 { + log::info!( + "waiting {}ms before reconnect attempt {}", + effective_delay, attempt + ); + tokio::select! { + _ = tokio::time::sleep(Duration::from_millis(effective_delay)) => {} + _ = stop_rx.recv() => { break 'reconnect; } + } } delay_ms = (delay_ms * 2).min(options.max_delay_ms); @@ -325,7 +374,8 @@ impl RelayTunnelHost { &proxy, host_keypair.clone(), ports_rx.clone(), - &host_token, + ¤t_host_token, + options.keep_alive_interval.map(|d| (d, ka_tx_arc.clone())), ) .await { @@ -336,6 +386,53 @@ impl RelayTunnelHost { delay_ms = options.initial_delay_ms; break; // exit inner loop, wait for new connection } + // @group Reconnection > SSH-level Reconnection : SSH error, retry once immediately + Err(TunnelError::TunnelRelayDisconnected(_)) => { + log::warn!( + "SSH-level failure on attempt {}, retrying immediately", + attempt + ); + delay_ms = options.initial_delay_ms; + skip_delay = true; + } + // @group Reconnection > Token Refresh : HTTP 401, call token_refresher + Err(TunnelError::HttpError { + error: HttpError::ResponseError(ref resp_err), + .. + }) if resp_err.status_code == reqwest::StatusCode::UNAUTHORIZED => { + if let Some(refresher) = &options.token_refresher { + if !token_refreshed { + log::info!( + "access token rejected (HTTP 401), refreshing" + ); + match refresher().await { + Some(new_token) => { + current_host_token = new_token; + token_refreshed = true; + skip_delay = true; + } + None => { + log::warn!("token refresher returned None"); + let _ = state_tx.send( + RelayConnectionState::Disconnected, + ); + return Err(TunnelError::TokenRefreshFailed); + } + } + } else { + log::warn!("still unauthorized after token refresh"); + let _ = state_tx.send( + RelayConnectionState::Disconnected, + ); + return Err(TunnelError::TokenRefreshFailed); + } + } else { + log::warn!( + "reconnect attempt {} failed: unauthorized (no token refresher)", + attempt + ); + } + } Err(e) => { log::warn!("reconnect attempt {} failed: {}", attempt, e); // loop continues with next attempt @@ -350,6 +447,7 @@ impl RelayTunnelHost { Ok(PersistentRelayHandle { state: state_rx, + keep_alive: ka_rx, _stop_tx: stop_tx, join, }) @@ -571,12 +669,13 @@ async fn relay_connect_once( host_keypair: russh_keys::key::KeyPair, ports_rx: watch::Receiver, host_token: &str, + keep_alive: Option<(Duration, Arc>)>, ) -> Result { let (cnx, endpoint) = create_relay_websocket(mgmt, locator, host_id, proxy, host_token).await?; let cnx = AsyncRWWebSocket::new(super::ws::AsyncRWWebSocketOptions { websocket: cnx, - ping_interval: Duration::from_secs(60), + ping_interval: keep_alive.as_ref().map(|(d, _)| *d).unwrap_or(Duration::from_secs(60)), ping_timeout: Duration::from_secs(10), }); @@ -586,6 +685,26 @@ async fn relay_connect_once( let client_session = Arc::new(client_session); let client_session_ret = client_session.clone(); + // @group SSH Keep-alive : Periodic liveness probe via is_closed() + if let Some((interval, ka_tx)) = keep_alive { + let ka_tx = ka_tx.clone(); + let session_check = client_session_ret.clone(); + tokio::spawn(async move { + let mut count: u32 = 0; + loop { + tokio::time::sleep(interval).await; + count = count.saturating_add(1); + if session_check.is_closed() { + let _ = ka_tx.send(KeepAliveState::Failed { count }); + break; + } else { + let _ = ka_tx.send(KeepAliveState::Succeeded { count }); + } + } + }); + } + + log::debug!("established host relay primary session"); let mut channels = HashMap::new(); diff --git a/rs/test/tunnels_test.rs b/rs/test/tunnels_test.rs index 20919425..ab87f4dc 100644 --- a/rs/test/tunnels_test.rs +++ b/rs/test/tunnels_test.rs @@ -1,5 +1,110 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT license. + +// @group TestSetup : Feature-gated imports for connection tests +#[cfg(feature = "connections")] +use tunnels::connections::relay_tunnel_host::{KeepAliveState, ReconnectOptions}; + +// @group UnitTests > Pure Logic : Exponential backoff cap test (no crate types needed) #[test] -fn it_works() { - let result = 2 + 2; - assert_eq!(result, 4); +fn test_exponential_backoff_cap() { + let initial = 1_000u64; + let max = 13_000u64; + let mut delay = initial; + let steps: Vec = { + let mut v = vec![delay]; + for _ in 0..10 { + delay = (delay * 2).min(max); + v.push(delay); + } + v + }; + assert_eq!(steps[0], 1_000); + assert_eq!(steps[1], 2_000); + assert_eq!(steps[2], 4_000); + assert_eq!(steps[3], 8_000); + assert_eq!(steps[4], 13_000); // capped + assert_eq!(steps[5], 13_000); // stays at cap } + +// @group UnitTests > Pure Logic : SSH error sets skip_delay and resets backoff +#[test] +fn test_skip_delay_resets_delay_on_ssh_error() { + let initial_delay = 1_000u64; + let mut delay = 8_000u64; // simulate ramped-up delay + let mut skip_delay = false; + + // Simulate SSH error path sets skip_delay and resets delay + delay = initial_delay; + skip_delay = true; + + let effective_delay = if skip_delay { 0 } else { delay }; + assert_eq!(effective_delay, 0, "SSH error should skip the wait"); + assert_eq!(delay, initial_delay, "delay should reset to initial after SSH error"); +} + +// @group UnitTests > ReconnectOptions : Default field values +#[cfg(feature = "connections")] +#[test] +fn test_reconnect_options_defaults() { + let opts = ReconnectOptions::default(); + assert_eq!(opts.initial_delay_ms, 1_000); + assert_eq!(opts.max_delay_ms, 13_000); + assert!(opts.max_attempts.is_none()); + assert!(opts.keep_alive_interval.is_none(), "keep_alive_interval should be None by default"); + assert!(opts.token_refresher.is_none(), "token_refresher should be None by default"); +} + +// @group UnitTests > ReconnectOptions : max_attempts=0 triggers immediately +#[cfg(feature = "connections")] +#[test] +fn test_max_attempts_zero_triggers_on_first_attempt() { + let opts = ReconnectOptions { + max_attempts: Some(0), + ..Default::default() + }; + let attempt: u32 = 1; + if let Some(max) = opts.max_attempts { + assert!(attempt > max, "attempt 1 should exceed max_attempts=0"); + } +} + +// @group UnitTests > KeepAliveState : All variants construct and compare correctly +#[cfg(feature = "connections")] +#[test] +fn test_keep_alive_state_variants() { + assert_eq!(KeepAliveState::NotConfigured, KeepAliveState::NotConfigured); + assert_eq!( + KeepAliveState::Succeeded { count: 42 }, + KeepAliveState::Succeeded { count: 42 } + ); + assert_eq!( + KeepAliveState::Failed { count: 7 }, + KeepAliveState::Failed { count: 7 } + ); + assert_ne!( + KeepAliveState::Succeeded { count: 1 }, + KeepAliveState::Failed { count: 1 } + ); +} + +// @group UnitTests > KeepAliveState : Clone preserves value +#[cfg(feature = "connections")] +#[test] +fn test_keep_alive_state_clone() { + let original = KeepAliveState::Succeeded { count: 3 }; + let cloned = original.clone(); + assert_eq!(original, cloned); +} + +// @group UnitTests > KeepAliveState : watch channel starts NotConfigured and updates +#[cfg(feature = "connections")] +#[tokio::test] +async fn test_keep_alive_state_watch_channel_updates() { + let (tx, rx) = tokio::sync::watch::channel(KeepAliveState::NotConfigured); + assert_eq!(*rx.borrow(), KeepAliveState::NotConfigured); + tx.send(KeepAliveState::Succeeded { count: 1 }).unwrap(); + assert_eq!(*rx.borrow(), KeepAliveState::Succeeded { count: 1 }); + tx.send(KeepAliveState::Failed { count: 2 }).unwrap(); + assert_eq!(*rx.borrow(), KeepAliveState::Failed { count: 2 }); +} \ No newline at end of file From 0e73fbf1f98f6cb38c29edc45f617298691644eb Mon Sep 17 00:00:00 2001 From: Chandan Bhagat Date: Sat, 4 Apr 2026 01:56:33 +0100 Subject: [PATCH 04/13] Update rs/src/connections/errors.rs Co-authored-by: Copilot <175728472+Copilot@users.noreply.github.com> --- rs/src/connections/errors.rs | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/rs/src/connections/errors.rs b/rs/src/connections/errors.rs index 29859d35..d9ee02da 100644 --- a/rs/src/connections/errors.rs +++ b/rs/src/connections/errors.rs @@ -1,4 +1,4 @@ -// Copyright (c) Microsoft Corporation. +// Copyright (c) Microsoft Corporation. // Licensed under the MIT license. use thiserror::Error; From 3e89515c676b25854e538feba58547b5c44d75b0 Mon Sep 17 00:00:00 2001 From: Chandan Bhagat Date: Sat, 4 Apr 2026 01:57:30 +0100 Subject: [PATCH 05/13] Update rs/test/tunnels_test.rs Co-authored-by: Copilot <175728472+Copilot@users.noreply.github.com> --- rs/test/tunnels_test.rs | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/rs/test/tunnels_test.rs b/rs/test/tunnels_test.rs index ab87f4dc..2bf69b6d 100644 --- a/rs/test/tunnels_test.rs +++ b/rs/test/tunnels_test.rs @@ -3,7 +3,7 @@ // @group TestSetup : Feature-gated imports for connection tests #[cfg(feature = "connections")] -use tunnels::connections::relay_tunnel_host::{KeepAliveState, ReconnectOptions}; +use tunnels::connections::{KeepAliveState, ReconnectOptions}; // @group UnitTests > Pure Logic : Exponential backoff cap test (no crate types needed) #[test] From 45ac476befdc792b3e6e4945083c4fe0bf8419f5 Mon Sep 17 00:00:00 2001 From: Chandan Bhagat Date: Sat, 4 Apr 2026 01:57:50 +0100 Subject: [PATCH 06/13] Update rs/test/tunnels_test.rs Co-authored-by: Copilot <175728472+Copilot@users.noreply.github.com> --- rs/test/tunnels_test.rs | 14 ++++++++------ 1 file changed, 8 insertions(+), 6 deletions(-) diff --git a/rs/test/tunnels_test.rs b/rs/test/tunnels_test.rs index 2bf69b6d..f7e7987e 100644 --- a/rs/test/tunnels_test.rs +++ b/rs/test/tunnels_test.rs @@ -55,18 +55,20 @@ fn test_reconnect_options_defaults() { assert!(opts.token_refresher.is_none(), "token_refresher should be None by default"); } -// @group UnitTests > ReconnectOptions : max_attempts=0 triggers immediately +// @group UnitTests > ReconnectOptions : max_attempts=0 is preserved in configuration #[cfg(feature = "connections")] #[test] -fn test_max_attempts_zero_triggers_on_first_attempt() { +fn test_max_attempts_zero_is_preserved_in_options() { let opts = ReconnectOptions { max_attempts: Some(0), ..Default::default() }; - let attempt: u32 = 1; - if let Some(max) = opts.max_attempts { - assert!(attempt > max, "attempt 1 should exceed max_attempts=0"); - } + + assert_eq!(opts.max_attempts, Some(0)); + assert_eq!(opts.initial_delay_ms, 1_000); + assert_eq!(opts.max_delay_ms, 13_000); + assert!(opts.keep_alive_interval.is_none(), "keep_alive_interval should remain None by default"); + assert!(opts.token_refresher.is_none(), "token_refresher should remain None by default"); } // @group UnitTests > KeepAliveState : All variants construct and compare correctly From 172cbbc29ce27a575ccd57529217fe0573450487 Mon Sep 17 00:00:00 2001 From: Chandan Bhagat Date: Sat, 4 Apr 2026 01:58:03 +0100 Subject: [PATCH 07/13] Update rs/test/tunnels_test.rs Co-authored-by: Copilot <175728472+Copilot@users.noreply.github.com> --- rs/test/tunnels_test.rs | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/rs/test/tunnels_test.rs b/rs/test/tunnels_test.rs index f7e7987e..e3b1c220 100644 --- a/rs/test/tunnels_test.rs +++ b/rs/test/tunnels_test.rs @@ -1,4 +1,4 @@ -// Copyright (c) Microsoft Corporation. +// Copyright (c) Microsoft Corporation. // Licensed under the MIT license. // @group TestSetup : Feature-gated imports for connection tests From 658829d50cffe6af9857f7c7e1e4719fb0b7d4fc Mon Sep 17 00:00:00 2001 From: Chandan Bhagat Date: Sat, 4 Apr 2026 01:58:18 +0100 Subject: [PATCH 08/13] Update rs/src/connections/relay_tunnel_host.rs Co-authored-by: Copilot <175728472+Copilot@users.noreply.github.com> --- rs/src/connections/relay_tunnel_host.rs | 6 +++++- 1 file changed, 5 insertions(+), 1 deletion(-) diff --git a/rs/src/connections/relay_tunnel_host.rs b/rs/src/connections/relay_tunnel_host.rs index 820848da..3e4c5521 100644 --- a/rs/src/connections/relay_tunnel_host.rs +++ b/rs/src/connections/relay_tunnel_host.rs @@ -329,7 +329,11 @@ impl RelayTunnelHost { Err(_) => log::warn!("relay task panicked"), } } - _ = stop_rx.recv() => { break 'reconnect; } + _ = stop_rx.recv() => { + current_join.abort(); + let _ = current_join.await; + break 'reconnect; + } } // Reconnect inner loop: retry with exponential back-off. From 8cc97fe723f148d7fdf1a7946d020e4ea3960b91 Mon Sep 17 00:00:00 2001 From: Chandan Bhagat Date: Sat, 4 Apr 2026 01:58:43 +0100 Subject: [PATCH 09/13] Update rs/src/connections/relay_tunnel_host.rs Co-authored-by: Copilot <175728472+Copilot@users.noreply.github.com> --- rs/src/connections/relay_tunnel_host.rs | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/rs/src/connections/relay_tunnel_host.rs b/rs/src/connections/relay_tunnel_host.rs index 3e4c5521..67c56a65 100644 --- a/rs/src/connections/relay_tunnel_host.rs +++ b/rs/src/connections/relay_tunnel_host.rs @@ -679,7 +679,7 @@ async fn relay_connect_once( create_relay_websocket(mgmt, locator, host_id, proxy, host_token).await?; let cnx = AsyncRWWebSocket::new(super::ws::AsyncRWWebSocketOptions { websocket: cnx, - ping_interval: keep_alive.as_ref().map(|(d, _)| *d).unwrap_or(Duration::from_secs(60)), + ping_interval: Duration::from_secs(60), ping_timeout: Duration::from_secs(10), }); From 442d8b70af9d23f75b26bc676c2cf4d623973ce7 Mon Sep 17 00:00:00 2001 From: Chandan Bhagat Date: Sat, 4 Apr 2026 01:59:09 +0100 Subject: [PATCH 10/13] Update rs/src/connections/errors.rs Co-authored-by: Copilot <175728472+Copilot@users.noreply.github.com> --- rs/src/connections/errors.rs | 1 - 1 file changed, 1 deletion(-) diff --git a/rs/src/connections/errors.rs b/rs/src/connections/errors.rs index d9ee02da..1a7cd62e 100644 --- a/rs/src/connections/errors.rs +++ b/rs/src/connections/errors.rs @@ -32,7 +32,6 @@ pub enum TunnelError { #[error("tunnel access token refresh failed")] TokenRefreshFailed, - #[error("proxy connection failed: {0}")] ProxyConnectionFailed(std::io::Error), From 45db5c6dc3b168cfca753c978f3c7e21b1403009 Mon Sep 17 00:00:00 2001 From: Chandan Bhagat Date: Sat, 4 Apr 2026 01:59:38 +0100 Subject: [PATCH 11/13] Update README.md Co-authored-by: Copilot <175728472+Copilot@users.noreply.github.com> --- README.md | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/README.md b/README.md index 0c7f73a4..284ec0a9 100644 --- a/README.md +++ b/README.md @@ -15,7 +15,7 @@ Dev tunnels allows developers to securely expose local web services to the Inter | Reconnection | ✅ | ✅ | ❌ | ❌ | ✅ | | SSH-level Reconnection | ✅ | ✅ | ❌ | ❌ | ✅ | | Automatic tunnel access token refresh | ✅ | ✅ | ❌ | ❌ | ✅ | -| Ssh Keep-alive | ✅ | ✅ | ❌ | ❌ | ✅ | +| SSH keep-alive | ✅ | ✅ | ❌ | ❌ | ✅ | ✅ - Supported 🚧 - In Progress From 83518cdfc65a4d932d2f5ff564e54400118b3210 Mon Sep 17 00:00:00 2001 From: Chandan Bhagat Date: Mon, 6 Jul 2026 19:20:48 +0100 Subject: [PATCH 12/13] Potential fix for pull request finding Co-authored-by: Copilot Autofix powered by AI <175728472+Copilot@users.noreply.github.com> --- rs/src/connections/relay_tunnel_host.rs | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/rs/src/connections/relay_tunnel_host.rs b/rs/src/connections/relay_tunnel_host.rs index aa20a4af..71126ad4 100644 --- a/rs/src/connections/relay_tunnel_host.rs +++ b/rs/src/connections/relay_tunnel_host.rs @@ -19,7 +19,7 @@ use crate::{ }, }; use async_trait::async_trait; -use futures::{future::BoxFuture, stream::FuturesUnordered, StreamExt, TryFutureExt}; +use futures::{future::BoxFuture, stream::FuturesUnordered, StreamExt}; use russh::{server::Server as ServerTrait, CryptoVec}; use russh_keys::PublicKeyBase64; use tokio::{ From 49ab16d6e5b81e92dfc4f14d42577fe2a547d2d5 Mon Sep 17 00:00:00 2001 From: Chandan Bhagat Date: Sat, 3 Oct 2026 20:53:37 +0100 Subject: [PATCH 13/13] fix(rust): address reconnection review findings --- README.md | 1 - rs/src/connections/errors.rs | 1 + rs/src/connections/relay_tunnel_host.rs | 416 +++++++++++++++--------- rs/test/tunnels_test.rs | 112 ------- rs/tests/reconnection.rs | 43 +++ 5 files changed, 308 insertions(+), 265 deletions(-) delete mode 100644 rs/test/tunnels_test.rs create mode 100644 rs/tests/reconnection.rs diff --git a/README.md b/README.md index 284ec0a9..303e4138 100644 --- a/README.md +++ b/README.md @@ -15,7 +15,6 @@ Dev tunnels allows developers to securely expose local web services to the Inter | Reconnection | ✅ | ✅ | ❌ | ❌ | ✅ | | SSH-level Reconnection | ✅ | ✅ | ❌ | ❌ | ✅ | | Automatic tunnel access token refresh | ✅ | ✅ | ❌ | ❌ | ✅ | -| SSH keep-alive | ✅ | ✅ | ❌ | ❌ | ✅ | ✅ - Supported 🚧 - In Progress diff --git a/rs/src/connections/errors.rs b/rs/src/connections/errors.rs index 87e05e4c..ab62eeb8 100644 --- a/rs/src/connections/errors.rs +++ b/rs/src/connections/errors.rs @@ -29,6 +29,7 @@ pub enum TunnelError { #[error("max reconnect attempts ({0}) exceeded")] MaxReconnectAttemptsExceeded(u32), + #[error("tunnel access token refresh failed")] TokenRefreshFailed, diff --git a/rs/src/connections/relay_tunnel_host.rs b/rs/src/connections/relay_tunnel_host.rs index 71126ad4..56e6b674 100644 --- a/rs/src/connections/relay_tunnel_host.rs +++ b/rs/src/connections/relay_tunnel_host.rs @@ -19,7 +19,7 @@ use crate::{ }, }; use async_trait::async_trait; -use futures::{future::BoxFuture, stream::FuturesUnordered, StreamExt}; +use futures_util::{future::BoxFuture, stream::FuturesUnordered, StreamExt}; use russh::{server::Server as ServerTrait, CryptoVec}; use russh_keys::PublicKeyBase64; use tokio::{ @@ -40,7 +40,14 @@ use super::{ /// sent. Shared by the host relay to each connected session. type PortMap = HashMap>; -// @group Reconnection : Types for automatic reconnection with exponential backoff +// WebSocket writes can arrive here as a single large frame. If we hand that +// whole buffer to russh and report it fully written, russh may queue data until +// the SSH channel window is exhausted before the caller gets another chance to +// poll and drain the other side. 32 KiB matches the SSH channel max packet size +// used by the relay, so each accepted write maps to at most one SSH packet. +// Reporting bounded chunks preserves normal AsyncWrite backpressure and lets +// large responses continue making progress. +const CHANNEL_WRITE_CHUNK_SIZE: usize = 32 * 1024; /// The connection state of a persistent relay host. #[derive(Debug, Clone, PartialEq, Eq)] @@ -58,19 +65,23 @@ pub enum RelayConnectionState { Disconnected, } -/// Observable state of the SSH keep-alive probing for a [`PersistentRelayHandle`]. +/// Observable state of the local SSH-session health check for a +/// [`PersistentRelayHandle`]. +/// +/// This check reports whether russh has already observed the session closing. +/// It does not send network traffic and therefore is not a keep-alive probe. #[derive(Debug, Clone, PartialEq, Eq)] -pub enum KeepAliveState { - /// Keep-alive is not configured (default). +pub enum SessionHealthState { + /// Session-health monitoring is not configured (default). NotConfigured, - /// The most recent keep-alive probe succeeded. - Succeeded { - /// Number of successful probes so far. + /// The session was open at the most recent check. + Open { + /// Number of checks that found the session open. count: u32, }, - /// The most recent keep-alive probe failed or timed out. - Failed { - /// Number of failed probes so far. + /// The session was closed at the most recent check. + Closed { + /// Number of checks performed, including the final closed check. count: u32, }, } @@ -84,9 +95,11 @@ pub struct ReconnectOptions { pub initial_delay_ms: u64, /// Upper bound on retry delay, in milliseconds. Default: 13 000 ms. pub max_delay_ms: u64, - /// Interval between SSH keep-alive probes. `None` (default) disables keep-alive - /// (WebSocket-level pings still run regardless). - pub keep_alive_interval: Option, + /// Interval between local checks of the russh session state. `None` + /// (default) disables session-health monitoring. This does not send an SSH + /// keep-alive or detect a half-open connection; WebSocket pings continue to + /// provide active liveness detection independently. + pub session_health_check_interval: Option, /// Async callback invoked when the access token is rejected (HTTP 401). /// Should return a fresh token, or `None` if a new token cannot be obtained. /// When `None` (default), unauthorized errors follow normal back-off. @@ -99,12 +112,52 @@ impl Default for ReconnectOptions { max_attempts: None, initial_delay_ms: 1_000, max_delay_ms: 13_000, - keep_alive_interval: None, + session_health_check_interval: None, token_refresher: None, } } } +#[derive(Debug)] +struct ReconnectBackoff { + initial_delay_ms: u64, + max_delay_ms: u64, + next_delay_ms: u64, +} + +impl ReconnectBackoff { + fn new(initial_delay_ms: u64, max_delay_ms: u64) -> Self { + let initial_delay_ms = initial_delay_ms.min(max_delay_ms); + Self { + initial_delay_ms, + max_delay_ms, + next_delay_ms: initial_delay_ms, + } + } + + fn delay_before_attempt(&self, skip_delay: bool) -> u64 { + if skip_delay { + 0 + } else { + self.next_delay_ms + } + } + + fn record_failure(&mut self, waited: bool) { + if waited { + self.next_delay_ms = self.next_delay_ms.saturating_mul(2).min(self.max_delay_ms); + } + } + + fn reset(&mut self) { + self.next_delay_ms = self.initial_delay_ms; + } +} + +fn reconnect_limit_exceeded(max_attempts: Option, attempt: u32) -> bool { + max_attempts.is_some_and(|max| attempt > max) +} + /// Handle returned by [`RelayTunnelHost::connect_persistent`]. /// /// Drop this value (or call [`PersistentRelayHandle::stop`]) to request a @@ -112,8 +165,8 @@ impl Default for ReconnectOptions { pub struct PersistentRelayHandle { /// Observe connection-state changes as they happen. pub state: watch::Receiver, - /// Observe keep-alive probe state changes as they happen. - pub keep_alive: watch::Receiver, + /// Observe local SSH-session health changes as they happen. + pub session_health: watch::Receiver, /// Dropping this sender signals the reconnect loop to exit. _stop_tx: mpsc::Sender<()>, join: JoinHandle>, @@ -259,48 +312,50 @@ impl RelayTunnelHost { /// reconnect if this happens, and they can reconnect using the same /// RelayTunnelHost. pub async fn connect(&mut self, host_token: &str) -> Result { - relay_connect_once( - &self.mgmt, - &self.locator, - self.host_id, - &self.proxy, - self.host_keypair.clone(), - self.ports_rx.clone(), + relay_connect_once(RelayConnectArgs { + mgmt: &self.mgmt, + locator: &self.locator, + host_id: self.host_id, + proxy: &self.proxy, + host_keypair: self.host_keypair.clone(), + ports_rx: self.ports_rx.clone(), host_token, - None, // keep-alive not configured for single connect - ) + // Session-health monitoring is only exposed by persistent connections. + session_health: None, + }) .await } /// Connects to the relay and automatically reconnects on disconnection. /// - /// Unlike [`connect`], this method retries indefinitely (or up to - /// `options.max_attempts` times) with exponential back-off. + /// Unlike [`RelayTunnelHost::connect`], this method retries indefinitely + /// (or up to `options.max_attempts` times) with exponential back-off. /// /// The first connection attempt is made eagerly so callers surface /// configuration errors immediately. Drop the returned /// [`PersistentRelayHandle`] (or call [`PersistentRelayHandle::stop`]) to /// request a clean shutdown. - // @group Reconnection : Persistent connection with automatic exponential backoff pub async fn connect_persistent( &mut self, host_token: String, options: ReconnectOptions, ) -> Result { // Fail-fast: establish the first connection eagerly. - let (ka_tx, ka_rx) = watch::channel(KeepAliveState::NotConfigured); - let ka_tx_arc = Arc::new(ka_tx); - - let initial_handle = relay_connect_once( - &self.mgmt, - &self.locator, - self.host_id, - &self.proxy, - self.host_keypair.clone(), - self.ports_rx.clone(), - &host_token, - options.keep_alive_interval.map(|d| (d, ka_tx_arc.clone())), - ) + let (health_tx, health_rx) = watch::channel(SessionHealthState::NotConfigured); + let health_tx = Arc::new(health_tx); + + let initial_handle = relay_connect_once(RelayConnectArgs { + mgmt: &self.mgmt, + locator: &self.locator, + host_id: self.host_id, + proxy: &self.proxy, + host_keypair: self.host_keypair.clone(), + ports_rx: self.ports_rx.clone(), + host_token: &host_token, + session_health: options + .session_health_check_interval + .map(|interval| (interval, health_tx.clone())), + }) .await?; let (state_tx, state_rx) = watch::channel(RelayConnectionState::Connected); @@ -314,44 +369,39 @@ impl RelayTunnelHost { let ports_rx = self.ports_rx.clone(); let join = tokio::spawn(async move { - let mut current_join = initial_handle.join; - let mut delay_ms = options.initial_delay_ms; - // @group Reconnection > Token Refresh : Track single-attempt token refresh per session + let mut current_handle = initial_handle; + let mut backoff = ReconnectBackoff::new(options.initial_delay_ms, options.max_delay_ms); let mut current_host_token = host_token; 'reconnect: loop { - // Wait for the current connection to finish or a stop signal. - tokio::select! { - r = &mut current_join => { - match r { - Ok(Ok(())) => log::debug!("relay connection ended cleanly"), - Ok(Err(e)) => log::warn!("relay connection ended with error: {}", e), - Err(_) => log::warn!("relay task panicked"), - } - } + let connection_result = tokio::select! { + result = &mut current_handle => result, _ = stop_rx.recv() => { - current_join.abort(); - let _ = current_join.await; - break 'reconnect; + let result = current_handle.close().await; + let _ = state_tx.send(RelayConnectionState::Disconnected); + return result; } + }; + + match connection_result { + Ok(()) => log::debug!("relay connection ended cleanly"), + Err(e) => log::warn!("relay connection ended with error: {}", e), } - // Reconnect inner loop: retry with exponential back-off. let mut attempt: u32 = 0; - // @group Reconnection > SSH-level Reconnection : Skip delay after SSH protocol failures let mut skip_delay = false; - // @group Reconnection > Token Refresh : Single refresh per reconnect session let mut token_refreshed = false; + loop { - attempt += 1; - if let Some(max) = options.max_attempts { - if attempt > max { - let _ = state_tx.send(RelayConnectionState::Disconnected); - return Err(TunnelError::MaxReconnectAttemptsExceeded(max)); - } + attempt = attempt.saturating_add(1); + if reconnect_limit_exceeded(options.max_attempts, attempt) { + let max = options.max_attempts.unwrap(); + let _ = state_tx.send(RelayConnectionState::Disconnected); + return Err(TunnelError::MaxReconnectAttemptsExceeded(max)); } - let effective_delay = if skip_delay { 0 } else { delay_ms }; + let effective_delay = backoff.delay_before_attempt(skip_delay); + let waited = effective_delay > 0; skip_delay = false; let _ = state_tx.send(RelayConnectionState::Reconnecting { attempt, @@ -361,7 +411,8 @@ impl RelayTunnelHost { if effective_delay > 0 { log::info!( "waiting {}ms before reconnect attempt {}", - effective_delay, attempt + effective_delay, + attempt ); tokio::select! { _ = tokio::time::sleep(Duration::from_millis(effective_delay)) => {} @@ -369,65 +420,61 @@ impl RelayTunnelHost { } } - delay_ms = (delay_ms * 2).min(options.max_delay_ms); - - match relay_connect_once( - &mgmt, - &locator, - host_id, - &proxy, - host_keypair.clone(), - ports_rx.clone(), - ¤t_host_token, - options.keep_alive_interval.map(|d| (d, ka_tx_arc.clone())), - ) - .await - { + let connect_result = tokio::select! { + result = relay_connect_once(RelayConnectArgs { + mgmt: &mgmt, + locator: &locator, + host_id, + proxy: &proxy, + host_keypair: host_keypair.clone(), + ports_rx: ports_rx.clone(), + host_token: ¤t_host_token, + session_health: options.session_health_check_interval + .map(|interval| (interval, health_tx.clone())), + }) => result, + _ = stop_rx.recv() => { break 'reconnect; } + }; + + match connect_result { Ok(handle) => { log::info!("reconnected to relay on attempt {}", attempt); let _ = state_tx.send(RelayConnectionState::Connected); - current_join = handle.join; - delay_ms = options.initial_delay_ms; - break; // exit inner loop, wait for new connection + current_handle = handle; + backoff.reset(); + break; } - // @group Reconnection > SSH-level Reconnection : SSH error, retry once immediately Err(TunnelError::TunnelRelayDisconnected(_)) => { log::warn!( "SSH-level failure on attempt {}, retrying immediately", attempt ); - delay_ms = options.initial_delay_ms; + backoff.reset(); skip_delay = true; } - // @group Reconnection > Token Refresh : HTTP 401, call token_refresher Err(TunnelError::HttpError { error: HttpError::ResponseError(ref resp_err), .. }) if resp_err.status_code == reqwest::StatusCode::UNAUTHORIZED => { if let Some(refresher) = &options.token_refresher { if !token_refreshed { - log::info!( - "access token rejected (HTTP 401), refreshing" - ); + log::info!("access token rejected (HTTP 401), refreshing"); match refresher().await { Some(new_token) => { current_host_token = new_token; token_refreshed = true; + backoff.reset(); skip_delay = true; } None => { log::warn!("token refresher returned None"); - let _ = state_tx.send( - RelayConnectionState::Disconnected, - ); + let _ = + state_tx.send(RelayConnectionState::Disconnected); return Err(TunnelError::TokenRefreshFailed); } } } else { log::warn!("still unauthorized after token refresh"); - let _ = state_tx.send( - RelayConnectionState::Disconnected, - ); + let _ = state_tx.send(RelayConnectionState::Disconnected); return Err(TunnelError::TokenRefreshFailed); } } else { @@ -435,11 +482,12 @@ impl RelayTunnelHost { "reconnect attempt {} failed: unauthorized (no token refresher)", attempt ); + backoff.record_failure(waited); } } Err(e) => { log::warn!("reconnect attempt {} failed: {}", attempt, e); - // loop continues with next attempt + backoff.record_failure(waited); } } } @@ -451,7 +499,7 @@ impl RelayTunnelHost { Ok(PersistentRelayHandle { state: state_rx, - keep_alive: ka_rx, + session_health: health_rx, _stop_tx: stop_tx, join, }) @@ -565,10 +613,6 @@ impl RelayTunnelHost { Server { config } } - fn host_public_keys_for_endpoint(&self) -> Vec { - encode_host_public_keys(&self.host_keypair) - } - async fn make_ssh_client( rw: impl AsyncRead + AsyncWrite + Unpin + Send + 'static, ) -> Result< @@ -604,39 +648,19 @@ impl RelayTunnelHost { } } -// @group Reconnection : Free helper functions backing connect() and connect_persistent() - async fn create_relay_websocket( mgmt: &TunnelManagementClient, locator: &TunnelLocator, host_id: Uuid, proxy: &Option, + host_public_keys: Vec, host_token: &str, -) -> Result< - ( - WebSocketStream>, - TunnelRelayTunnelEndpoint, - ), - TunnelError, -> { +) -> Result<(WebSocketStream>, TunnelEndpoint), TunnelError> { + let requested_endpoint = make_relay_endpoint(host_id, host_public_keys); let endpoint = mgmt .update_tunnel_relay_endpoints( locator, - &TunnelRelayTunnelEndpoint { - base: TunnelEndpoint { - id: Some(format!("{}-relay", host_id)), - connection_mode: TunnelConnectionMode::TunnelRelay, - host_id: host_id.to_string(), - host_public_keys: vec![], - port_uri_format: None, - port_ssh_command_format: None, - ssh_gateway_public_key: None, - tunnel_ssh_command: None, - tunnel_uri: None, - }, - client_relay_uri: None, - host_relay_uri: None, - }, + &requested_endpoint, &TunnelRequestOptions { authorization: Some(Authorization::Tunnel(host_token.to_string())), ..TunnelRequestOptions::default() @@ -649,6 +673,7 @@ async fn create_relay_websocket( })?; let url = endpoint + .tunnel_relay_tunnel_endpoint .host_relay_uri .as_deref() .ok_or(TunnelError::MissingHostEndpoint)?; @@ -672,18 +697,48 @@ async fn create_relay_websocket( Ok((cnx, endpoint)) } -async fn relay_connect_once( - mgmt: &TunnelManagementClient, - locator: &TunnelLocator, +fn make_relay_endpoint(host_id: Uuid, host_public_keys: Vec) -> TunnelEndpoint { + TunnelEndpoint { + id: Some(format!("{}-relay", host_id)), + connection_mode: TunnelConnectionMode::TunnelRelay, + host_id: host_id.to_string(), + host_public_keys, + port_uri_format: None, + port_ssh_command_format: None, + ssh_gateway_public_key: None, + tunnel_ssh_command: None, + tunnel_uri: None, + local_network_tunnel_endpoint: Default::default(), + tunnel_relay_tunnel_endpoint: Default::default(), + } +} + +struct RelayConnectArgs<'a> { + mgmt: &'a TunnelManagementClient, + locator: &'a TunnelLocator, host_id: Uuid, - proxy: &Option, + proxy: &'a Option, host_keypair: russh_keys::key::KeyPair, ports_rx: watch::Receiver, - host_token: &str, - keep_alive: Option<(Duration, Arc>)>, -) -> Result { + host_token: &'a str, + session_health: Option<(Duration, Arc>)>, +} + +async fn relay_connect_once(args: RelayConnectArgs<'_>) -> Result { + let RelayConnectArgs { + mgmt, + locator, + host_id, + proxy, + host_keypair, + ports_rx, + host_token, + session_health, + } = args; + + let host_public_keys = encode_host_public_keys(&host_keypair); let (cnx, endpoint) = - create_relay_websocket(mgmt, locator, host_id, proxy, host_token).await?; + create_relay_websocket(mgmt, locator, host_id, proxy, host_public_keys, host_token).await?; let cnx = AsyncRWWebSocket::new(super::ws::AsyncRWWebSocketOptions { websocket: cnx, ping_interval: Duration::from_secs(60), @@ -696,9 +751,7 @@ async fn relay_connect_once( let client_session = Arc::new(client_session); let client_session_ret = client_session.clone(); - // @group SSH Keep-alive : Periodic liveness probe via is_closed() - if let Some((interval, ka_tx)) = keep_alive { - let ka_tx = ka_tx.clone(); + let session_health_join = session_health.map(|(interval, health_tx)| { let session_check = client_session_ret.clone(); tokio::spawn(async move { let mut count: u32 = 0; @@ -706,15 +759,14 @@ async fn relay_connect_once( tokio::time::sleep(interval).await; count = count.saturating_add(1); if session_check.is_closed() { - let _ = ka_tx.send(KeepAliveState::Failed { count }); + let _ = health_tx.send(SessionHealthState::Closed { count }); break; } else { - let _ = ka_tx.send(KeepAliveState::Succeeded { count }); + let _ = health_tx.send(SessionHealthState::Open { count }); } } - }); - } - + }) + }); log::debug!("established host relay primary session"); @@ -759,6 +811,7 @@ async fn relay_connect_once( endpoint, join, session: client_session_ret, + session_health_join, }) } @@ -1433,6 +1486,7 @@ pub struct RelayHandle { endpoint: TunnelEndpoint, session: Arc>, join: JoinHandle>, + session_health_join: Option>, } impl RelayHandle { @@ -1442,12 +1496,16 @@ impl RelayHandle { } /// Closes the tunnel and waits for all associated tasks to end. - pub async fn close(self) -> Result<(), TunnelError> { + pub async fn close(mut self) -> Result<(), TunnelError> { let result = self .session .disconnect(russh::Disconnect::ByApplication, "disconnect", "en") .await; self.join.await.ok(); + if let Some(join) = self.session_health_join.take() { + join.abort(); + let _ = join.await; + } result.map_err(TunnelError::TunnelRelayDisconnected) } } @@ -1456,11 +1514,16 @@ impl std::future::Future for RelayHandle { type Output = Result<(), TunnelError>; fn poll(mut self: Pin<&mut Self>, cx: &mut std::task::Context<'_>) -> Poll { match std::future::Future::poll(Pin::new(&mut self.join), cx) { - Poll::Ready(r) => Poll::Ready(match r { - Ok(Ok(_)) => Ok(()), - Ok(Err(e)) => Err(TunnelError::TunnelRelayDisconnected(e)), - Err(_) => Ok(()), - }), + Poll::Ready(r) => { + if let Some(join) = self.session_health_join.take() { + join.abort(); + } + Poll::Ready(match r { + Ok(Ok(_)) => Ok(()), + Ok(Err(e)) => Err(TunnelError::TunnelRelayDisconnected(e)), + Err(_) => Ok(()), + }) + } Poll::Pending => Poll::Pending, } } @@ -1481,8 +1544,8 @@ mod tests { fn encodes_rsa_public_key_in_canonical_ssh_wire_form() { // Generate an RSA keypair with a non-default hash variant — this is // what `RelayTunnelHost::new` does (SHA2_512). - let keypair = KeyPair::generate_rsa(2048, SignatureHash::SHA2_512) - .expect("generate rsa keypair"); + let keypair = + KeyPair::generate_rsa(2048, SignatureHash::SHA2_512).expect("generate rsa keypair"); let advertised = encode_host_public_keys(&keypair); assert_eq!(advertised.len(), 1, "expected exactly one host public key"); @@ -1496,5 +1559,54 @@ mod tests { .public_key_base64(); assert_eq!(advertised[0], public_key_b64); } -} + #[test] + fn relay_endpoint_advertises_host_public_keys() { + let host_id = Uuid::new_v4(); + let host_public_keys = vec!["test-host-key".to_string()]; + + let endpoint = make_relay_endpoint(host_id, host_public_keys.clone()); + + assert_eq!(endpoint.id, Some(format!("{}-relay", host_id))); + assert_eq!(endpoint.host_id, host_id.to_string()); + assert!(matches!( + endpoint.connection_mode, + TunnelConnectionMode::TunnelRelay + )); + assert_eq!(endpoint.host_public_keys, host_public_keys); + } + + #[test] + fn reconnect_backoff_doubles_after_waited_failures_and_caps() { + let mut backoff = ReconnectBackoff::new(1_000, 13_000); + + let mut delays = Vec::new(); + for _ in 0..6 { + delays.push(backoff.delay_before_attempt(false)); + backoff.record_failure(true); + } + + assert_eq!(delays, vec![1_000, 2_000, 4_000, 8_000, 13_000, 13_000]); + } + + #[test] + fn immediate_retry_does_not_inflate_next_backoff() { + let mut backoff = ReconnectBackoff::new(1_000, 13_000); + backoff.record_failure(true); + assert_eq!(backoff.delay_before_attempt(false), 2_000); + + backoff.reset(); + assert_eq!(backoff.delay_before_attempt(true), 0); + backoff.record_failure(false); + + assert_eq!(backoff.delay_before_attempt(false), 1_000); + } + + #[test] + fn reconnect_attempt_limit_uses_actual_attempt_count() { + assert!(!reconnect_limit_exceeded(None, u32::MAX)); + assert!(reconnect_limit_exceeded(Some(0), 1)); + assert!(!reconnect_limit_exceeded(Some(3), 3)); + assert!(reconnect_limit_exceeded(Some(3), 4)); + } +} diff --git a/rs/test/tunnels_test.rs b/rs/test/tunnels_test.rs deleted file mode 100644 index e3b1c220..00000000 --- a/rs/test/tunnels_test.rs +++ /dev/null @@ -1,112 +0,0 @@ -// Copyright (c) Microsoft Corporation. -// Licensed under the MIT license. - -// @group TestSetup : Feature-gated imports for connection tests -#[cfg(feature = "connections")] -use tunnels::connections::{KeepAliveState, ReconnectOptions}; - -// @group UnitTests > Pure Logic : Exponential backoff cap test (no crate types needed) -#[test] -fn test_exponential_backoff_cap() { - let initial = 1_000u64; - let max = 13_000u64; - let mut delay = initial; - let steps: Vec = { - let mut v = vec![delay]; - for _ in 0..10 { - delay = (delay * 2).min(max); - v.push(delay); - } - v - }; - assert_eq!(steps[0], 1_000); - assert_eq!(steps[1], 2_000); - assert_eq!(steps[2], 4_000); - assert_eq!(steps[3], 8_000); - assert_eq!(steps[4], 13_000); // capped - assert_eq!(steps[5], 13_000); // stays at cap -} - -// @group UnitTests > Pure Logic : SSH error sets skip_delay and resets backoff -#[test] -fn test_skip_delay_resets_delay_on_ssh_error() { - let initial_delay = 1_000u64; - let mut delay = 8_000u64; // simulate ramped-up delay - let mut skip_delay = false; - - // Simulate SSH error path sets skip_delay and resets delay - delay = initial_delay; - skip_delay = true; - - let effective_delay = if skip_delay { 0 } else { delay }; - assert_eq!(effective_delay, 0, "SSH error should skip the wait"); - assert_eq!(delay, initial_delay, "delay should reset to initial after SSH error"); -} - -// @group UnitTests > ReconnectOptions : Default field values -#[cfg(feature = "connections")] -#[test] -fn test_reconnect_options_defaults() { - let opts = ReconnectOptions::default(); - assert_eq!(opts.initial_delay_ms, 1_000); - assert_eq!(opts.max_delay_ms, 13_000); - assert!(opts.max_attempts.is_none()); - assert!(opts.keep_alive_interval.is_none(), "keep_alive_interval should be None by default"); - assert!(opts.token_refresher.is_none(), "token_refresher should be None by default"); -} - -// @group UnitTests > ReconnectOptions : max_attempts=0 is preserved in configuration -#[cfg(feature = "connections")] -#[test] -fn test_max_attempts_zero_is_preserved_in_options() { - let opts = ReconnectOptions { - max_attempts: Some(0), - ..Default::default() - }; - - assert_eq!(opts.max_attempts, Some(0)); - assert_eq!(opts.initial_delay_ms, 1_000); - assert_eq!(opts.max_delay_ms, 13_000); - assert!(opts.keep_alive_interval.is_none(), "keep_alive_interval should remain None by default"); - assert!(opts.token_refresher.is_none(), "token_refresher should remain None by default"); -} - -// @group UnitTests > KeepAliveState : All variants construct and compare correctly -#[cfg(feature = "connections")] -#[test] -fn test_keep_alive_state_variants() { - assert_eq!(KeepAliveState::NotConfigured, KeepAliveState::NotConfigured); - assert_eq!( - KeepAliveState::Succeeded { count: 42 }, - KeepAliveState::Succeeded { count: 42 } - ); - assert_eq!( - KeepAliveState::Failed { count: 7 }, - KeepAliveState::Failed { count: 7 } - ); - assert_ne!( - KeepAliveState::Succeeded { count: 1 }, - KeepAliveState::Failed { count: 1 } - ); -} - -// @group UnitTests > KeepAliveState : Clone preserves value -#[cfg(feature = "connections")] -#[test] -fn test_keep_alive_state_clone() { - let original = KeepAliveState::Succeeded { count: 3 }; - let cloned = original.clone(); - assert_eq!(original, cloned); -} - -// @group UnitTests > KeepAliveState : watch channel starts NotConfigured and updates -#[cfg(feature = "connections")] -#[tokio::test] -async fn test_keep_alive_state_watch_channel_updates() { - let (tx, rx) = tokio::sync::watch::channel(KeepAliveState::NotConfigured); - assert_eq!(*rx.borrow(), KeepAliveState::NotConfigured); - tx.send(KeepAliveState::Succeeded { count: 1 }).unwrap(); - assert_eq!(*rx.borrow(), KeepAliveState::Succeeded { count: 1 }); - tx.send(KeepAliveState::Failed { count: 2 }).unwrap(); - assert_eq!(*rx.borrow(), KeepAliveState::Failed { count: 2 }); -} \ No newline at end of file diff --git a/rs/tests/reconnection.rs b/rs/tests/reconnection.rs new file mode 100644 index 00000000..2d152cdb --- /dev/null +++ b/rs/tests/reconnection.rs @@ -0,0 +1,43 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT license. + +#![cfg(feature = "connections")] + +use tunnels::connections::{ReconnectOptions, SessionHealthState}; + +#[test] +fn reconnect_options_have_expected_defaults() { + let options = ReconnectOptions::default(); + + assert_eq!(options.initial_delay_ms, 1_000); + assert_eq!(options.max_delay_ms, 13_000); + assert_eq!(options.max_attempts, None); + assert_eq!(options.session_health_check_interval, None); + assert!(options.token_refresher.is_none()); +} + +#[test] +fn max_attempts_zero_is_preserved() { + let options = ReconnectOptions { + max_attempts: Some(0), + ..Default::default() + }; + + assert_eq!(options.max_attempts, Some(0)); +} + +#[test] +fn session_health_states_compare_by_variant_and_count() { + assert_eq!( + SessionHealthState::NotConfigured, + SessionHealthState::NotConfigured + ); + assert_eq!( + SessionHealthState::Open { count: 2 }, + SessionHealthState::Open { count: 2 } + ); + assert_ne!( + SessionHealthState::Open { count: 2 }, + SessionHealthState::Closed { count: 2 } + ); +}