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 url::Url;
8
9use crate::{Client, Error};
10
11#[derive(Clone, Debug, clap::Args, serde::Serialize, serde::Deserialize)]
13#[serde(default, deny_unknown_fields)]
14#[non_exhaustive]
15pub struct Backoff {
16 #[arg(
18 id = "backoff-initial",
19 long,
20 default_value = "1s",
21 env = "MOQ_BACKOFF_INITIAL",
22 value_parser = humantime::parse_duration,
23 )]
24 #[serde(with = "humantime_serde")]
25 pub initial: Duration,
26
27 #[arg(id = "backoff-multiplier", long, default_value_t = 2, env = "MOQ_BACKOFF_MULTIPLIER")]
29 pub multiplier: u32,
30
31 #[arg(
33 id = "backoff-max",
34 long,
35 default_value = "30s",
36 env = "MOQ_BACKOFF_MAX",
37 value_parser = humantime::parse_duration,
38 )]
39 #[serde(with = "humantime_serde")]
40 pub max: Duration,
41
42 #[arg(
47 id = "backoff-timeout",
48 long,
49 default_value = "5m",
50 env = "MOQ_BACKOFF_TIMEOUT",
51 value_parser = humantime::parse_duration,
52 )]
53 #[serde(with = "humantime_serde")]
54 pub timeout: Duration,
55}
56
57impl Default for Backoff {
58 fn default() -> Self {
59 Self {
60 initial: Duration::from_secs(1),
61 multiplier: 2,
62 max: Duration::from_secs(30),
63 timeout: Duration::from_secs(300),
64 }
65 }
66}
67
68impl Backoff {
69 pub fn linger(&self) -> Duration {
75 match self.timeout.is_zero() {
76 true => Duration::MAX,
77 false => self.timeout.saturating_add(Duration::from_secs(1)),
78 }
79 }
80}
81
82#[derive(Clone, Copy, Debug, PartialEq, Eq)]
84#[non_exhaustive]
85pub enum Status {
86 Connected,
88 Disconnected,
90}
91
92#[derive(Default)]
97struct State {
98 status: Option<Status>,
100 version: Option<Version>,
102 error: Option<Error>,
104 session: Option<moq_net::Session>,
107}
108
109#[derive(Clone)]
114pub struct ConnectionStatsReader {
115 state: kio::Consumer<State>,
116}
117
118impl ConnectionStatsReader {
119 pub fn stats(&self) -> Option<moq_net::ConnectionStats> {
121 self.state.read().session.as_ref().map(moq_net::Session::stats)
122 }
123}
124
125pub struct Reconnect {
135 abort: tokio::task::AbortHandle,
136 state: kio::Consumer<State>,
137 send_bandwidth: BandwidthConsumer,
139 recv_bandwidth: BandwidthConsumer,
141 last_reported: Option<Status>,
143}
144
145impl Reconnect {
146 pub(crate) fn new(client: Client, url: Url, backoff: Backoff) -> Self {
147 let producer = kio::Producer::<State>::default();
148 let state = producer.consume();
149
150 let send_bw = BandwidthProducer::new();
153 let recv_bw = BandwidthProducer::new();
154 let send_bandwidth = send_bw.consume();
155 let recv_bandwidth = recv_bw.consume();
156
157 let task = tokio::spawn(async move {
158 if let Err(err) = Self::run(&producer, &send_bw, &recv_bw, client, url, backoff).await {
159 tracing::error!(%err, "reconnect loop exited");
160 if let Ok(mut state) = producer.write() {
161 state.error = Some(err);
162 }
163 }
164 });
166 Self {
167 abort: task.abort_handle(),
168 state,
169 send_bandwidth,
170 recv_bandwidth,
171 last_reported: None,
172 }
173 }
174
175 async fn run(
176 state: &kio::Producer<State>,
177 send_bw: &BandwidthProducer,
178 recv_bw: &BandwidthProducer,
179 client: Client,
180 url: Url,
181 backoff: Backoff,
182 ) -> crate::Result<()> {
183 let mut delay = backoff.initial;
184 let mut retry_start = tokio::time::Instant::now();
185 let mut last_error: Option<Error> = None;
186
187 loop {
188 if !backoff.timeout.is_zero() && retry_start.elapsed() > backoff.timeout {
189 let timeout = backoff.timeout;
190 let msg = match last_error {
191 Some(err) => format!("reconnect timed out after {timeout:?}: {err}"),
192 None => format!("reconnect timed out after {timeout:?}"),
193 };
194 return Err(Error::Reconnect(msg));
195 }
196
197 tracing::info!(%url, "connecting");
198
199 match client.connect(url.clone()).await {
200 Ok(session) => {
201 tracing::info!(%url, "connected");
202 if let Ok(mut state) = state.write() {
203 state.status = Some(Status::Connected);
204 state.version = Some(session.version());
205 state.session = Some(session.clone());
206 }
207
208 let connected = tokio::time::Instant::now();
209 let closed = run_session(send_bw, recv_bw, &session).await;
212 if let Ok(mut state) = state.write() {
213 state.status = Some(Status::Disconnected);
214 state.version = None;
215 state.session = None;
216 }
217 let _ = send_bw.set(None);
219 let _ = recv_bw.set(None);
220
221 if connected.elapsed() >= backoff.initial {
222 tracing::warn!(%url, "session closed, reconnecting");
225 delay = backoff.initial;
226 retry_start = tokio::time::Instant::now();
227 last_error = None;
228 } else {
229 if let Err(err) = closed {
234 let err = Error::from(err);
235 tracing::warn!(%url, %err, "session severed immediately, retrying");
236 last_error = Some(err);
237 } else {
238 tracing::warn!(%url, "session severed immediately, retrying");
239 }
240 }
241 }
242 Err(err) => {
243 if err.is_auth() {
244 return Err(err);
245 }
246 last_error = Some(err);
247 }
248 }
249
250 tracing::warn!(%url, ?delay, "reconnecting after backoff");
251 tokio::time::sleep(delay).await;
252 delay = std::cmp::min(delay * backoff.multiplier, backoff.max);
253 }
254 }
255
256 pub fn poll_status(&mut self, waiter: &kio::Waiter) -> Poll<crate::Result<Status>> {
261 let last = self.last_reported;
262 let status = match ready!(self.state.poll(waiter, |state| match state.status {
263 Some(status) if Some(status) != last => Poll::Ready(status),
264 _ => Poll::Pending,
265 })) {
266 Ok(status) => status,
267 Err(state) => return Poll::Ready(Err(terminal(&state))),
268 };
269
270 self.last_reported = Some(status);
271 Poll::Ready(Ok(status))
272 }
273
274 pub async fn status(&mut self) -> crate::Result<Status> {
280 kio::wait(|waiter| self.poll_status(waiter)).await
281 }
282
283 pub fn connected(&self) -> bool {
288 self.state.read().status == Some(Status::Connected)
289 }
290
291 pub fn version(&self) -> Option<Version> {
296 self.state.read().version
297 }
298
299 pub fn send_bandwidth(&self) -> BandwidthConsumer {
306 self.send_bandwidth.clone()
307 }
308
309 pub fn recv_bandwidth(&self) -> BandwidthConsumer {
313 self.recv_bandwidth.clone()
314 }
315
316 pub fn poll_closed(&self, waiter: &kio::Waiter) -> Poll<crate::Result<()>> {
321 ready!(self.state.poll_closed(waiter));
322 Poll::Ready(match &self.state.read().error {
323 Some(err) => Err(err.clone()),
324 None => Ok(()),
325 })
326 }
327
328 pub async fn closed(&self) -> crate::Result<()> {
330 kio::wait(|waiter| self.poll_closed(waiter)).await
331 }
332
333 pub fn stats(&self) -> ConnectionStatsReader {
337 ConnectionStatsReader {
338 state: self.state.clone(),
339 }
340 }
341}
342
343async fn run_session(
351 send_bw: &BandwidthProducer,
352 recv_bw: &BandwidthProducer,
353 session: &moq_net::Session,
354) -> Result<(), moq_net::Error> {
355 let mut send = session.send_bandwidth();
356 let mut recv = session.recv_bandwidth();
357 let closed = session.closed();
358 tokio::pin!(closed);
359
360 let err = kio::wait(|waiter| {
361 poll_forward(&mut send, send_bw, waiter);
362 poll_forward(&mut recv, recv_bw, waiter);
363 waiter.poll_future(closed.as_mut())
364 })
365 .await;
366
367 Err(err)
368}
369
370fn poll_forward(bw: &mut Option<BandwidthConsumer>, out: &BandwidthProducer, waiter: &kio::Waiter) {
378 loop {
379 let Some(consumer) = bw.as_mut() else { return };
380 let Poll::Ready(res) = consumer.poll_changed(waiter) else {
381 return;
382 };
383 match res {
384 Ok(rate) => {
385 let _ = out.set(rate);
386 }
387 Err(_) => {
388 *bw = None;
389 return;
390 }
391 }
392 }
393}
394
395impl Drop for Reconnect {
396 fn drop(&mut self) {
397 self.abort.abort();
398 }
399}
400
401fn terminal(state: &State) -> Error {
403 match &state.error {
404 Some(err) => err.clone(),
405 None => Error::Reconnect("reconnect stopped".to_string()),
406 }
407}
408
409#[cfg(test)]
410mod tests {
411 use super::*;
412
413 #[test]
414 fn test_backoff_default() {
415 let backoff = Backoff::default();
416 assert_eq!(backoff.initial, Duration::from_secs(1));
417 assert_eq!(backoff.multiplier, 2);
418 assert_eq!(backoff.max, Duration::from_secs(30));
419 assert_eq!(backoff.timeout, Duration::from_secs(300));
420 }
421
422 #[test]
425 fn test_backoff_linger() {
426 let backoff = Backoff::default();
427 assert_eq!(backoff.linger(), backoff.timeout + Duration::from_secs(1));
428
429 let unlimited = Backoff {
430 timeout: Duration::ZERO,
431 ..Backoff::default()
432 };
433 assert_eq!(unlimited.linger(), Duration::MAX);
434 }
435
436 #[test]
437 fn poll_forward_mirrors_until_the_source_closes() {
438 let src = BandwidthProducer::new();
439 let out = BandwidthProducer::new();
440 let out_rx = out.consume();
441 let waiter = kio::Waiter::noop();
442
443 let mut bw = Some(src.consume());
445 poll_forward(&mut bw, &out, &waiter);
446 assert_eq!(out_rx.peek(), None);
447 assert!(bw.is_some());
448
449 src.set(Some(3_000)).unwrap();
451 poll_forward(&mut bw, &out, &waiter);
452 assert_eq!(out_rx.peek(), Some(3_000));
453
454 src.set(None).unwrap();
457 poll_forward(&mut bw, &out, &waiter);
458 assert_eq!(out_rx.peek(), None);
459 assert!(bw.is_some());
460
461 src.set(Some(9_000)).unwrap();
465 poll_forward(&mut bw, &out, &waiter);
466 assert_eq!(out_rx.peek(), Some(9_000));
467
468 src.abort(moq_net::Error::Cancel).unwrap();
470 poll_forward(&mut bw, &out, &waiter);
471 assert!(bw.is_none());
472 }
473}