1use 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
14pub struct GoogleVerifier {
17 keys: HashMap<String, DecodingKey>,
18 validation: Validation,
19 usable_until: Timestamp,
20}
21
22impl GoogleVerifier {
23 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 validation.validate_exp = false;
89 validation.leeway = 0;
90 Ok(Self {
91 keys,
92 validation,
93 usable_until,
94 })
95 }
96
97 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 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
144pub(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 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}