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::task::Poll;
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(kio::Shared<Vec<Request>>);
271
272 impl Log {
273 fn events(&self) -> Vec<Event> {
274 self.0.read().iter().map(|r| r.event.clone()).collect()
275 }
276
277 fn revalidates(&self) -> usize {
278 self.events()
279 .iter()
280 .filter(|event| **event == Event::Revalidate)
281 .count()
282 }
283
284 async fn until(&self, mut done: impl FnMut(&[Request]) -> bool + Unpin) {
286 self.0
287 .wait(|log| if done(log) { Poll::Ready(()) } else { Poll::Pending })
288 .await;
289 }
290
291 async fn end(&self) -> Request {
293 let is_end = |r: &Request| matches!(r.event, Event::End { .. });
294 self.until(|log| log.iter().any(is_end)).await;
295 self.0.read().iter().find(|r| is_end(r)).cloned().unwrap()
296 }
297 }
298
299 impl wiremock::Match for Log {
300 fn matches(&self, received: &Received) -> bool {
301 self.0.lock().push(received.body_json().unwrap());
302 true
303 }
304 }
305
306 async fn server(log: Log, respond: impl Fn(&Request) -> ResponseTemplate + Send + Sync + 'static) -> MockServer {
307 let server = MockServer::start().await;
308 Mock::given(method("POST"))
309 .and(path("/"))
310 .and(log)
311 .respond_with(move |received: &Received| respond(&received.body_json().unwrap()))
312 .mount(&server)
313 .await;
314 server
315 }
316
317 fn client(server: &MockServer) -> Client {
318 Client::new(server.uri().parse().unwrap(), None).unwrap()
319 }
320
321 fn grant(expires_in: Option<Duration>, revalidate: Option<Duration>) -> Grant {
324 let mut grant = Grant::new(patterns(&["**"]), Patterns::new());
325 grant.expires = expires_in.map(|d| SystemTime::now() + d);
326 grant.revalidate = revalidate;
327 grant
328 }
329
330 #[tokio::test]
331 async fn connect_admits_and_end_follows_the_close() {
332 let log = Log::default();
333 let server = server(log.clone(), |_| {
334 ResponseTemplate::new(200).set_body_json(grant(None, None))
335 })
336 .await;
337
338 let consumer = client(&server).connect(request()).await.unwrap();
339 assert_eq!(consumer.grant().publish, patterns(&["**"]));
340 assert_eq!(log.events(), [Event::Connect]);
341
342 consumer.close("disconnected", Bytes { sent: 7, received: 11 });
343
344 let end = log.end().await;
345 assert_eq!(end.id, "0123");
346 match end.event {
347 Event::End { reason, bytes, .. } => {
348 assert_eq!(reason, Reason::Session("disconnected".into()));
349 assert_eq!(bytes, Bytes { sent: 7, received: 11 });
350 }
351 other => panic!("expected an end, got {other:?}"),
352 }
353 }
354
355 #[tokio::test]
356 async fn a_bare_drop_ends_as_dropped() {
357 let log = Log::default();
358 let server = server(log.clone(), |_| {
359 ResponseTemplate::new(200).set_body_json(grant(None, None))
360 })
361 .await;
362
363 let consumer = client(&server).connect(request()).await.unwrap();
364 drop(consumer);
365
366 match log.end().await.event {
367 Event::End {
368 reason: Reason::Dropped,
369 bytes,
370 ..
371 } => assert_eq!(bytes, Bytes::default()),
372 other => panic!("expected a dropped end, got {other:?}"),
373 }
374 }
375
376 #[tokio::test]
377 async fn refusals_and_outages_refuse_at_connect() {
378 for status in [401, 403, 400, 404, 408, 429, 500, 503] {
379 let server = server(Log::default(), move |_| ResponseTemplate::new(status)).await;
380 let err = client(&server).connect(request()).await.unwrap_err();
381 match status {
382 401 | 403 => assert!(matches!(err, Error::Refused), "{status}: {err}"),
383 _ => assert!(matches!(err, Error::Unavailable(_)), "{status}: {err}"),
384 }
385 }
386
387 let garbage = server(Log::default(), |_| ResponseTemplate::new(200).set_body_string("nope")).await;
388 let err = client(&garbage).connect(request()).await.unwrap_err();
389 assert!(matches!(err, Error::Unavailable(_)), "{err}");
390
391 let nothing = server(Log::default(), |_| {
392 ResponseTemplate::new(200).set_body_json(Grant::default())
393 })
394 .await;
395 let err = client(¬hing).connect(request()).await.unwrap_err();
396 assert!(matches!(err, Error::Refused), "{err}");
397 }
398
399 #[tokio::test]
400 async fn revalidate_runs_on_cadence_and_applies_the_reply() {
401 let log = Log::default();
402 let server = server(log.clone(), |request| {
403 let mut grant = grant(Some(Duration::from_secs(3600)), Some(Duration::from_secs(1)));
404 if request.event == Event::Revalidate {
405 grant.tier = Some("moved".into());
406 }
407 ResponseTemplate::new(200).set_body_json(grant)
408 })
409 .await;
410
411 let mut consumer = client(&server).connect(request()).await.unwrap();
412 let fresh = tokio::time::timeout(Duration::from_secs(3), consumer.changed())
413 .await
414 .expect("a re-check within the cadence")
415 .unwrap();
416 assert_eq!(fresh.tier.as_deref(), Some("moved"));
417 assert_eq!(log.events()[..2], [Event::Connect, Event::Revalidate]);
418 }
419
420 #[tokio::test]
421 async fn a_refusal_on_recheck_revokes() {
422 for answer in [
423 ResponseTemplate::new(401),
424 ResponseTemplate::new(403),
425 ResponseTemplate::new(200).set_body_json(Grant::default()),
426 ] {
427 let server = server(Log::default(), {
428 let answer = answer.clone();
429 move |request| match request.event {
430 Event::Connect => ResponseTemplate::new(200)
431 .set_body_json(grant(Some(Duration::from_secs(3600)), Some(Duration::from_secs(1)))),
432 _ => answer.clone(),
433 }
434 })
435 .await;
436
437 let consumer = client(&server).connect(request()).await.unwrap();
438 let reason = tokio::time::timeout(Duration::from_secs(3), consumer.closed())
439 .await
440 .expect("revoked within the cadence");
441 assert_eq!(reason, Reason::Refused);
442 }
443 }
444
445 #[tokio::test]
446 async fn an_invalid_grant_on_recheck_revokes() {
447 let unbounded = grant(None, Some(Duration::from_secs(1)));
448 let server = server(Log::default(), move |request| match request.event {
449 Event::Connect => ResponseTemplate::new(200)
450 .set_body_json(grant(Some(Duration::from_secs(3600)), Some(Duration::from_secs(1)))),
451 _ => ResponseTemplate::new(200).set_body_json(unbounded.clone()),
452 })
453 .await;
454
455 let consumer = client(&server).connect(request()).await.unwrap();
456 let reason = tokio::time::timeout(Duration::from_secs(3), consumer.closed())
457 .await
458 .expect("revoked within the cadence");
459 assert_eq!(reason, Reason::Invalid);
460 }
461
462 #[tokio::test]
463 async fn a_grant_within_clock_skew_stays_live() {
464 let server = server(Log::default(), |_| {
465 let mut grant = Grant::new(patterns(&["**"]), Patterns::new());
466 grant.expires = Some(SystemTime::now() - Duration::from_secs(1));
467 ResponseTemplate::new(200).set_body_json(grant)
468 })
469 .await;
470
471 let consumer = client(&server).connect(request()).await.unwrap();
472 tokio::time::sleep(Duration::from_millis(500)).await;
473 assert!(
474 tokio::time::timeout(Duration::from_millis(100), consumer.closed())
475 .await
476 .is_err(),
477 "still live inside the skew window"
478 );
479
480 let reason = tokio::time::timeout(crate::grant::CLOCK_SKEW + Duration::from_secs(1), consumer.closed())
481 .await
482 .expect("expired once the skew window ended");
483 assert_eq!(reason, Reason::Expired);
484 }
485
486 #[tokio::test]
487 async fn an_outage_keeps_the_grant_until_expires() {
488 let log = Log::default();
489 let server = server(log.clone(), |request| match request.event {
490 Event::Connect => ResponseTemplate::new(200)
491 .set_body_json(grant(Some(Duration::from_secs(3)), Some(Duration::from_secs(1)))),
492 _ => ResponseTemplate::new(503),
493 })
494 .await;
495
496 let consumer = client(&server).connect(request()).await.unwrap();
497 tokio::time::sleep(Duration::from_millis(1500)).await;
498 assert!(log.revalidates() >= 1, "re-checks happened");
499 assert_eq!(
500 consumer.grant().publish,
501 patterns(&["**"]),
502 "the grant stands through the outage"
503 );
504
505 let reason = tokio::time::timeout(Duration::from_secs(5), consumer.closed())
506 .await
507 .expect("expired");
508 assert_eq!(reason, Reason::Expired);
509 assert!(matches!(
510 log.end().await.event,
511 Event::End {
512 reason: Reason::Expired,
513 ..
514 }
515 ));
516 }
517
518 #[tokio::test]
519 async fn expiry_fires_while_a_recheck_is_stalled() {
520 let server = server(Log::default(), |request| match request.event {
521 Event::Connect => ResponseTemplate::new(200)
522 .set_body_json(grant(Some(Duration::from_secs(3)), Some(Duration::from_secs(1)))),
523 _ => ResponseTemplate::new(200)
524 .set_body_json(grant(Some(Duration::from_secs(3600)), Some(Duration::from_secs(60))))
525 .set_delay(Duration::from_secs(30)),
526 })
527 .await;
528
529 let consumer = client(&server).connect(request()).await.unwrap();
530 let reason = tokio::time::timeout(Duration::from_secs(5), consumer.closed())
531 .await
532 .expect("expired while the re-check was in flight");
533 assert_eq!(reason, Reason::Expired);
534 }
535
536 #[tokio::test]
537 async fn a_close_is_reported_while_a_recheck_is_stalled() {
538 let log = Log::default();
539 let server = server(log.clone(), |request| match request.event {
540 Event::Connect => ResponseTemplate::new(200)
541 .set_body_json(grant(Some(Duration::from_secs(3600)), Some(Duration::from_secs(1)))),
542 Event::Revalidate => ResponseTemplate::new(200)
543 .set_body_json(grant(Some(Duration::from_secs(3600)), Some(Duration::from_secs(60))))
544 .set_delay(Duration::from_secs(30)),
545 Event::End { .. } => ResponseTemplate::new(200),
546 })
547 .await;
548
549 let consumer = client(&server).connect(request()).await.unwrap();
550 tokio::time::sleep(Duration::from_millis(1500)).await;
551 assert!(log.events().contains(&Event::Revalidate), "the re-check is in flight");
552 consumer.close("disconnected", Bytes::default());
553 assert!(matches!(
554 log.end().await.event,
555 Event::End {
556 reason: Reason::Session(_),
557 ..
558 }
559 ));
560 }
561
562 #[test]
563 fn url_schemes_are_checked_at_construction() {
564 assert!(Client::new("http://127.0.0.1:4440/".parse().unwrap(), None).is_ok());
565 assert!(Client::new("http://localhost:4440/".parse().unwrap(), None).is_ok());
566 assert!(Client::new("http://[::1]:4440/".parse().unwrap(), None).is_ok());
567 assert!(matches!(
568 Client::new("http://auth.example/".parse().unwrap(), None),
569 Err(Error::InsecureUrl(_))
570 ));
571 assert!(Client::new("https://auth.example/".parse().unwrap(), None).is_ok());
572 assert!(matches!(
573 Client::new("ftp://auth.example/".parse().unwrap(), None),
574 Err(Error::InvalidUrl(_))
575 ));
576 #[cfg(unix)]
577 assert!(Client::new("unix:///run/moq-auth.sock".parse().unwrap(), None).is_ok());
578 }
579
580 #[tokio::test]
583 async fn a_nudge_while_idle_posts_at_once() {
584 let log = Log::default();
585 let server = server(log.clone(), |_| {
586 ResponseTemplate::new(200)
587 .set_body_json(grant(Some(Duration::from_secs(3600)), Some(Duration::from_secs(3600))))
588 })
589 .await;
590
591 let consumer = client(&server).connect(request()).await.unwrap();
592 assert_eq!(log.events(), [Event::Connect]);
593 consumer.revalidate();
594 tokio::time::timeout(
595 Duration::from_secs(2),
596 log.until(|log| log.iter().any(|r| r.event == Event::Revalidate)),
597 )
598 .await
599 .expect("a nudge while idle POSTs at once");
600 }
601
602 #[tokio::test]
603 async fn a_nudge_during_inflight_posts_once_more_when_the_reply_lands() {
604 let log = Log::default();
605 let server = server(log.clone(), |request| match request.event {
606 Event::Connect => ResponseTemplate::new(200)
607 .set_body_json(grant(Some(Duration::from_secs(3600)), Some(Duration::from_secs(3600)))),
608 Event::Revalidate => ResponseTemplate::new(200)
609 .set_body_json(grant(Some(Duration::from_secs(3600)), Some(Duration::from_secs(3600))))
610 .set_delay(Duration::from_millis(400)),
611 Event::End { .. } => ResponseTemplate::new(200),
612 })
613 .await;
614
615 let consumer = client(&server).connect(request()).await.unwrap();
616 consumer.revalidate();
617 tokio::time::timeout(
618 Duration::from_secs(2),
619 log.until(|log| log.iter().any(|r| r.event == Event::Revalidate)),
620 )
621 .await
622 .expect("the first re-check is in flight");
623
624 consumer.revalidate();
625 consumer.revalidate();
626 tokio::time::timeout(
629 Duration::from_secs(3),
630 log.until(|log| log.iter().filter(|r| r.event == Event::Revalidate).count() >= 2),
631 )
632 .await
633 .expect("the in-flight nudge POSTs once more when the reply lands");
634 consumer.close("disconnected", Bytes::default());
635 log.end().await;
636 assert_eq!(
637 log.revalidates(),
638 2,
639 "a burst during an in-flight re-check is one extra POST"
640 );
641 }
642
643 #[test]
644 fn backoff_grows_and_stays_bounded() {
645 let cadence = Duration::from_secs(30);
646 let first = backoff(1, cadence);
647 assert!(
648 first >= Duration::from_millis(750) && first <= Duration::from_millis(1250),
649 "{first:?}"
650 );
651 let later = backoff(10, cadence);
652 assert!(later <= cadence.mul_f64(1.25), "{later:?}");
653 assert!(backoff(40, Duration::from_secs(3600)) <= BACKOFF_MAX.mul_f64(1.25));
654 }
655}