From 302da130c57633b7497aea4473de20ebc1feb8da Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Jo=C3=A3o=20Lucas=20de=20Oliveira=20Lopes?= Date: Tue, 6 Oct 2026 14:40:23 +0000 Subject: [PATCH 1/2] fix(lifecycle): recheck host admission before queued dials --- src/client/lifecycle.rs | 45 ++-- src/client/lifecycle/admission_tests.rs | 54 +++- src/types/connect_admission.rs | 23 +- tests/api-consumer/src/lifecycle/admission.rs | 27 ++ tests/connect_admission.rs | 234 +++++++++++++++++- 5 files changed, 363 insertions(+), 20 deletions(-) diff --git a/src/client/lifecycle.rs b/src/client/lifecycle.rs index a58ab72a6..d0728c044 100644 --- a/src/client/lifecycle.rs +++ b/src/client/lifecycle.rs @@ -898,7 +898,12 @@ impl Client { let pause_generation = self.pause_generation.load(Ordering::SeqCst); let delay = admission.delay(); if !self - .wait_for_connect_admission(delay, pause_generation, &shutdown) + .wait_for_connect_admission( + admission.as_ref(), + delay, + pause_generation, + &shutdown, + ) .await { if self.is_paused() @@ -1147,7 +1152,8 @@ impl Client { /// ordinary teardown wake. The generation catches a rapid pause/resume. async fn wait_for_connect_admission( &self, - delay: Duration, + admission: &dyn crate::ConnectAdmission, + mut delay: Duration, pause_generation: u64, shutdown: &wacore::runtime::ShutdownSignal, ) -> bool { @@ -1157,23 +1163,32 @@ impl Client { && !self.is_paused() && self.pause_generation.load(Ordering::SeqCst) == pause_generation }; - if delay.is_zero() { - return can_dial(); - } - let sleep = self.runtime.sleep(delay).fuse(); - let stopped = wacore::runtime::wait_for_shutdown(shutdown).fuse(); - futures::pin_mut!(sleep, stopped); loop { - // Register before checking: pause/resume and supervision-stop - // notifications must not be lost between the check and select. - let changed = self.session_state_notifier.listen(); + if !delay.is_zero() { + let sleep = self.runtime.sleep(delay).fuse(); + let stopped = wacore::runtime::wait_for_shutdown(shutdown).fuse(); + futures::pin_mut!(sleep, stopped); + loop { + // Register before checking; unrelated wakes retain this timer. + let changed = self.session_state_notifier.listen(); + if !can_dial() { + return false; + } + futures::select! { + _ = sleep => break, + _ = stopped => return false, + _ = changed.fuse() => {} + } + } + } if !can_dial() { return false; } - futures::select! { - _ = sleep => return can_dial(), - _ = stopped => return false, - _ = changed.fuse() => {} + // Keep the reservation and pause generation across extensions. + // Calling delay() again would book another slot in host budgets. + delay = admission.recheck(); + if delay.is_zero() { + return can_dial(); } } } diff --git a/src/client/lifecycle/admission_tests.rs b/src/client/lifecycle/admission_tests.rs index 3361a0453..f622144bf 100644 --- a/src/client/lifecycle/admission_tests.rs +++ b/src/client/lifecycle/admission_tests.rs @@ -10,6 +10,10 @@ impl ConnectAdmission for Wait { } async fn client() -> Arc { + client_with_policy(Wait).await +} + +async fn client_with_policy(policy: impl ConnectAdmission + 'static) -> Arc { Client::builder() .with_runtime(crate::runtime_impl::TokioRuntime) .with_persistence_manager(Arc::new( @@ -19,13 +23,61 @@ async fn client() -> Arc { )) .with_http_client(crate::test_utils::MockHttpClient) .with_transport_factory(crate::transport::mock::MockTransportFactory::new()) - .with_connect_admission(Wait) + .with_connect_admission(policy) .build() .await .unwrap() .into_client() } +struct Extend(Arc); +impl ConnectAdmission for Extend { + fn delay(&self) -> Duration { + Duration::ZERO + } + fn recheck(&self) -> Duration { + self.0.fetch_add(1, Ordering::SeqCst); + Duration::from_secs(900) + } +} + +#[tokio::test(start_paused = true)] +async fn admission_extension_retains_timer_on_unrelated_wakes_and_cancels_on_stop() { + let checks = Arc::new(std::sync::atomic::AtomicUsize::new(0)); + let client = client_with_policy(Extend(checks.clone())).await; + client.auto_reconnect_errors.store(5, Ordering::Relaxed); + client + .backoff_reset_suppressed + .store(true, Ordering::Relaxed); + let runner = client.clone(); + let run = tokio::spawn(async move { runner.run().await }); + crate::test_utils::wait_for_notifier_listeners(&client.session_state_notifier, 1).await; + for _ in 0..2 { + tokio::time::advance(Duration::from_secs(300)).await; + client.notify_connection_shutdown(); + client.notify_session_state(); + tokio::task::yield_now().await; + crate::test_utils::wait_for_notifier_listeners(&client.session_state_notifier, 1).await; + assert_eq!(checks.load(Ordering::SeqCst), 1); + assert!(!client.is_connecting.load(Ordering::Relaxed)); + assert_eq!(client.stats().reconnects, 0); + assert_eq!(client.stats().reconnect_errors, 5); + assert!(client.backoff_reset_suppressed.load(Ordering::Relaxed)); + } + tokio::time::advance(Duration::from_secs(300)).await; + tokio::task::yield_now().await; + assert_eq!( + checks.load(Ordering::SeqCst), + 2, + "wakes must not restart the timer" + ); + client.stop_supervision_loop(); + assert!(matches!(run.await.unwrap(), RunCompletionReason::Stopped)); + assert_eq!(client.stats().reconnects, 0); + assert_eq!(client.stats().reconnect_errors, 5); + client.shutdown().await; +} + #[tokio::test] async fn admission_stop_preserves_stopped_reason() { let client = client().await; diff --git a/src/types/connect_admission.rs b/src/types/connect_admission.rs index 080e5cfa3..290d94f45 100644 --- a/src/types/connect_admission.rs +++ b/src/types/connect_admission.rs @@ -2,11 +2,13 @@ use std::time::Duration; /// Optional host-owned pacing of [`Client::run`](crate::Client::run) dials. /// -/// The run loop consults this policy once before starting each connect attempt, +/// The run loop calls [`delay`](Self::delay) once before each connect attempt, /// including the first and forced reconnects, after its ordinary reconnect /// backoff and pause gate. It serves the returned delay **before** calling /// [`Client::connect`](crate::Client::connect), outside the transport/version -/// connect timeouts. Zero proceeds without sleeping or consulting again. +/// connect timeouts. It then calls [`recheck`](Self::recheck), including when +/// the initial delay is zero, so a host can extend a wait after learning of a +/// shared cooldown. Existing policies need only implement `delay`. /// Direct calls to `connect()` do not consult this policy. /// /// Shutdown or a pause (even one already resumed) abandons the wait without a @@ -17,7 +19,7 @@ use std::time::Duration; /// cancellation/refund callback; hosts must account for abandoned reservations. /// The SDK still owns WhatsApp's backoff, stable reset, and rate-limit penalties. /// -/// This synchronous callback runs inline. It must return promptly without I/O, +/// These synchronous callbacks run inline. They must return promptly without I/O, /// blocking, or awaiting. Shutdown cannot interrupt a blocking callback; panics /// propagate to the caller driving the run loop, with no recovery guarantee. /// Share a policy across builders to coordinate a host-wide budget; no limiter @@ -50,4 +52,19 @@ use std::time::Duration; pub trait ConnectAdmission: wacore::sync_marker::MaybeSendSync { /// Reserve a run-loop attempt and return how long to delay it. fn delay(&self) -> Duration; + + /// Recheck an existing reservation before dialing, without reserving again. + /// + /// Called after the initial delay, even if zero, and after every positive + /// extension returned here. Return zero to proceed, or a positive duration + /// to wait before checking again. The default preserves one-shot policies. + /// Extensions remain outside connect timeouts and are cancelled by shutdown, + /// supervision stop, or pause just like the initial wait. Resume reserves a + /// new attempt through `delay`. + /// + /// This is a synchronous snapshot, not atomic exclusion with the dial: + /// a host update after the final zero cannot revoke that permission. + fn recheck(&self) -> Duration { + Duration::ZERO + } } diff --git a/tests/api-consumer/src/lifecycle/admission.rs b/tests/api-consumer/src/lifecycle/admission.rs index 79b7763c9..26f5887ce 100644 --- a/tests/api-consumer/src/lifecycle/admission.rs +++ b/tests/api-consumer/src/lifecycle/admission.rs @@ -17,6 +17,21 @@ impl ConnectAdmission for Policy { } } +pub struct RecheckingPolicy(pub Policy); +impl ConnectAdmission for RecheckingPolicy { + fn delay(&self) -> Duration { + self.0.delay() + } + fn recheck(&self) -> Duration { + // Reading the same host state remains valid with wasm's local Rc. + #[cfg(target_arch = "wasm32")] + let calls = self.0.0.get(); + #[cfg(not(target_arch = "wasm32"))] + let calls = self.0.0.load(std::sync::atomic::Ordering::Relaxed); + Duration::from_millis(calls as u64) + } +} + pub fn builder(policy: Arc) -> whatsapp_rust::ClientBuilder { Client::builder().with_connect_admission_arc(policy) } @@ -40,5 +55,17 @@ fn shared_policy_is_object_safe() { let policy = Policy(std::sync::atomic::AtomicUsize::new(0)); let boxed: Box = Box::new(policy); assert_eq!(boxed.delay(), Duration::ZERO); + assert_eq!(boxed.recheck(), Duration::ZERO); let _ = builder(Arc::from(boxed)); } + +#[cfg(not(target_arch = "wasm32"))] +#[test] +fn rechecking_policy_is_object_safe() { + let policy = Policy(std::sync::atomic::AtomicUsize::new(0)); + let shared: Arc = Arc::new(RecheckingPolicy(policy)); + assert_eq!(shared.delay(), Duration::ZERO); + assert_eq!(shared.recheck(), Duration::from_millis(1)); + assert_eq!(shared.recheck(), Duration::from_millis(1)); + let _ = builder(shared); +} diff --git a/tests/connect_admission.rs b/tests/connect_admission.rs index 472c5082f..d254d2c51 100644 --- a/tests/connect_admission.rs +++ b/tests/connect_admission.rs @@ -72,6 +72,13 @@ struct Fixture { } impl Fixture { async fn new(delays: Option>, fallback: Duration) -> Self { + Self::with_policy(delays, fallback, None).await + } + async fn with_policy( + delays: Option>, + fallback: Duration, + policy: Option>, + ) -> Self { let decisions = Arc::new(AtomicUsize::new(0)); let dials = Arc::new(AtomicUsize::new(0)); let order = Arc::new(Mutex::new(Vec::new())); @@ -94,7 +101,9 @@ impl Fixture { entered: enter_tx, release: release_rx, }); - if let Some(delays) = delays { + if let Some(policy) = policy { + builder = builder.with_connect_admission_arc(policy); + } else if let Some(delays) = delays { // The Arc setter also proves a shared trait object works publicly. builder = builder.with_connect_admission_arc(Arc::new(Policy { calls: decisions.clone(), @@ -149,6 +158,229 @@ async fn scheduled() { } } +struct CooldownPolicy { + initial_delay: Duration, + deadline: Mutex, + reservations: AtomicUsize, + rechecks: AtomicUsize, +} +impl CooldownPolicy { + fn new(initial_delay: Duration) -> Arc { + Arc::new(Self { + initial_delay, + deadline: Mutex::new(tokio::time::Instant::now()), + reservations: AtomicUsize::new(0), + rechecks: AtomicUsize::new(0), + }) + } + fn hold_for(&self, duration: Duration) { + *self.deadline.lock().unwrap() = tokio::time::Instant::now() + duration; + } +} +impl ConnectAdmission for CooldownPolicy { + fn delay(&self) -> Duration { + self.reservations.fetch_add(1, Ordering::SeqCst); + self.initial_delay + } + fn recheck(&self) -> Duration { + self.rechecks.fetch_add(1, Ordering::SeqCst); + self.deadline + .lock() + .unwrap() + .saturating_duration_since(tokio::time::Instant::now()) + } +} + +#[tokio::test(start_paused = true)] +async fn connect_admission_rechecks_cooldown_without_reserving_again() { + let policy = CooldownPolicy::new(Duration::from_secs(10)); + let f = Fixture::with_policy(None, Duration::ZERO, Some(policy.clone())).await; + let run = f.run(); + scheduled().await; + tokio::time::advance(Duration::from_secs(5)).await; + policy.hold_for(Duration::from_secs(25)); // A push-back moves admission to t=30. + tokio::time::advance(Duration::from_secs(5)).await; + scheduled().await; + f.no_attempt(); + assert_eq!(policy.rechecks.load(Ordering::SeqCst), 1); + tokio::time::advance(Duration::from_secs(19)).await; + scheduled().await; + f.no_attempt(); + policy.hold_for(Duration::from_secs(11)); // Another push-back moves it to t=40. + tokio::time::advance(Duration::from_secs(1)).await; + scheduled().await; + f.no_attempt(); + assert_eq!(policy.rechecks.load(Ordering::SeqCst), 2); + tokio::time::advance(Duration::from_secs(9)).await; + scheduled().await; + f.no_attempt(); + tokio::time::advance(Duration::from_secs(1)).await; + scheduled().await; + assert_eq!(f.dials.load(Ordering::SeqCst), 1); + assert_eq!(policy.reservations.load(Ordering::SeqCst), 1); + assert_eq!(policy.rechecks.load(Ordering::SeqCst), 3); + f.stop(run).await; +} + +#[tokio::test(start_paused = true)] +async fn connect_admission_zero_reservation_rechecks_and_keeps_full_transport_budget() { + let policy = CooldownPolicy::new(Duration::ZERO); + policy.hold_for(Duration::from_secs(60)); + let f = Fixture::with_policy(None, Duration::ZERO, Some(policy.clone())).await; + f.client.set_auto_reconnect(false); + let run = f.run(); + scheduled().await; + f.no_attempt(); + tokio::time::advance(Duration::from_secs(59)).await; + scheduled().await; + f.no_attempt(); + tokio::time::advance(Duration::from_secs(1)).await; + scheduled().await; + assert_eq!(f.dials.load(Ordering::SeqCst), 1); + tokio::time::advance(Duration::from_secs(19)).await; + scheduled().await; + assert!(!run.is_finished()); + tokio::time::advance(Duration::from_secs(1)).await; + assert!( + matches!(run.await.unwrap(), RunCompletionReason::AutoReconnectDisabled { + connect_error: Some(ConnectError::Timeout { stage: whatsapp_rust::ConnectStage::Transport, timeout }), .. + } if timeout == Duration::from_secs(20)) + ); + assert_eq!(policy.reservations.load(Ordering::SeqCst), 1); + assert_eq!(policy.rechecks.load(Ordering::SeqCst), 2); + f.client.shutdown().await; +} + +#[tokio::test(start_paused = true)] +async fn connect_admission_shutdown_cancels_cooldown_extension() { + let policy = CooldownPolicy::new(Duration::from_secs(10)); + policy.hold_for(Duration::from_secs(900)); + let f = Fixture::with_policy(None, Duration::ZERO, Some(policy.clone())).await; + let run = f.run(); + scheduled().await; + tokio::time::advance(Duration::from_secs(10)).await; + scheduled().await; + assert_eq!(policy.rechecks.load(Ordering::SeqCst), 1); + f.stop(run).await; + f.no_attempt(); + assert_eq!(policy.reservations.load(Ordering::SeqCst), 1); +} + +#[tokio::test(start_paused = true)] +async fn connect_admission_pause_resume_abandons_cooldown_extension() { + for rapid in [false, true] { + let policy = CooldownPolicy::new(Duration::from_secs(10)); + policy.hold_for(Duration::from_secs(900)); + let f = Fixture::with_policy(None, Duration::ZERO, Some(policy.clone())).await; + let run = f.run(); + scheduled().await; + tokio::time::advance(Duration::from_secs(10)).await; + scheduled().await; + assert_eq!(policy.rechecks.load(Ordering::SeqCst), 1); + f.client.pause().await; + if !rapid { + scheduled().await; + tokio::time::advance(Duration::from_secs(1000)).await; + scheduled().await; + f.no_attempt(); + assert_eq!(policy.reservations.load(Ordering::SeqCst), 1); + } + policy.hold_for(Duration::ZERO); + f.client.resume(); + scheduled().await; + f.no_attempt(); + assert_eq!(policy.reservations.load(Ordering::SeqCst), 2); + tokio::time::advance(Duration::from_secs(10)).await; + scheduled().await; + assert_eq!(f.dials.load(Ordering::SeqCst), 1); + assert_eq!(f.client.stats().reconnects, 0); + f.release.send(()).await.unwrap(); + scheduled().await; + assert_eq!(f.client.stats().reconnect_errors, 1); + f.stop(run).await; + } +} + +#[tokio::test(start_paused = true)] +async fn connect_admission_legacy_constant_delay_still_dials_once() { + let f = Fixture::new(Some(vec![]), Duration::from_millis(50)).await; + let run = f.run(); + f.admitted.recv().await.unwrap(); + tokio::time::advance(Duration::from_millis(50)).await; + scheduled().await; + assert_eq!(f.dials.load(Ordering::SeqCst), 1); + assert_eq!(f.decisions.load(Ordering::SeqCst), 1); + f.stop(run).await; +} + +#[tokio::test(start_paused = true)] +async fn connect_admission_rechecks_normal_and_forced_reconnects() { + for forced in [false, true] { + let policy = CooldownPolicy::new(Duration::ZERO); + let f = Fixture::with_policy(None, Duration::ZERO, Some(policy.clone())).await; + let run = f.run(); + f.entered.recv().await.unwrap(); + policy.hold_for(Duration::from_secs(60)); + if forced { + f.client.reconnect_immediately().await; + } + f.release.send(()).await.unwrap(); + scheduled().await; + tokio::time::advance(Duration::from_millis(1101)).await; + scheduled().await; + assert_eq!(policy.reservations.load(Ordering::SeqCst), 2); + assert_eq!(policy.rechecks.load(Ordering::SeqCst), 2); + assert_eq!(f.dials.load(Ordering::SeqCst), 1); + assert_eq!(f.client.stats().reconnects, 0); + assert_eq!(f.client.stats().reconnect_errors, u32::from(!forced)); + tokio::time::advance(Duration::from_millis(58899)).await; + scheduled().await; + assert_eq!(f.dials.load(Ordering::SeqCst), 2); + assert_eq!(f.client.stats().reconnects, 1); + assert_eq!(policy.reservations.load(Ordering::SeqCst), 2); + f.stop(run).await; + } +} + +struct StopOnRecheck { + client: Mutex>, + extension: Duration, +} +impl ConnectAdmission for StopOnRecheck { + fn delay(&self) -> Duration { + Duration::ZERO + } + fn recheck(&self) -> Duration { + self.client + .lock() + .unwrap() + .upgrade() + .unwrap() + .signal_shutdown_sync(); + self.extension + } +} + +#[tokio::test(start_paused = true)] +async fn connect_admission_checks_shutdown_after_inline_recheck() { + for extension in [Duration::ZERO, Duration::from_secs(900)] { + let policy = Arc::new(StopOnRecheck { + client: Mutex::new(std::sync::Weak::new()), + extension, + }); + let f = Fixture::with_policy(None, Duration::ZERO, Some(policy.clone())).await; + *policy.client.lock().unwrap() = Arc::downgrade(&f.client); + let run = f.run(); + scheduled().await; + assert!(run.is_finished()); + assert!(matches!( + run.await.unwrap(), + RunCompletionReason::ShutdownRequested + )); + f.no_attempt(); + } +} + #[tokio::test(start_paused = true)] async fn connect_admission_first_and_forced_reconnect_order() { let f = Fixture::new(Some(vec![Duration::ZERO; 3]), Duration::from_secs(900)).await; From 3f01a56dcab59ebae1e092a2eef8bfeaef8e55bb Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Jo=C3=A3o=20Lucas=20de=20Oliveira=20Lopes?= Date: Tue, 6 Oct 2026 14:41:02 +0000 Subject: [PATCH 2/2] test(lifecycle): use imported atomic counter type --- src/client/lifecycle/admission_tests.rs | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/src/client/lifecycle/admission_tests.rs b/src/client/lifecycle/admission_tests.rs index f622144bf..631b3223b 100644 --- a/src/client/lifecycle/admission_tests.rs +++ b/src/client/lifecycle/admission_tests.rs @@ -30,7 +30,7 @@ async fn client_with_policy(policy: impl ConnectAdmission + 'static) -> Arc); +struct Extend(Arc); impl ConnectAdmission for Extend { fn delay(&self) -> Duration { Duration::ZERO @@ -43,7 +43,7 @@ impl ConnectAdmission for Extend { #[tokio::test(start_paused = true)] async fn admission_extension_retains_timer_on_unrelated_wakes_and_cancels_on_stop() { - let checks = Arc::new(std::sync::atomic::AtomicUsize::new(0)); + let checks = Arc::new(AtomicUsize::new(0)); let client = client_with_policy(Extend(checks.clone())).await; client.auto_reconnect_errors.store(5, Ordering::Relaxed); client