1use std::collections::HashMap;
12use std::sync::{Arc, Mutex};
13use std::time::{Duration, Instant, SystemTime, UNIX_EPOCH};
14
15use base64::Engine;
16use base64::engine::general_purpose::URL_SAFE_NO_PAD;
17use serde::Deserialize;
18use serde_json::Value;
19
20use crate::error::{Error, Result};
21
22pub const ASSERTION_HEADER: &str = "Cf-Access-Jwt-Assertion";
23const KEY_CACHE_TTL: Duration = Duration::from_secs(3600);
24pub const REFETCH_MIN: Duration = Duration::from_secs(10);
26const LEEWAY_SECS: f64 = 30.0;
27const MAX_JWKS_BYTES: u64 = 1 << 20;
28const MAX_TOKEN_BYTES: usize = 16 * 1024;
29
30pub type JwksFetcher = Arc<dyn Fn(&str) -> std::result::Result<Vec<u8>, String> + Send + Sync>;
32
33#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize)]
35pub struct Identity {
36 pub email: Option<String>,
38 pub sub: String,
39 pub common_name: Option<String>,
41}
42
43impl Identity {
44 pub fn name(&self) -> &str {
46 self.email
47 .as_deref()
48 .or(self.common_name.as_deref())
49 .unwrap_or(&self.sub)
50 }
51
52 pub fn is_service_token(&self) -> bool {
53 self.email.is_none() && self.common_name.is_some()
54 }
55}
56
57#[derive(Debug, Clone, PartialEq, Eq)]
59pub struct Denied(pub String);
60
61impl std::fmt::Display for Denied {
62 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
63 f.write_str(&self.0)
64 }
65}
66
67fn deny<T>(msg: impl Into<String>) -> std::result::Result<T, Denied> {
68 Err(Denied(msg.into()))
69}
70
71#[derive(Clone)]
72struct RsaKey {
73 n: Vec<u8>,
74 e: Vec<u8>,
75}
76
77#[derive(Default)]
78struct KeyCache {
79 keys: HashMap<String, RsaKey>,
80 expires: Option<Instant>,
81 last_fetch: Option<Instant>,
82}
83
84pub struct AccessValidator {
86 issuer: String,
87 audience: String,
88 certs_url: String,
89 fetcher: JwksFetcher,
90 cache: Mutex<KeyCache>,
91}
92
93impl std::fmt::Debug for AccessValidator {
94 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
95 f.debug_struct("AccessValidator")
96 .field("issuer", &self.issuer)
97 .field("audience", &self.audience)
98 .finish()
99 }
100}
101
102impl AccessValidator {
103 pub fn new(team_domain: &str, audience: &str) -> Result<Self> {
107 let issuer = normalize_team_domain(team_domain)?;
108 let audience = audience.trim();
109 if audience.is_empty() {
110 return Err(Error::invalid(
111 "Cloudflare Access team domain and audience are both required",
112 ));
113 }
114 Ok(AccessValidator {
115 certs_url: format!("{issuer}/cdn-cgi/access/certs"),
116 issuer,
117 audience: audience.to_string(),
118 fetcher: Arc::new(fetch_https),
119 cache: Mutex::new(KeyCache::default()),
120 })
121 }
122
123 pub fn with_fetcher(mut self, fetcher: JwksFetcher) -> Self {
125 self.fetcher = fetcher;
126 self
127 }
128
129 pub fn issuer(&self) -> &str {
130 &self.issuer
131 }
132
133 pub fn audience(&self) -> &str {
134 &self.audience
135 }
136
137 pub fn certs_url(&self) -> &str {
138 &self.certs_url
139 }
140
141 pub fn validate(&self, token: &str) -> std::result::Result<Identity, Denied> {
144 self.validate_at(token, unix_now())
145 }
146
147 fn validate_at(&self, token: &str, now: f64) -> std::result::Result<Identity, Denied> {
148 let token = token.trim();
149 if token.is_empty() {
150 return deny("missing assertion");
151 }
152 if token.len() > MAX_TOKEN_BYTES {
153 return deny("assertion too large");
154 }
155 let mut parts = token.split('.');
156 let (Some(h64), Some(p64), Some(s64), None) =
157 (parts.next(), parts.next(), parts.next(), parts.next())
158 else {
159 return deny("assertion is not a compact JWS");
160 };
161 let header: Header = decode_json(h64, "header")?;
162 if header.alg != "RS256" {
163 return deny(format!("unexpected signing algorithm {:?}", header.alg));
164 }
165 if header.crit.is_some() {
166 return deny("unsupported critical header");
167 }
168 let kid = match header.kid.as_deref() {
169 Some(k) if !k.is_empty() => k,
170 _ => return deny("assertion has no key id"),
171 };
172 let sig = URL_SAFE_NO_PAD
173 .decode(s64)
174 .map_err(|_| Denied("signature is not base64url".into()))?;
175 let key = self.key(kid)?;
176 let signed = &token[..h64.len() + 1 + p64.len()];
177 ring::signature::RsaPublicKeyComponents {
178 n: &key.n,
179 e: &key.e,
180 }
181 .verify(
182 &ring::signature::RSA_PKCS1_2048_8192_SHA256,
183 signed.as_bytes(),
184 &sig,
185 )
186 .map_err(|_| Denied("bad signature".into()))?;
187
188 let c: Claims = decode_json(p64, "claims")?;
189 if c.iss.as_deref() != Some(self.issuer.as_str()) {
190 return deny(format!("issuer {:?} is not {:?}", c.iss, self.issuer));
191 }
192 let aud_ok = match &c.aud {
193 Some(Value::String(a)) => *a == self.audience,
194 Some(Value::Array(a)) => a.iter().any(|v| v.as_str() == Some(&self.audience)),
195 _ => false,
196 };
197 if !aud_ok {
198 return deny("audience does not match");
199 }
200 match c.exp {
201 None => return deny("assertion has no expiry"),
202 Some(exp) if now > exp + LEEWAY_SECS => return deny("assertion expired"),
203 _ => {}
204 }
205 if c.nbf.is_some_and(|nbf| now + LEEWAY_SECS < nbf) {
206 return deny("assertion not yet valid");
207 }
208 if c.iat.is_some_and(|iat| now + LEEWAY_SECS < iat) {
209 return deny("assertion issued in the future");
210 }
211 let nonempty = |s: Option<String>| s.filter(|s| !s.is_empty());
212 Ok(Identity {
213 email: nonempty(c.email),
214 sub: c.sub.unwrap_or_default(),
215 common_name: nonempty(c.common_name),
216 })
217 }
218
219 fn key(&self, kid: &str) -> std::result::Result<RsaKey, Denied> {
220 let mut c = self.cache.lock().unwrap_or_else(|p| p.into_inner());
223 let now = Instant::now();
224 let fresh = c.expires.is_some_and(|t| now < t);
225 if let Some(k) = c.keys.get(kid).filter(|_| fresh) {
226 return Ok(k.clone());
227 }
228 if c.last_fetch
229 .is_some_and(|t| now.duration_since(t) < REFETCH_MIN)
230 {
231 return deny(if fresh {
232 format!("signing key {kid:?} is not published")
233 } else {
234 "Access certs are unavailable".to_string()
235 });
236 }
237 c.last_fetch = Some(now);
238 let keys = (self.fetcher)(&self.certs_url)
239 .map_err(|e| Denied(format!("fetch Access certs: {e}")))
240 .and_then(|b| parse_jwks(&b))?;
241 c.keys = keys;
242 c.expires = Some(now + KEY_CACHE_TTL);
243 c.keys
244 .get(kid)
245 .cloned()
246 .ok_or_else(|| Denied(format!("signing key {kid:?} is not published")))
247 }
248}
249
250#[derive(Deserialize)]
251struct Header {
252 #[serde(default)]
253 alg: String,
254 kid: Option<String>,
255 crit: Option<Value>,
256}
257
258#[derive(Deserialize)]
259struct Claims {
260 iss: Option<String>,
261 aud: Option<Value>,
262 exp: Option<f64>,
263 nbf: Option<f64>,
264 iat: Option<f64>,
265 sub: Option<String>,
266 email: Option<String>,
267 common_name: Option<String>,
268}
269
270fn decode_json<T: serde::de::DeserializeOwned>(
271 part: &str,
272 what: &str,
273) -> std::result::Result<T, Denied> {
274 let bytes = URL_SAFE_NO_PAD
275 .decode(part)
276 .map_err(|_| Denied(format!("{what} is not base64url")))?;
277 serde_json::from_slice(&bytes).map_err(|e| Denied(format!("{what}: {e}")))
278}
279
280fn parse_jwks(body: &[u8]) -> std::result::Result<HashMap<String, RsaKey>, Denied> {
281 #[derive(Deserialize)]
282 struct Doc {
283 keys: Vec<Jwk>,
284 }
285 #[derive(Deserialize)]
286 struct Jwk {
287 #[serde(default)]
288 kty: String,
289 #[serde(default)]
290 kid: String,
291 #[serde(default)]
292 n: String,
293 #[serde(default)]
294 e: String,
295 }
296 let doc: Doc =
297 serde_json::from_slice(body).map_err(|e| Denied(format!("decode Access certs: {e}")))?;
298 let mut keys = HashMap::new();
299 for k in doc.keys {
300 if k.kty != "RSA" || k.kid.is_empty() {
301 continue;
302 }
303 let bad = || Denied(format!("Access signing key {:?} is malformed", k.kid));
304 let n = URL_SAFE_NO_PAD.decode(&k.n).map_err(|_| bad())?;
305 let e = URL_SAFE_NO_PAD.decode(&k.e).map_err(|_| bad())?;
306 if n.is_empty() || e.is_empty() || e.len() > 4 {
307 return Err(bad());
308 }
309 keys.insert(k.kid, RsaKey { n, e });
310 }
311 if keys.is_empty() {
312 return deny("Access certs document contains no RSA signing keys");
313 }
314 Ok(keys)
315}
316
317pub fn normalize_team_domain(team_domain: &str) -> Result<String> {
319 let t = team_domain.trim();
320 if t.is_empty() {
321 return Err(Error::invalid(
322 "Cloudflare Access team domain and audience are both required",
323 ));
324 }
325 let with_scheme = if t.contains("://") {
326 t.to_string()
327 } else {
328 format!("https://{t}")
329 };
330 let bad = || Error::invalid(format!("invalid Cloudflare Access team domain {t:?}"));
331 let (scheme, rest) = with_scheme.split_once("://").ok_or_else(bad)?;
332 let scheme = scheme.to_ascii_lowercase();
333 if rest.contains(['?', '#', '@']) || !(scheme == "https" || scheme == "http") {
334 return Err(bad());
335 }
336 let (authority, path) = match rest.find('/') {
337 Some(i) => (&rest[..i], &rest[i..]),
338 None => (rest, ""),
339 };
340 let host = if let Some(v6) = authority.strip_prefix('[') {
341 v6.split(']').next().unwrap_or_default()
342 } else {
343 authority.split(':').next().unwrap_or_default()
344 };
345 if host.is_empty() {
346 return Err(bad());
347 }
348 if scheme != "https" && !matches!(host, "localhost" | "127.0.0.1" | "::1") {
349 return Err(Error::invalid(
350 "Cloudflare Access team domain must use https",
351 ));
352 }
353 Ok(format!(
354 "{scheme}://{authority}{}",
355 path.trim_end_matches('/')
356 ))
357}
358
359fn unix_now() -> f64 {
360 SystemTime::now()
361 .duration_since(UNIX_EPOCH)
362 .map(|d| d.as_secs_f64())
363 .unwrap_or(0.0)
364}
365
366fn fetch_https(url: &str) -> std::result::Result<Vec<u8>, String> {
367 let agent: ureq::Agent = ureq::Agent::config_builder()
368 .timeout_global(Some(Duration::from_secs(5)))
369 .user_agent(concat!("isb/", env!("CARGO_PKG_VERSION")))
370 .build()
371 .into();
372 let mut resp = agent.get(url).call().map_err(|e| e.to_string())?;
373 resp.body_mut()
374 .with_config()
375 .limit(MAX_JWKS_BYTES)
376 .read_to_vec()
377 .map_err(|e| e.to_string())
378}
379
380#[cfg(any(test, feature = "test-support"))]
382#[doc(hidden)]
383pub mod tests {
384 use super::*;
385 use ring::signature::{RSA_PKCS1_SHA256, RsaKeyPair};
386 use serde_json::json;
387 use std::sync::atomic::{AtomicUsize, Ordering};
388
389 pub const TEAM: &str = "https://team.cloudflareaccess.com";
390 pub const AUD: &str = "aud-tag-123";
391 pub const KID: &str = "test-kid";
392
393 fn keypair() -> RsaKeyPair {
394 RsaKeyPair::from_pkcs8(include_bytes!("testdata/access_test_key.pk8")).unwrap()
395 }
396
397 pub fn jwks() -> Vec<u8> {
398 let kp = keypair();
399 let p = ring::rsa::PublicKeyComponents::<Vec<u8>>::from(kp.public());
400 let b = |x: &[u8]| URL_SAFE_NO_PAD.encode(x);
401 let (n, e) = (&p.n, &p.e);
402 serde_json::to_vec(&json!({"keys": [
403 {"kty": "EC", "kid": "ignored", "crv": "P-256"},
404 {"kty": "RSA", "kid": KID, "alg": "RS256", "use": "sig", "n": b(n), "e": b(e)},
405 ]}))
406 .unwrap()
407 }
408
409 pub fn sign(header: &Value, claims: &Value) -> String {
410 let enc = |v: &Value| URL_SAFE_NO_PAD.encode(serde_json::to_vec(v).unwrap());
411 let input = format!("{}.{}", enc(header), enc(claims));
412 let kp = keypair();
413 let mut sig = vec![0u8; kp.public().modulus_len()];
414 kp.sign(
415 &RSA_PKCS1_SHA256,
416 &ring::rand::SystemRandom::new(),
417 input.as_bytes(),
418 &mut sig,
419 )
420 .unwrap();
421 format!("{input}.{}", URL_SAFE_NO_PAD.encode(sig))
422 }
423
424 pub fn claims() -> Value {
425 let now = unix_now() as i64;
426 json!({"iss": TEAM, "aud": [AUD], "exp": now + 300, "iat": now, "nbf": now,
427 "sub": "user-1", "email": "alice@example.com"})
428 }
429
430 pub fn header() -> Value {
431 json!({"alg": "RS256", "kid": KID, "typ": "JWT"})
432 }
433
434 pub fn validator() -> (AccessValidator, Arc<AtomicUsize>) {
435 let fetches = Arc::new(AtomicUsize::new(0));
436 let f = fetches.clone();
437 let v = AccessValidator::new("team.cloudflareaccess.com/", AUD)
438 .unwrap()
439 .with_fetcher(Arc::new(move |url: &str| {
440 assert_eq!(
441 url,
442 "https://team.cloudflareaccess.com/cdn-cgi/access/certs"
443 );
444 f.fetch_add(1, Ordering::SeqCst);
445 Ok(jwks())
446 }));
447 (v, fetches)
448 }
449
450 #[cfg(test)]
451 fn with(mut v: Value, k: &str, x: Value) -> Value {
452 v[k] = x;
453 v
454 }
455
456 #[test]
457 fn valid_token_yields_identity_and_caches_keys() {
458 let (v, fetches) = validator();
459 let id = v.validate(&sign(&header(), &claims())).unwrap();
460 assert_eq!(id.name(), "alice@example.com");
461 assert_eq!(id.sub, "user-1");
462 v.validate(&sign(&header(), &claims())).unwrap();
463 assert_eq!(fetches.load(Ordering::SeqCst), 1, "keys cached");
464 v.validate(&sign(&header(), &with(claims(), "aud", json!(AUD))))
466 .unwrap();
467 }
468
469 #[test]
470 fn service_token_identity() {
471 let (v, _) = validator();
472 let mut c = claims();
473 c.as_object_mut().unwrap().remove("email");
474 c["sub"] = json!("");
475 c["common_name"] = json!("abc.access");
476 let id = v.validate(&sign(&header(), &c)).unwrap();
477 assert!(id.is_service_token());
478 assert_eq!(id.name(), "abc.access");
479 }
480
481 #[test]
482 fn rejects_bad_claims() {
483 let (v, _) = validator();
484 let now = unix_now() as i64;
485 let cases = [
486 ("wrong aud", with(claims(), "aud", json!(["other"]))),
487 (
488 "wrong iss",
489 with(claims(), "iss", json!("https://evil.cloudflareaccess.com")),
490 ),
491 ("expired", with(claims(), "exp", json!(now - 31))),
492 ("no exp", {
493 let mut c = claims();
494 c.as_object_mut().unwrap().remove("exp");
495 c
496 }),
497 ("nbf future", with(claims(), "nbf", json!(now + 120))),
498 ("iat future", with(claims(), "iat", json!(now + 120))),
499 ];
500 for (what, c) in cases {
501 assert!(v.validate(&sign(&header(), &c)).is_err(), "{what} accepted");
502 }
503 v.validate(&sign(&header(), &with(claims(), "exp", json!(now - 5))))
505 .unwrap();
506 }
507
508 #[test]
509 fn rejects_bad_signature() {
510 let (v, _) = validator();
511 let good = sign(&header(), &claims());
512 let mut parts: Vec<&str> = good.split('.').collect();
513 let forged = URL_SAFE_NO_PAD.encode(
514 serde_json::to_vec(&with(claims(), "email", json!("mallory@example.com"))).unwrap(),
515 );
516 parts[1] = &forged;
517 let e = v.validate(&parts.join(".")).unwrap_err();
518 assert_eq!(e.0, "bad signature");
519 assert!(v.validate(&format!("{good}x")).is_err());
520 assert!(v.validate("a.b").is_err());
521 assert!(v.validate("").is_err());
522 }
523
524 #[test]
525 fn rejects_other_algorithms() {
526 let (v, fetches) = validator();
527 let enc = |x: &Value| URL_SAFE_NO_PAD.encode(serde_json::to_vec(x).unwrap());
528 let none = format!(
529 "{}.{}.",
530 enc(&json!({"alg": "none", "kid": KID})),
531 enc(&claims())
532 );
533 assert!(v.validate(&none).unwrap_err().0.contains("algorithm"));
534 let hs = sign(&with(header(), "alg", json!("HS256")), &claims());
536 assert!(v.validate(&hs).unwrap_err().0.contains("algorithm"));
537 let crit = sign(&with(header(), "crit", json!(["b64"])), &claims());
538 assert!(v.validate(&crit).is_err());
539 assert_eq!(
540 fetches.load(Ordering::SeqCst),
541 0,
542 "rejected before any fetch"
543 );
544 }
545
546 #[test]
547 fn unknown_kid_refetches_at_most_once_per_interval() {
548 let (v, fetches) = validator();
549 v.validate(&sign(&header(), &claims())).unwrap();
550 for _ in 0..50 {
553 let t = sign(&with(header(), "kid", json!("bogus")), &claims());
554 assert!(v.validate(&t).is_err());
555 }
556 assert_eq!(fetches.load(Ordering::SeqCst), 1);
557 v.cache.lock().unwrap().last_fetch = Some(Instant::now() - REFETCH_MIN * 2);
559 let t = sign(&with(header(), "kid", json!("bogus")), &claims());
560 assert!(v.validate(&t).unwrap_err().0.contains("not published"));
561 assert_eq!(fetches.load(Ordering::SeqCst), 2);
562 v.validate(&sign(&header(), &claims())).unwrap();
564 assert_eq!(fetches.load(Ordering::SeqCst), 2);
565 }
566
567 #[test]
568 fn fetch_failure_fails_closed() {
569 let v = AccessValidator::new(TEAM, AUD)
570 .unwrap()
571 .with_fetcher(Arc::new(|_: &str| Err("offline".to_string())));
572 let e = v.validate(&sign(&header(), &claims())).unwrap_err();
573 assert!(e.0.contains("offline"), "{e}");
574 }
575
576 #[test]
578 #[ignore = "network"]
579 fn fetches_a_real_jwks() {
580 let body = fetch_https("https://www.googleapis.com/oauth2/v3/certs").unwrap();
581 assert!(!parse_jwks(&body).unwrap().is_empty());
582 }
583
584 #[test]
585 fn team_domain_normalization() {
586 let n = |s: &str| normalize_team_domain(s);
587 assert_eq!(n("team.cloudflareaccess.com").unwrap(), TEAM);
588 assert_eq!(n(" https://team.cloudflareaccess.com/ ").unwrap(), TEAM);
589 assert_eq!(
590 n("http://127.0.0.1:8080/").unwrap(),
591 "http://127.0.0.1:8080"
592 );
593 assert_eq!(n("http://[::1]:9").unwrap(), "http://[::1]:9");
594 assert!(n("http://team.cloudflareaccess.com").is_err());
595 assert!(n("https://team.example?x=1").is_err());
596 assert!(n("ftp://team.example").is_err());
597 assert!(n("https://").is_err());
598 assert!(n("").is_err());
599 assert!(AccessValidator::new(TEAM, " ").is_err());
600 }
601}