Skip to main content

isb_server/server/
access.rs

1//! Cloudflare Access JWT validation.
2//!
3//! Behind a tunnel, Access forwards every authenticated request with a signed
4//! assertion in `Cf-Access-Jwt-Assertion`. Checking it at the origin means a
5//! request that reaches the loopback port some other way (a second tunnel, a
6//! local process) is still refused. RS256 only; keys come from the team's
7//! JWKS, cached for an hour, refetched when an unknown `kid` shows up, and
8//! refetched at most once per [`REFETCH_MIN`] so a flood of made-up key ids
9//! cannot be turned into a flood of requests to Cloudflare.
10
11use 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);
24/// The least time between two JWKS fetches.
25pub 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
30/// Fetches the JWKS document at a URL. Injectable so tests run offline.
31pub type JwksFetcher = Arc<dyn Fn(&str) -> std::result::Result<Vec<u8>, String> + Send + Sync>;
32
33/// The verified caller behind an Access assertion.
34#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize)]
35pub struct Identity {
36    /// Users have an email; service tokens do not.
37    pub email: Option<String>,
38    pub sub: String,
39    /// A service token's client id.
40    pub common_name: Option<String>,
41}
42
43impl Identity {
44    /// Email, else service-token common name, else subject.
45    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/// Why an assertion was refused. Logged, never sent to the client.
58#[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
84/// Verifies Access assertions for one application audience.
85pub 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    /// `team_domain` is `team.cloudflareaccess.com` (https:// is implied), and
104    /// must be https unless it is loopback (for tests). `audience` is the
105    /// Access application's AUD tag.
106    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    /// Replace the JWKS fetcher.
124    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    /// Verify an assertion: signature, `iss`, `aud`, `exp` (required), `nbf`
142    /// and `iat`, with 30s of leeway for clock skew.
143    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        // Held across the fetch on purpose: concurrent misses wait for one
221        // fetch instead of each starting their own.
222        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
317/// Normalize a team domain into the issuer string Access puts in `iss`.
318pub 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// Signing helpers for other crates' tests too (the `test-support` feature).
381#[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        // A single audience string is accepted too.
465        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        // Inside the leeway is fine.
504        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        // HS256 is the RSA-public-key-as-HMAC-secret confusion attack.
535        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        // The cache is fresh but this kid is new: one refetch is allowed only
551        // after REFETCH_MIN since the last fetch.
552        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        // Pretend the last fetch was long ago: an unknown kid refetches once.
558        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        // A known kid still works without fetching.
563        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    /// The real HTTPS path (rustls, webpki roots). Run with `--ignored`.
577    #[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}