use std::ops::Deref;
use openidconnect::core::CoreProviderMetadata;
use openidconnect::reqwest;
use openidconnect::{ClientId, ClientSecret, IssuerUrl, RedirectUrl};
use super::oidc_generic::{OidcError, OidcProvider};
pub const GOOGLE_ISSUER: &str = "https://accounts.google.com";
pub struct GoogleProvider<S>(OidcProvider<S>);
impl<S> GoogleProvider<S> {
pub fn from_provider_metadata(
metadata: CoreProviderMetadata,
client_id: ClientId,
client_secret: Option<ClientSecret>,
redirect_uri: RedirectUrl,
flows: S,
) -> Self {
Self(OidcProvider::from_provider_metadata(
metadata,
client_id,
client_secret,
redirect_uri,
flows,
))
}
pub async fn discover(
client_id: ClientId,
client_secret: Option<ClientSecret>,
redirect_uri: RedirectUrl,
flows: S,
http: &reqwest::Client,
) -> Result<Self, OidcError> {
let issuer = IssuerUrl::new(GOOGLE_ISSUER.to_owned())
.map_err(|e| OidcError::Config(format!("invalid Google issuer URL: {e}")))?;
Ok(Self(
OidcProvider::discover(issuer, client_id, client_secret, redirect_uri, flows, http)
.await?,
))
}
pub fn into_inner(self) -> OidcProvider<S> {
self.0
}
}
impl<S> Deref for GoogleProvider<S> {
type Target = OidcProvider<S>;
fn deref(&self) -> &Self::Target {
&self.0
}
}
#[cfg(test)]
mod tests {
use super::*;
use chrono::{Duration, Utc};
use openidconnect::core::{
CoreIdToken, CoreIdTokenClaims, CoreIdTokenFields,
CoreJsonWebKeySet, CoreJwsSigningAlgorithm, CoreResponseType, CoreRsaPrivateSigningKey,
CoreSubjectIdentifierType, CoreTokenResponse, CoreTokenType,
};
use openidconnect::{
AccessToken, AuthUrl, Audience, AuthorizationCode, ClientId, ClientSecret,
EmptyAdditionalClaims, EmptyAdditionalProviderMetadata, EmptyExtraTokenFields,
EndUserEmail, EndUserName, IssuerUrl, JsonWebKeyId, JsonWebKeySetUrl, LanguageTag,
LocalizedClaim, Nonce, PrivateSigningKey, RedirectUrl, ResponseTypes, StandardClaims,
SubjectIdentifier, TokenUrl,
};
use openidconnect::reqwest;
use wiremock::matchers::{method, path};
use wiremock::{Mock, MockServer, ResponseTemplate};
use crate::providers::oidc_generic::{MemoryOidcFlowStore, OidcCallback, OidcFlowStore};
const TEST_RSA_PEM: &str = concat!(
"-----BEGIN RSA PRIVATE KEY-----\n",
"MIIEowIBAAKCAQEAsRMj0YYjy7du6v1gWyKSTJx3YjBzZTG0XotRP0IaObw0k+68\n",
"30dXadjL5jVhSWNdcg9OyMyTGWfdNqfdrS6ppBqlQNgjZJdloIqL9zOLBZrDm7G4\n",
"+qN4KeZ4/5TyEilq2zOHHGFEzXpOq/UxqVnm3J4fhjqCNaS2nKd7HVVXGBQQ+4+F\n",
"dVT+MyJXemw5maz2F/h324TQi6XoUPEwUddxBwLQFSOlzWnHYMc4/lcyZJ8MpTXC\n",
"MPe/YJFNtb9CaikKUdf8x4mzwH7usSf8s2d6R4dQITzKrjrEJ0u3w3eGkBBapoMV\n",
"FBGPjP3Haz5FsVtHc5VEN3FZVIDF6HrbJH1C4QIDAQABAoIBAHSS3izM+3nc7Bel\n",
"8S5uRxRKmcm5je6b11u6qiVUFkHWJmMRc6QmqmSThkCq+b4/vUAe1cYZ7+l02Exo\n",
"HOcrZiEULaDP6hUKGqyjKVv3wdlRtt8kFFxlC/HBufzAiNDuFVvzw0oquwnvMCXC\n",
"yQvtlK+/JY/PqvM32cSt+b4o9apySsHqAtdsoHHohK82jsQqIfCi1v8XYV/xRBJB\n",
"cQMCaA0Ls3tFpmJv3JdikyyQxio4kZ5tswghC63znCp1iL+qDq1wjjKzjick9MDb\n",
"Qzb95X09QQP201l1FPWN7Kbhj4ybg6PJGz/VHQcvILcBCoYIc0UY/OMSBt9VN9yD\n",
"wr1WlbECgYEA37difsTMcLmUEN57sicFe1q4lxH6eqnUBjmoKBflx4oMIIyRnfjF\n",
"Jwsu9yIiBkJfBCP85nl2tZdcV0wfZLf6amxB/KMtdfW6r8eoTDzE472OYxSIg1F5\n",
"dI4qn2nBI0Dou0g58xj+Kv0iLaym0pxtyJkSg/rxZGwKb9a+x5WAs50CgYEAyqC0\n",
"NcZs2BRIiT5kEOF6+MeUvarbKh1mangKHKcTdXRrvoJ+Z5izm7FifBixo/79MYpt\n",
"0VofW0IzYKtAI9KZDq2JcozEbZ+lt/ZPH5QEXO4T39QbDoAG8BbOmEP7l+6m+7QO\n",
"PiQ0WSNjDnwk3W7Zihgg31DH7hyxsxQCapKLcxUCgYAwERXPiPcoDSd8DGFlYK7z\n",
"1wUsKEe6DT0p7T9tBd1v5wA+ChXLbETn46Y+oQ3QbHg/yn+vAU/5KkFD3G4uVL0w\n",
"Gnx/DIxa+OYYmHxXjQL8r6ClNycxl9LRsS4FPFKsAWk/u///dFI/6E1spNjfDY8k\n",
"94ab5tHwsqn3Z5tsBHo3nQKBgFUmxbSXh2Qi2fy6+GhTqU7k6G/wXhvLsR9rBKzX\n",
"1YiVfTXZNu+oL0ptd/q4keZeIN7x0oaY/fZm0pp8PP8Q4HtXmBxIZb+/yG+Pld6q\n",
"YE8BSd7VDu3ABapdm0JHx3Iou4mpOBcLNeiDw3vx1bgsfkTXMPFHzE0XR+H+tak9\n",
"nlalAoGBALAmAF7WBGdOt43Rj8hPaKOM/ahj+6z3CNwVreToNsVBHoyNmiO8q7MC\n",
"+tRo4jgdrzk1pzs66OIHfbx5P1mXKPtgPZhvI5omAY8WqXEgeNqSL1Ksp6LZ2ql/\n",
"ouZns5xwKc9+aRL+GWoAGNzwzcjE8cP52sBy/r0rYXTs/sZo5kgV\n",
"-----END RSA PRIVATE KEY-----\n",
);
const TEST_KID: &str = "cheers-test-key";
const CLIENT_ID: &str = "test-client.apps.googleusercontent.com";
const CLIENT_SECRET: &str = "test-secret";
const REDIRECT_URI: &str = "https://app.example/auth/callback/google";
fn signing_key() -> CoreRsaPrivateSigningKey {
CoreRsaPrivateSigningKey::from_pem(
TEST_RSA_PEM,
Some(JsonWebKeyId::new(TEST_KID.into())),
)
.expect("test RSA PEM parses")
}
fn dummy_http() -> reqwest::Client {
reqwest::ClientBuilder::new()
.redirect(reqwest::redirect::Policy::none())
.build()
.expect("reqwest builds")
}
enum NameClaim {
UnTagged(&'static str),
LocalizedOnly(&'static [(&'static str, &'static str)]),
}
fn build_id_token(
issuer: &str,
nonce: &Nonce,
email_verified: bool,
name: NameClaim,
) -> CoreIdToken {
let now = Utc::now();
let mut std_claims =
StandardClaims::new(SubjectIdentifier::new("user-1234567890".to_owned()))
.set_email(Some(EndUserEmail::new("alice@example.com".to_owned())))
.set_email_verified(Some(email_verified));
match name {
NameClaim::UnTagged(s) => {
let mut lc: LocalizedClaim<EndUserName> = LocalizedClaim::default();
lc.insert(None, EndUserName::new(s.to_owned()));
std_claims = std_claims.set_name(Some(lc));
}
NameClaim::LocalizedOnly(entries) => {
let mut lc: LocalizedClaim<EndUserName> = LocalizedClaim::default();
for (tag, value) in entries {
lc.insert(
Some(LanguageTag::new((*tag).to_owned())),
EndUserName::new((*value).to_owned()),
);
}
std_claims = std_claims.set_name(Some(lc));
}
}
let claims = CoreIdTokenClaims::new(
IssuerUrl::new(issuer.to_owned()).expect("test issuer URL parses"),
vec![Audience::new(CLIENT_ID.to_owned())],
now + Duration::seconds(600),
now,
std_claims,
EmptyAdditionalClaims {},
)
.set_nonce(Some(nonce.clone()));
CoreIdToken::new(
claims,
&signing_key(),
CoreJwsSigningAlgorithm::RsaSsaPkcs1V15Sha256,
None,
None,
)
.expect("ID token signs cleanly")
}
async fn mount_discovery_and_jwks(server: &MockServer, base: &str) {
let metadata = CoreProviderMetadata::new(
IssuerUrl::new(base.to_owned()).unwrap(),
AuthUrl::new(format!("{base}/o/oauth2/auth")).unwrap(),
JsonWebKeySetUrl::new(format!("{base}/jwks")).unwrap(),
vec![ResponseTypes::new(vec![CoreResponseType::Code])],
vec![CoreSubjectIdentifierType::Public],
vec![CoreJwsSigningAlgorithm::RsaSsaPkcs1V15Sha256],
EmptyAdditionalProviderMetadata {},
)
.set_token_endpoint(Some(TokenUrl::new(format!("{base}/token")).unwrap()));
Mock::given(method("GET"))
.and(path("/.well-known/openid-configuration"))
.respond_with(ResponseTemplate::new(200).set_body_json(&metadata))
.mount(server)
.await;
let jwks = CoreJsonWebKeySet::new(vec![signing_key().as_verification_key()]);
Mock::given(method("GET"))
.and(path("/jwks"))
.respond_with(ResponseTemplate::new(200).set_body_json(&jwks))
.mount(server)
.await;
}
async fn peek_stashed_nonce(
provider: &GoogleProvider<MemoryOidcFlowStore>,
csrf_state_secret: &str,
) -> Nonce {
let st = provider
.flows()
.take(csrf_state_secret)
.await
.expect("store ok")
.expect("flow stashed");
let nonce = st.nonce().clone();
provider
.flows()
.put(csrf_state_secret, st)
.await
.expect("re-put");
nonce
}
async fn build_provider_via_discovery(
server: &MockServer,
http: &reqwest::Client,
) -> GoogleProvider<MemoryOidcFlowStore> {
let issuer = IssuerUrl::new(server.uri()).expect("issuer URL parses");
let inner = OidcProvider::discover(
issuer,
ClientId::new(CLIENT_ID.into()),
Some(ClientSecret::new(CLIENT_SECRET.into())),
RedirectUrl::new(REDIRECT_URI.into()).expect("redirect URL parses"),
MemoryOidcFlowStore::new(),
http,
)
.await
.expect("discovery succeeds");
GoogleProvider(inner)
}
async fn mount_token_endpoint(server: &MockServer, id_token: CoreIdToken) {
let token_response = CoreTokenResponse::new(
AccessToken::new("test-access-token".to_owned()),
CoreTokenType::Bearer,
CoreIdTokenFields::new(Some(id_token), EmptyExtraTokenFields {}),
);
Mock::given(method("POST"))
.and(path("/token"))
.respond_with(ResponseTemplate::new(200).set_body_json(&token_response))
.mount(server)
.await;
}
#[tokio::test]
async fn discover_then_finish_round_trip_extracts_claims() {
let http = dummy_http();
let server = MockServer::start().await;
let base = server.uri();
mount_discovery_and_jwks(&server, &base).await;
let provider = build_provider_via_discovery(&server, &http).await;
let now_seconds = Utc::now().timestamp();
let begin = provider.begin(now_seconds).await.expect("begin succeeds");
let nonce = peek_stashed_nonce(&provider, begin.csrf_state.secret()).await;
let id_token =
build_id_token(&base, &nonce, true, NameClaim::UnTagged("Alice Anderson"));
mount_token_endpoint(&server, id_token).await;
let verified = provider
.finish(
OidcCallback::new(
AuthorizationCode::new("auth-code-xyz".into()),
begin.csrf_state,
),
&http,
now_seconds,
)
.await
.expect("finish round-trip");
assert_eq!(verified.issuer, base);
assert_eq!(verified.subject, "user-1234567890");
assert_eq!(verified.email.as_deref(), Some("alice@example.com"));
assert_eq!(verified.email_verified, Some(true));
assert_eq!(verified.name.as_deref(), Some("Alice Anderson"));
}
#[tokio::test]
async fn finish_falls_back_to_localized_name_when_untagged_missing() {
let http = dummy_http();
let server = MockServer::start().await;
let base = server.uri();
mount_discovery_and_jwks(&server, &base).await;
let provider = build_provider_via_discovery(&server, &http).await;
let now_seconds = Utc::now().timestamp();
let begin = provider.begin(now_seconds).await.expect("begin");
let nonce = peek_stashed_nonce(&provider, begin.csrf_state.secret()).await;
let id_token = build_id_token(
&base,
&nonce,
false,
NameClaim::LocalizedOnly(&[("en", "Bob Localized")]),
);
mount_token_endpoint(&server, id_token).await;
let verified = provider
.finish(
OidcCallback::new(
AuthorizationCode::new("auth-code-2".into()),
begin.csrf_state,
),
&http,
now_seconds,
)
.await
.expect("finish round-trip");
assert_eq!(verified.email_verified, Some(false));
assert_eq!(verified.name.as_deref(), Some("Bob Localized"));
}
#[test]
fn google_issuer_constant_is_the_canonical_url() {
assert_eq!(GOOGLE_ISSUER, "https://accounts.google.com");
let _ = IssuerUrl::new(GOOGLE_ISSUER.to_owned())
.expect("Google issuer URL parses");
}
#[tokio::test]
async fn from_provider_metadata_derefs_to_oidc_provider() {
const FIXTURE: &str = r#"{
"issuer": "https://accounts.google.com",
"authorization_endpoint": "https://accounts.google.com/o/oauth2/auth",
"token_endpoint": "https://accounts.google.com/o/oauth2/token",
"jwks_uri": "https://accounts.google.com/o/oauth2/jwks",
"response_types_supported": ["code"],
"subject_types_supported": ["public"],
"id_token_signing_alg_values_supported": ["RS256"]
}"#;
let metadata: CoreProviderMetadata =
serde_json::from_str(FIXTURE).expect("fixture parses");
let provider = GoogleProvider::from_provider_metadata(
metadata,
ClientId::new(CLIENT_ID.into()),
Some(ClientSecret::new(CLIENT_SECRET.into())),
RedirectUrl::new(REDIRECT_URI.into()).unwrap(),
MemoryOidcFlowStore::new(),
);
let begin = provider.begin(1_000).await.expect("begin");
assert_eq!(
begin.authorize_url.host_str(),
Some("accounts.google.com"),
"Google-baked authorization endpoint should still drive the begin URL"
);
let inner = provider.into_inner();
assert_eq!(inner.flow_ttl_seconds(), 300);
}
}