Skip to main content

wowlab_centrifuge/client/
mod.rs

1// #t(file: rust_alloc_in_loop) client reconnect loop builds channel names per subscription
2
3use std::{
4    sync::{
5        Arc,
6        atomic::{AtomicBool, AtomicU32, Ordering},
7    },
8    time::{Duration, Instant},
9};
10
11use prost::Message;
12use tokio::sync::{RwLock, mpsc, oneshot};
13use tokio_util::sync::CancellationToken;
14use wowlab_types::{constants::MS_PER_SECOND_U64, sensitive::Sensitive, sim::FastMap};
15
16use crate::{
17    backoff::Backoff,
18    error::Error,
19    proto,
20    subscription::{SubscriptionInner, SubscriptionState},
21    transport::Transport,
22    types::{ClientEvent, ConnectResult},
23};
24
25mod commands;
26mod config;
27mod push;
28#[cfg(test)]
29mod tests;
30
31pub use config::ClientConfig;
32
33const MAX_SERVER_PING_DELAY: Duration = Duration::from_secs(10);
34const EVENT_CHANNEL_BUFFER: usize = 64;
35const COMMAND_CHANNEL_BUFFER: usize = 32;
36const BACKOFF_RESET_THRESHOLD: Duration = Duration::from_secs(30);
37const TOKEN_REFRESH_TIMEOUT: Duration = Duration::from_secs(10);
38
39fn check_reply(reply: proto::Reply) -> Result<proto::Reply, Error> {
40    if let Some(err) = reply.error {
41        return Err(Error::from_proto(err));
42    }
43
44    Ok(reply)
45}
46
47fn format_ttl(ttl: u32, expires: bool) -> String {
48    if expires && ttl > 0 {
49        format!("{ttl}s")
50    } else {
51        "none".to_string()
52    }
53}
54
55fn ttl_to_instant(ttl: u32, expires: bool) -> Option<Instant> {
56    (expires && ttl > 0).then(|| {
57        let ttl_ms = (u64::from(ttl) * MS_PER_SECOND_U64).min(i32::MAX as u64);
58
59        Instant::now() + Duration::from_millis(ttl_ms)
60    })
61}
62
63fn log_connection_lost(error: &Error) {
64    if error.is_benign_disconnect() {
65        tracing::debug!(%error, "Connection lost");
66    } else {
67        tracing::warn!(%error, "Connection lost");
68    }
69}
70
71fn log_disconnected() {
72    tracing::info!("Disconnected");
73}
74
75fn log_reconnect(delay: Duration) {
76    tracing::info!(?delay, "Scheduling reconnect");
77}
78
79fn log_no_ping() {
80    tracing::warn!("No ping from server; disconnecting");
81}
82
83fn log_token_refresh_failed(error: &Error) {
84    tracing::warn!(%error, "Token refresh failed");
85}
86
87fn log_command_send_failed(error: &Error) {
88    tracing::error!(%error, "Failed to send command");
89}
90
91/// Lifecycle state of a Centrifugo connection.
92#[derive(Clone, Copy, Debug, Default, Eq, PartialEq)]
93#[non_exhaustive]
94pub enum ClientState {
95    #[default]
96    Disconnected,
97    Connecting,
98    Connected,
99}
100
101type PendingRequest = oneshot::Sender<Result<proto::Reply, Error>>;
102
103struct ClientInner {
104    config: ClientConfig,
105    state: ClientState,
106    subscriptions: FastMap<String, SubscriptionInner>,
107    pending: FastMap<u32, PendingRequest>,
108    next_id: AtomicU32,
109    event_tx: mpsc::Sender<ClientEvent>,
110    cmd_tx: Option<mpsc::Sender<CommandRequest>>,
111    server_ping: Option<Duration>,
112    send_pong: bool,
113    refresh_required: bool,
114}
115
116struct CommandRequest {
117    id: u32,
118    cmd: proto::Command,
119}
120
121struct PendingRefresh {
122    id: u32,
123    token: Sensitive<String>,
124    deadline: tokio::time::Instant,
125}
126
127struct DisconnectAdvice {
128    code: u32,
129    reason: String,
130    reconnect: bool,
131}
132
133enum ConnectionExit {
134    Shutdown,
135    ServerDisconnect(DisconnectAdvice),
136}
137
138struct RunGuard(Arc<AtomicBool>);
139
140impl Drop for RunGuard {
141    fn drop(&mut self) {
142        self.0.store(false, Ordering::Release);
143    }
144}
145
146/// Centrifugo realtime client with auto-reconnect; clones share state.
147#[derive(Clone)]
148pub struct Client {
149    inner: Arc<RwLock<ClientInner>>,
150    shutdown: CancellationToken,
151    run_active: Arc<AtomicBool>,
152}
153
154impl Client {
155    /// Creates a disconnected client with shared connection state.
156    #[must_use]
157    pub fn new(config: ClientConfig) -> Self {
158        let (event_tx, _) = mpsc::channel(EVENT_CHANNEL_BUFFER);
159
160        Self {
161            inner: Arc::new(RwLock::new(ClientInner {
162                config,
163                state: ClientState::Disconnected,
164                subscriptions: FastMap::default(),
165                pending: FastMap::default(),
166                next_id: AtomicU32::new(1),
167                event_tx,
168                cmd_tx: None,
169                server_ping: None,
170                send_pong: false,
171                refresh_required: false,
172            })),
173            shutdown: CancellationToken::new(),
174            run_active: Arc::new(AtomicBool::new(false)),
175        }
176    }
177
178    /// Replaces the client event receiver and returns the new receiver; install before [`Self::connect`].
179    pub async fn events(&self) -> mpsc::Receiver<ClientEvent> {
180        let (tx, rx) = mpsc::channel(EVENT_CHANNEL_BUFFER);
181
182        self.inner.write().await.event_tx = tx;
183
184        rx
185    }
186
187    /// Returns the current connection lifecycle state.
188    pub async fn state(&self) -> ClientState {
189        self.inner.read().await.state
190    }
191
192    /// Returns whether the client completed the Centrifugo handshake.
193    pub async fn is_connected(&self) -> bool {
194        self.inner.read().await.state == ClientState::Connected
195    }
196
197    /// Spawns a background task that maintains the connection with auto-reconnect.
198    pub fn connect(&self) {
199        let client = self.clone();
200
201        tokio::spawn(async move {
202            client.run().await;
203        });
204    }
205
206    /// Maintains the connection and reconnects until [`Self::disconnect`] is called.
207    pub async fn run(&self) {
208        if self
209            .run_active
210            .compare_exchange(false, true, Ordering::AcqRel, Ordering::Acquire)
211            .is_err()
212        {
213            tracing::debug!("Connection loop is already running");
214
215            return;
216        }
217
218        let _run_guard = RunGuard(Arc::clone(&self.run_active));
219        let mut backoff = {
220            let inner = self.inner.read().await;
221
222            Backoff::new(
223                inner.config.min_reconnect_delay,
224                inner.config.max_reconnect_delay,
225            )
226        };
227
228        loop {
229            if self.shutdown.is_cancelled() {
230                return;
231            }
232
233            self.inner.write().await.state = ClientState::Connecting;
234            self.emit_client_event(ClientEvent::Connecting).await;
235
236            let start = Instant::now();
237
238            let disconnect = match self.connect_and_run().await {
239                Ok(ConnectionExit::Shutdown) => {
240                    log_disconnected();
241
242                    DisconnectAdvice {
243                        code: 0,
244                        reason: "Disconnected".to_string(),
245                        reconnect: false,
246                    }
247                }
248                Ok(ConnectionExit::ServerDisconnect(disconnect)) => {
249                    log_disconnected();
250
251                    disconnect
252                }
253                Err(error) => {
254                    log_connection_lost(&error);
255
256                    if error.requires_token_refresh() {
257                        self.inner.write().await.refresh_required = true;
258                    }
259
260                    self.emit_client_event(ClientEvent::Error(error.to_string()))
261                        .await;
262
263                    DisconnectAdvice {
264                        code: 0,
265                        reason: "Connection closed".to_string(),
266                        reconnect: true,
267                    }
268                }
269            };
270
271            let disconnected_event = ClientEvent::Disconnected {
272                code: disconnect.code,
273                reason: disconnect.reason,
274                reconnect: disconnect.reconnect,
275            };
276
277            let mut inner = self.inner.write().await;
278
279            inner.state = ClientState::Disconnected;
280            inner.cmd_tx = None;
281
282            inner.pending.drain().for_each(|(_, reply_tx)| {
283                let _ = reply_tx.send(Err(Error::connection_closed()));
284            });
285
286            inner
287                .subscriptions
288                .values_mut()
289                .for_each(|sub| sub.state = SubscriptionState::Unsubscribed);
290            drop(inner);
291            self.emit_client_event(disconnected_event).await;
292
293            if self.shutdown.is_cancelled() || !disconnect.reconnect {
294                return;
295            }
296
297            if start.elapsed() > BACKOFF_RESET_THRESHOLD {
298                backoff.reset();
299            }
300
301            let delay = backoff.next_delay();
302
303            log_reconnect(delay);
304
305            tokio::select! {
306                () = tokio::time::sleep(delay) => {}
307                () = self.shutdown.cancelled() => return,
308            }
309
310            tokio::task::yield_now().await;
311        }
312    }
313
314    async fn emit_client_event(&self, event: ClientEvent) {
315        let event_tx = self.inner.read().await.event_tx.clone();
316
317        tokio::select! {
318            biased;
319            () = self.shutdown.cancelled() => {
320                if let Err(error) = event_tx.try_send(event) {
321                    tracing::debug!(%error, "Client event receiver unavailable during shutdown");
322                }
323            }
324            result = event_tx.send(event.clone()) => {
325                if let Err(error) = result {
326                    tracing::debug!(%error, "Client event receiver closed");
327                }
328            }
329        }
330    }
331
332    async fn emit_subscription_event(&self, channel: &str, event: crate::types::SubscriptionEvent) {
333        let event_tx = {
334            let inner = self.inner.read().await;
335
336            inner
337                .subscriptions
338                .get(channel)
339                .map(|sub| sub.event_tx.clone())
340        };
341
342        let Some(event_tx) = event_tx else {
343            return;
344        };
345
346        tokio::select! {
347            biased;
348            () = self.shutdown.cancelled() => {}
349            result = event_tx.send(event) => {
350                if let Err(error) = result {
351                    tracing::debug!(%channel, %error, "Subscription event receiver closed");
352                }
353            }
354        }
355    }
356
357    async fn connection_config(&self) -> Result<Option<ClientConfig>, Error> {
358        let mut config = self.inner.read().await.config.clone();
359        let refresh_required = self.inner.read().await.refresh_required;
360
361        if !refresh_required {
362            return Ok(Some(config));
363        }
364
365        let Some(ref get_token) = config.get_token else {
366            return Ok(Some(config));
367        };
368
369        let token_result = tokio::select! {
370            biased;
371            () = self.shutdown.cancelled() => return Ok(None),
372            result = tokio::time::timeout(TOKEN_REFRESH_TIMEOUT, get_token()) => {
373                result.map_err(|_elapsed| Error::protocol("Token callback timed out"))?
374            }
375        };
376
377        match token_result {
378            Ok(new_token) => {
379                config.token = new_token.clone();
380                let mut inner = self.inner.write().await;
381
382                inner.config.token = new_token;
383                inner.refresh_required = false;
384            }
385            Err(error) => {
386                tracing::error!(%error, "Token refresh failed");
387
388                return Err(error);
389            }
390        }
391
392        Ok(Some(config))
393    }
394
395    async fn establish_connection(
396        &self,
397        config: &ClientConfig,
398    ) -> Result<(Transport, proto::ConnectResult), Error> {
399        let mut transport = Transport::connect(&config.url).await?;
400        let subs = self.build_recovery_subs().await;
401        let connect_req = proto::ConnectRequest {
402            token: config.token.expose().clone(),
403            name: config.name.clone(),
404            version: config.version.clone(),
405            data: config.data.clone().unwrap_or_default(),
406            headers: config.headers.clone().into_iter().collect(),
407            subs: subs.into_iter().collect(),
408            ..Default::default()
409        };
410        let cmd = proto::Command {
411            id: 1,
412            connect: Some(connect_req),
413            ..Default::default()
414        };
415        let reply = check_reply(transport.send_command(cmd).await?)?;
416        let connect_result = reply
417            .connect
418            .ok_or_else(|| Error::protocol("No connect result"))?;
419
420        Ok((transport, connect_result))
421    }
422
423    async fn connect_and_run(&self) -> Result<ConnectionExit, Error> {
424        let Some(config) = self.connection_config().await? else {
425            return Ok(ConnectionExit::Shutdown);
426        };
427
428        let (mut transport, connect_result) = self.establish_connection(&config).await?;
429
430        let server_ping =
431            (connect_result.ping > 0).then(|| Duration::from_secs(u64::from(connect_result.ping)));
432        let send_pong = connect_result.pong;
433
434        let refresh_at = ttl_to_instant(connect_result.ttl, connect_result.expires);
435        let ttl_display = format_ttl(connect_result.ttl, connect_result.expires);
436        let ping_display =
437            server_ping.map_or_else(|| "disabled".to_string(), |d| format!("{}s", d.as_secs()));
438
439        tracing::info!(ping = %ping_display, token_ttl = %ttl_display, "Connected");
440
441        let server_subs = connect_result.subs.clone();
442
443        self.process_server_subs(server_subs).await;
444
445        let (cmd_tx, mut cmd_rx) = mpsc::channel::<CommandRequest>(COMMAND_CHANNEL_BUFFER);
446
447        let connected_event = ClientEvent::Connected(ConnectResult::from(connect_result));
448
449        let mut inner = self.inner.write().await;
450
451        inner.state = ClientState::Connected;
452        inner.cmd_tx = Some(cmd_tx);
453        inner.server_ping = server_ping;
454        inner.send_pong = send_pong;
455        drop(inner);
456        self.emit_client_event(connected_event).await;
457
458        self.resubscribe_all().await;
459
460        let ping_timeout = server_ping.map(|p| p + MAX_SERVER_PING_DELAY);
461        let mut last_data = Instant::now();
462        let mut refresh_at = refresh_at;
463        let mut pending_refresh: Option<PendingRefresh> = None;
464
465        let exit = loop {
466            let timeout_future = async {
467                if let Some(timeout) = ping_timeout {
468                    tokio::time::sleep_until((last_data + timeout).into()).await;
469                } else {
470                    std::future::pending::<()>().await;
471                }
472            };
473
474            let refresh_future = async {
475                if let Some(at) = refresh_at {
476                    tokio::time::sleep_until(at.into()).await;
477                } else {
478                    std::future::pending::<()>().await;
479                }
480            };
481
482            let refresh_timeout_future = async {
483                if let Some(refresh) = &pending_refresh {
484                    tokio::time::sleep_until(refresh.deadline).await;
485                } else {
486                    std::future::pending::<()>().await;
487                }
488            };
489
490            tokio::select! {
491                () = self.shutdown.cancelled() => {
492                    break ConnectionExit::Shutdown;
493                }
494
495                () = timeout_future => {
496                    log_no_ping();
497                    return Err(Error::no_ping());
498                }
499
500                () = refresh_future => {
501                    match self.begin_refresh(&mut transport).await {
502                        Ok(Some(refresh)) => {
503                            pending_refresh = Some(refresh);
504                            refresh_at = None;
505                        }
506                        Ok(None) => refresh_at = None,
507                        Err(error) => {
508                            log_token_refresh_failed(&error);
509                            self.inner.write().await.refresh_required = true;
510
511                            return Err(error);
512                        }
513                    }
514                }
515
516                () = refresh_timeout_future => {
517                    self.inner.write().await.refresh_required = true;
518
519                    return Err(Error::protocol("Token refresh response timed out"));
520                }
521
522                Some(req) = cmd_rx.recv() => {
523                    let mut cmd = req.cmd;
524                    cmd.id = req.id;
525
526                    let data = cmd.encode_length_delimited_to_vec();
527                    if let Err(error) = transport.send_raw(data).await {
528                        log_command_send_failed(&error);
529
530                        return Err(error);
531                    }
532                }
533
534                result = transport.read_message() => {
535                    let reply = result?;
536                    last_data = Instant::now();
537
538                    if pending_refresh.as_ref().is_some_and(|refresh| refresh.id == reply.id) {
539                        let refresh = pending_refresh.take().ok_or_else(|| {
540                            Error::protocol("Refresh response arrived without a pending refresh")
541                        })?;
542
543                        match self.finish_refresh(refresh, reply).await {
544                            Ok(new_refresh_at) => refresh_at = new_refresh_at,
545                            Err(error) => {
546                                log_token_refresh_failed(&error);
547                                self.inner.write().await.refresh_required = true;
548
549                                return Err(error);
550                            }
551                        }
552                    } else if let Some(disconnect) = self.handle_reply(reply, &mut transport).await? {
553                        break ConnectionExit::ServerDisconnect(disconnect);
554                    }
555                }
556            }
557        };
558
559        transport.close().await;
560
561        Ok(exit)
562    }
563
564    async fn begin_refresh(
565        &self,
566        transport: &mut Transport,
567    ) -> Result<Option<PendingRefresh>, Error> {
568        let get_token = {
569            let inner = self.inner.read().await;
570
571            inner.config.get_token.clone()
572        };
573
574        let Some(get_token) = get_token else {
575            tracing::warn!("Token refresh requested but no get_token callback configured");
576
577            return Ok(None);
578        };
579
580        let new_token = tokio::select! {
581            biased;
582            () = self.shutdown.cancelled() => return Err(Error::connection_closed()),
583            result = tokio::time::timeout(TOKEN_REFRESH_TIMEOUT, get_token()) => {
584                result
585                    .map_err(|_elapsed| Error::protocol("Token callback timed out"))??
586            }
587        };
588
589        let refresh_req = proto::RefreshRequest {
590            token: new_token.expose().clone(),
591        };
592
593        let id = self
594            .inner
595            .read()
596            .await
597            .next_id
598            .fetch_add(1, Ordering::Relaxed);
599        let cmd = proto::Command {
600            id,
601            refresh: Some(refresh_req),
602            ..Default::default()
603        };
604
605        let data = cmd.encode_length_delimited_to_vec();
606
607        self.inner.write().await.refresh_required = true;
608        transport.send_raw(data).await?;
609
610        Ok(Some(PendingRefresh {
611            id,
612            token: new_token,
613            deadline: tokio::time::Instant::now() + TOKEN_REFRESH_TIMEOUT,
614        }))
615    }
616
617    async fn finish_refresh(
618        &self,
619        refresh: PendingRefresh,
620        reply: proto::Reply,
621    ) -> Result<Option<Instant>, Error> {
622        let reply = check_reply(reply)?;
623
624        let result = reply
625            .refresh
626            .ok_or_else(|| Error::protocol("No refresh result"))?;
627
628        {
629            let mut inner = self.inner.write().await;
630
631            inner.config.token = refresh.token;
632            inner.refresh_required = false;
633        };
634
635        let ttl_display = format_ttl(result.ttl, result.expires);
636
637        tracing::info!(next_ttl = %ttl_display, "Token refreshed");
638
639        Ok(ttl_to_instant(result.ttl, result.expires))
640    }
641}
642impl std::fmt::Debug for Client {
643    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
644        f.debug_struct("Client")
645            .field("inner", &"<ClientInner>")
646            .field("shutdown", &"<CancellationToken>")
647            .finish()
648    }
649}