diff --git a/README.md b/README.md index 65925810..303e4138 100644 --- a/README.md +++ b/README.md @@ -12,10 +12,9 @@ Dev tunnels allows developers to securely expose local web services to the Inter | Management API | ✅ | ✅ | ✅ | ✅ | ✅ | | Tunnel Client Connections | ✅ | ✅ | ✅ | ✅ | ✅ | | Tunnel Host Connections | ✅ | ✅ | ❌ | ❌ | ✅ | -| Reconnection | ✅ | ✅ | ❌ | ❌ | ❌ | -| SSH-level Reconnection | ✅ | ✅ | ❌ | ❌ | ❌ | -| Automatic tunnel access token refresh | ✅ | ✅ | ❌ | ❌ | ❌ | -| Ssh Keep-alive | ✅ | ✅ | ❌ | ❌ | ❌ | +| Reconnection | ✅ | ✅ | ❌ | ❌ | ✅ | +| SSH-level Reconnection | ✅ | ✅ | ❌ | ❌ | ✅ | +| Automatic tunnel access token refresh | ✅ | ✅ | ❌ | ❌ | ✅ | ✅ - Supported 🚧 - In Progress diff --git a/rs/src/connections/errors.rs b/rs/src/connections/errors.rs index 2b29fb01..ab62eeb8 100644 --- a/rs/src/connections/errors.rs +++ b/rs/src/connections/errors.rs @@ -27,6 +27,12 @@ pub enum TunnelError { #[error("port {0} already exists in the relay")] PortAlreadyExists(u32), + #[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 c3fddaa9..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_util::{stream::FuturesUnordered, StreamExt}; +use futures_util::{future::BoxFuture, stream::FuturesUnordered, StreamExt}; use russh::{server::Server as ServerTrait, CryptoVec}; use russh_keys::PublicKeyBase64; use tokio::{ @@ -49,6 +49,137 @@ type PortMap = HashMap>; // 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)] +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, +} + +/// 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 SessionHealthState { + /// Session-health monitoring is not configured (default). + NotConfigured, + /// The session was open at the most recent check. + Open { + /// Number of checks that found the session open. + count: u32, + }, + /// The session was closed at the most recent check. + Closed { + /// Number of checks performed, including the final closed check. + count: u32, + }, +} + +/// 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, + /// 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. + pub token_refresher: Option BoxFuture<'static, Option> + Send + Sync>>, +} + +impl Default for ReconnectOptions { + fn default() -> Self { + Self { + max_attempts: None, + initial_delay_ms: 1_000, + max_delay_ms: 13_000, + 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 +/// clean shutdown of the reconnect loop. +pub struct PersistentRelayHandle { + /// Observe connection-state changes as they happen. + pub state: 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>, +} + +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. @@ -181,65 +312,196 @@ 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(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, + // Session-health monitoring is only exposed by persistent connections. + session_health: None, + }) + .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 [`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. + pub async fn connect_persistent( + &mut self, + host_token: String, + options: ReconnectOptions, + ) -> Result { + // Fail-fast: establish the first connection eagerly. + 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?; - 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 { - 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); + 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 { + let connection_result = tokio::select! { + result = &mut current_handle => result, + _ = stop_rx.recv() => { + 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), + } + + let mut attempt: u32 = 0; + let mut skip_delay = false; + let mut token_refreshed = false; + + loop { + 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 = backoff.delay_before_attempt(skip_delay); + let waited = effective_delay > 0; + 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; } + } + } + + 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_handle = handle; + backoff.reset(); + break; + } + Err(TunnelError::TunnelRelayDisconnected(_)) => { + log::warn!( + "SSH-level failure on attempt {}, retrying immediately", + attempt + ); + backoff.reset(); + skip_delay = true; + } + 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; + backoff.reset(); + 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 + ); + backoff.record_failure(waited); } - }, - }, - else => break, + } + Err(e) => { + log::warn!("reconnect attempt {} failed: {}", attempt, e); + backoff.record_failure(waited); + } + } } } - client_session - .disconnect(russh::Disconnect::ByApplication, "going away", "en") - .await - .ok(); - - log::debug!("disconnected primary session after EOF"); - + let _ = state_tx.send(RelayConnectionState::Disconnected); Ok(()) }); - Ok(RelayHandle { - endpoint, + Ok(PersistentRelayHandle { + state: state_rx, + session_health: health_rx, + _stop_tx: stop_tx, join, - session: client_session_ret, }) } @@ -351,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< @@ -388,63 +646,173 @@ impl RelayTunnelHost { Ok((session, rx)) } +} - async fn create_websocket( - &self, - host_token: &str, - ) -> Result<(WebSocketStream>, TunnelEndpoint), TunnelError> { - let endpoint = self - .mgmt - .update_tunnel_relay_endpoints( - &self.locator, - &TunnelEndpoint { - id: Some(format!("{}-relay", self.host_id)), - connection_mode: TunnelConnectionMode::TunnelRelay, - host_id: self.host_id.to_string(), - host_public_keys: self.host_public_keys_for_endpoint(), - 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(), - }, - &TunnelRequestOptions { - authorization: Some(Authorization::Tunnel(host_token.to_string())), - ..TunnelRequestOptions::default() +async fn create_relay_websocket( + mgmt: &TunnelManagementClient, + locator: &TunnelLocator, + host_id: Uuid, + proxy: &Option, + host_public_keys: Vec, + host_token: &str, +) -> Result<(WebSocketStream>, TunnelEndpoint), TunnelError> { + let requested_endpoint = make_relay_endpoint(host_id, host_public_keys); + let endpoint = mgmt + .update_tunnel_relay_endpoints( + locator, + &requested_endpoint, + &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 + .tunnel_relay_tunnel_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)) +} + +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: &'a Option, + host_keypair: russh_keys::key::KeyPair, + ports_rx: watch::Receiver, + 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_public_keys, 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(); + + 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; + loop { + tokio::time::sleep(interval).await; + count = count.saturating_add(1); + if session_check.is_closed() { + let _ = health_tx.send(SessionHealthState::Closed { count }); + break; + } else { + let _ = health_tx.send(SessionHealthState::Open { count }); + } + } + }) + }); + + 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 - .tunnel_relay_tunnel_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, + session_health_join, + }) } /// Encodes the host's public key(s) for advertisement in a TunnelEndpoint. @@ -1118,6 +1486,7 @@ pub struct RelayHandle { endpoint: TunnelEndpoint, session: Arc>, join: JoinHandle>, + session_health_join: Option>, } impl RelayHandle { @@ -1127,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) } } @@ -1141,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, } } @@ -1166,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"); @@ -1181,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 20919425..00000000 --- a/rs/test/tunnels_test.rs +++ /dev/null @@ -1,5 +0,0 @@ -#[test] -fn it_works() { - let result = 2 + 2; - assert_eq!(result, 4); -} 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 } + ); +}