Skip to main content

tollgate_server/
google.rs

1//! Google service-account ID tokens. Key discovery stays off the HTTP handlers;
2//! an unknown key or an expired key set fails closed until refresh succeeds.
3
4use jiff::Timestamp;
5use jsonwebtoken::{Algorithm, DecodingKey, Validation, decode, decode_header};
6use serde::Deserialize;
7use sha2::{Digest, Sha256};
8use std::collections::HashMap;
9use tollgate_auth::{CredentialVerifier, Verified};
10use tollgate_core::Principal;
11
12use crate::security::SecurityError;
13
14/// Fixed issuer, RS256 only, explicit audience. `sub` is the stable Google
15/// account identifier; mutable email addresses do not grant a role.
16pub struct GoogleVerifier {
17    keys: HashMap<String, DecodingKey>,
18    validation: Validation,
19    usable_until: Timestamp,
20}
21
22impl GoogleVerifier {
23    /// Builds a verifier for ID tokens issued to `audience`, from Google's
24    /// published JSON Web Key Set.
25    ///
26    /// `usable_until` is when the key set stops being trusted. Every
27    /// [`Verified`] this verifier returns expires at the earlier of the
28    /// token's `exp` and this instant, so a stale key set cannot extend a
29    /// token's validity. The caller compares that expiry against its own
30    /// clock; this verifier reads no clock.
31    ///
32    /// A token verifies only with an RS256 signature under a key named by its
33    /// `kid`, a Google issuer (`https://accounts.google.com` or
34    /// `accounts.google.com`), exactly this audience, a nonempty `sub`, no
35    /// `nbf`, and `iat` before `exp`.
36    ///
37    /// # Errors
38    ///
39    /// Returns a [`SecurityError`] for an empty audience, malformed JSON, an
40    /// empty key set, a duplicate `kid`, or any key that is not an identified
41    /// RSA `RS256` signing key.
42    pub fn from_jwks(
43        audience: &str,
44        jwks: &[u8],
45        usable_until: Timestamp,
46    ) -> Result<Self, SecurityError> {
47        #[derive(Deserialize)]
48        struct Keys {
49            keys: Vec<Key>,
50        }
51        #[derive(Deserialize)]
52        struct Key {
53            kid: String,
54            kty: String,
55            alg: String,
56            #[serde(rename = "use")]
57            usage: String,
58            n: String,
59            e: String,
60        }
61        if audience.is_empty() {
62            return Err(SecurityError("Google audience must not be empty"));
63        }
64        let parsed: Keys = serde_json::from_slice(jwks)
65            .map_err(|_| SecurityError("invalid Google signing key set"))?;
66        let mut keys = HashMap::new();
67        for key in parsed.keys {
68            if key.kid.is_empty() || key.kty != "RSA" || key.alg != "RS256" || key.usage != "sig" {
69                return Err(SecurityError(
70                    "Google signing key is not an identified RS256 signing key",
71                ));
72            }
73            let decoded = DecodingKey::from_rsa_components(&key.n, &key.e)
74                .map_err(|_| SecurityError("invalid Google RSA key"))?;
75            if keys.insert(key.kid, decoded).is_some() {
76                return Err(SecurityError("duplicate Google signing key identifier"));
77            }
78        }
79        if keys.is_empty() {
80            return Err(SecurityError("Google signing key set is empty"));
81        }
82        let mut validation = Validation::new(Algorithm::RS256);
83        validation.set_issuer(&["https://accounts.google.com", "accounts.google.com"]);
84        validation.set_audience(&[audience]);
85        validation.set_required_spec_claims(&["iss", "sub", "aud", "exp"]);
86        // Return expiry as verified evidence; the owning server compares it to
87        // its explicit Clock exactly once, with no hidden library clock/skew.
88        validation.validate_exp = false;
89        validation.leeway = 0;
90        Ok(Self {
91            keys,
92            validation,
93            usable_until,
94        })
95    }
96
97    /// The [`Principal`] a token with this `sub` claim verifies as: a
98    /// domain-separated SHA-256 digest of the subject, truncated to 128 bits.
99    ///
100    /// Map a service account's numeric unique ID through this to give it a
101    /// role with [`SecurityPolicy::with_bearer`]. The email claim plays no
102    /// part.
103    ///
104    /// [`SecurityPolicy::with_bearer`]: crate::security::SecurityPolicy::with_bearer
105    pub fn principal(subject: &str) -> Principal {
106        let mut digest = Sha256::new();
107        digest.update(b"tollgate-control:google-sub:");
108        digest.update(subject.as_bytes());
109        let digest = digest.finalize();
110        let mut principal = [0; 16];
111        principal.copy_from_slice(&digest[..16]);
112        Principal(u128::from_be_bytes(principal))
113    }
114}
115
116impl CredentialVerifier for GoogleVerifier {
117    fn verify(&self, credential: &[u8]) -> Option<Verified> {
118        #[derive(Clone, Deserialize)]
119        struct Claims {
120            sub: String,
121            exp: i64,
122            iat: i64,
123            nbf: Option<i64>,
124        }
125        let token = std::str::from_utf8(credential).ok()?;
126        let header = decode_header(token).ok()?;
127        if header.alg != Algorithm::RS256 {
128            return None;
129        }
130        let key = self.keys.get(header.kid.as_ref()?)?;
131        let claims = decode::<Claims>(token, key, &self.validation).ok()?.claims;
132        // Google's service-account tokens carry iat/exp, not nbf. Do not
133        // silently accept a future-validity restriction this seam cannot carry.
134        if claims.sub.is_empty() || claims.nbf.is_some() || claims.iat >= claims.exp {
135            return None;
136        }
137        let expiry = Timestamp::from_second(claims.exp)
138            .ok()?
139            .min(self.usable_until);
140        Some(Verified::until(Self::principal(&claims.sub), expiry))
141    }
142}
143
144/// The loader fixes Google's published endpoint; the private parameter lets
145/// transport tests exercise this same bounded fetch against a local fixture.
146/// Diagnostics never contain a token or response body.
147pub(crate) async fn fetch_keys(
148    endpoint: &str,
149) -> Result<(Vec<u8>, std::time::Duration), SecurityError> {
150    let client = reqwest::Client::builder()
151        .use_rustls_tls()
152        .no_proxy()
153        .redirect(reqwest::redirect::Policy::none())
154        .connect_timeout(std::time::Duration::from_secs(2))
155        .timeout(std::time::Duration::from_secs(5))
156        .build()
157        .map_err(|_| SecurityError("cannot configure signing-key client"))?;
158    let mut response = client
159        .get(endpoint)
160        .send()
161        .await
162        .map_err(|_| SecurityError("Google signing-key fetch failed"))?;
163    if !response.status().is_success() {
164        return Err(SecurityError("Google signing-key endpoint refused refresh"));
165    }
166    let lifetime = key_cache_lifetime(response.headers())?;
167    // Much larger than Google's rotating RSA set; this is a transport envelope,
168    // not a limit on the number of accounts or trusted service identities.
169    let mut bytes = Vec::new();
170    while let Some(chunk) = response
171        .chunk()
172        .await
173        .map_err(|_| SecurityError("Google signing-key response interrupted"))?
174    {
175        if bytes.len() + chunk.len() > 1024 * 1024 {
176            return Err(SecurityError("Google signing-key response exceeds 1 MiB"));
177        }
178        bytes.extend_from_slice(&chunk);
179    }
180    Ok((bytes, lifetime))
181}
182
183fn key_cache_lifetime(
184    headers: &reqwest::header::HeaderMap,
185) -> Result<std::time::Duration, SecurityError> {
186    let mut maximum: Option<u64> = None;
187    for value in headers.get_all(reqwest::header::CACHE_CONTROL) {
188        let value = value
189            .to_str()
190            .map_err(|_| SecurityError("invalid signing-key cache policy"))?;
191        for directive in value.split(',').map(str::trim) {
192            if directive.eq_ignore_ascii_case("no-store")
193                || directive.eq_ignore_ascii_case("no-cache")
194            {
195                return Ok(std::time::Duration::ZERO);
196            }
197            if let Some((name, value)) = directive.split_once('=')
198                && name.trim().eq_ignore_ascii_case("max-age")
199            {
200                let seconds: u64 = value
201                    .trim()
202                    .trim_matches('"')
203                    .parse()
204                    .map_err(|_| SecurityError("invalid signing-key max-age"))?;
205                maximum = Some(maximum.map_or(seconds, |old| old.min(seconds)));
206            }
207        }
208    }
209    let age = headers
210        .get(reqwest::header::AGE)
211        .map(|value| {
212            value
213                .to_str()
214                .ok()
215                .and_then(|value| value.parse::<u64>().ok())
216                .ok_or(SecurityError("invalid signing-key age"))
217        })
218        .transpose()?
219        .unwrap_or(0);
220    Ok(std::time::Duration::from_secs(
221        maximum.unwrap_or(300).saturating_sub(age).min(3600),
222    ))
223}
224
225#[cfg(test)]
226mod tests {
227    use super::*;
228
229    #[test]
230    fn every_signing_key_must_independently_name_an_rs256_signature_key() {
231        let fixture: serde_json::Value =
232            serde_json::from_str(include_str!("../tests/fixtures/google-tokens.json")).unwrap();
233        let good = fixture["jwks"].clone();
234        let until = Timestamp::from_second(1_800_000_000).unwrap();
235        let decode = |audience: &str, value: &serde_json::Value| {
236            GoogleVerifier::from_jwks(audience, &serde_json::to_vec(value).unwrap(), until)
237        };
238        assert!(decode("audience", &good).is_ok());
239        assert!(decode("", &good).is_err());
240        for (field, value) in [
241            ("kid", ""),
242            ("kty", "EC"),
243            ("alg", "RS512"),
244            ("use", "enc"),
245            ("n", "@"),
246            ("e", "@"),
247        ] {
248            let mut invalid = good.clone();
249            invalid["keys"][0][field] = value.into();
250            assert!(decode("audience", &invalid).is_err(), "field={field}");
251        }
252        let key = good["keys"][0].clone();
253        assert!(decode("audience", &serde_json::json!({"keys":[key.clone(),key]})).is_err());
254        assert!(decode("audience", &serde_json::json!({"keys":[]})).is_err());
255        assert!(GoogleVerifier::from_jwks("audience", b"{", until).is_err());
256    }
257
258    #[tokio::test]
259    async fn signing_key_transport_enforces_status_cache_policy_and_complete_body_bounds() {
260        use axum::http::{HeaderMap, StatusCode, header};
261        use std::sync::{Arc, Mutex};
262        use std::time::Duration;
263        let response = Arc::new(Mutex::new((
264            StatusCode::OK,
265            HeaderMap::new(),
266            Vec::<u8>::new(),
267        )));
268        let replies = response.clone();
269        let app = axum::Router::new().route(
270            "/certs",
271            axum::routing::get(move || {
272                let replies = replies.clone();
273                async move { replies.lock().unwrap().clone() }
274            }),
275        );
276        let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
277        let endpoint = format!("http://{}/certs", listener.local_addr().unwrap());
278        let (stop, stopped) = tokio::sync::oneshot::channel();
279        let server = tokio::spawn(async move {
280            axum::serve(listener, app)
281                .with_graceful_shutdown(async {
282                    stopped.await.unwrap();
283                })
284                .await
285                .unwrap();
286        });
287        let mut headers = HeaderMap::new();
288        headers.insert(header::CACHE_CONTROL, "max-age=600".parse().unwrap());
289        headers.insert(header::AGE, "100".parse().unwrap());
290        for length in [0, 1, 1024, 1024 * 1024, 1024 * 1024 + 1, 2 * 1024 * 1024] {
291            let body = vec![b'k'; length];
292            *response.lock().unwrap() = (StatusCode::OK, headers.clone(), body.clone());
293            let result = fetch_keys(&endpoint).await;
294            if length <= 1024 * 1024 {
295                let (received, lifetime) = result.unwrap();
296                assert_eq!(received, body);
297                assert_eq!(lifetime, Duration::from_secs(500));
298            } else {
299                assert_eq!(
300                    result.unwrap_err(),
301                    SecurityError("Google signing-key response exceeds 1 MiB")
302                );
303            }
304        }
305        for status in [
306            StatusCode::UNAUTHORIZED,
307            StatusCode::SERVICE_UNAVAILABLE,
308            StatusCode::FOUND,
309        ] {
310            *response.lock().unwrap() = (status, HeaderMap::new(), b"refused".to_vec());
311            assert_eq!(
312                fetch_keys(&endpoint).await.unwrap_err(),
313                SecurityError("Google signing-key endpoint refused refresh")
314            );
315        }
316        headers.insert(header::CACHE_CONTROL, "max-age=invalid".parse().unwrap());
317        *response.lock().unwrap() = (StatusCode::OK, headers, b"invalid cache".to_vec());
318        assert!(fetch_keys(&endpoint).await.is_err());
319        stop.send(()).unwrap();
320        server.await.unwrap();
321
322        use tokio::io::{AsyncReadExt, AsyncWriteExt};
323        let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
324        let endpoint = format!("http://{}/certs", listener.local_addr().unwrap());
325        let server = tokio::spawn(async move {
326            let (mut stream, _) = listener.accept().await.unwrap();
327            let mut request = Vec::new();
328            while !request.ends_with(b"\r\n\r\n") {
329                request.push(stream.read_u8().await.unwrap());
330            }
331            stream
332                .write_all(
333                    b"HTTP/1.1 200 OK\r\nContent-Length: 100\r\nConnection: close\r\n\r\ntruncated",
334                )
335                .await
336                .unwrap();
337        });
338        assert_eq!(
339            fetch_keys(&endpoint).await.unwrap_err(),
340            SecurityError("Google signing-key response interrupted")
341        );
342        server.await.unwrap();
343    }
344    #[test]
345    fn signing_keys_never_outlive_the_issuer_cache_policy_or_one_hour() {
346        for (policy, age, seconds) in [
347            ("public, max-age=30000", "0", 3600),
348            ("max-age=600", "100", 500),
349            ("max-age=10", "20", 0),
350            ("no-cache, max-age=600", "0", 0),
351            ("no-store", "0", 0),
352            ("max-age=\"100\"", "0", 100),
353        ] {
354            let mut headers = reqwest::header::HeaderMap::new();
355            headers.insert(reqwest::header::CACHE_CONTROL, policy.parse().unwrap());
356            headers.insert(reqwest::header::AGE, age.parse().unwrap());
357            assert_eq!(key_cache_lifetime(&headers).unwrap().as_secs(), seconds);
358        }
359        assert_eq!(
360            key_cache_lifetime(&reqwest::header::HeaderMap::new())
361                .unwrap()
362                .as_secs(),
363            300
364        );
365        let mut invalid = reqwest::header::HeaderMap::new();
366        invalid.insert(
367            reqwest::header::CACHE_CONTROL,
368            "max-age=invalid".parse().unwrap(),
369        );
370        assert!(key_cache_lifetime(&invalid).is_err());
371    }
372}