1use std::future::Future;
2use std::pin::Pin;
3use std::time::{Duration, Instant};
4
5use url::Url;
6
7use crate::lease::{self, Reason};
8use crate::{Bytes, Error, Event, Grant, Request};
9
10const TIMEOUT: Duration = Duration::from_secs(10);
12
13const BACKOFF_MAX: Duration = Duration::from_secs(60);
15
16#[derive(Clone)]
23pub struct Client {
24 http: reqwest::Client,
25 url: Url,
26}
27
28impl Client {
29 pub fn new(url: Url, tls: Option<rustls::ClientConfig>) -> crate::Result<Self> {
35 let builder = reqwest::Client::builder().timeout(TIMEOUT);
36
37 let (builder, url) = match url.scheme() {
38 "http" => {
39 let loopback = match url.host() {
40 Some(url::Host::Ipv4(ip)) => ip.is_loopback(),
41 Some(url::Host::Ipv6(ip)) => ip.is_loopback(),
42 Some(url::Host::Domain(host)) => host == "localhost",
43 None => false,
44 };
45 if !loopback {
46 return Err(Error::InsecureUrl(url.to_string()));
47 }
48 (builder, url)
49 }
50 "https" => match tls {
51 Some(tls) => (builder.use_preconfigured_tls(tls), url),
52 None => (builder, url),
53 },
54 #[cfg(unix)]
55 "unix" => {
56 let path = url.to_file_path().map_err(|()| Error::InvalidUrl(url.to_string()))?;
57 let target = Url::parse("http://localhost/").expect("a constant URL parses");
59 (builder.unix_socket(path), target)
60 }
61 _ => return Err(Error::InvalidUrl(url.to_string())),
62 };
63
64 Ok(Self {
65 http: builder.build()?,
66 url,
67 })
68 }
69
70 pub async fn connect(&self, request: Request) -> crate::Result<lease::Consumer> {
73 let mut request = request;
74 request.event = Event::Connect;
75
76 let grant = self.post(&request).await?;
77 let (producer, consumer) = lease::Producer::new(grant.clone());
78
79 let driver = Driver {
80 client: self.clone(),
81 request,
82 producer: Some(producer),
83 started: Instant::now(),
84 };
85 tokio::spawn(driver.run(grant));
86
87 Ok(consumer)
88 }
89
90 async fn post(&self, request: &Request) -> crate::Result<Grant> {
94 let response = self.http.post(self.url.clone()).json(request).send().await?;
95 let status = response.status();
96 if status == reqwest::StatusCode::UNAUTHORIZED || status == reqwest::StatusCode::FORBIDDEN {
97 return Err(Error::Refused);
98 }
99 if !status.is_success() {
100 return Err(Error::Unavailable(format!("auth server answered {status}")));
101 }
102 let grant: Grant = response.json().await?;
103 grant.validate().map_err(|err| match err {
104 Error::UselessGrant => Error::Refused,
106 other => other,
107 })?;
108 Ok(grant)
109 }
110}
111
112struct Driver {
114 client: Client,
115 request: Request,
116 producer: Option<lease::Producer>,
117 started: Instant,
118}
119
120impl Driver {
121 async fn run(mut self, grant: Grant) {
122 let (reason, bytes) = self.drive(grant).await;
123
124 let mut request = self.request.clone();
125 request.event = Event::End {
126 reason,
127 duration: self.started.elapsed(),
128 bytes,
129 };
130 if let Err(err) = self.client.post_end(&request).await {
132 tracing::warn!(id = %request.id, %err, "failed to report the session end");
133 }
134 }
135
136 async fn drive(&mut self, mut grant: Grant) -> (Reason, Bytes) {
138 let producer = self
139 .producer
140 .take()
141 .expect("the driver owns the producer until it ends");
142 let mut failures = 0u32;
143 let mut next = grant.revalidate.map(|cadence| Instant::now() + cadence);
144 let mut inflight: Option<Pin<Box<dyn Future<Output = crate::Result<Grant>> + Send>>> = None;
147 let mut pending = false;
150
151 loop {
152 let expires = grant.expires.map(crate::grant::until);
153 let revalidate = async {
154 match next {
155 Some(at) => tokio::time::sleep_until(at.into()).await,
156 None => std::future::pending().await,
157 }
158 };
159 let expire = async {
160 match expires {
161 Some(after) => tokio::time::sleep(after).await,
162 None => std::future::pending().await,
163 }
164 };
165 let reply = async {
166 match inflight.as_mut() {
167 Some(request) => request.await,
168 None => std::future::pending().await,
169 }
170 };
171
172 tokio::select! {
173 ended = producer.closed() => return ended,
174 () = expire => return producer.finish(Reason::Expired, Bytes::default()),
175 result = reply => {
176 inflight = None;
177 match result {
178 Ok(fresh) => {
179 failures = 0;
180 next = fresh.revalidate.map(|cadence| Instant::now() + cadence);
181 producer.update(fresh.clone());
182 grant = fresh;
183 }
184 Err(Error::Refused) => return producer.finish(Reason::Refused, Bytes::default()),
185 Err(Error::GrantExpired | Error::UnboundedRevalidate | Error::ZeroRevalidate) => {
186 return producer.finish(Reason::Invalid, Bytes::default());
187 }
188 Err(err) => {
189 failures += 1;
191 let delay = backoff(failures, grant.revalidate.unwrap_or(BACKOFF_MAX));
192 tracing::warn!(id = %self.request.id, %err, ?delay, "auth revalidation failed; retrying");
193 next = Some(Instant::now() + delay);
194 }
195 }
196 if pending {
197 pending = false;
198 next = None;
199 inflight = Some(self.post_revalidate());
200 }
201 }
202 () = revalidate => {
203 next = None;
205 inflight = Some(self.post_revalidate());
206 }
207 () = producer.revalidate_requested() => {
208 if inflight.is_some() {
209 pending = true;
210 } else {
211 next = None;
212 inflight = Some(self.post_revalidate());
213 }
214 }
215 }
216 }
217 }
218
219 fn post_revalidate(&self) -> Pin<Box<dyn Future<Output = crate::Result<Grant>> + Send>> {
220 let client = self.client.clone();
221 let mut request = self.request.clone();
222 request.event = Event::Revalidate;
223 Box::pin(async move { client.post(&request).await })
224 }
225}
226
227impl Client {
228 async fn post_end(&self, request: &Request) -> crate::Result<()> {
229 self.http
230 .post(self.url.clone())
231 .json(request)
232 .send()
233 .await?
234 .error_for_status()?;
235 Ok(())
236 }
237}
238
239fn backoff(failures: u32, cadence: Duration) -> Duration {
242 use rand::RngExt;
243 let base = Duration::from_secs(1) * 2u32.saturating_pow(failures.saturating_sub(1).min(16));
244 let base = base.min(cadence).min(BACKOFF_MAX);
245 let jitter = rand::rng().random_range(0.75..=1.25);
246 base.mul_f64(jitter)
247}
248
249#[cfg(test)]
250mod tests {
251 use super::*;
252 use moq_pattern::Patterns;
253 use std::sync::{Arc, Mutex};
254 use std::time::SystemTime;
255 use wiremock::matchers::{method, path};
256 use wiremock::{Mock, MockServer, Request as Received, ResponseTemplate};
257
258 fn patterns(texts: &[&str]) -> Patterns {
259 texts.iter().map(|text| text.parse().unwrap()).collect()
260 }
261
262 fn request() -> Request {
263 let mut request = Request::new("relay-1", crate::Transport::Quic, "/demo/room");
264 request.id = "0123".into();
265 request
266 }
267
268 #[derive(Clone, Default)]
270 struct Log(Arc<Mutex<Vec<Request>>>);
271
272 impl Log {
273 fn events(&self) -> Vec<Event> {
274 self.0.lock().unwrap().iter().map(|r| r.event.clone()).collect()
275 }
276
277 fn last(&self) -> Request {
278 self.0.lock().unwrap().last().cloned().unwrap()
279 }
280 }
281
282 impl wiremock::Match for Log {
283 fn matches(&self, received: &Received) -> bool {
284 self.0.lock().unwrap().push(received.body_json().unwrap());
285 true
286 }
287 }
288
289 async fn server(log: Log, respond: impl Fn(&Request) -> ResponseTemplate + Send + Sync + 'static) -> MockServer {
290 let server = MockServer::start().await;
291 Mock::given(method("POST"))
292 .and(path("/"))
293 .and(log)
294 .respond_with(move |received: &Received| respond(&received.body_json().unwrap()))
295 .mount(&server)
296 .await;
297 server
298 }
299
300 fn client(server: &MockServer) -> Client {
301 Client::new(server.uri().parse().unwrap(), None).unwrap()
302 }
303
304 fn grant(expires_in: Option<Duration>, revalidate: Option<Duration>) -> Grant {
307 let mut grant = Grant::new(patterns(&["**"]), Patterns::new());
308 grant.expires = expires_in.map(|d| SystemTime::now() + d);
309 grant.revalidate = revalidate;
310 grant
311 }
312
313 async fn settle() {
314 tokio::time::sleep(Duration::from_millis(50)).await;
315 }
316
317 #[tokio::test]
318 async fn connect_admits_and_end_follows_the_close() {
319 let log = Log::default();
320 let server = server(log.clone(), |_| {
321 ResponseTemplate::new(200).set_body_json(grant(None, None))
322 })
323 .await;
324
325 let consumer = client(&server).connect(request()).await.unwrap();
326 assert_eq!(consumer.grant().publish, patterns(&["**"]));
327 assert_eq!(log.events(), [Event::Connect]);
328
329 consumer.close("disconnected", Bytes { sent: 7, received: 11 });
330 settle().await;
331
332 let end = log.last();
333 assert_eq!(end.id, "0123");
334 match end.event {
335 Event::End { reason, bytes, .. } => {
336 assert_eq!(reason, Reason::Session("disconnected".into()));
337 assert_eq!(bytes, Bytes { sent: 7, received: 11 });
338 }
339 other => panic!("expected an end, got {other:?}"),
340 }
341 }
342
343 #[tokio::test]
344 async fn a_bare_drop_ends_as_dropped() {
345 let log = Log::default();
346 let server = server(log.clone(), |_| {
347 ResponseTemplate::new(200).set_body_json(grant(None, None))
348 })
349 .await;
350
351 let consumer = client(&server).connect(request()).await.unwrap();
352 drop(consumer);
353 settle().await;
354
355 match log.last().event {
356 Event::End {
357 reason: Reason::Dropped,
358 bytes,
359 ..
360 } => assert_eq!(bytes, Bytes::default()),
361 other => panic!("expected a dropped end, got {other:?}"),
362 }
363 }
364
365 #[tokio::test]
366 async fn refusals_and_outages_refuse_at_connect() {
367 for status in [401, 403, 400, 404, 408, 429, 500, 503] {
368 let server = server(Log::default(), move |_| ResponseTemplate::new(status)).await;
369 let err = client(&server).connect(request()).await.unwrap_err();
370 match status {
371 401 | 403 => assert!(matches!(err, Error::Refused), "{status}: {err}"),
372 _ => assert!(matches!(err, Error::Unavailable(_)), "{status}: {err}"),
373 }
374 }
375
376 let garbage = server(Log::default(), |_| ResponseTemplate::new(200).set_body_string("nope")).await;
377 let err = client(&garbage).connect(request()).await.unwrap_err();
378 assert!(matches!(err, Error::Unavailable(_)), "{err}");
379
380 let nothing = server(Log::default(), |_| {
381 ResponseTemplate::new(200).set_body_json(Grant::default())
382 })
383 .await;
384 let err = client(¬hing).connect(request()).await.unwrap_err();
385 assert!(matches!(err, Error::Refused), "{err}");
386 }
387
388 #[tokio::test]
389 async fn revalidate_runs_on_cadence_and_applies_the_reply() {
390 let log = Log::default();
391 let server = server(log.clone(), |request| {
392 let mut grant = grant(Some(Duration::from_secs(3600)), Some(Duration::from_secs(1)));
393 if request.event == Event::Revalidate {
394 grant.tier = Some("moved".into());
395 }
396 ResponseTemplate::new(200).set_body_json(grant)
397 })
398 .await;
399
400 let mut consumer = client(&server).connect(request()).await.unwrap();
401 let fresh = tokio::time::timeout(Duration::from_secs(3), consumer.changed())
402 .await
403 .expect("a re-check within the cadence")
404 .unwrap();
405 assert_eq!(fresh.tier.as_deref(), Some("moved"));
406 assert_eq!(log.events()[..2], [Event::Connect, Event::Revalidate]);
407 }
408
409 #[tokio::test]
410 async fn a_refusal_on_recheck_revokes() {
411 for answer in [
412 ResponseTemplate::new(401),
413 ResponseTemplate::new(403),
414 ResponseTemplate::new(200).set_body_json(Grant::default()),
415 ] {
416 let server = server(Log::default(), {
417 let answer = answer.clone();
418 move |request| match request.event {
419 Event::Connect => ResponseTemplate::new(200)
420 .set_body_json(grant(Some(Duration::from_secs(3600)), Some(Duration::from_secs(1)))),
421 _ => answer.clone(),
422 }
423 })
424 .await;
425
426 let consumer = client(&server).connect(request()).await.unwrap();
427 let reason = tokio::time::timeout(Duration::from_secs(3), consumer.closed())
428 .await
429 .expect("revoked within the cadence");
430 assert_eq!(reason, Reason::Refused);
431 }
432 }
433
434 #[tokio::test]
435 async fn an_invalid_grant_on_recheck_revokes() {
436 let unbounded = grant(None, Some(Duration::from_secs(1)));
437 let server = server(Log::default(), move |request| match request.event {
438 Event::Connect => ResponseTemplate::new(200)
439 .set_body_json(grant(Some(Duration::from_secs(3600)), Some(Duration::from_secs(1)))),
440 _ => ResponseTemplate::new(200).set_body_json(unbounded.clone()),
441 })
442 .await;
443
444 let consumer = client(&server).connect(request()).await.unwrap();
445 let reason = tokio::time::timeout(Duration::from_secs(3), consumer.closed())
446 .await
447 .expect("revoked within the cadence");
448 assert_eq!(reason, Reason::Invalid);
449 }
450
451 #[tokio::test]
452 async fn a_grant_within_clock_skew_stays_live() {
453 let server = server(Log::default(), |_| {
454 let mut grant = Grant::new(patterns(&["**"]), Patterns::new());
455 grant.expires = Some(SystemTime::now() - Duration::from_secs(1));
456 ResponseTemplate::new(200).set_body_json(grant)
457 })
458 .await;
459
460 let consumer = client(&server).connect(request()).await.unwrap();
461 tokio::time::sleep(Duration::from_millis(500)).await;
462 assert!(
463 tokio::time::timeout(Duration::from_millis(100), consumer.closed())
464 .await
465 .is_err(),
466 "still live inside the skew window"
467 );
468
469 let reason = tokio::time::timeout(crate::grant::CLOCK_SKEW + Duration::from_secs(1), consumer.closed())
470 .await
471 .expect("expired once the skew window ended");
472 assert_eq!(reason, Reason::Expired);
473 }
474
475 #[tokio::test]
476 async fn an_outage_keeps_the_grant_until_expires() {
477 let log = Log::default();
478 let server = server(log.clone(), |request| match request.event {
479 Event::Connect => ResponseTemplate::new(200)
480 .set_body_json(grant(Some(Duration::from_secs(3)), Some(Duration::from_secs(1)))),
481 _ => ResponseTemplate::new(503),
482 })
483 .await;
484
485 let consumer = client(&server).connect(request()).await.unwrap();
486 tokio::time::sleep(Duration::from_millis(1500)).await;
487 assert!(
488 log.events().iter().filter(|e| **e == Event::Revalidate).count() >= 1,
489 "re-checks happened"
490 );
491 assert_eq!(
492 consumer.grant().publish,
493 patterns(&["**"]),
494 "the grant stands through the outage"
495 );
496
497 let reason = tokio::time::timeout(Duration::from_secs(5), consumer.closed())
498 .await
499 .expect("expired");
500 assert_eq!(reason, Reason::Expired);
501 settle().await;
502 assert!(matches!(
503 log.last().event,
504 Event::End {
505 reason: Reason::Expired,
506 ..
507 }
508 ));
509 }
510
511 #[tokio::test]
512 async fn expiry_fires_while_a_recheck_is_stalled() {
513 let server = server(Log::default(), |request| match request.event {
514 Event::Connect => ResponseTemplate::new(200)
515 .set_body_json(grant(Some(Duration::from_secs(3)), Some(Duration::from_secs(1)))),
516 _ => ResponseTemplate::new(200)
517 .set_body_json(grant(Some(Duration::from_secs(3600)), Some(Duration::from_secs(60))))
518 .set_delay(Duration::from_secs(30)),
519 })
520 .await;
521
522 let consumer = client(&server).connect(request()).await.unwrap();
523 let reason = tokio::time::timeout(Duration::from_secs(5), consumer.closed())
524 .await
525 .expect("expired while the re-check was in flight");
526 assert_eq!(reason, Reason::Expired);
527 }
528
529 #[tokio::test]
530 async fn a_close_is_reported_while_a_recheck_is_stalled() {
531 let log = Log::default();
532 let server = server(log.clone(), |request| match request.event {
533 Event::Connect => ResponseTemplate::new(200)
534 .set_body_json(grant(Some(Duration::from_secs(3600)), Some(Duration::from_secs(1)))),
535 Event::Revalidate => ResponseTemplate::new(200)
536 .set_body_json(grant(Some(Duration::from_secs(3600)), Some(Duration::from_secs(60))))
537 .set_delay(Duration::from_secs(30)),
538 Event::End { .. } => ResponseTemplate::new(200),
539 })
540 .await;
541
542 let consumer = client(&server).connect(request()).await.unwrap();
543 tokio::time::sleep(Duration::from_millis(1500)).await;
544 assert!(log.events().contains(&Event::Revalidate), "the re-check is in flight");
545 consumer.close("disconnected", Bytes::default());
546 settle().await;
547 assert!(matches!(
548 log.last().event,
549 Event::End {
550 reason: Reason::Session(_),
551 ..
552 }
553 ));
554 }
555
556 #[test]
557 fn url_schemes_are_checked_at_construction() {
558 assert!(Client::new("http://127.0.0.1:4440/".parse().unwrap(), None).is_ok());
559 assert!(Client::new("http://localhost:4440/".parse().unwrap(), None).is_ok());
560 assert!(Client::new("http://[::1]:4440/".parse().unwrap(), None).is_ok());
561 assert!(matches!(
562 Client::new("http://auth.example/".parse().unwrap(), None),
563 Err(Error::InsecureUrl(_))
564 ));
565 assert!(Client::new("https://auth.example/".parse().unwrap(), None).is_ok());
566 assert!(matches!(
567 Client::new("ftp://auth.example/".parse().unwrap(), None),
568 Err(Error::InvalidUrl(_))
569 ));
570 #[cfg(unix)]
571 assert!(Client::new("unix:///run/moq-auth.sock".parse().unwrap(), None).is_ok());
572 }
573
574 #[tokio::test]
575 async fn a_nudge_while_idle_posts_at_once() {
576 let log = Log::default();
577 let server = server(log.clone(), |_| {
578 ResponseTemplate::new(200)
579 .set_body_json(grant(Some(Duration::from_secs(3600)), Some(Duration::from_secs(3600))))
580 })
581 .await;
582
583 let consumer = client(&server).connect(request()).await.unwrap();
584 assert_eq!(log.events(), [Event::Connect]);
585 consumer.revalidate();
586 tokio::time::timeout(Duration::from_secs(2), async {
587 loop {
588 if log.events().contains(&Event::Revalidate) {
589 break;
590 }
591 tokio::time::sleep(Duration::from_millis(10)).await;
592 }
593 })
594 .await
595 .expect("a nudge while idle POSTs at once");
596 }
597
598 #[tokio::test]
599 async fn a_nudge_during_backoff_posts_at_once() {
600 let log = Log::default();
601 let server = server(log.clone(), |request| match request.event {
602 Event::Connect => ResponseTemplate::new(200)
603 .set_body_json(grant(Some(Duration::from_secs(3600)), Some(Duration::from_secs(3600)))),
604 _ => ResponseTemplate::new(503),
605 })
606 .await;
607
608 let consumer = client(&server).connect(request()).await.unwrap();
609 consumer.revalidate();
610 tokio::time::timeout(Duration::from_secs(2), async {
611 loop {
612 if log.events().contains(&Event::Revalidate) {
613 break;
614 }
615 tokio::time::sleep(Duration::from_millis(10)).await;
616 }
617 })
618 .await
619 .expect("the first re-check ran");
620 settle().await;
621 let before = log.events().iter().filter(|event| **event == Event::Revalidate).count();
622
623 consumer.revalidate();
624 tokio::time::timeout(Duration::from_secs(2), async {
625 loop {
626 if log.events().iter().filter(|event| **event == Event::Revalidate).count() > before {
627 break;
628 }
629 tokio::time::sleep(Duration::from_millis(10)).await;
630 }
631 })
632 .await
633 .expect("a nudge during backoff POSTs at once");
634 }
635
636 #[tokio::test]
637 async fn a_nudge_during_inflight_posts_once_more_when_the_reply_lands() {
638 let log = Log::default();
639 let server = server(log.clone(), |request| match request.event {
640 Event::Connect => ResponseTemplate::new(200)
641 .set_body_json(grant(Some(Duration::from_secs(3600)), Some(Duration::from_secs(3600)))),
642 Event::Revalidate => ResponseTemplate::new(200)
643 .set_body_json(grant(Some(Duration::from_secs(3600)), Some(Duration::from_secs(3600))))
644 .set_delay(Duration::from_millis(400)),
645 Event::End { .. } => ResponseTemplate::new(200),
646 })
647 .await;
648
649 let consumer = client(&server).connect(request()).await.unwrap();
650 consumer.revalidate();
651 tokio::time::timeout(Duration::from_secs(2), async {
652 loop {
653 if log.events().contains(&Event::Revalidate) {
654 break;
655 }
656 tokio::time::sleep(Duration::from_millis(10)).await;
657 }
658 })
659 .await
660 .expect("the first re-check is in flight");
661
662 consumer.revalidate();
663 consumer.revalidate();
664 tokio::time::timeout(Duration::from_secs(3), async {
665 loop {
666 if log.events().iter().filter(|event| **event == Event::Revalidate).count() >= 2 {
667 break;
668 }
669 tokio::time::sleep(Duration::from_millis(10)).await;
670 }
671 })
672 .await
673 .expect("the in-flight nudge POSTs once more when the reply lands");
674 settle().await;
675 assert_eq!(
676 log.events().iter().filter(|event| **event == Event::Revalidate).count(),
677 2,
678 "a burst during an in-flight re-check is one extra POST"
679 );
680 }
681
682 #[test]
683 fn backoff_grows_and_stays_bounded() {
684 let cadence = Duration::from_secs(30);
685 let first = backoff(1, cadence);
686 assert!(
687 first >= Duration::from_millis(750) && first <= Duration::from_millis(1250),
688 "{first:?}"
689 );
690 let later = backoff(10, cadence);
691 assert!(later <= cadence.mul_f64(1.25), "{later:?}");
692 assert!(backoff(40, Duration::from_secs(3600)) <= BACKOFF_MAX.mul_f64(1.25));
693 }
694}