Skip to main content

moq_auth/
client.rs

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
10/// Every request is bounded so a hung server refuses rather than parks the session.
11const TIMEOUT: Duration = Duration::from_secs(10);
12
13/// The longest a failed re-check waits before trying again.
14const BACKOFF_MAX: Duration = Duration::from_secs(60);
15
16/// The HTTP side of the contract: one JSON POST per event to an auth server.
17///
18/// `connect` admits a session and hands back the [`lease::Consumer`] it holds; a task
19/// behind it re-POSTs `revalidate` on the grant's cadence, applies each reply, revokes
20/// when the server refuses, answers an invalid grant, or the grant expires, and POSTs
21/// `end` when the session closes.
22#[derive(Clone)]
23pub struct Client {
24	http: reqwest::Client,
25	url: Url,
26}
27
28impl Client {
29	/// A client for the server at `url`.
30	///
31	/// `http://` is accepted for a loopback host only; `https://` presents `tls`, the
32	/// caller's client identity and roots; `unix://` speaks HTTP over the socket at
33	/// the URL's path. Anything else is refused here rather than at the first session.
34	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				// The socket is the transport; the request target is the server's root.
58				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	/// Admit a session: POST `connect`, validate the reply, and return the lease the
71	/// session holds. The session reports totals through [`lease::Consumer::close`].
72	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	/// One POST: a 2xx with a valid grant admits, a 401 or 403 refuses, a 2xx
91	/// whose grant fails validation is that error, and everything else is an
92	/// outage the caller decides about.
93	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			// An empty grant is a refusal, not a malformed answer.
105			Error::UselessGrant => Error::Refused,
106			other => other,
107		})?;
108		Ok(grant)
109	}
110}
111
112/// The task behind a lease: re-checks on cadence and reports the end.
113struct 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		// The session is already gone; nothing to do with a failure but say so.
131		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	/// Re-check until the lease ends, returning why it did and the totals the session reported.
137	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		// The re-check in flight, kept out of the select so expiry and the session's
145		// close are still polled while a stalled server holds the reply.
146		let mut inflight: Option<Pin<Box<dyn Future<Output = crate::Result<Grant>> + Send>>> = None;
147		// A nudge while a re-check is in flight: POST once more when the reply lands,
148		// so a ban set after this request left is not missed until the next cadence.
149		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							// Evidence of nothing: the grant stands until `expires`.
190							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					// One re-check at a time; the reply schedules the next.
204					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
239/// Exponential backoff from one second, capped by the cadence and [`BACKOFF_MAX`],
240/// jittered by up to a quarter so a fleet does not retry in lockstep.
241fn 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	/// Records every request body the server saw, in order.
269	#[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		/// Wait until the requests seen so far satisfy `done`.
285		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		/// Wait for the background `end` POST and return it.
292		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	/// The wire carries whole seconds, so a cadence under one second would serialize
322	/// as zero and be refused; these tests run on a one-second cadence.
323	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(&nothing).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	/// Backoff leaves the driver in this same state (nothing in flight, a timer armed),
581	/// so this covers a nudge during backoff too.
582	#[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		// `changed` jumps to the latest grant epoch, so two replies can arrive as one
627		// observation. The request log does not coalesce.
628		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}