wowlab_centrifuge/client/
mod.rs1use 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#[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#[derive(Clone)]
148pub struct Client {
149 inner: Arc<RwLock<ClientInner>>,
150 shutdown: CancellationToken,
151 run_active: Arc<AtomicBool>,
152}
153
154impl Client {
155 #[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 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 pub async fn state(&self) -> ClientState {
189 self.inner.read().await.state
190 }
191
192 pub async fn is_connected(&self) -> bool {
194 self.inner.read().await.state == ClientState::Connected
195 }
196
197 pub fn connect(&self) {
199 let client = self.clone();
200
201 tokio::spawn(async move {
202 client.run().await;
203 });
204 }
205
206 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}