@@ -4,8 +4,10 @@ use alloy::{
44} ;
55use async_trait:: async_trait;
66use blocksense_config:: WebsocketReconnectConfig ;
7- use std:: sync:: Arc ;
8- use tokio:: time:: { sleep, Duration } ;
7+ use std:: { future:: Future , sync:: Arc } ;
8+ #[ cfg( not( test) ) ]
9+ use tokio:: time:: sleep;
10+ use tokio:: time:: Duration ;
911use tracing:: { info, warn} ;
1012
1113#[ async_trait]
@@ -60,6 +62,30 @@ pub struct ResilientWsConnect {
6062 network : Arc < String > ,
6163}
6264
65+ #[ cfg( test) ]
66+ mod backoff_recorder {
67+ use std:: { future:: Future , sync:: Arc , time:: Duration } ;
68+ use tokio:: sync:: Mutex ;
69+
70+ tokio:: task_local! {
71+ static RECORDER : Arc <Mutex <Vec <Duration >>>;
72+ }
73+
74+ // Run fut with RECORDER set as task local var for the current task.
75+ pub ( super ) async fn with_recorder < F , R > ( recorder : Arc < Mutex < Vec < Duration > > > , fut : F ) -> R
76+ where
77+ F : Future < Output = R > ,
78+ {
79+ RECORDER . scope ( recorder, fut) . await
80+ }
81+
82+ pub ( super ) async fn record_backoff ( delay : Duration ) {
83+ if let Ok ( recorder) = RECORDER . try_with ( |r| r. clone ( ) ) {
84+ recorder. lock ( ) . await . push ( delay) ;
85+ }
86+ }
87+ }
88+
6389impl ResilientWsConnect {
6490 pub fn new (
6591 inner : WsConnect ,
@@ -90,6 +116,17 @@ impl PubSubConnect for ResilientWsConnect {
90116 }
91117
92118 async fn try_reconnect ( & self ) -> TransportResult < ConnectionHandle > {
119+ let inner = self . inner . clone ( ) ;
120+ self . reconnect_loop ( || PubSubConnect :: connect ( & inner) ) . await
121+ }
122+ }
123+
124+ impl ResilientWsConnect {
125+ async fn reconnect_loop < F , Fut > ( & self , mut connect : F ) -> TransportResult < ConnectionHandle >
126+ where
127+ F : FnMut ( ) -> Fut + Send ,
128+ Fut : Future < Output = TransportResult < ConnectionHandle > > + Send ,
129+ {
93130 self . metrics . on_disconnect ( ) . await ;
94131
95132 warn ! (
@@ -105,7 +142,7 @@ impl PubSubConnect for ResilientWsConnect {
105142 attempt = attempt. saturating_add ( 1 ) ;
106143 self . metrics . on_attempt ( ) . await ;
107144
108- match PubSubConnect :: connect ( & self . inner ) . await {
145+ match connect ( ) . await {
109146 Ok ( handle) => {
110147 self . metrics . on_success ( ) . await ;
111148 info ! (
@@ -126,9 +163,201 @@ impl PubSubConnect for ResilientWsConnect {
126163 error = %err,
127164 "WS reconnect attempt failed; will retry"
128165 ) ;
129- sleep ( delay) . await ;
166+ self . wait_before_retry ( delay) . await ;
130167 }
131168 }
132169 }
133170 }
171+
172+ #[ cfg( not( test) ) ]
173+ async fn wait_before_retry ( & self , delay : Duration ) {
174+ sleep ( delay) . await ;
175+ }
176+
177+ #[ cfg( test) ]
178+ async fn wait_before_retry ( & self , delay : Duration ) {
179+ backoff_recorder:: record_backoff ( delay) . await ;
180+ }
181+ }
182+
183+ #[ cfg( test) ]
184+ mod tests {
185+ use super :: * ;
186+ use alloy:: transports:: TransportErrorKind ;
187+ use async_trait:: async_trait;
188+ use std:: {
189+ collections:: VecDeque ,
190+ sync:: {
191+ atomic:: { AtomicUsize , Ordering } ,
192+ Arc ,
193+ } ,
194+ } ;
195+ use tokio:: sync:: Mutex ;
196+
197+ #[ derive( Clone ) ]
198+ struct MockConnector {
199+ outcomes : Arc < Mutex < VecDeque < MockOutcome > > > ,
200+ attempts : Arc < AtomicUsize > ,
201+ }
202+
203+ #[ derive( Clone ) ]
204+ enum MockOutcome {
205+ Ok ,
206+ Err ( & ' static str ) ,
207+ }
208+
209+ impl MockConnector {
210+ fn new ( outcomes : Vec < MockOutcome > ) -> Self {
211+ Self {
212+ outcomes : Arc :: new ( Mutex :: new ( outcomes. into ( ) ) ) ,
213+ attempts : Arc :: new ( AtomicUsize :: new ( 0 ) ) ,
214+ }
215+ }
216+
217+ fn next ( & self ) -> impl Future < Output = TransportResult < ConnectionHandle > > + Send + ' static {
218+ let outcomes = Arc :: clone ( & self . outcomes ) ;
219+ let attempts = Arc :: clone ( & self . attempts ) ;
220+ async move {
221+ attempts. fetch_add ( 1 , Ordering :: SeqCst ) ;
222+ let outcome = outcomes
223+ . lock ( )
224+ . await
225+ . pop_front ( )
226+ . expect ( "mock connector exhausted" ) ;
227+ match outcome {
228+ MockOutcome :: Ok => {
229+ let ( handle, _iface) = ConnectionHandle :: new ( ) ;
230+ Ok ( handle)
231+ }
232+ MockOutcome :: Err ( msg) => Err ( TransportErrorKind :: custom_str ( msg) ) ,
233+ }
234+ }
235+ }
236+
237+ fn attempts ( & self ) -> usize {
238+ self . attempts . load ( Ordering :: SeqCst )
239+ }
240+
241+ async fn remaining ( & self ) -> usize {
242+ self . outcomes . lock ( ) . await . len ( )
243+ }
244+ }
245+
246+ #[ derive( Default ) ]
247+ struct RecordingMetrics {
248+ disconnects : AtomicUsize ,
249+ attempts : AtomicUsize ,
250+ successes : AtomicUsize ,
251+ order : Mutex < Vec < & ' static str > > ,
252+ }
253+
254+ impl RecordingMetrics {
255+ fn counts ( & self ) -> ( usize , usize , usize ) {
256+ (
257+ self . disconnects . load ( Ordering :: SeqCst ) ,
258+ self . attempts . load ( Ordering :: SeqCst ) ,
259+ self . successes . load ( Ordering :: SeqCst ) ,
260+ )
261+ }
262+
263+ async fn events ( & self ) -> Vec < & ' static str > {
264+ self . order . lock ( ) . await . clone ( )
265+ }
266+ }
267+
268+ #[ async_trait]
269+ impl WsReconnectMetrics for RecordingMetrics {
270+ async fn on_disconnect ( & self ) {
271+ self . disconnects . fetch_add ( 1 , Ordering :: SeqCst ) ;
272+ self . order . lock ( ) . await . push ( "disconnect" ) ;
273+ }
274+
275+ async fn on_attempt ( & self ) {
276+ self . attempts . fetch_add ( 1 , Ordering :: SeqCst ) ;
277+ self . order . lock ( ) . await . push ( "attempt" ) ;
278+ }
279+
280+ async fn on_success ( & self ) {
281+ self . successes . fetch_add ( 1 , Ordering :: SeqCst ) ;
282+ self . order . lock ( ) . await . push ( "success" ) ;
283+ }
284+ }
285+
286+ #[ test]
287+ fn backoff_delay_scales_and_caps ( ) {
288+ let cfg = WebsocketReconnectConfig {
289+ initial_backoff_ms : 100 ,
290+ backoff_multiplier : 2.0 ,
291+ max_backoff_ms : 350 ,
292+ } ;
293+ let policy = WsReconnectPolicy :: from_config ( Some ( & cfg) ) ;
294+
295+ assert_eq ! ( policy. backoff_delay( 0 ) , Duration :: ZERO ) ;
296+ assert_eq ! ( policy. backoff_delay( 1 ) , Duration :: from_millis( 100 ) ) ;
297+ assert_eq ! ( policy. backoff_delay( 2 ) , Duration :: from_millis( 200 ) ) ;
298+ assert_eq ! ( policy. backoff_delay( 3 ) , Duration :: from_millis( 350 ) ) ;
299+ assert_eq ! ( policy. backoff_delay( 4 ) , Duration :: from_millis( 350 ) ) ;
300+ }
301+
302+ #[ tokio:: test]
303+ async fn reconnect_loop_retries_metrics_and_stops_after_success ( ) {
304+ let cfg = WebsocketReconnectConfig {
305+ initial_backoff_ms : 100 ,
306+ backoff_multiplier : 2.0 ,
307+ max_backoff_ms : 350 ,
308+ } ;
309+ let policy = WsReconnectPolicy :: from_config ( Some ( & cfg) ) ;
310+ let metrics = Arc :: new ( RecordingMetrics :: default ( ) ) ;
311+ let connector = MockConnector :: new ( vec ! [
312+ MockOutcome :: Err ( "fail-1" ) ,
313+ MockOutcome :: Err ( "fail-2" ) ,
314+ MockOutcome :: Ok ,
315+ MockOutcome :: Err ( "unused" ) ,
316+ ] ) ;
317+
318+ let recorded_delays = Arc :: new ( Mutex :: new ( Vec :: new ( ) ) ) ;
319+ let resilient = ResilientWsConnect :: new (
320+ WsConnect :: new ( "ws://example.invalid" ) ,
321+ policy. clone ( ) ,
322+ metrics. clone ( ) ,
323+ "testnet" ,
324+ ) ;
325+
326+ let expected_delays = [ policy. backoff_delay ( 1 ) , policy. backoff_delay ( 2 ) ] ;
327+
328+ backoff_recorder:: with_recorder ( recorded_delays. clone ( ) , async {
329+ let connector_for_task = connector. clone ( ) ;
330+ let resilient_for_task = resilient. clone ( ) ;
331+ let join = tokio:: spawn ( backoff_recorder:: with_recorder (
332+ recorded_delays. clone ( ) ,
333+ async move {
334+ resilient_for_task
335+ . reconnect_loop ( move || connector_for_task. next ( ) )
336+ . await
337+ } ,
338+ ) ) ;
339+
340+ let result = join. await . expect ( "task completed" ) ;
341+ let handle = result. expect ( "reconnect eventually succeeds" ) ;
342+ handle. shutdown ( ) ;
343+ } )
344+ . await ;
345+
346+ assert_eq ! ( connector. attempts( ) , 3 ) ;
347+ assert_eq ! ( connector. remaining( ) . await , 1 ) ;
348+
349+ let recorded = recorded_delays. lock ( ) . await . clone ( ) ;
350+ assert_eq ! ( recorded, expected_delays. to_vec( ) ) ;
351+
352+ let ( disconnects, attempts, successes) = metrics. counts ( ) ;
353+ assert_eq ! ( disconnects, 1 ) ;
354+ assert_eq ! ( attempts, 3 ) ;
355+ assert_eq ! ( successes, 1 ) ;
356+
357+ let events = metrics. events ( ) . await ;
358+ assert_eq ! (
359+ events,
360+ vec![ "disconnect" , "attempt" , "attempt" , "attempt" , "success" ]
361+ ) ;
362+ }
134363}
0 commit comments