From 4e9e9788f5c974208e80d3cc6572cd77467e87cf Mon Sep 17 00:00:00 2001 From: Arth Srivastava Date: Thu, 8 Oct 2026 19:24:26 +0530 Subject: [PATCH 1/2] feat(rivetkit-client): configurable reconnection backoff Add BackoffConfig for tunable exponential backoff with jitter support in the Rust client SDK. Users can configure initial/max delays, multiplier, retry limits, and jitter factor via ClientConfig builder methods or per-connection overrides. - Replace hardcoded 1s/30s reconnect policy with configurable backoff - Add disable_reconnect() to prevent retries after connection drops - Add connect_with_backoff() for per-connection override - Enforce max_retries in the reconnect loop - Reset backoff state on successful reconnection - Add rand dependency (workspace) for jitter randomization - Add 17 tests (6 unit + 11 integration) covering progression, capping, retry limits, jitter bounds, normalization, and end-to-end reconnect behavior --- Cargo.lock | 171 ++++---- rivetkit-rust/packages/client/Cargo.toml | 1 + rivetkit-rust/packages/client/src/backoff.rs | 315 ++++++++++++++- rivetkit-rust/packages/client/src/client.rs | 42 +- .../packages/client/src/connection.rs | 19 +- rivetkit-rust/packages/client/src/handle.rs | 16 + rivetkit-rust/packages/client/src/lib.rs | 1 + .../packages/client/tests/backoff.rs | 366 ++++++++++++++++++ 8 files changed, 831 insertions(+), 100 deletions(-) create mode 100644 rivetkit-rust/packages/client/tests/backoff.rs diff --git a/Cargo.lock b/Cargo.lock index bc82c45852..8ef50ea3e0 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -1320,7 +1320,7 @@ checksum = "2a2330da5de22e8a3cb63252ce2abb30116bf5265e89c0e01bc17015ce30a476" [[package]] name = "datacenter" -version = "2.3.19" +version = "2.3.20" dependencies = [ "anyhow", "futures-util", @@ -1379,7 +1379,7 @@ dependencies = [ [[package]] name = "depot" -version = "2.3.19" +version = "2.3.20" dependencies = [ "anyhow", "async-channel", @@ -1422,7 +1422,7 @@ dependencies = [ [[package]] name = "depot-client-embedded" -version = "2.3.19" +version = "2.3.20" dependencies = [ "anyhow", "async-trait", @@ -1702,7 +1702,7 @@ checksum = "c34f04666d835ff5d62e058c3995147c06f42fe86ff053337632bca83e42702d" [[package]] name = "epoxy" -version = "2.3.19" +version = "2.3.20" dependencies = [ "anyhow", "axum 0.8.4", @@ -1748,7 +1748,7 @@ dependencies = [ [[package]] name = "epoxy-protocol" -version = "2.3.19" +version = "2.3.20" dependencies = [ "anyhow", "serde", @@ -2100,7 +2100,7 @@ dependencies = [ [[package]] name = "gasoline" -version = "2.3.19" +version = "2.3.20" dependencies = [ "anyhow", "async-stream", @@ -2151,7 +2151,7 @@ dependencies = [ [[package]] name = "gasoline-macros" -version = "2.3.19" +version = "2.3.20" dependencies = [ "proc-macro2", "quote", @@ -2160,7 +2160,7 @@ dependencies = [ [[package]] name = "gasoline-runtime" -version = "2.3.19" +version = "2.3.20" dependencies = [ "anyhow", "epoxy", @@ -3412,7 +3412,7 @@ dependencies = [ [[package]] name = "namespace" -version = "2.3.19" +version = "2.3.20" dependencies = [ "anyhow", "epoxy", @@ -3991,7 +3991,7 @@ dependencies = [ [[package]] name = "pegboard" -version = "2.3.19" +version = "2.3.20" dependencies = [ "anyhow", "base64 0.22.1", @@ -4049,7 +4049,7 @@ dependencies = [ [[package]] name = "pegboard-envoy" -version = "2.3.19" +version = "2.3.20" dependencies = [ "anyhow", "async-trait", @@ -4098,7 +4098,7 @@ dependencies = [ [[package]] name = "pegboard-gateway" -version = "2.3.19" +version = "2.3.20" dependencies = [ "anyhow", "async-trait", @@ -4132,7 +4132,7 @@ dependencies = [ [[package]] name = "pegboard-gateway2" -version = "2.3.19" +version = "2.3.20" dependencies = [ "anyhow", "async-trait", @@ -4167,7 +4167,7 @@ dependencies = [ [[package]] name = "pegboard-gateway3" -version = "2.3.19" +version = "2.3.20" dependencies = [ "anyhow", "async-trait", @@ -4203,7 +4203,7 @@ dependencies = [ [[package]] name = "pegboard-outbound" -version = "2.3.19" +version = "2.3.20" dependencies = [ "anyhow", "epoxy", @@ -4230,7 +4230,7 @@ dependencies = [ [[package]] name = "pegboard-runner" -version = "2.3.19" +version = "2.3.20" dependencies = [ "anyhow", "async-trait", @@ -5102,7 +5102,7 @@ dependencies = [ [[package]] name = "rivet-actor-runtime-socket-protocol" -version = "2.3.19" +version = "2.3.20" dependencies = [ "anyhow", "serde", @@ -5114,7 +5114,7 @@ dependencies = [ [[package]] name = "rivet-api-builder" -version = "2.3.19" +version = "2.3.20" dependencies = [ "anyhow", "axum 0.8.4", @@ -5157,7 +5157,7 @@ dependencies = [ [[package]] name = "rivet-api-peer" -version = "2.3.19" +version = "2.3.20" dependencies = [ "anyhow", "axum 0.8.4", @@ -5193,7 +5193,7 @@ dependencies = [ [[package]] name = "rivet-api-public" -version = "2.3.19" +version = "2.3.20" dependencies = [ "anyhow", "axum 0.8.4", @@ -5233,7 +5233,7 @@ dependencies = [ [[package]] name = "rivet-api-public-openapi-gen" -version = "2.3.19" +version = "2.3.20" dependencies = [ "rivet-api-public", "serde_json", @@ -5242,7 +5242,7 @@ dependencies = [ [[package]] name = "rivet-api-types" -version = "2.3.19" +version = "2.3.20" dependencies = [ "anyhow", "gasoline", @@ -5258,7 +5258,7 @@ dependencies = [ [[package]] name = "rivet-api-util" -version = "2.3.19" +version = "2.3.20" dependencies = [ "anyhow", "axum 0.8.4", @@ -5316,7 +5316,7 @@ dependencies = [ [[package]] name = "rivet-auth" -version = "2.3.19" +version = "2.3.20" dependencies = [ "anyhow", "gasoline", @@ -5332,7 +5332,7 @@ dependencies = [ [[package]] name = "rivet-auth-jwt" -version = "2.3.19" +version = "2.3.20" dependencies = [ "anyhow", "base64 0.22.1", @@ -5365,7 +5365,7 @@ dependencies = [ [[package]] name = "rivet-auth-policy" -version = "2.3.19" +version = "2.3.20" dependencies = [ "anyhow", "rivet-error", @@ -5376,7 +5376,7 @@ dependencies = [ [[package]] name = "rivet-bootstrap" -version = "2.3.19" +version = "2.3.20" dependencies = [ "datacenter", "depot", @@ -5399,7 +5399,7 @@ dependencies = [ [[package]] name = "rivet-build-meta" -version = "2.3.19" +version = "2.3.20" dependencies = [ "anyhow", "epoxy-protocol", @@ -5413,7 +5413,7 @@ dependencies = [ [[package]] name = "rivet-cache" -version = "2.3.19" +version = "2.3.20" dependencies = [ "anyhow", "futures-util", @@ -5439,7 +5439,7 @@ dependencies = [ [[package]] name = "rivet-cache-purge" -version = "2.3.19" +version = "2.3.20" dependencies = [ "anyhow", "futures-util", @@ -5456,14 +5456,14 @@ dependencies = [ [[package]] name = "rivet-cache-result" -version = "2.3.19" +version = "2.3.20" dependencies = [ "rivet-util", ] [[package]] name = "rivet-cli" -version = "2.3.19" +version = "2.3.20" dependencies = [ "anstyle", "anyhow", @@ -5484,7 +5484,7 @@ dependencies = [ [[package]] name = "rivet-config" -version = "2.3.19" +version = "2.3.20" dependencies = [ "anyhow", "chrono", @@ -5507,7 +5507,7 @@ dependencies = [ [[package]] name = "rivet-config-schema-gen" -version = "2.3.19" +version = "2.3.20" dependencies = [ "rivet-config", "schemars 0.8.22", @@ -5516,7 +5516,7 @@ dependencies = [ [[package]] name = "rivet-data" -version = "2.3.19" +version = "2.3.20" dependencies = [ "anyhow", "gasoline", @@ -5530,7 +5530,7 @@ dependencies = [ [[package]] name = "rivet-depot-client" -version = "2.3.19" +version = "2.3.20" dependencies = [ "anyhow", "async-trait", @@ -5559,11 +5559,11 @@ dependencies = [ [[package]] name = "rivet-depot-client-types" -version = "2.3.19" +version = "2.3.20" [[package]] name = "rivet-depot-protocol" -version = "2.3.19" +version = "2.3.20" dependencies = [ "anyhow", "serde", @@ -5574,7 +5574,7 @@ dependencies = [ [[package]] name = "rivet-dynamic-config" -version = "2.3.19" +version = "2.3.20" dependencies = [ "anyhow", "futures-util", @@ -5592,7 +5592,7 @@ dependencies = [ [[package]] name = "rivet-engine" -version = "2.3.19" +version = "2.3.20" dependencies = [ "anyhow", "async-trait", @@ -5675,7 +5675,7 @@ dependencies = [ [[package]] name = "rivet-env" -version = "2.3.19" +version = "2.3.20" dependencies = [ "lazy_static", "uuid", @@ -5683,7 +5683,7 @@ dependencies = [ [[package]] name = "rivet-envoy-client" -version = "2.3.19" +version = "2.3.20" dependencies = [ "anyhow", "futures-util", @@ -5715,7 +5715,7 @@ dependencies = [ [[package]] name = "rivet-envoy-protocol" -version = "2.3.19" +version = "2.3.20" dependencies = [ "anyhow", "hex", @@ -5731,7 +5731,7 @@ dependencies = [ [[package]] name = "rivet-error" -version = "2.3.19" +version = "2.3.20" dependencies = [ "anyhow", "indoc", @@ -5743,7 +5743,7 @@ dependencies = [ [[package]] name = "rivet-error-macros" -version = "2.3.19" +version = "2.3.20" dependencies = [ "indoc", "proc-macro2", @@ -5754,7 +5754,7 @@ dependencies = [ [[package]] name = "rivet-guard" -version = "2.3.19" +version = "2.3.20" dependencies = [ "anyhow", "async-trait", @@ -5819,7 +5819,7 @@ dependencies = [ [[package]] name = "rivet-guard-core" -version = "2.3.19" +version = "2.3.20" dependencies = [ "anyhow", "async-trait", @@ -5870,7 +5870,7 @@ dependencies = [ [[package]] name = "rivet-logs" -version = "2.3.19" +version = "2.3.20" dependencies = [ "anyhow", "chrono", @@ -5884,7 +5884,7 @@ dependencies = [ [[package]] name = "rivet-metrics" -version = "2.3.19" +version = "2.3.20" dependencies = [ "lazy_static", "prometheus", @@ -5892,7 +5892,7 @@ dependencies = [ [[package]] name = "rivet-metrics-server" -version = "2.3.19" +version = "2.3.20" dependencies = [ "anyhow", "console-subscriber", @@ -5910,7 +5910,7 @@ dependencies = [ [[package]] name = "rivet-outbound-guard" -version = "2.3.19" +version = "2.3.20" dependencies = [ "anyhow", "ipnet", @@ -5924,7 +5924,7 @@ dependencies = [ [[package]] name = "rivet-perf" -version = "2.3.19" +version = "2.3.20" dependencies = [ "prometheus", "tokio", @@ -5934,7 +5934,7 @@ dependencies = [ [[package]] name = "rivet-pools" -version = "2.3.19" +version = "2.3.20" dependencies = [ "anyhow", "clickhouse", @@ -5967,7 +5967,7 @@ dependencies = [ [[package]] name = "rivet-postgres-util" -version = "2.3.19" +version = "2.3.20" dependencies = [ "anyhow", "rustls", @@ -5978,7 +5978,7 @@ dependencies = [ [[package]] name = "rivet-profiling" -version = "2.3.19" +version = "2.3.20" dependencies = [ "anyhow", "gasoline", @@ -5993,7 +5993,7 @@ dependencies = [ [[package]] name = "rivet-runner-protocol" -version = "2.3.19" +version = "2.3.20" dependencies = [ "anyhow", "gasoline", @@ -6011,7 +6011,7 @@ dependencies = [ [[package]] name = "rivet-runtime" -version = "2.3.19" +version = "2.3.20" dependencies = [ "anyhow", "console-subscriber", @@ -6040,7 +6040,7 @@ dependencies = [ [[package]] name = "rivet-service-manager" -version = "2.3.19" +version = "2.3.20" dependencies = [ "anyhow", "chrono", @@ -6057,7 +6057,7 @@ dependencies = [ [[package]] name = "rivet-telemetry" -version = "2.3.19" +version = "2.3.20" dependencies = [ "anyhow", "rivet-config", @@ -6068,7 +6068,7 @@ dependencies = [ [[package]] name = "rivet-term" -version = "2.3.19" +version = "2.3.20" dependencies = [ "console", "derive_builder 0.12.0", @@ -6081,7 +6081,7 @@ dependencies = [ [[package]] name = "rivet-test-deps" -version = "2.3.19" +version = "2.3.20" dependencies = [ "anyhow", "futures-util", @@ -6099,7 +6099,7 @@ dependencies = [ [[package]] name = "rivet-test-deps-docker" -version = "2.3.19" +version = "2.3.20" dependencies = [ "anyhow", "portpicker", @@ -6116,7 +6116,7 @@ dependencies = [ [[package]] name = "rivet-test-envoy" -version = "2.3.19" +version = "2.3.20" dependencies = [ "anyhow", "async-stream", @@ -6132,7 +6132,7 @@ dependencies = [ [[package]] name = "rivet-tracing-utils" -version = "2.3.19" +version = "2.3.20" dependencies = [ "futures-util", "lazy_static", @@ -6142,7 +6142,7 @@ dependencies = [ [[package]] name = "rivet-types" -version = "2.3.19" +version = "2.3.20" dependencies = [ "anyhow", "gasoline", @@ -6159,7 +6159,7 @@ dependencies = [ [[package]] name = "rivet-universaldb-commit" -version = "2.3.19" +version = "2.3.20" dependencies = [ "anyhow", "serde", @@ -6171,7 +6171,7 @@ dependencies = [ [[package]] name = "rivet-ups-protocol" -version = "2.3.19" +version = "2.3.20" dependencies = [ "anyhow", "base64 0.22.1", @@ -6184,7 +6184,7 @@ dependencies = [ [[package]] name = "rivet-util" -version = "2.3.19" +version = "2.3.20" dependencies = [ "anyhow", "async-trait", @@ -6212,7 +6212,7 @@ dependencies = [ [[package]] name = "rivet-util-id" -version = "2.3.19" +version = "2.3.20" dependencies = [ "serde", "thiserror 1.0.69", @@ -6223,7 +6223,7 @@ dependencies = [ [[package]] name = "rivet-util-serde" -version = "2.3.19" +version = "2.3.20" dependencies = [ "indexmap 2.14.0", "serde", @@ -6233,7 +6233,7 @@ dependencies = [ [[package]] name = "rivet-version-management" -version = "2.3.19" +version = "2.3.20" dependencies = [ "anyhow", "futures-util", @@ -6253,7 +6253,7 @@ dependencies = [ [[package]] name = "rivet-workflow-worker" -version = "2.3.19" +version = "2.3.20" dependencies = [ "anyhow", "datacenter", @@ -6270,7 +6270,7 @@ dependencies = [ [[package]] name = "rivetkit" -version = "2.3.19" +version = "2.3.20" dependencies = [ "anyhow", "async-trait", @@ -6299,7 +6299,7 @@ dependencies = [ [[package]] name = "rivetkit-actor-persist" -version = "2.3.19" +version = "2.3.20" dependencies = [ "anyhow", "serde", @@ -6311,7 +6311,7 @@ dependencies = [ [[package]] name = "rivetkit-client" -version = "2.3.19" +version = "2.3.20" dependencies = [ "anyhow", "axum 0.8.4", @@ -6324,6 +6324,7 @@ dependencies = [ "opentelemetry_sdk", "parking_lot", "portpicker", + "rand 0.8.5", "reqwest 0.12.22", "rivetkit-client-protocol", "scc", @@ -6345,7 +6346,7 @@ dependencies = [ [[package]] name = "rivetkit-client-protocol" -version = "2.3.19" +version = "2.3.20" dependencies = [ "anyhow", "serde", @@ -6357,7 +6358,7 @@ dependencies = [ [[package]] name = "rivetkit-core" -version = "2.3.19" +version = "2.3.20" dependencies = [ "anyhow", "async-trait", @@ -6420,7 +6421,7 @@ dependencies = [ [[package]] name = "rivetkit-engine-process" -version = "2.3.19" +version = "2.3.20" dependencies = [ "anyhow", "libc", @@ -6437,7 +6438,7 @@ dependencies = [ [[package]] name = "rivetkit-inspector-protocol" -version = "2.3.19" +version = "2.3.20" dependencies = [ "anyhow", "serde", @@ -6449,7 +6450,7 @@ dependencies = [ [[package]] name = "rivetkit-napi" -version = "2.3.19" +version = "2.3.20" dependencies = [ "anyhow", "async-trait", @@ -6475,7 +6476,7 @@ dependencies = [ [[package]] name = "rivetkit-shared-types" -version = "2.3.19" +version = "2.3.20" dependencies = [ "serde", "serde_json", @@ -6483,7 +6484,7 @@ dependencies = [ [[package]] name = "rivetkit-wasm" -version = "2.3.19" +version = "2.3.20" dependencies = [ "anyhow", "console_error_panic_hook", @@ -7667,7 +7668,7 @@ dependencies = [ [[package]] name = "test-snapshot-gen" -version = "2.3.19" +version = "2.3.20" dependencies = [ "anyhow", "async-trait", @@ -8450,7 +8451,7 @@ checksum = "ebc1c04c71510c7f702b52b7c350734c9ff1295c464a03335b00bb84fc54f853" [[package]] name = "universaldb" -version = "2.3.19" +version = "2.3.20" dependencies = [ "anyhow", "async-trait", @@ -8488,7 +8489,7 @@ dependencies = [ [[package]] name = "universalpubsub" -version = "2.3.19" +version = "2.3.20" dependencies = [ "anyhow", "async-trait", diff --git a/rivetkit-rust/packages/client/Cargo.toml b/rivetkit-rust/packages/client/Cargo.toml index d655cffd17..fd4e97f7c1 100644 --- a/rivetkit-rust/packages/client/Cargo.toml +++ b/rivetkit-rust/packages/client/Cargo.toml @@ -17,6 +17,7 @@ opentelemetry = { version = "0.28", default-features = false, features = ["trace opentelemetry-http.workspace = true opentelemetry_sdk = { version = "0.28", default-features = false, features = ["trace"] } parking_lot.workspace = true +rand = { workspace = true } reqwest = { version = "0.12.12", default-features = false, features = ["json", "charset", "http2", "macos-system-configuration", "rustls-tls-native-roots", "rustls-tls-webpki-roots"] } rivetkit-client-protocol.workspace = true scc.workspace = true diff --git a/rivetkit-rust/packages/client/src/backoff.rs b/rivetkit-rust/packages/client/src/backoff.rs index 14117660fc..84bd043948 100644 --- a/rivetkit-rust/packages/client/src/backoff.rs +++ b/rivetkit-rust/packages/client/src/backoff.rs @@ -1,24 +1,325 @@ -use std::{cmp, time::Duration}; +use std::time::Duration; +/// Configuration for reconnection exponential backoff and retry policies. +#[derive(Debug, Clone, PartialEq)] +pub struct BackoffConfig { + /// Initial delay before the first reconnect attempt. Default: 1 second. + pub initial_delay: Duration, + /// Maximum delay between reconnect attempts. Default: 30 seconds. + pub max_delay: Duration, + /// Multiplier applied to the delay after each attempt. Default: 2.0. + pub multiplier: f64, + /// Maximum number of retries after the initial connection attempt. + /// For example, `Some(3)` means: 1 initial attempt + up to 3 retries = 4 total attempts. + /// `None` indicates infinite retries (default). + pub max_retries: Option, + /// Jitter factor between 0.0 and 1.0 to randomize delays and prevent thundering herds. + /// Default: 0.0 (disabled). + pub jitter_factor: f64, +} + +impl Default for BackoffConfig { + fn default() -> Self { + Self { + initial_delay: Duration::from_secs(1), + max_delay: Duration::from_secs(30), + multiplier: 2.0, + max_retries: None, + jitter_factor: 0.0, + } + } +} + +impl BackoffConfig { + /// Normalizes the configuration to ensure invariants: + /// - Ensures `max_delay >= initial_delay` + /// - Clamps `multiplier` to at least `1.0` (or `1.0` if `NaN`) + /// - Clamps `jitter_factor` to `0.0..=1.0` (or `0.0` if `NaN`) + pub fn normalize(&mut self) { + if self.max_delay < self.initial_delay { + self.max_delay = self.initial_delay; + } + if self.multiplier.is_nan() || self.multiplier < 1.0 { + self.multiplier = 1.0; + } + if self.jitter_factor.is_nan() { + self.jitter_factor = 0.0; + } else { + self.jitter_factor = self.jitter_factor.clamp(0.0, 1.0); + } + } + + /// Creates a new `BackoffConfig` with the given initial and max delays. + /// + /// If `max_delay < initial_delay`, `max_delay` is raised to `initial_delay`. + pub fn new(initial_delay: Duration, max_delay: Duration) -> Self { + let max_delay = max_delay.max(initial_delay); + Self { + initial_delay, + max_delay, + ..Default::default() + } + } + + /// Sets the initial delay before the first retry attempt. + /// + /// If `max_delay < delay`, `max_delay` is automatically raised to match `delay`. + pub fn initial_delay(mut self, delay: Duration) -> Self { + self.initial_delay = delay; + if self.max_delay < delay { + self.max_delay = delay; + } + self + } + + /// Sets the maximum ceiling delay between retry attempts. + /// + /// Clamped to be at least `initial_delay`. + pub fn max_delay(mut self, delay: Duration) -> Self { + self.max_delay = delay.max(self.initial_delay); + self + } + + /// Sets the exponential multiplier. + /// + /// Values less than `1.0` are silently clamped to `1.0` (ensuring delays never shrink across retries). + pub fn multiplier(mut self, multiplier: f64) -> Self { + self.multiplier = if multiplier.is_nan() { + 1.0 + } else { + multiplier.max(1.0) + }; + self + } + + /// Sets the maximum number of retries after the initial connection attempt. + /// For example, `Some(3)` allows 1 initial attempt + 3 retries = 4 total. + /// Pass `None` for indefinite retries. + pub fn max_retries(mut self, max_retries: Option) -> Self { + self.max_retries = max_retries; + self + } + + /// Enables or disables standard jitter (uses 20% jitter if enabled). + pub fn jitter(mut self, enabled: bool) -> Self { + self.jitter_factor = if enabled { 0.2 } else { 0.0 }; + self + } + + /// Sets a custom jitter factor. + /// + /// Values outside `[0.0, 1.0]` are silently clamped into `0.0..=1.0` (e.g. values `< 0.0` + /// become `0.0`, and values `> 1.0` become `1.0`). + pub fn jitter_factor(mut self, factor: f64) -> Self { + self.jitter_factor = if factor.is_nan() { + 0.0 + } else { + factor.clamp(0.0, 1.0) + }; + self + } +} + +/// Exponential backoff calculator with optional jitter and attempt limits. +#[derive(Debug, Clone)] pub struct Backoff { - max_delay: Duration, + config: BackoffConfig, delay: Duration, + attempt: usize, } impl Backoff { + /// Creates a backoff with the given initial and max delay using default settings. pub fn new(initial: Duration, max_delay: Duration) -> Self { + Self::from_config(BackoffConfig::new(initial, max_delay)) + } + + /// Creates a backoff from a `BackoffConfig`, normalizing any invalid fields. + pub fn from_config(mut config: BackoffConfig) -> Self { + config.normalize(); + let delay = config.initial_delay; Self { - max_delay, - delay: initial, + config, + delay, + attempt: 0, } } + /// Returns a reference to the active `BackoffConfig`. + pub fn config(&self) -> &BackoffConfig { + &self.config + } + + /// Returns the number of retry attempts made so far. + pub fn attempt(&self) -> usize { + self.attempt + } + + /// Returns whether additional retries are permitted under the configured `max_retries`. + pub fn can_retry(&self) -> bool { + match self.config.max_retries { + Some(max) => self.attempt < max, + None => true, + } + } + + /// Returns the current base delay for this attempt before stepping. pub fn delay(&self) -> Duration { self.delay } - pub async fn tick(&mut self) { - tokio::time::sleep(self.delay).await; - self.delay = cmp::min(self.delay * 2, self.max_delay); + /// Advances the backoff state by one attempt and returns the computed sleep duration. + /// Returns `None` if `can_retry()` is false. + /// + /// **Note:** The internal attempt counter and delay are advanced immediately, + /// before the caller performs the actual sleep. If using `step()` directly + /// (rather than `tick()`), be aware that cancellation after `step()` but + /// before sleeping will still have advanced the backoff state. + pub fn step(&mut self) -> Option { + if !self.can_retry() { + return None; + } + + let base = self.delay; + self.attempt += 1; + + let next_secs = + (base.as_secs_f64() * self.config.multiplier).min(self.config.max_delay.as_secs_f64()); + self.delay = Duration::from_secs_f64(next_secs); + + let sleep_duration = if self.config.jitter_factor > 0.0 { + let jitter_offset = (rand::random::() * 2.0 - 1.0) * self.config.jitter_factor; + let factor = (1.0 + jitter_offset).max(0.0); + Duration::from_secs_f64((base.as_secs_f64() * factor).max(0.0)) + } else { + base + }; + + Some(sleep_duration) + } + + /// Waits for the backoff delay. Returns `true` if waited, or `false` if max retries exceeded. + pub async fn tick(&mut self) -> bool { + let Some(duration) = self.step() else { + return false; + }; + tokio::time::sleep(duration).await; + true + } + + /// Resets the backoff state back to the initial delay and attempt 0. + pub fn reset(&mut self) { + self.delay = self.config.initial_delay; + self.attempt = 0; + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn test_default_progression() { + let mut backoff = Backoff::new(Duration::from_secs(1), Duration::from_secs(4)); + assert_eq!(backoff.delay(), Duration::from_secs(1)); + assert_eq!(backoff.attempt(), 0); + assert!(backoff.can_retry()); + + let step1 = backoff.step().expect("step 1"); + assert_eq!(step1, Duration::from_secs(1)); + assert_eq!(backoff.delay(), Duration::from_secs(2)); + assert_eq!(backoff.attempt(), 1); + + let step2 = backoff.step().expect("step 2"); + assert_eq!(step2, Duration::from_secs(2)); + assert_eq!(backoff.delay(), Duration::from_secs(4)); + assert_eq!(backoff.attempt(), 2); + + let step3 = backoff.step().expect("step 3"); + assert_eq!(step3, Duration::from_secs(4)); + // Capped at max_delay + assert_eq!(backoff.delay(), Duration::from_secs(4)); + assert_eq!(backoff.attempt(), 3); + } + + #[test] + fn test_max_retries() { + let config = BackoffConfig::new(Duration::from_millis(100), Duration::from_secs(1)) + .max_retries(Some(2)); + let mut backoff = Backoff::from_config(config); + + assert!(backoff.can_retry()); + assert!(backoff.step().is_some()); // attempt 1 + assert!(backoff.can_retry()); + assert!(backoff.step().is_some()); // attempt 2 + assert!(!backoff.can_retry()); + assert!(backoff.step().is_none()); // attempt 3 blocked + } + + #[test] + fn test_reset() { + let mut backoff = Backoff::new(Duration::from_secs(1), Duration::from_secs(10)); + backoff.step(); + backoff.step(); + assert_eq!(backoff.attempt(), 2); + assert_eq!(backoff.delay(), Duration::from_secs(4)); + + backoff.reset(); + assert_eq!(backoff.attempt(), 0); + assert_eq!(backoff.delay(), Duration::from_secs(1)); + } + + #[test] + fn test_jitter_bounds() { + let config = BackoffConfig::new(Duration::from_millis(1000), Duration::from_secs(10)) + .jitter_factor(0.2); + let mut backoff = Backoff::from_config(config); + + let mut seen_values = std::collections::HashSet::new(); + for _ in 0..50 { + let dur = backoff.step().expect("step"); + // jitter_factor=0.2 gives base * [0.8, 1.2], so 800ms..1200ms + assert!( + dur >= Duration::from_millis(800) && dur <= Duration::from_millis(1200), + "jittered duration {dur:?} outside [800ms, 1200ms]" + ); + seen_values.insert(dur.as_millis()); + backoff.reset(); + } + // Verify jitter is actually producing varying values + assert!( + seen_values.len() > 1, + "jitter should produce varying durations, but all {len} iterations returned the same value", + len = seen_values.len() + ); + } + + #[test] + fn test_clamping_behavior() { + let config = BackoffConfig::default() + .multiplier(0.5) // should clamp to 1.0 + .jitter_factor(-0.5); // should clamp to 0.0 + assert_eq!(config.multiplier, 1.0); + assert_eq!(config.jitter_factor, 0.0); + + let config2 = BackoffConfig::default().jitter_factor(5.0); // should clamp to 1.0 + assert_eq!(config2.jitter_factor, 1.0); + } + + #[test] + fn test_direct_struct_normalization() { + // Directly instantiate struct bypassing builder methods + let unvalidated = BackoffConfig { + initial_delay: Duration::from_secs(10), + max_delay: Duration::from_secs(1), // invalid: max < initial + multiplier: -50.0, // invalid: < 1.0 + max_retries: None, + jitter_factor: 500.0, // invalid: > 1.0 + }; + + let backoff = Backoff::from_config(unvalidated); + assert_eq!(backoff.config().max_delay, Duration::from_secs(10)); // normalized to initial_delay + assert_eq!(backoff.config().multiplier, 1.0); // normalized to 1.0 + assert_eq!(backoff.config().jitter_factor, 1.0); // normalized to 1.0 } } diff --git a/rivetkit-rust/packages/client/src/client.rs b/rivetkit-rust/packages/client/src/client.rs index 79a8a622b4..91e863a131 100644 --- a/rivetkit-rust/packages/client/src/client.rs +++ b/rivetkit-rust/packages/client/src/client.rs @@ -48,6 +48,7 @@ pub struct ClientConfig { pub headers: Option>, pub max_input_size: Option, pub disable_metadata_lookup: bool, + pub reconnect_backoff: Option, } impl ClientConfig { @@ -62,6 +63,7 @@ impl ClientConfig { headers: None, max_input_size: None, disable_metadata_lookup: false, + reconnect_backoff: None, } } @@ -116,12 +118,39 @@ impl ClientConfig { self.disable_metadata_lookup = disable; self } + + pub fn reconnect_backoff(mut self, mut backoff: crate::backoff::BackoffConfig) -> Self { + backoff.normalize(); + self.reconnect_backoff = Some(backoff); + self + } + + pub fn reconnect_delays( + mut self, + initial: std::time::Duration, + max: std::time::Duration, + ) -> Self { + self.reconnect_backoff = Some(crate::backoff::BackoffConfig::new(initial, max)); + self + } + + /// Disables automatic reconnection after a connection drops. + /// + /// The initial connection attempt via `connect()` will still be made. + /// This only prevents retries if that initial connection (or a later + /// established connection) fails or disconnects. + pub fn disable_reconnect(mut self) -> Self { + self.reconnect_backoff = + Some(crate::backoff::BackoffConfig::default().max_retries(Some(0))); + self + } } pub struct Client { remote_manager: RemoteManager, encoding_kind: EncodingKind, transport_kind: TransportKind, + reconnect_backoff: crate::backoff::BackoffConfig, shutdown_tx: Arc>, } @@ -131,6 +160,7 @@ impl Clone for Client { remote_manager: self.remote_manager.clone(), encoding_kind: self.encoding_kind, transport_kind: self.transport_kind, + reconnect_backoff: self.reconnect_backoff.clone(), shutdown_tx: self.shutdown_tx.clone(), } } @@ -141,6 +171,7 @@ impl std::fmt::Debug for Client { f.debug_struct("Client") .field("encoding_kind", &self.encoding_kind) .field("transport_kind", &self.transport_kind) + .field("reconnect_backoff", &self.reconnect_backoff) .finish_non_exhaustive() } } @@ -157,10 +188,14 @@ impl Client { config.disable_metadata_lookup, ); + let mut reconnect_backoff = config.reconnect_backoff.unwrap_or_default(); + reconnect_backoff.normalize(); + Self { remote_manager, encoding_kind: config.encoding, transport_kind: config.transport, + reconnect_backoff, shutdown_tx: Arc::new(tokio::sync::broadcast::channel(1).0), } } @@ -170,16 +205,15 @@ impl Client { } fn create_handle(&self, params: Option, query: ActorQuery) -> ActorHandle { - let handle = ActorHandle::new( + ActorHandle::new( self.remote_manager.clone(), params, query, self.shutdown_tx.clone(), self.transport_kind, self.encoding_kind, - ); - - handle + self.reconnect_backoff.clone(), + ) } pub fn get(&self, name: &str, key: ActorKey, opts: GetOptions) -> Result { diff --git a/rivetkit-rust/packages/client/src/connection.rs b/rivetkit-rust/packages/client/src/connection.rs index fdf4807b9f..552608b94a 100644 --- a/rivetkit-rust/packages/client/src/connection.rs +++ b/rivetkit-rust/packages/client/src/connection.rs @@ -7,7 +7,6 @@ use std::fmt::Debug; use std::ops::Deref; use std::sync::atomic::{AtomicBool, AtomicU64, Ordering}; use std::sync::{Arc, Weak}; -use std::time::Duration; use tokio::sync::{broadcast, oneshot, watch, Mutex}; use crate::{ @@ -137,6 +136,7 @@ pub struct ActorConnectionInner { watch::Receiver, ), disconnection_rx: Mutex>>, + backoff_config: crate::backoff::BackoffConfig, } impl ActorConnectionInner { @@ -146,6 +146,7 @@ impl ActorConnectionInner { transport_kind: TransportKind, encoding_kind: EncodingKind, parameters: Option, + backoff_config: crate::backoff::BackoffConfig, ) -> ActorConnection { Arc::new(Self { remote_manager, @@ -169,6 +170,7 @@ impl ActorConnectionInner { dc_watch: watch::channel(false), status_watch: watch::channel(ConnectionStatus::Idle), disconnection_rx: Mutex::new(None), + backoff_config, }) } @@ -703,7 +705,7 @@ pub fn start_connection( 'keepalive: loop { debug!("Attempting to reconnect"); - let mut backoff = Backoff::new(Duration::from_secs(1), Duration::from_secs(30)); + let mut backoff = Backoff::from_config(conn.backoff_config.clone()); let mut retry_attempt = 0; 'retry: loop { retry_attempt += 1; @@ -719,14 +721,23 @@ pub fn start_connection( } if attempt.did_open { + backoff.reset(); break 'retry; } let mut dc_rx = conn.dc_watch.0.subscribe(); tokio::select! { - _ = backoff.tick() => {}, - _ = dc_rx.wait_for(|x| *x == true) => { + waited = backoff.tick() => { + if !waited { + debug!( + "Max reconnect attempts ({}) reached, stopping connection", + retry_attempt + ); + break 'keepalive; + } + }, + _ = dc_rx.wait_for(|x| *x) => { break 'keepalive; } _ = shutdown_rx.recv() => { diff --git a/rivetkit-rust/packages/client/src/handle.rs b/rivetkit-rust/packages/client/src/handle.rs index ff1b3701d3..c2d640572c 100644 --- a/rivetkit-rust/packages/client/src/handle.rs +++ b/rivetkit-rust/packages/client/src/handle.rs @@ -357,6 +357,7 @@ pub struct ActorHandle { client_shutdown_tx: Arc>, transport_kind: crate::TransportKind, encoding_kind: EncodingKind, + reconnect_backoff: crate::backoff::BackoffConfig, } impl ActorHandle { @@ -367,6 +368,7 @@ impl ActorHandle { client_shutdown_tx: Arc>, transport_kind: TransportKind, encoding_kind: EncodingKind, + reconnect_backoff: crate::backoff::BackoffConfig, ) -> Self { let handle = ActorHandleStateless::new( remote_manager.clone(), @@ -383,16 +385,30 @@ impl ActorHandle { client_shutdown_tx, transport_kind, encoding_kind, + reconnect_backoff, } } + pub fn reconnect_backoff(&self) -> &crate::backoff::BackoffConfig { + &self.reconnect_backoff + } + pub fn connect(&self) -> ActorConnection { + self.connect_with_backoff(self.reconnect_backoff.clone()) + } + + pub fn connect_with_backoff( + &self, + mut backoff: crate::backoff::BackoffConfig, + ) -> ActorConnection { + backoff.normalize(); let conn = ActorConnectionInner::new( self.remote_manager.clone(), self.query.clone(), self.transport_kind, self.encoding_kind, self.params.clone(), + backoff, ); let rx = self.client_shutdown_tx.subscribe(); diff --git a/rivetkit-rust/packages/client/src/lib.rs b/rivetkit-rust/packages/client/src/lib.rs index 114b3bcaac..bebf41d67e 100644 --- a/rivetkit-rust/packages/client/src/lib.rs +++ b/rivetkit-rust/packages/client/src/lib.rs @@ -14,6 +14,7 @@ pub mod handle; pub mod protocol; mod remote_manager; +pub use backoff::{Backoff, BackoffConfig}; pub use client::{ Client, ClientConfig, CreateOptions, GetOptions, GetOrCreateOptions, GetWithIdOptions, }; diff --git a/rivetkit-rust/packages/client/tests/backoff.rs b/rivetkit-rust/packages/client/tests/backoff.rs new file mode 100644 index 0000000000..d965eed3a6 --- /dev/null +++ b/rivetkit-rust/packages/client/tests/backoff.rs @@ -0,0 +1,366 @@ +use rivetkit_client::{Backoff, BackoffConfig, Client, ClientConfig}; +use std::time::Duration; + +#[test] +fn default_backoff_config_matches_standard_policy() { + let config = BackoffConfig::default(); + assert_eq!(config.initial_delay, Duration::from_secs(1)); + assert_eq!(config.max_delay, Duration::from_secs(30)); + assert_eq!(config.multiplier, 2.0); + assert_eq!(config.max_retries, None); + assert_eq!(config.jitter_factor, 0.0); +} + +#[test] +fn backoff_builder_methods_work_as_expected() { + let config = BackoffConfig::new(Duration::from_millis(500), Duration::from_secs(10)) + .multiplier(1.5) + .max_retries(Some(5)) + .jitter(true); + + assert_eq!(config.initial_delay, Duration::from_millis(500)); + assert_eq!(config.max_delay, Duration::from_secs(10)); + assert_eq!(config.multiplier, 1.5); + assert_eq!(config.max_retries, Some(5)); + assert_eq!(config.jitter_factor, 0.2); + + let custom_jitter = config.jitter_factor(0.35); + assert_eq!(custom_jitter.jitter_factor, 0.35); + + // Test silent clamping on invalid builder inputs + let clamped_multiplier = BackoffConfig::default().multiplier(0.1); + assert_eq!(clamped_multiplier.multiplier, 1.0); + + let clamped_jitter_low = BackoffConfig::default().jitter_factor(-1.0); + assert_eq!(clamped_jitter_low.jitter_factor, 0.0); + + let clamped_jitter_high = BackoffConfig::default().jitter_factor(3.0); + assert_eq!(clamped_jitter_high.jitter_factor, 1.0); +} + +#[test] +fn backoff_exponential_growth_and_capping() { + let config = + BackoffConfig::new(Duration::from_millis(100), Duration::from_millis(600)).multiplier(2.0); + let mut backoff = Backoff::from_config(config); + + assert_eq!(backoff.attempt(), 0); + assert_eq!(backoff.delay(), Duration::from_millis(100)); + + // Attempt 1: yields 100ms, advances next base to 200ms + let dur1 = backoff.step().expect("step 1"); + assert_eq!(dur1, Duration::from_millis(100)); + assert_eq!(backoff.attempt(), 1); + assert_eq!(backoff.delay(), Duration::from_millis(200)); + + // Attempt 2: yields 200ms, advances next base to 400ms + let dur2 = backoff.step().expect("step 2"); + assert_eq!(dur2, Duration::from_millis(200)); + assert_eq!(backoff.attempt(), 2); + assert_eq!(backoff.delay(), Duration::from_millis(400)); + + // Attempt 3: yields 400ms, advances next base capped at 600ms + let dur3 = backoff.step().expect("step 3"); + assert_eq!(dur3, Duration::from_millis(400)); + assert_eq!(backoff.attempt(), 3); + assert_eq!(backoff.delay(), Duration::from_millis(600)); + + // Attempt 4: yields 600ms, remains capped at 600ms + let dur4 = backoff.step().expect("step 4"); + assert_eq!(dur4, Duration::from_millis(600)); + assert_eq!(backoff.attempt(), 4); + assert_eq!(backoff.delay(), Duration::from_millis(600)); +} + +#[test] +fn backoff_respects_max_retries() { + let config = BackoffConfig::new(Duration::from_millis(10), Duration::from_millis(100)) + .max_retries(Some(3)); + let mut backoff = Backoff::from_config(config); + + assert!(backoff.can_retry()); + assert!(backoff.step().is_some()); // attempt 1 + assert!(backoff.can_retry()); + assert!(backoff.step().is_some()); // attempt 2 + assert!(backoff.can_retry()); + assert!(backoff.step().is_some()); // attempt 3 + + // At 3 attempts, max is reached + assert!(!backoff.can_retry()); + assert!(backoff.step().is_none()); + assert_eq!(backoff.attempt(), 3); +} + +#[test] +fn backoff_reset_restores_initial_state() { + let config = BackoffConfig::new(Duration::from_millis(100), Duration::from_millis(800)); + let mut backoff = Backoff::from_config(config); + + backoff.step(); + backoff.step(); + backoff.step(); + assert_eq!(backoff.attempt(), 3); + assert!(backoff.delay() > Duration::from_millis(100)); + + backoff.reset(); + assert_eq!(backoff.attempt(), 0); + assert_eq!(backoff.delay(), Duration::from_millis(100)); + assert!(backoff.can_retry()); +} + +#[test] +fn backoff_with_jitter_stays_within_bounds() { + let config = + BackoffConfig::new(Duration::from_millis(1000), Duration::from_secs(10)).jitter_factor(0.2); // +/- 20% + let mut backoff = Backoff::from_config(config); + + let mut seen_values = std::collections::HashSet::new(); + for _ in 0..50 { + let dur = backoff.step().expect("step"); + // 1000ms +/- 20% = [800ms, 1200ms] + assert!( + dur >= Duration::from_millis(800) && dur <= Duration::from_millis(1200), + "duration {dur:?} outside jitter range [800ms, 1200ms]" + ); + seen_values.insert(dur.as_millis()); + backoff.reset(); + } + + assert!( + seen_values.len() > 1, + "expected multiple distinct jitter durations, got only {}", + seen_values.len() + ); +} + +#[test] +fn client_config_reconnect_builder_integration() { + let client_config = ClientConfig::new("http://127.0.0.1:6420") + .reconnect_delays(Duration::from_millis(250), Duration::from_secs(15)); + + let backoff = client_config + .reconnect_backoff + .as_ref() + .expect("backoff set"); + assert_eq!(backoff.initial_delay, Duration::from_millis(250)); + assert_eq!(backoff.max_delay, Duration::from_secs(15)); + + let disabled_config = ClientConfig::new("http://127.0.0.1:6420").disable_reconnect(); + let disabled_backoff = disabled_config + .reconnect_backoff + .as_ref() + .expect("backoff set"); + assert_eq!(disabled_backoff.max_retries, Some(0)); + + let client = Client::new(client_config); + let handle = client + .get("test", vec!["key1".to_string()], Default::default()) + .expect("handle"); + assert_eq!( + handle.reconnect_backoff().initial_delay, + Duration::from_millis(250) + ); + assert_eq!( + handle.reconnect_backoff().max_delay, + Duration::from_secs(15) + ); +} + +#[test] +fn direct_struct_construction_is_normalized_on_backoff_creation() { + let custom = BackoffConfig { + initial_delay: Duration::from_secs(5), + max_delay: Duration::from_millis(500), // invalid: max < initial + multiplier: 0.2, // invalid: < 1.0 + max_retries: Some(2), + jitter_factor: 10.0, // invalid: > 1.0 + }; + + let mut backoff = Backoff::from_config(custom); + assert_eq!(backoff.config().max_delay, Duration::from_secs(5)); + assert_eq!(backoff.config().multiplier, 1.0); + assert_eq!(backoff.config().jitter_factor, 1.0); + + // Ensure step() works reliably with normalized config + let dur1 = backoff.step().expect("step 1"); + assert!(dur1 <= Duration::from_secs(10)); +} + +#[tokio::test] +async fn integration_max_retries_stops_reconnect_loop() { + use axum::{http::StatusCode, routing::any, Router}; + use rivetkit_client::GetOrCreateOptions; + use std::sync::{ + atomic::{AtomicUsize, Ordering}, + Arc, + }; + use tokio::{net::TcpListener, time::sleep}; + + let attempts = Arc::new(AtomicUsize::new(0)); + let app = Router::new().route( + "/gateway/{actor_id}/connect", + any({ + let attempts = attempts.clone(); + move || { + attempts.fetch_add(1, Ordering::SeqCst); + async { StatusCode::INTERNAL_SERVER_ERROR } + } + }), + ); + + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let addr = listener.local_addr().unwrap(); + let server = tokio::spawn(async move { + axum::serve(listener, app).await.unwrap(); + }); + + let client_config = ClientConfig::new(format!("http://{addr}")) + .disable_metadata_lookup(true) + .reconnect_backoff( + BackoffConfig::default() + .initial_delay(Duration::from_millis(5)) + .max_delay(Duration::from_millis(20)) + .max_retries(Some(2)), // 1 initial + 2 retries = 3 total attempts + ); + + let client = Client::new(client_config); + let actor = client + .get_or_create( + "test-actor", + vec!["key1".to_string()], + GetOrCreateOptions::default(), + ) + .unwrap(); + + let _conn = actor.connect(); + + // Wait for reconnect attempts to complete (5ms + 10ms + buffer) + sleep(Duration::from_millis(150)).await; + + // Assert exactly 3 connection attempts: initial + 2 retries + assert_eq!( + attempts.load(Ordering::SeqCst), + 3, + "expected exactly 3 attempts (1 initial + 2 retries)" + ); + + // Ensure no further reconnection attempts occur + sleep(Duration::from_millis(100)).await; + assert_eq!(attempts.load(Ordering::SeqCst), 3); + + server.abort(); +} + +#[tokio::test] +async fn integration_disable_reconnect_attempts_once() { + use axum::{http::StatusCode, routing::any, Router}; + use rivetkit_client::GetOrCreateOptions; + use std::sync::{ + atomic::{AtomicUsize, Ordering}, + Arc, + }; + use tokio::{net::TcpListener, time::sleep}; + + let attempts = Arc::new(AtomicUsize::new(0)); + let app = Router::new().route( + "/gateway/{actor_id}/connect", + any({ + let attempts = attempts.clone(); + move || { + attempts.fetch_add(1, Ordering::SeqCst); + async { StatusCode::INTERNAL_SERVER_ERROR } + } + }), + ); + + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let addr = listener.local_addr().unwrap(); + let server = tokio::spawn(async move { + axum::serve(listener, app).await.unwrap(); + }); + + let client_config = ClientConfig::new(format!("http://{addr}")) + .disable_metadata_lookup(true) + .disable_reconnect(); + + let client = Client::new(client_config); + let actor = client + .get_or_create( + "test-actor", + vec!["key1".to_string()], + GetOrCreateOptions::default(), + ) + .unwrap(); + + let _conn = actor.connect(); + + // Wait and verify only the initial attempt occurred + sleep(Duration::from_millis(150)).await; + assert_eq!( + attempts.load(Ordering::SeqCst), + 1, + "disable_reconnect should allow the initial attempt but zero retries" + ); + + server.abort(); +} + +#[tokio::test] +async fn integration_handle_connect_with_backoff_override() { + use axum::{http::StatusCode, routing::any, Router}; + use rivetkit_client::GetOrCreateOptions; + use std::sync::{ + atomic::{AtomicUsize, Ordering}, + Arc, + }; + use tokio::{net::TcpListener, time::sleep}; + + let attempts = Arc::new(AtomicUsize::new(0)); + let app = Router::new().route( + "/gateway/{actor_id}/connect", + any({ + let attempts = attempts.clone(); + move || { + attempts.fetch_add(1, Ordering::SeqCst); + async { StatusCode::INTERNAL_SERVER_ERROR } + } + }), + ); + + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let addr = listener.local_addr().unwrap(); + let server = tokio::spawn(async move { + axum::serve(listener, app).await.unwrap(); + }); + + // Client has disable_reconnect by default + let client = Client::new( + ClientConfig::new(format!("http://{addr}")) + .disable_metadata_lookup(true) + .disable_reconnect(), + ); + + let actor = client + .get_or_create( + "test-actor", + vec!["key1".to_string()], + GetOrCreateOptions::default(), + ) + .unwrap(); + + // Per-connection override allowing 1 retry (2 attempts total) + let _conn = actor.connect_with_backoff( + BackoffConfig::default() + .initial_delay(Duration::from_millis(5)) + .max_retries(Some(1)), + ); + + sleep(Duration::from_millis(150)).await; + assert_eq!( + attempts.load(Ordering::SeqCst), + 2, + "connect_with_backoff override should permit 2 attempts (1 initial + 1 retry)" + ); + + server.abort(); +} From 03476e01af5d31f1d71cadc7ac03f6b9667beb94 Mon Sep 17 00:00:00 2001 From: Arth Srivastava Date: Thu, 8 Oct 2026 20:51:12 +0530 Subject: [PATCH 2/2] fix(rivetkit-client): handle reconnect timing, jitter clamping, and keepalive termination --- rivetkit-rust/packages/client/src/backoff.rs | 114 +-------- .../packages/client/src/connection.rs | 26 ++ .../packages/client/tests/backoff.rs | 229 ++++++++++++++++++ 3 files changed, 258 insertions(+), 111 deletions(-) diff --git a/rivetkit-rust/packages/client/src/backoff.rs b/rivetkit-rust/packages/client/src/backoff.rs index 84bd043948..c309e34d1d 100644 --- a/rivetkit-rust/packages/client/src/backoff.rs +++ b/rivetkit-rust/packages/client/src/backoff.rs @@ -190,7 +190,9 @@ impl Backoff { let sleep_duration = if self.config.jitter_factor > 0.0 { let jitter_offset = (rand::random::() * 2.0 - 1.0) * self.config.jitter_factor; let factor = (1.0 + jitter_offset).max(0.0); - Duration::from_secs_f64((base.as_secs_f64() * factor).max(0.0)) + let max_secs = self.config.max_delay.as_secs_f64(); + let jittered = (base.as_secs_f64() * factor).clamp(0.0, max_secs); + Duration::from_secs_f64(jittered) } else { base }; @@ -213,113 +215,3 @@ impl Backoff { self.attempt = 0; } } - -#[cfg(test)] -mod tests { - use super::*; - - #[test] - fn test_default_progression() { - let mut backoff = Backoff::new(Duration::from_secs(1), Duration::from_secs(4)); - assert_eq!(backoff.delay(), Duration::from_secs(1)); - assert_eq!(backoff.attempt(), 0); - assert!(backoff.can_retry()); - - let step1 = backoff.step().expect("step 1"); - assert_eq!(step1, Duration::from_secs(1)); - assert_eq!(backoff.delay(), Duration::from_secs(2)); - assert_eq!(backoff.attempt(), 1); - - let step2 = backoff.step().expect("step 2"); - assert_eq!(step2, Duration::from_secs(2)); - assert_eq!(backoff.delay(), Duration::from_secs(4)); - assert_eq!(backoff.attempt(), 2); - - let step3 = backoff.step().expect("step 3"); - assert_eq!(step3, Duration::from_secs(4)); - // Capped at max_delay - assert_eq!(backoff.delay(), Duration::from_secs(4)); - assert_eq!(backoff.attempt(), 3); - } - - #[test] - fn test_max_retries() { - let config = BackoffConfig::new(Duration::from_millis(100), Duration::from_secs(1)) - .max_retries(Some(2)); - let mut backoff = Backoff::from_config(config); - - assert!(backoff.can_retry()); - assert!(backoff.step().is_some()); // attempt 1 - assert!(backoff.can_retry()); - assert!(backoff.step().is_some()); // attempt 2 - assert!(!backoff.can_retry()); - assert!(backoff.step().is_none()); // attempt 3 blocked - } - - #[test] - fn test_reset() { - let mut backoff = Backoff::new(Duration::from_secs(1), Duration::from_secs(10)); - backoff.step(); - backoff.step(); - assert_eq!(backoff.attempt(), 2); - assert_eq!(backoff.delay(), Duration::from_secs(4)); - - backoff.reset(); - assert_eq!(backoff.attempt(), 0); - assert_eq!(backoff.delay(), Duration::from_secs(1)); - } - - #[test] - fn test_jitter_bounds() { - let config = BackoffConfig::new(Duration::from_millis(1000), Duration::from_secs(10)) - .jitter_factor(0.2); - let mut backoff = Backoff::from_config(config); - - let mut seen_values = std::collections::HashSet::new(); - for _ in 0..50 { - let dur = backoff.step().expect("step"); - // jitter_factor=0.2 gives base * [0.8, 1.2], so 800ms..1200ms - assert!( - dur >= Duration::from_millis(800) && dur <= Duration::from_millis(1200), - "jittered duration {dur:?} outside [800ms, 1200ms]" - ); - seen_values.insert(dur.as_millis()); - backoff.reset(); - } - // Verify jitter is actually producing varying values - assert!( - seen_values.len() > 1, - "jitter should produce varying durations, but all {len} iterations returned the same value", - len = seen_values.len() - ); - } - - #[test] - fn test_clamping_behavior() { - let config = BackoffConfig::default() - .multiplier(0.5) // should clamp to 1.0 - .jitter_factor(-0.5); // should clamp to 0.0 - assert_eq!(config.multiplier, 1.0); - assert_eq!(config.jitter_factor, 0.0); - - let config2 = BackoffConfig::default().jitter_factor(5.0); // should clamp to 1.0 - assert_eq!(config2.jitter_factor, 1.0); - } - - #[test] - fn test_direct_struct_normalization() { - // Directly instantiate struct bypassing builder methods - let unvalidated = BackoffConfig { - initial_delay: Duration::from_secs(10), - max_delay: Duration::from_secs(1), // invalid: max < initial - multiplier: -50.0, // invalid: < 1.0 - max_retries: None, - jitter_factor: 500.0, // invalid: > 1.0 - }; - - let backoff = Backoff::from_config(unvalidated); - assert_eq!(backoff.config().max_delay, Duration::from_secs(10)); // normalized to initial_delay - assert_eq!(backoff.config().multiplier, 1.0); // normalized to 1.0 - assert_eq!(backoff.config().jitter_factor, 1.0); // normalized to 1.0 - } -} diff --git a/rivetkit-rust/packages/client/src/connection.rs b/rivetkit-rust/packages/client/src/connection.rs index 552608b94a..776d3f7121 100644 --- a/rivetkit-rust/packages/client/src/connection.rs +++ b/rivetkit-rust/packages/client/src/connection.rs @@ -722,6 +722,32 @@ pub fn start_connection( if attempt.did_open { backoff.reset(); + + // After a successful connection that later closed, check + // if the retry budget allows another reconnect cycle. + // This is how disable_reconnect() (max_retries=0) stops + // reconnection after the initial connection drops. + if !backoff.can_retry() { + break 'keepalive; + } + + let mut dc_rx = conn.dc_watch.0.subscribe(); + + tokio::select! { + waited = backoff.tick() => { + if !waited { + break 'keepalive; + } + } + _ = dc_rx.wait_for(|x| *x) => { + break 'keepalive; + } + _ = shutdown_rx.recv() => { + debug!("Received shutdown signal, stopping connection attempts"); + break 'keepalive; + } + } + break 'retry; } diff --git a/rivetkit-rust/packages/client/tests/backoff.rs b/rivetkit-rust/packages/client/tests/backoff.rs index d965eed3a6..6ef4a58f59 100644 --- a/rivetkit-rust/packages/client/tests/backoff.rs +++ b/rivetkit-rust/packages/client/tests/backoff.rs @@ -133,6 +133,31 @@ fn backoff_with_jitter_stays_within_bounds() { ); } +#[test] +fn backoff_with_jitter_clamped_to_max_delay() { + let config = BackoffConfig::new(Duration::from_millis(500), Duration::from_millis(500)) + .multiplier(2.0) + .jitter_factor(0.5); // base is 500ms, +/- 50% jitter would reach 750ms without clamping + let mut backoff = Backoff::from_config(config); + + let mut seen_values = std::collections::HashSet::new(); + for _ in 0..100 { + let dur = backoff.step().expect("step"); + assert!( + dur <= Duration::from_millis(500), + "jittered duration {dur:?} exceeded max_delay of 500ms" + ); + seen_values.insert(dur.as_millis()); + backoff.reset(); + } + + assert!( + seen_values.len() > 1, + "expected multiple distinct jitter durations below cap, got only {}", + seen_values.len() + ); +} + #[test] fn client_config_reconnect_builder_integration() { let client_config = ClientConfig::new("http://127.0.0.1:6420") @@ -364,3 +389,207 @@ async fn integration_handle_connect_with_backoff_override() { server.abort(); } + +#[tokio::test] +async fn integration_disable_reconnect_stops_after_connection_closes() { + use axum::{ + extract::ws::{Message as AxumWsMessage, WebSocket, WebSocketUpgrade}, + routing::any, + Router, + }; + use futures_util::SinkExt; + use rivetkit_client::GetOrCreateOptions; + use rivetkit_client_protocol as wire; + use std::sync::{ + atomic::{AtomicUsize, Ordering}, + Arc, + }; + use tokio::{net::TcpListener, time::sleep}; + use vbare::OwnedVersionedData; + + let connect_count = Arc::new(AtomicUsize::new(0)); + let connect_count_handler = connect_count.clone(); + + let app = Router::new().route( + "/gateway/{actor_id}/connect", + any(move |ws: WebSocketUpgrade| { + let count = connect_count_handler.clone(); + async move { + count.fetch_add(1, Ordering::SeqCst); + ws.protocols(["rivet"]) + .on_upgrade(|mut socket: WebSocket| async move { + let payload = wire::versioned::ToClient::wrap_latest(wire::ToClient { + body: wire::ToClientBody::Init(wire::Init { + actor_id: "test-actor".to_owned(), + connection_id: "conn-1".to_owned(), + }), + }) + .serialize_with_embedded_version(wire::PROTOCOL_VERSION) + .unwrap(); + socket + .send(AxumWsMessage::Binary(payload.into())) + .await + .unwrap(); + // Close connection from server side + let _ = socket.close().await; + }) + } + }), + ); + + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let addr = listener.local_addr().unwrap(); + let server = tokio::spawn(async move { + axum::serve(listener, app).await.unwrap(); + }); + + let client_config = ClientConfig::new(format!("http://{addr}")) + .disable_metadata_lookup(true) + .disable_reconnect(); + + let client = Client::new(client_config); + let actor = client + .get_or_create( + "test-actor", + vec!["key1".to_string()], + GetOrCreateOptions::default(), + ) + .unwrap(); + + let _conn = actor.connect(); + + // Wait for the connection to establish, close, and verify no reconnection occurs + sleep(Duration::from_millis(200)).await; + + assert_eq!( + connect_count.load(Ordering::SeqCst), + 1, + "disable_reconnect should not reconnect after established connection closes" + ); + + server.abort(); +} + +#[tokio::test] +async fn integration_reconnect_after_healthy_close_respects_initial_delay() { + use axum::{ + extract::ws::{Message as AxumWsMessage, WebSocket, WebSocketUpgrade}, + routing::any, + Router, + }; + use futures_util::SinkExt; + use rivetkit_client::GetOrCreateOptions; + use rivetkit_client_protocol as wire; + use std::sync::{ + atomic::{AtomicUsize, Ordering}, + Arc, + }; + use tokio::{ + net::TcpListener, + sync::mpsc, + time::{sleep, timeout, Instant}, + }; + use vbare::OwnedVersionedData; + + let attempts = Arc::new(AtomicUsize::new(0)); + let attempts_handler = attempts.clone(); + + let (closed_tx, mut closed_rx) = mpsc::unbounded_channel::(); + let closed_tx = Arc::new(tokio::sync::Mutex::new(Some(closed_tx))); + + let (reconnect_tx, mut reconnect_rx) = mpsc::unbounded_channel::(); + let reconnect_tx = Arc::new(reconnect_tx); + + let app = Router::new().route( + "/gateway/{actor_id}/connect", + any(move |ws: WebSocketUpgrade| { + let attempts = attempts_handler.clone(); + let closed_tx = closed_tx.clone(); + let reconnect_tx = reconnect_tx.clone(); + async move { + let attempt_num = attempts.fetch_add(1, Ordering::SeqCst); + if attempt_num == 0 { + ws.protocols(["rivet"]) + .on_upgrade(move |mut socket: WebSocket| async move { + let payload = wire::versioned::ToClient::wrap_latest(wire::ToClient { + body: wire::ToClientBody::Init(wire::Init { + actor_id: "test-actor".to_owned(), + connection_id: "conn-1".to_owned(), + }), + }) + .serialize_with_embedded_version(wire::PROTOCOL_VERSION) + .unwrap(); + socket + .send(AxumWsMessage::Binary(payload.into())) + .await + .unwrap(); + + // Server closes the connection and records timestamp + let _ = socket.close().await; + if let Some(tx) = closed_tx.lock().await.take() { + let _ = tx.send(Instant::now()); + } + }) + } else { + let _ = reconnect_tx.send(Instant::now()); + ws.protocols(["rivet"]) + .on_upgrade(|_socket: WebSocket| async move {}) + } + } + }), + ); + + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let addr = listener.local_addr().unwrap(); + let server = tokio::spawn(async move { + axum::serve(listener, app).await.unwrap(); + }); + + let client_config = ClientConfig::new(format!("http://{addr}")) + .disable_metadata_lookup(true) + .reconnect_backoff( + BackoffConfig::default() + .initial_delay(Duration::from_millis(100)) + .max_delay(Duration::from_millis(500)) + .jitter_factor(0.0), // zero jitter for deterministic delay assertion + ); + + let client = Client::new(client_config); + let actor = client + .get_or_create( + "test-actor", + vec!["key1".to_string()], + GetOrCreateOptions::default(), + ) + .unwrap(); + + let _conn = actor.connect(); + + // Wait for the first connection to close + let closed_at = timeout(Duration::from_secs(2), closed_rx.recv()) + .await + .expect("first connection did not close in time") + .expect("channel closed"); + + // At 25ms, reconnect attempt must NOT have happened yet (initial_delay is 100ms) + sleep(Duration::from_millis(25)).await; + assert_eq!( + attempts.load(Ordering::SeqCst), + 1, + "reconnect must not happen immediately after healthy connection closes" + ); + + // Wait for reconnect attempt to occur + let reconnected_at = timeout(Duration::from_secs(2), reconnect_rx.recv()) + .await + .expect("reconnect did not occur within timeout") + .expect("channel closed"); + + let elapsed = reconnected_at.duration_since(closed_at); + assert!( + elapsed >= Duration::from_millis(80), + "reconnect happened too fast ({elapsed:?}), expected >= 80ms for 100ms initial_delay" + ); + + server.abort(); +}