1use std::task::{Poll, ready};
2use std::time::Duration;
3
4use moq_net::Version;
5use moq_net::bandwidth::{Consumer as BandwidthConsumer, Producer as BandwidthProducer};
6use moq_net::kio;
7use rand::RngExt;
8use url::Url;
9
10use crate::{Client, Error, RedactedUrl};
11
12#[derive(Clone, Debug, clap::Args, serde::Serialize, serde::Deserialize)]
23#[serde(default, deny_unknown_fields)]
24#[non_exhaustive]
25pub struct Backoff {
26 #[arg(
28 id = "backoff-initial",
29 long,
30 default_value = "1s",
31 env = "MOQ_BACKOFF_INITIAL",
32 value_parser = humantime::parse_duration,
33 )]
34 #[serde(with = "humantime_serde")]
35 pub initial: Duration,
36
37 #[arg(id = "backoff-multiplier", long, default_value_t = 2, env = "MOQ_BACKOFF_MULTIPLIER")]
39 pub multiplier: u32,
40
41 #[arg(
43 id = "backoff-max",
44 long,
45 default_value = "5s",
46 env = "MOQ_BACKOFF_MAX",
47 value_parser = humantime::parse_duration,
48 )]
49 #[serde(with = "humantime_serde")]
50 pub max: Duration,
51
52 #[arg(
57 id = "backoff-timeout",
58 long,
59 default_value = "10s",
60 env = "MOQ_BACKOFF_TIMEOUT",
61 value_parser = humantime::parse_duration,
62 )]
63 #[serde(with = "humantime_serde")]
64 pub timeout: Duration,
65}
66
67impl Default for Backoff {
68 fn default() -> Self {
69 Self {
70 initial: Duration::from_secs(1),
71 multiplier: 2,
72 max: Duration::from_secs(5),
73 timeout: Duration::from_secs(10),
74 }
75 }
76}
77
78impl Backoff {
79 pub(crate) fn validate(&self) -> crate::Result<()> {
90 match self.initial.is_zero() || self.multiplier == 0 || self.max.is_zero() {
91 true => Err(crate::Error::BackoffUnpaced),
92 false => Ok(()),
93 }
94 }
95
96 fn next_delay(&self, delay: Duration) -> Duration {
98 delay.saturating_mul(self.multiplier.max(1)).min(self.max)
99 }
100
101 pub fn linger(&self) -> Duration {
107 match self.timeout.is_zero() {
108 true => Duration::MAX,
109 false => self.timeout.saturating_add(Duration::from_secs(1)),
110 }
111 }
112}
113
114fn deadline_from(backoff: &Backoff) -> Option<tokio::time::Instant> {
117 match backoff.timeout.is_zero() {
118 true => None,
119 false => Some(tokio::time::Instant::now() + backoff.timeout),
120 }
121}
122
123#[derive(Clone, Copy, Debug, PartialEq, Eq)]
125#[non_exhaustive]
126pub enum Status {
127 Connected,
129 Disconnected,
131}
132
133#[derive(Default)]
138struct State {
139 status: Option<Status>,
141 presence: moq_net::stats::Presence,
144 version: Option<Version>,
146 error: Option<Error>,
149 session: Option<moq_net::Session>,
152}
153
154#[non_exhaustive]
156pub struct ConnectionSnapshot {
157 pub stats: moq_net::ConnectionStats,
159 pub version: Version,
161}
162
163#[derive(Clone)]
168pub struct ConnectionStatsReader {
169 state: kio::Consumer<State>,
170 last_presence: moq_net::stats::Presence,
171}
172
173impl ConnectionStatsReader {
174 pub fn presence(&self) -> moq_net::stats::Presence {
178 self.state.read().presence
179 }
180
181 pub fn poll_presence(&mut self, waiter: &kio::Waiter) -> Poll<crate::Result<moq_net::stats::Presence>> {
183 let last = self.last_presence;
184 let presence = match ready!(self.state.poll(waiter, |state| match state.presence {
185 presence if presence != last => Poll::Ready(presence),
186 _ => Poll::Pending,
187 })) {
188 Ok(presence) => presence,
189 Err(state) => return Poll::Ready(Err(terminal(&state))),
190 };
191
192 self.last_presence = presence;
193 Poll::Ready(Ok(presence))
194 }
195
196 pub async fn presence_changed(&mut self) -> crate::Result<moq_net::stats::Presence> {
201 kio::wait(|waiter| self.poll_presence(waiter)).await
202 }
203
204 pub fn stats(&self) -> Option<moq_net::ConnectionStats> {
206 self.state.read().session.as_ref().map(moq_net::Session::stats)
207 }
208
209 pub fn snapshot(&self) -> Option<ConnectionSnapshot> {
211 let state = self.state.read();
212 let session = state.session.as_ref()?;
213 Some(ConnectionSnapshot {
214 stats: session.stats(),
215 version: session.version(),
216 })
217 }
218}
219
220pub struct Reconnect {
234 abort: tokio::task::AbortHandle,
235 state: kio::Consumer<State>,
236 send_bandwidth: BandwidthConsumer,
238 recv_bandwidth: BandwidthConsumer,
240 last_reported: Option<Status>,
242}
243
244impl Reconnect {
245 pub(crate) fn new(client: Client, url: Url, backoff: Backoff) -> Self {
246 let producer = kio::Producer::<State>::default();
247 let state = producer.consume();
248
249 let send_bw = BandwidthProducer::new();
252 let recv_bw = BandwidthProducer::new();
253 let send_bandwidth = send_bw.consume();
254 let recv_bandwidth = recv_bw.consume();
255
256 let task = tokio::spawn(async move {
257 if let Err(err) = Self::run(&producer, &send_bw, &recv_bw, client, url, backoff).await {
258 tracing::error!(%err, "reconnect loop exited");
259 if let Ok(mut state) = producer.write() {
260 state.error = Some(err);
261 }
262 }
263 });
265 Self {
266 abort: task.abort_handle(),
267 state,
268 send_bandwidth,
269 recv_bandwidth,
270 last_reported: None,
271 }
272 }
273
274 async fn run(
275 state: &kio::Producer<State>,
276 send_bw: &BandwidthProducer,
277 recv_bw: &BandwidthProducer,
278 client: Client,
279 url: Url,
280 backoff: Backoff,
281 ) -> crate::Result<()> {
282 let mut delay = backoff.initial;
286 let mut deadline = deadline_from(&backoff);
287 let mut last_error: Option<Error> = None;
288
289 let url_log = RedactedUrl::new(&url);
292
293 loop {
294 tracing::info!(url = %url_log, "connecting");
295
296 match client.connect(url.clone()).await {
297 Ok(session) => {
298 tracing::info!(url = %url_log, "connected");
299 if let Ok(mut state) = state.write() {
300 state.presence.sessions += 1;
301 state.status = Some(Status::Connected);
302 state.version = Some(session.version());
303 state.session = Some(session.clone());
304 }
305
306 let connected = tokio::time::Instant::now();
307 let closed = run_session(send_bw, recv_bw, &session).await;
310 if let Ok(mut state) = state.write() {
311 state.presence.sessions_closed += 1;
312 state.status = Some(Status::Disconnected);
313 state.version = None;
314 state.session = None;
315 }
316 let _ = send_bw.set(None);
318 let _ = recv_bw.set(None);
319
320 if connected.elapsed() >= backoff.initial {
321 tracing::warn!(url = %url_log, "session closed, reconnecting");
324 delay = backoff.initial;
325 deadline = deadline_from(&backoff);
326 last_error = None;
327 } else {
328 if let Err(err) = closed {
333 let err = Error::from(err);
334 tracing::warn!(url = %url_log, %err, "session severed immediately, retrying");
335 last_error = Some(err);
336 } else {
337 tracing::warn!(url = %url_log, "session severed immediately, retrying");
338 }
339 }
340 }
341 Err(err) => {
342 if err.is_auth() {
347 return Err(err);
348 }
349 if let Some(status) = err.status()
350 && !crate::error::status_retryable(status)
351 {
352 return Err(err);
353 }
354 last_error = Some(err);
355 }
356 }
357
358 let now = tokio::time::Instant::now();
359 if deadline.is_some_and(|deadline| now >= deadline) {
360 let timeout = backoff.timeout;
361 let msg = match last_error {
362 Some(err) => format!("reconnect timed out after {timeout:?}: {err}"),
363 None => format!("reconnect timed out after {timeout:?}"),
364 };
365 return Err(Error::Reconnect(msg));
366 }
367
368 let mut wait = delay.mul_f64(0.5 + rand::rng().random::<f64>() / 2.0);
371 if let Some(deadline) = deadline {
372 wait = wait.min(deadline - now);
373 }
374 delay = backoff.next_delay(delay);
375
376 tracing::warn!(url = %url_log, ?wait, "reconnecting after backoff");
377 tokio::time::sleep(wait).await;
378 }
379 }
380
381 pub fn poll_status(&mut self, waiter: &kio::Waiter) -> Poll<crate::Result<Status>> {
386 let last = self.last_reported;
387 let status = match ready!(self.state.poll(waiter, |state| match state.status {
388 Some(status) if Some(status) != last => Poll::Ready(status),
389 _ => Poll::Pending,
390 })) {
391 Ok(status) => status,
392 Err(state) => return Poll::Ready(Err(terminal(&state))),
393 };
394
395 self.last_reported = Some(status);
396 Poll::Ready(Ok(status))
397 }
398
399 pub async fn status(&mut self) -> crate::Result<Status> {
405 kio::wait(|waiter| self.poll_status(waiter)).await
406 }
407
408 pub fn connected(&self) -> bool {
413 self.state.read().status == Some(Status::Connected)
414 }
415
416 pub fn version(&self) -> Option<Version> {
421 self.state.read().version
422 }
423
424 pub fn send_bandwidth(&self) -> BandwidthConsumer {
431 self.send_bandwidth.clone()
432 }
433
434 pub fn recv_bandwidth(&self) -> BandwidthConsumer {
438 self.recv_bandwidth.clone()
439 }
440
441 pub fn poll_closed(&self, waiter: &kio::Waiter) -> Poll<crate::Result<()>> {
447 ready!(self.state.poll_closed(waiter));
448 Poll::Ready(match &self.state.read().error {
449 Some(err) => Err(err.clone()),
450 None => Ok(()),
451 })
452 }
453
454 pub async fn closed(&self) -> crate::Result<()> {
456 kio::wait(|waiter| self.poll_closed(waiter)).await
457 }
458
459 pub fn stats(&self) -> ConnectionStatsReader {
463 ConnectionStatsReader {
464 state: self.state.clone(),
465 last_presence: moq_net::stats::Presence::default(),
466 }
467 }
468}
469
470async fn run_session(
478 send_bw: &BandwidthProducer,
479 recv_bw: &BandwidthProducer,
480 session: &moq_net::Session,
481) -> Result<(), moq_net::Error> {
482 let mut send = session.send_bandwidth();
483 let mut recv = session.recv_bandwidth();
484 let closed = session.closed();
485 tokio::pin!(closed);
486
487 let err = kio::wait(|waiter| {
488 poll_forward(&mut send, send_bw, waiter);
489 poll_forward(&mut recv, recv_bw, waiter);
490 waiter.poll_future(closed.as_mut())
491 })
492 .await;
493
494 Err(err)
495}
496
497fn poll_forward(bw: &mut Option<BandwidthConsumer>, out: &BandwidthProducer, waiter: &kio::Waiter) {
505 loop {
506 let Some(consumer) = bw.as_mut() else { return };
507 let Poll::Ready(res) = consumer.poll_changed(waiter) else {
508 return;
509 };
510 match res {
511 Ok(rate) => {
512 let _ = out.set(rate);
513 }
514 Err(_) => {
515 *bw = None;
516 return;
517 }
518 }
519 }
520}
521
522impl Drop for Reconnect {
523 fn drop(&mut self) {
524 self.abort.abort();
525 }
526}
527
528fn terminal(state: &State) -> Error {
530 match &state.error {
531 Some(err) => err.clone(),
532 None => Error::Reconnect("reconnect stopped".to_string()),
533 }
534}
535
536#[cfg(test)]
537mod tests {
538 #[tokio::test]
539 async fn snapshot_uses_one_live_session() {
540 let mut config = crate::ServerConfig {
541 bind: Some("[::]:0".into()),
542 ..Default::default()
543 };
544 config.tls.generate = vec!["localhost".into()];
545 let mut server = config.init().unwrap();
546 let url = format!("moqt://localhost:{}", server.local_addr().unwrap().port())
547 .parse()
548 .unwrap();
549 let mut config = crate::ClientConfig::default();
550 config.tls.disable_verify = Some(true);
551 let client = config.init().unwrap();
552 let (accepted, connected) = tokio::time::timeout(Duration::from_secs(10), async {
553 tokio::join!(
554 async { server.accept().await.unwrap().ok().await.unwrap() },
555 client.connect(url)
556 )
557 })
558 .await
559 .unwrap();
560 let session = connected.unwrap();
561 let version = session.version();
562 let producer = kio::Producer::<State>::default();
563 let reader = ConnectionStatsReader {
564 state: producer.consume(),
565 last_presence: Default::default(),
566 };
567 assert_eq!(reader.presence().active(), 0);
568 assert!(reader.snapshot().is_none());
569 {
570 let mut state = producer.write().ok().unwrap();
571 state.presence.sessions = 2;
572 state.presence.sessions_closed = 1;
573 state.session = Some(session);
574 }
575 assert_eq!(reader.presence().sessions, 2);
576 assert_eq!(reader.presence().active(), 1);
577 assert_eq!(reader.snapshot().unwrap().version, version);
579 producer.write().ok().unwrap().session = None;
580 assert!(reader.snapshot().is_none());
581 drop(accepted);
582 }
583
584 #[tokio::test]
587 async fn presence_change_survives_coalesced_status() {
588 let producer = kio::Producer::<State>::default();
589 let mut reader = ConnectionStatsReader {
590 state: producer.consume(),
591 last_presence: Default::default(),
592 };
593 {
594 let mut state = producer.write().ok().unwrap();
595 state.presence.sessions = 1;
596 state.status = Some(Status::Connected);
597 }
598 {
599 let mut state = producer.write().ok().unwrap();
600 state.presence.sessions_closed = 1;
601 state.status = Some(Status::Disconnected);
602 }
603 let presence = reader.presence_changed().await.unwrap();
604 assert_eq!((presence.sessions, presence.sessions_closed), (1, 1));
605 assert_eq!(presence.active(), 0);
606 }
607
608 #[test]
612 fn backoff_rejects_an_unpaced_retry() {
613 assert!(Backoff::default().validate().is_ok());
614
615 for bad in [
616 Backoff {
617 initial: Duration::ZERO,
618 ..Default::default()
619 },
620 Backoff {
621 multiplier: 0,
622 ..Default::default()
623 },
624 Backoff {
625 max: Duration::ZERO,
626 ..Default::default()
627 },
628 ] {
629 assert!(
630 matches!(bad.validate(), Err(crate::Error::BackoffUnpaced)),
631 "{bad:?} should be rejected"
632 );
633 }
634
635 let forever = Backoff {
638 timeout: Duration::ZERO,
639 multiplier: 1,
640 ..Default::default()
641 };
642 assert!(forever.validate().is_ok());
643 }
644
645 #[test]
646 fn backoff_growth_saturates_before_applying_the_cap() {
647 let backoff = Backoff {
648 multiplier: u32::MAX,
649 ..Default::default()
650 };
651 assert_eq!(backoff.next_delay(Duration::MAX), backoff.max);
652 }
653
654 use super::*;
655
656 #[test]
657 fn test_backoff_default() {
658 let backoff = Backoff::default();
659 assert_eq!(backoff.initial, Duration::from_secs(1));
660 assert_eq!(backoff.multiplier, 2);
661 assert_eq!(backoff.max, Duration::from_secs(5));
662 assert_eq!(backoff.timeout, Duration::from_secs(10));
663 }
664
665 #[test]
668 fn test_backoff_linger() {
669 let backoff = Backoff::default();
670 assert_eq!(backoff.linger(), backoff.timeout + Duration::from_secs(1));
671
672 let unlimited = Backoff {
673 timeout: Duration::ZERO,
674 ..Backoff::default()
675 };
676 assert_eq!(unlimited.linger(), Duration::MAX);
677 }
678
679 #[test]
680 fn poll_forward_mirrors_until_the_source_closes() {
681 let src = BandwidthProducer::new();
682 let out = BandwidthProducer::new();
683 let out_rx = out.consume();
684 let waiter = kio::Waiter::noop();
685
686 let mut bw = Some(src.consume());
688 poll_forward(&mut bw, &out, &waiter);
689 assert_eq!(out_rx.peek(), None);
690 assert!(bw.is_some());
691
692 src.set(Some(3_000)).unwrap();
694 poll_forward(&mut bw, &out, &waiter);
695 assert_eq!(out_rx.peek(), Some(3_000));
696
697 src.set(None).unwrap();
700 poll_forward(&mut bw, &out, &waiter);
701 assert_eq!(out_rx.peek(), None);
702 assert!(bw.is_some());
703
704 src.set(Some(9_000)).unwrap();
708 poll_forward(&mut bw, &out, &waiter);
709 assert_eq!(out_rx.peek(), Some(9_000));
710
711 src.abort(moq_net::Error::Cancel).unwrap();
713 poll_forward(&mut bw, &out, &waiter);
714 assert!(bw.is_none());
715 }
716}