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 version: Option<Version>,
143 error: Option<Error>,
146 session: Option<moq_net::Session>,
149}
150
151#[non_exhaustive]
153pub struct ConnectionSnapshot {
154 pub stats: moq_net::ConnectionStats,
156 pub version: Version,
158}
159
160#[derive(Clone)]
165pub struct ConnectionStatsReader {
166 state: kio::Consumer<State>,
167}
168
169impl ConnectionStatsReader {
170 pub fn stats(&self) -> Option<moq_net::ConnectionStats> {
172 self.state.read().session.as_ref().map(moq_net::Session::stats)
173 }
174
175 pub fn snapshot(&self) -> Option<ConnectionSnapshot> {
177 let state = self.state.read();
178 let session = state.session.as_ref()?;
179 Some(ConnectionSnapshot {
180 stats: session.stats(),
181 version: session.version(),
182 })
183 }
184}
185
186pub struct Reconnect {
200 abort: tokio::task::AbortHandle,
201 state: kio::Consumer<State>,
202 send_bandwidth: BandwidthConsumer,
204 recv_bandwidth: BandwidthConsumer,
206 last_reported: Option<Status>,
208}
209
210impl Reconnect {
211 pub(crate) fn new(client: Client, url: Url, backoff: Backoff) -> Self {
212 let producer = kio::Producer::<State>::default();
213 let state = producer.consume();
214
215 let send_bw = BandwidthProducer::new();
218 let recv_bw = BandwidthProducer::new();
219 let send_bandwidth = send_bw.consume();
220 let recv_bandwidth = recv_bw.consume();
221
222 let task = tokio::spawn(async move {
223 if let Err(err) = Self::run(&producer, &send_bw, &recv_bw, client, url, backoff).await {
224 tracing::error!(%err, "reconnect loop exited");
225 if let Ok(mut state) = producer.write() {
226 state.error = Some(err);
227 }
228 }
229 });
231 Self {
232 abort: task.abort_handle(),
233 state,
234 send_bandwidth,
235 recv_bandwidth,
236 last_reported: None,
237 }
238 }
239
240 async fn run(
241 state: &kio::Producer<State>,
242 send_bw: &BandwidthProducer,
243 recv_bw: &BandwidthProducer,
244 client: Client,
245 url: Url,
246 backoff: Backoff,
247 ) -> crate::Result<()> {
248 let mut delay = backoff.initial;
252 let mut deadline = deadline_from(&backoff);
253 let mut last_error: Option<Error> = None;
254
255 let url_log = RedactedUrl::new(&url);
258
259 loop {
260 tracing::info!(url = %url_log, "connecting");
261
262 match client.connect(url.clone()).await {
263 Ok(session) => {
264 tracing::info!(url = %url_log, "connected");
265 if let Ok(mut state) = state.write() {
266 state.status = Some(Status::Connected);
267 state.version = Some(session.version());
268 state.session = Some(session.clone());
269 }
270
271 let connected = tokio::time::Instant::now();
272 let closed = run_session(send_bw, recv_bw, &session).await;
275 if let Ok(mut state) = state.write() {
276 state.status = Some(Status::Disconnected);
277 state.version = None;
278 state.session = None;
279 }
280 let _ = send_bw.set(None);
282 let _ = recv_bw.set(None);
283
284 if connected.elapsed() >= backoff.initial {
285 tracing::warn!(url = %url_log, "session closed, reconnecting");
288 delay = backoff.initial;
289 deadline = deadline_from(&backoff);
290 last_error = None;
291 } else {
292 if let Err(err) = closed {
297 let err = Error::from(err);
298 tracing::warn!(url = %url_log, %err, "session severed immediately, retrying");
299 last_error = Some(err);
300 } else {
301 tracing::warn!(url = %url_log, "session severed immediately, retrying");
302 }
303 }
304 }
305 Err(err) => {
306 if err.is_auth() {
311 return Err(err);
312 }
313 if let Some(status) = err.status()
314 && !crate::error::status_retryable(status)
315 {
316 return Err(err);
317 }
318 last_error = Some(err);
319 }
320 }
321
322 let now = tokio::time::Instant::now();
323 if deadline.is_some_and(|deadline| now >= deadline) {
324 let timeout = backoff.timeout;
325 let msg = match last_error {
326 Some(err) => format!("reconnect timed out after {timeout:?}: {err}"),
327 None => format!("reconnect timed out after {timeout:?}"),
328 };
329 return Err(Error::Reconnect(msg));
330 }
331
332 let mut wait = delay.mul_f64(0.5 + rand::rng().random::<f64>() / 2.0);
335 if let Some(deadline) = deadline {
336 wait = wait.min(deadline - now);
337 }
338 delay = backoff.next_delay(delay);
339
340 tracing::warn!(url = %url_log, ?wait, "reconnecting after backoff");
341 tokio::time::sleep(wait).await;
342 }
343 }
344
345 pub fn poll_status(&mut self, waiter: &kio::Waiter) -> Poll<crate::Result<Status>> {
350 let last = self.last_reported;
351 let status = match ready!(self.state.poll(waiter, |state| match state.status {
352 Some(status) if Some(status) != last => Poll::Ready(status),
353 _ => Poll::Pending,
354 })) {
355 Ok(status) => status,
356 Err(state) => return Poll::Ready(Err(terminal(&state))),
357 };
358
359 self.last_reported = Some(status);
360 Poll::Ready(Ok(status))
361 }
362
363 pub async fn status(&mut self) -> crate::Result<Status> {
369 kio::wait(|waiter| self.poll_status(waiter)).await
370 }
371
372 pub fn connected(&self) -> bool {
377 self.state.read().status == Some(Status::Connected)
378 }
379
380 pub fn version(&self) -> Option<Version> {
385 self.state.read().version
386 }
387
388 pub fn send_bandwidth(&self) -> BandwidthConsumer {
395 self.send_bandwidth.clone()
396 }
397
398 pub fn recv_bandwidth(&self) -> BandwidthConsumer {
402 self.recv_bandwidth.clone()
403 }
404
405 pub fn poll_closed(&self, waiter: &kio::Waiter) -> Poll<crate::Result<()>> {
411 ready!(self.state.poll_closed(waiter));
412 Poll::Ready(match &self.state.read().error {
413 Some(err) => Err(err.clone()),
414 None => Ok(()),
415 })
416 }
417
418 pub async fn closed(&self) -> crate::Result<()> {
420 kio::wait(|waiter| self.poll_closed(waiter)).await
421 }
422
423 pub fn stats(&self) -> ConnectionStatsReader {
427 ConnectionStatsReader {
428 state: self.state.clone(),
429 }
430 }
431}
432
433async fn run_session(
441 send_bw: &BandwidthProducer,
442 recv_bw: &BandwidthProducer,
443 session: &moq_net::Session,
444) -> Result<(), moq_net::Error> {
445 let mut send = session.send_bandwidth();
446 let mut recv = session.recv_bandwidth();
447 let closed = session.closed();
448 tokio::pin!(closed);
449
450 let err = kio::wait(|waiter| {
451 poll_forward(&mut send, send_bw, waiter);
452 poll_forward(&mut recv, recv_bw, waiter);
453 waiter.poll_future(closed.as_mut())
454 })
455 .await;
456
457 Err(err)
458}
459
460fn poll_forward(bw: &mut Option<BandwidthConsumer>, out: &BandwidthProducer, waiter: &kio::Waiter) {
468 loop {
469 let Some(consumer) = bw.as_mut() else { return };
470 let Poll::Ready(res) = consumer.poll_changed(waiter) else {
471 return;
472 };
473 match res {
474 Ok(rate) => {
475 let _ = out.set(rate);
476 }
477 Err(_) => {
478 *bw = None;
479 return;
480 }
481 }
482 }
483}
484
485impl Drop for Reconnect {
486 fn drop(&mut self) {
487 self.abort.abort();
488 }
489}
490
491fn terminal(state: &State) -> Error {
493 match &state.error {
494 Some(err) => err.clone(),
495 None => Error::Reconnect("reconnect stopped".to_string()),
496 }
497}
498
499#[cfg(test)]
500mod tests {
501 #[tokio::test]
502 async fn snapshot_uses_one_live_session() {
503 let mut config = crate::ServerConfig {
504 bind: Some("[::]:0".into()),
505 ..Default::default()
506 };
507 config.tls.generate = vec!["localhost".into()];
508 let mut server = config.init().unwrap();
509 let url = format!("moqt://localhost:{}", server.local_addr().unwrap().port())
510 .parse()
511 .unwrap();
512 let mut config = crate::ClientConfig::default();
513 config.tls.disable_verify = Some(true);
514 let client = config.init().unwrap();
515 let (accepted, connected) = tokio::time::timeout(Duration::from_secs(10), async {
516 tokio::join!(
517 async { server.accept().await.unwrap().ok().await.unwrap() },
518 client.connect(url)
519 )
520 })
521 .await
522 .unwrap();
523 let session = connected.unwrap();
524 let version = session.version();
525 let producer = kio::Producer::<State>::default();
526 let reader = ConnectionStatsReader {
527 state: producer.consume(),
528 };
529 assert!(reader.snapshot().is_none());
530 producer.write().ok().unwrap().session = Some(session);
531 assert_eq!(reader.snapshot().unwrap().version, version);
533 producer.write().ok().unwrap().session = None;
534 assert!(reader.snapshot().is_none());
535 drop(accepted);
536 }
537
538 #[test]
542 fn backoff_rejects_an_unpaced_retry() {
543 assert!(Backoff::default().validate().is_ok());
544
545 for bad in [
546 Backoff {
547 initial: Duration::ZERO,
548 ..Default::default()
549 },
550 Backoff {
551 multiplier: 0,
552 ..Default::default()
553 },
554 Backoff {
555 max: Duration::ZERO,
556 ..Default::default()
557 },
558 ] {
559 assert!(
560 matches!(bad.validate(), Err(crate::Error::BackoffUnpaced)),
561 "{bad:?} should be rejected"
562 );
563 }
564
565 let forever = Backoff {
568 timeout: Duration::ZERO,
569 multiplier: 1,
570 ..Default::default()
571 };
572 assert!(forever.validate().is_ok());
573 }
574
575 #[test]
576 fn backoff_growth_saturates_before_applying_the_cap() {
577 let backoff = Backoff {
578 multiplier: u32::MAX,
579 ..Default::default()
580 };
581 assert_eq!(backoff.next_delay(Duration::MAX), backoff.max);
582 }
583
584 use super::*;
585
586 #[test]
587 fn test_backoff_default() {
588 let backoff = Backoff::default();
589 assert_eq!(backoff.initial, Duration::from_secs(1));
590 assert_eq!(backoff.multiplier, 2);
591 assert_eq!(backoff.max, Duration::from_secs(5));
592 assert_eq!(backoff.timeout, Duration::from_secs(10));
593 }
594
595 #[test]
598 fn test_backoff_linger() {
599 let backoff = Backoff::default();
600 assert_eq!(backoff.linger(), backoff.timeout + Duration::from_secs(1));
601
602 let unlimited = Backoff {
603 timeout: Duration::ZERO,
604 ..Backoff::default()
605 };
606 assert_eq!(unlimited.linger(), Duration::MAX);
607 }
608
609 #[test]
610 fn poll_forward_mirrors_until_the_source_closes() {
611 let src = BandwidthProducer::new();
612 let out = BandwidthProducer::new();
613 let out_rx = out.consume();
614 let waiter = kio::Waiter::noop();
615
616 let mut bw = Some(src.consume());
618 poll_forward(&mut bw, &out, &waiter);
619 assert_eq!(out_rx.peek(), None);
620 assert!(bw.is_some());
621
622 src.set(Some(3_000)).unwrap();
624 poll_forward(&mut bw, &out, &waiter);
625 assert_eq!(out_rx.peek(), Some(3_000));
626
627 src.set(None).unwrap();
630 poll_forward(&mut bw, &out, &waiter);
631 assert_eq!(out_rx.peek(), None);
632 assert!(bw.is_some());
633
634 src.set(Some(9_000)).unwrap();
638 poll_forward(&mut bw, &out, &waiter);
639 assert_eq!(out_rx.peek(), Some(9_000));
640
641 src.abort(moq_net::Error::Cancel).unwrap();
643 poll_forward(&mut bw, &out, &waiter);
644 assert!(bw.is_none());
645 }
646}