Skip to main content

soma_auth/
google.rs

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";
17/// Google ID tokens can also carry this bare-form issuer (no scheme) —
18/// accepted alongside [`GOOGLE_ISSUER`] to preserve pre-extraction behavior
19/// (`google.rs`'s `verify_id_token` used to check both forms directly).
20const 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    /// Regression test for the alt-issuer fix: pre-extraction `google.rs`
274    /// accepted ID tokens carrying either the `https://accounts.google.com`
275    /// form OR the bare `accounts.google.com` form. Prove the bare form is
276    /// still accepted after the `OidcVerifier` extraction, not just that it
277    /// isn't rejected.
278    #[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}