1use std::time::Duration;
2
3use async_trait::async_trait;
4use reqwest::Url;
5use tracing::debug;
6
7use crate::error::AuthError;
8use crate::oauth_provider::{AuthorizeUrlRequest, OAuthProvider, ProviderExchange};
9use crate::oidc::OidcVerifier;
10use crate::provider_http::build_authorize_url;
11use crate::util::fingerprint;
12
13const GOOGLE_AUTHORIZE_ENDPOINT: &str = "https://accounts.google.com/o/oauth2/v2/auth";
14const GOOGLE_TOKEN_ENDPOINT: &str = "https://oauth2.googleapis.com/token";
15const GOOGLE_JWKS_ENDPOINT: &str = "https://www.googleapis.com/oauth2/v3/certs";
16const GOOGLE_ISSUER: &str = "https://accounts.google.com";
17const GOOGLE_ISSUER_ALT: &str = "accounts.google.com";
21const GOOGLE_HTTP_TIMEOUT: Duration = Duration::from_secs(30);
22
23#[derive(Clone)]
24pub struct GoogleProvider {
25 pub client_id: String,
26 pub client_secret: String,
27 pub redirect_uri: Url,
28 pub scopes: Vec<String>,
29 pub http: reqwest::Client,
30 authorize_endpoint: Url,
31 token_endpoint: Url,
32 verifier: OidcVerifier,
33}
34
35impl std::fmt::Debug for GoogleProvider {
36 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
37 f.debug_struct("GoogleProvider")
38 .field("client_id", &self.client_id)
39 .field("redirect_uri", &self.redirect_uri)
40 .field("scopes", &self.scopes)
41 .finish_non_exhaustive()
42 }
43}
44
45impl GoogleProvider {
46 pub fn new(
47 client_id: String,
48 client_secret: String,
49 redirect_uri: Url,
50 ) -> Result<Self, AuthError> {
51 crate::provider_http::install_rustls_default_once();
52 let http = reqwest::Client::builder()
53 .timeout(GOOGLE_HTTP_TIMEOUT)
54 .build()
55 .map_err(|error| {
56 AuthError::Storage(format!("build google oauth http client: {error}"))
57 })?;
58 let authorize_endpoint = Url::parse(GOOGLE_AUTHORIZE_ENDPOINT).map_err(|error| {
59 AuthError::Config(format!("parse google authorize endpoint: {error}"))
60 })?;
61 let token_endpoint = Url::parse(GOOGLE_TOKEN_ENDPOINT)
62 .map_err(|error| AuthError::Config(format!("parse google token endpoint: {error}")))?;
63 let jwks_endpoint = Url::parse(GOOGLE_JWKS_ENDPOINT)
64 .map_err(|error| AuthError::Config(format!("parse google jwks endpoint: {error}")))?;
65 let verifier = OidcVerifier::new(
66 "google",
67 GOOGLE_ISSUER.to_string(),
68 jwks_endpoint,
69 http.clone(),
70 )
71 .with_alt_issuer(GOOGLE_ISSUER_ALT);
72
73 Ok(Self {
74 client_id,
75 client_secret,
76 redirect_uri,
77 scopes: vec![
78 "openid".to_string(),
79 "email".to_string(),
80 "profile".to_string(),
81 ],
82 http,
83 authorize_endpoint,
84 token_endpoint,
85 verifier,
86 })
87 }
88
89 #[cfg(test)]
90 #[must_use]
91 pub fn with_endpoints(mut self, authorize_endpoint: Url, token_endpoint: Url) -> Self {
92 self.authorize_endpoint = authorize_endpoint;
93 self.token_endpoint = token_endpoint;
94 self
95 }
96
97 #[cfg(test)]
98 #[must_use]
99 pub fn with_jwks_endpoint(mut self, jwks_endpoint: Url) -> Self {
100 self.verifier = self.verifier.with_jwks_endpoint(jwks_endpoint);
101 self
102 }
103}
104
105#[async_trait]
106impl OAuthProvider for GoogleProvider {
107 fn provider_id(&self) -> &'static str {
108 "google"
109 }
110
111 fn callback_path(&self) -> &str {
112 self.redirect_uri.path()
113 }
114
115 fn authorize_url(&self, request: &AuthorizeUrlRequest) -> Result<Url, AuthError> {
116 let scope = self.scopes.join(" ");
117 let url = build_authorize_url(
118 &self.authorize_endpoint,
119 &self.client_id,
120 &self.redirect_uri,
121 &self.scopes,
122 request,
123 &[
124 ("access_type", "offline"),
125 ("include_granted_scopes", "true"),
126 ],
127 );
128 debug!(
129 provider = "google",
130 oauth_state_id = %fingerprint(&request.state),
131 scope = %scope,
132 redirect_uri = %self.redirect_uri,
133 "oauth upstream authorize URL constructed"
134 );
135 Ok(url)
136 }
137
138 async fn exchange_code(
139 &self,
140 code: &str,
141 code_verifier: &str,
142 ) -> Result<ProviderExchange, AuthError> {
143 self.verifier
144 .exchange_code(
145 &self.http,
146 &self.token_endpoint,
147 &self.client_id,
148 &self.client_secret,
149 &self.redirect_uri,
150 code,
151 code_verifier,
152 )
153 .await
154 }
155
156 async fn refresh(&self, refresh_token: &str) -> Result<ProviderExchange, AuthError> {
157 self.verifier
158 .refresh(
159 &self.http,
160 &self.token_endpoint,
161 &self.client_id,
162 &self.client_secret,
163 refresh_token,
164 )
165 .await
166 }
167}
168
169#[cfg(test)]
170mod tests {
171 use base64::Engine;
172 use base64::engine::general_purpose::URL_SAFE_NO_PAD;
173 use jsonwebtoken::{Algorithm, EncodingKey, Header, encode};
174 use rsa::RsaPrivateKey;
175 use rsa::pkcs8::EncodePrivateKey;
176 use rsa::rand_core::{TryCryptoRng, TryRng, UnwrapErr};
177 use rsa::traits::PublicKeyParts;
178 use serde_json::json;
179 use std::sync::OnceLock;
180 use url::Url;
181 use wiremock::matchers::{method, path};
182 use wiremock::{Mock, MockServer, ResponseTemplate};
183
184 use super::{AuthorizeUrlRequest, GoogleProvider};
185 use crate::oauth_provider::OAuthProvider;
186
187 #[test]
188 fn google_authorize_url_includes_offline_access_prompt_and_pkce() {
189 let provider = test_google_provider();
190 let request = sample_request();
191 let url = provider.authorize_url(&request).unwrap();
192 assert!(url.as_str().contains("access_type=offline"));
193 assert!(url.as_str().contains("prompt=consent"));
194 assert!(url.as_str().contains("code_challenge="));
195 }
196
197 #[test]
198 fn google_authorize_url_omits_prompt_when_consent_not_forced() {
199 let provider = test_google_provider();
200 let mut request = sample_request();
201 request.force_consent = false;
202 let url = provider.authorize_url(&request).unwrap();
203 assert!(url.as_str().contains("access_type=offline"));
204 assert!(!url.as_str().contains("prompt="));
205 }
206
207 #[tokio::test]
208 async fn google_exchange_parses_subject_and_refresh_token() {
209 let provider = mocked_google_provider().await;
210 let token = provider.exchange_code("code", "verifier").await.unwrap();
211 assert_eq!(token.subject, "google-subject-123");
212 assert_eq!(token.refresh_token.as_deref(), Some("refresh-token"));
213 }
214
215 #[tokio::test]
216 async fn google_exchange_rejects_unsigned_id_tokens() {
217 let provider = mocked_google_provider_with_id_token(test_id_token()).await;
218 let error = provider
219 .exchange_code("code", "verifier")
220 .await
221 .unwrap_err();
222 assert!(
223 error.to_string().contains("verify google id_token"),
224 "unexpected error: {error}"
225 );
226 }
227
228 #[tokio::test]
229 async fn google_exchange_rejects_wrong_audience_in_id_token() {
230 let provider =
231 mocked_google_provider_with_id_token(signed_test_id_token("other-client", false, true))
232 .await;
233 let error = provider
234 .exchange_code("code", "verifier")
235 .await
236 .unwrap_err();
237 assert!(
238 error.to_string().contains("invalid google id_token"),
239 "unexpected error: {error}"
240 );
241 }
242
243 #[tokio::test]
244 async fn google_exchange_rejects_expired_id_token() {
245 let provider =
246 mocked_google_provider_with_id_token(signed_test_id_token("client-id", true, true))
247 .await;
248 let error = provider
249 .exchange_code("code", "verifier")
250 .await
251 .unwrap_err();
252 assert!(
253 error.to_string().contains("invalid google id_token"),
254 "unexpected error: {error}"
255 );
256 }
257
258 #[tokio::test]
259 async fn google_exchange_rejects_wrong_issuer_in_id_token() {
260 let provider =
261 mocked_google_provider_with_id_token(signed_test_id_token("client-id", false, false))
262 .await;
263 let error = provider
264 .exchange_code("code", "verifier")
265 .await
266 .unwrap_err();
267 assert!(
268 error.to_string().contains("issuer"),
269 "unexpected error: {error}"
270 );
271 }
272
273 #[tokio::test]
279 async fn google_exchange_accepts_bare_form_issuer_in_id_token() {
280 let provider = mocked_google_provider_with_id_token(signed_test_id_token_with_issuer(
281 "client-id",
282 "accounts.google.com",
283 ))
284 .await;
285 let exchange = provider.exchange_code("code", "verifier").await;
286 assert!(
287 exchange.is_ok(),
288 "bare-form issuer must be accepted: {exchange:?}"
289 );
290 assert_eq!(exchange.unwrap().subject, "google-subject-123");
291 }
292
293 #[tokio::test]
294 async fn google_exchange_reuses_cached_jwks() {
295 let server = MockServer::start().await;
296 Mock::given(method("POST"))
297 .and(path("/token"))
298 .respond_with(ResponseTemplate::new(200).set_body_json(json!({
299 "access_token": "google-access-token",
300 "refresh_token": "refresh-token",
301 "expires_in": 3600,
302 "id_token": signed_test_id_token("client-id", false, true),
303 })))
304 .mount(&server)
305 .await;
306 Mock::given(method("GET"))
307 .and(path("/certs"))
308 .respond_with(
309 ResponseTemplate::new(200)
310 .insert_header("Cache-Control", "public, max-age=3600")
311 .set_body_json(test_jwks()),
312 )
313 .mount(&server)
314 .await;
315
316 let provider = test_google_provider()
317 .with_endpoints(
318 server.uri().parse::<Url>().unwrap(),
319 server.uri().parse::<Url>().unwrap().join("/token").unwrap(),
320 )
321 .with_jwks_endpoint(server.uri().parse::<Url>().unwrap().join("/certs").unwrap());
322
323 provider.exchange_code("code-1", "verifier").await.unwrap();
324 provider.exchange_code("code-2", "verifier").await.unwrap();
325
326 let requests = server.received_requests().await.unwrap();
327 let jwks_requests = requests
328 .iter()
329 .filter(|request| request.url.path() == "/certs")
330 .count();
331 assert_eq!(jwks_requests, 1);
332 }
333
334 #[tokio::test]
335 async fn google_exchange_succeeds_on_first_jwks_fetch_with_no_pre_seeded_cache() {
336 let server = MockServer::start().await;
337 Mock::given(method("POST"))
338 .and(path("/token"))
339 .respond_with(ResponseTemplate::new(200).set_body_json(json!({
340 "access_token": "google-access-token",
341 "refresh_token": "refresh-token",
342 "expires_in": 3600,
343 "id_token": signed_test_id_token("client-id", false, true),
344 })))
345 .mount(&server)
346 .await;
347 Mock::given(method("GET"))
348 .and(path("/certs"))
349 .respond_with(ResponseTemplate::new(200).set_body_json(test_jwks()))
350 .mount(&server)
351 .await;
352
353 let provider = test_google_provider()
354 .with_endpoints(
355 server.uri().parse::<Url>().unwrap(),
356 server.uri().parse::<Url>().unwrap().join("/token").unwrap(),
357 )
358 .with_jwks_endpoint(server.uri().parse::<Url>().unwrap().join("/certs").unwrap());
359
360 let exchange = provider.exchange_code("code", "verifier").await.unwrap();
361 assert_eq!(exchange.subject, "google-subject-123");
362
363 let requests = server.received_requests().await.unwrap();
364 let jwks_requests = requests
365 .iter()
366 .filter(|request| request.url.path() == "/certs")
367 .count();
368 assert_eq!(jwks_requests, 1);
369 }
370
371 fn test_google_provider() -> GoogleProvider {
372 GoogleProvider::new(
373 "client-id".to_string(),
374 "client-secret".to_string(),
375 Url::parse("https://lab.example.com/auth/google/callback").unwrap(),
376 )
377 .unwrap()
378 }
379
380 async fn mocked_google_provider() -> MockedGoogleProvider {
381 mocked_google_provider_with_id_token(signed_test_id_token("client-id", false, true)).await
382 }
383
384 struct MockedGoogleProvider {
385 provider: GoogleProvider,
386 _server: MockServer,
387 }
388
389 impl std::ops::Deref for MockedGoogleProvider {
390 type Target = GoogleProvider;
391
392 fn deref(&self) -> &Self::Target {
393 &self.provider
394 }
395 }
396
397 async fn mocked_google_provider_with_id_token(id_token: String) -> MockedGoogleProvider {
398 let server = MockServer::start().await;
399 Mock::given(method("POST"))
400 .and(path("/token"))
401 .respond_with(ResponseTemplate::new(200).set_body_json(json!({
402 "access_token": "google-access-token",
403 "refresh_token": "refresh-token",
404 "expires_in": 3600,
405 "id_token": id_token,
406 })))
407 .mount(&server)
408 .await;
409 Mock::given(method("GET"))
410 .and(path("/certs"))
411 .respond_with(ResponseTemplate::new(200).set_body_json(test_jwks()))
412 .mount(&server)
413 .await;
414
415 let provider = test_google_provider()
416 .with_endpoints(
417 server.uri().parse::<Url>().unwrap(),
418 server.uri().parse::<Url>().unwrap().join("/token").unwrap(),
419 )
420 .with_jwks_endpoint(server.uri().parse::<Url>().unwrap().join("/certs").unwrap());
421
422 MockedGoogleProvider {
423 provider,
424 _server: server,
425 }
426 }
427
428 fn sample_request() -> AuthorizeUrlRequest {
429 AuthorizeUrlRequest {
430 state: "state-123".to_string(),
431 code_challenge: "challenge".to_string(),
432 code_challenge_method: "S256".to_string(),
433 force_consent: true,
434 }
435 }
436
437 fn test_id_token() -> String {
438 let header = URL_SAFE_NO_PAD.encode(br#"{"alg":"none","typ":"JWT"}"#);
439 let payload = URL_SAFE_NO_PAD.encode(br#"{"sub":"google-subject-123"}"#);
440 format!("{header}.{payload}.")
441 }
442
443 fn signed_test_id_token(client_id: &str, expired: bool, valid_issuer: bool) -> String {
444 let issuer = if valid_issuer {
445 "https://accounts.google.com"
446 } else {
447 "https://evil.example.com"
448 };
449 signed_test_id_token_with_issuer_and_expiry(client_id, issuer, expired)
450 }
451
452 fn signed_test_id_token_with_issuer(client_id: &str, issuer: &str) -> String {
453 signed_test_id_token_with_issuer_and_expiry(client_id, issuer, false)
454 }
455
456 fn signed_test_id_token_with_issuer_and_expiry(
457 client_id: &str,
458 issuer: &str,
459 expired: bool,
460 ) -> String {
461 let claims = json!({
462 "iss": issuer,
463 "aud": client_id,
464 "sub": "google-subject-123",
465 "email": "user@example.com",
466 "iat": (unix_now() - 10) as usize,
467 "exp": if expired { (unix_now() - 3600) as usize } else { (unix_now() + 3600) as usize },
468 });
469 let mut header = Header::new(Algorithm::RS256);
470 header.kid = Some("test-kid".to_string());
471 encode(&header, &claims, &test_encoding_key()).unwrap()
472 }
473
474 fn test_jwks() -> serde_json::Value {
475 let key = test_rsa_key();
476 let public_key = key.to_public_key();
477 json!({
478 "keys": [{
479 "kid": "test-kid",
480 "alg": "RS256",
481 "kty": "RSA",
482 "use": "sig",
483 "n": URL_SAFE_NO_PAD.encode(public_key.n_bytes()),
484 "e": URL_SAFE_NO_PAD.encode(public_key.e_bytes()),
485 }]
486 })
487 }
488
489 fn test_rsa_key() -> &'static RsaPrivateKey {
490 static TEST_RSA_KEY: OnceLock<RsaPrivateKey> = OnceLock::new();
491 TEST_RSA_KEY.get_or_init(|| {
492 let mut rng = UnwrapErr(TestRng);
493 RsaPrivateKey::new(&mut rng, 2048).unwrap()
494 })
495 }
496
497 fn test_encoding_key() -> EncodingKey {
498 let pem = test_rsa_key().to_pkcs8_pem(Default::default()).unwrap();
499 EncodingKey::from_rsa_pem(pem.as_bytes()).unwrap()
500 }
501
502 fn unix_now() -> i64 {
503 std::time::SystemTime::now()
504 .duration_since(std::time::UNIX_EPOCH)
505 .unwrap()
506 .as_secs() as i64
507 }
508
509 struct TestRng;
510
511 impl TryRng for TestRng {
512 type Error = getrandom::Error;
513
514 fn try_next_u32(&mut self) -> Result<u32, Self::Error> {
515 let mut bytes = [0u8; 4];
516 getrandom::fill(&mut bytes)?;
517 Ok(u32::from_le_bytes(bytes))
518 }
519
520 fn try_next_u64(&mut self) -> Result<u64, Self::Error> {
521 let mut bytes = [0u8; 8];
522 getrandom::fill(&mut bytes)?;
523 Ok(u64::from_le_bytes(bytes))
524 }
525
526 fn try_fill_bytes(&mut self, dst: &mut [u8]) -> Result<(), Self::Error> {
527 getrandom::fill(dst)
528 }
529 }
530
531 impl TryCryptoRng for TestRng {}
532}