use std::collections::HashMap;
use std::sync::Mutex;
use async_trait::async_trait;
use openidconnect::core::{
CoreAuthenticationFlow, CoreClient, CoreIdTokenClaims, CoreProviderMetadata,
};
use openidconnect::reqwest;
use openidconnect::{
AuthorizationCode, ClientId, ClientSecret, CsrfToken, EndpointMaybeSet, EndpointNotSet,
EndpointSet, IssuerUrl, Nonce, PkceCodeChallenge, PkceCodeVerifier, RedirectUrl, Scope,
TokenResponse,
};
pub const DEFAULT_FLOW_TTL_SECONDS: i64 = 5 * 60;
pub const DEFAULT_SCOPES: &[&str] = &["openid", "email", "profile"];
type Client = CoreClient<
EndpointSet,
EndpointNotSet,
EndpointNotSet,
EndpointNotSet,
EndpointMaybeSet,
EndpointMaybeSet,
>;
#[derive(Debug, thiserror::Error)]
#[non_exhaustive]
pub enum OidcError {
#[error("oidc discovery: {0}")]
Discovery(String),
#[error("oidc config: {0}")]
Config(String),
#[error("oidc http: {0}")]
Http(String),
#[error("unknown or already-consumed oidc flow")]
UnknownFlow,
#[error("oidc flow expired")]
FlowExpired,
#[error("oidc csrf state mismatch")]
StateMismatch,
#[error("token response carried no id_token")]
MissingIdToken,
#[error("id token verification: {0}")]
IdToken(String),
#[error("oidc flow store: {0}")]
Store(String),
}
#[derive(Debug, Clone, PartialEq, Eq)]
#[non_exhaustive]
pub struct VerifiedIdToken {
pub issuer: String,
pub subject: String,
pub email: Option<String>,
pub email_verified: Option<bool>,
pub name: Option<String>,
}
pub struct OidcFlowState {
csrf_token: CsrfToken,
nonce: Nonce,
pkce_verifier: PkceCodeVerifier,
expires_at: i64,
}
impl OidcFlowState {
pub fn from_parts(
csrf_token: CsrfToken,
nonce: Nonce,
pkce_verifier: PkceCodeVerifier,
expires_at: i64,
) -> Self {
Self {
csrf_token,
nonce,
pkce_verifier,
expires_at,
}
}
pub fn expires_at(&self) -> i64 {
self.expires_at
}
pub fn is_expired_at(&self, now: i64) -> bool {
self.expires_at <= now
}
pub fn csrf_token(&self) -> &CsrfToken {
&self.csrf_token
}
pub fn nonce(&self) -> &Nonce {
&self.nonce
}
pub fn pkce_verifier(&self) -> &PkceCodeVerifier {
&self.pkce_verifier
}
pub fn into_parts(self) -> (CsrfToken, Nonce, PkceCodeVerifier, i64) {
(
self.csrf_token,
self.nonce,
self.pkce_verifier,
self.expires_at,
)
}
}
impl std::fmt::Debug for OidcFlowState {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("OidcFlowState")
.field("csrf_token", &"<redacted>")
.field("nonce", &"<redacted>")
.field("pkce_verifier", &"<redacted>")
.field("expires_at", &self.expires_at)
.finish()
}
}
#[async_trait]
pub trait OidcFlowStore: Send + Sync {
async fn put(&self, id: &str, state: OidcFlowState) -> Result<(), String>;
async fn take(&self, id: &str) -> Result<Option<OidcFlowState>, String>;
}
#[derive(Default)]
pub struct MemoryOidcFlowStore {
inner: Mutex<HashMap<String, OidcFlowState>>,
}
impl MemoryOidcFlowStore {
pub fn new() -> Self {
Self::default()
}
pub fn gc(&self, now: i64) {
self.inner
.lock()
.unwrap()
.retain(|_, s| s.expires_at > now);
}
pub fn len(&self) -> usize {
self.inner.lock().unwrap().len()
}
pub fn is_empty(&self) -> bool {
self.inner.lock().unwrap().is_empty()
}
}
#[async_trait]
impl OidcFlowStore for MemoryOidcFlowStore {
async fn put(&self, id: &str, state: OidcFlowState) -> Result<(), String> {
self.inner.lock().unwrap().insert(id.to_owned(), state);
Ok(())
}
async fn take(&self, id: &str) -> Result<Option<OidcFlowState>, String> {
Ok(self.inner.lock().unwrap().remove(id))
}
}
#[derive(Debug)]
#[non_exhaustive]
pub struct OidcBegin {
pub authorize_url: openidconnect::url::Url,
pub csrf_state: CsrfToken,
}
#[derive(Debug)]
#[non_exhaustive]
pub struct OidcCallback {
pub code: AuthorizationCode,
pub state: CsrfToken,
}
impl OidcCallback {
pub fn new(code: AuthorizationCode, state: CsrfToken) -> Self {
Self { code, state }
}
}
pub struct OidcProvider<S> {
client: Client,
scopes: Vec<Scope>,
flow_ttl_seconds: i64,
flows: S,
}
impl<S> OidcProvider<S> {
pub fn from_provider_metadata(
metadata: CoreProviderMetadata,
client_id: ClientId,
client_secret: Option<ClientSecret>,
redirect_uri: RedirectUrl,
flows: S,
) -> Self {
let client = CoreClient::from_provider_metadata(metadata, client_id, client_secret)
.set_redirect_uri(redirect_uri);
Self {
client,
scopes: DEFAULT_SCOPES
.iter()
.map(|s| Scope::new((*s).to_owned()))
.collect(),
flow_ttl_seconds: DEFAULT_FLOW_TTL_SECONDS,
flows,
}
}
pub async fn discover(
issuer: IssuerUrl,
client_id: ClientId,
client_secret: Option<ClientSecret>,
redirect_uri: RedirectUrl,
flows: S,
http: &reqwest::Client,
) -> Result<Self, OidcError> {
let metadata = CoreProviderMetadata::discover_async(issuer, http)
.await
.map_err(|e| OidcError::Discovery(format!("{e}")))?;
Ok(Self::from_provider_metadata(
metadata,
client_id,
client_secret,
redirect_uri,
flows,
))
}
pub fn with_scopes<I, V>(mut self, scopes: I) -> Self
where
I: IntoIterator<Item = V>,
V: Into<String>,
{
self.scopes = scopes.into_iter().map(|s| Scope::new(s.into())).collect();
self
}
pub fn with_flow_ttl_seconds(mut self, ttl: i64) -> Self {
self.flow_ttl_seconds = ttl;
self
}
pub fn scopes(&self) -> &[Scope] {
&self.scopes
}
pub fn flow_ttl_seconds(&self) -> i64 {
self.flow_ttl_seconds
}
pub fn flows(&self) -> &S {
&self.flows
}
}
impl<S: OidcFlowStore> OidcProvider<S> {
pub async fn begin(&self, now: i64) -> Result<OidcBegin, OidcError> {
let (pkce_challenge, pkce_verifier) = PkceCodeChallenge::new_random_sha256();
let mut req = self.client.authorize_url(
CoreAuthenticationFlow::AuthorizationCode,
CsrfToken::new_random,
Nonce::new_random,
);
for scope in &self.scopes {
req = req.add_scope(scope.clone());
}
let (authorize_url, csrf_state, nonce) = req.set_pkce_challenge(pkce_challenge).url();
let flow_state = OidcFlowState {
csrf_token: csrf_state.clone(),
nonce,
pkce_verifier,
expires_at: now.saturating_add(self.flow_ttl_seconds),
};
let id = csrf_state.secret().to_owned();
self.flows
.put(&id, flow_state)
.await
.map_err(OidcError::Store)?;
Ok(OidcBegin {
authorize_url,
csrf_state,
})
}
pub async fn finish(
&self,
callback: OidcCallback,
http: &reqwest::Client,
now: i64,
) -> Result<VerifiedIdToken, OidcError> {
let flow_state = self
.flows
.take(callback.state.secret())
.await
.map_err(OidcError::Store)?
.ok_or(OidcError::UnknownFlow)?;
if flow_state.csrf_token.secret() != callback.state.secret() {
return Err(OidcError::StateMismatch);
}
if flow_state.is_expired_at(now) {
return Err(OidcError::FlowExpired);
}
let token_response = self
.client
.exchange_code(callback.code)
.map_err(|e| OidcError::Config(format!("{e}")))?
.set_pkce_verifier(flow_state.pkce_verifier)
.request_async(http)
.await
.map_err(|e| OidcError::Http(format!("{e}")))?;
let id_token = token_response.id_token().ok_or(OidcError::MissingIdToken)?;
let verifier = self.client.id_token_verifier();
let claims: &CoreIdTokenClaims = id_token
.claims(&verifier, &flow_state.nonce)
.map_err(|e| OidcError::IdToken(format!("{e}")))?;
Ok(extract(claims))
}
}
fn extract(c: &CoreIdTokenClaims) -> VerifiedIdToken {
let name = c
.name()
.and_then(|n| n.get(None).or_else(|| n.iter().next().map(|(_, v)| v)))
.map(|v| v.as_str().to_owned());
VerifiedIdToken {
issuer: c.issuer().as_str().to_owned(),
subject: c.subject().as_str().to_owned(),
email: c.email().map(|e| e.as_str().to_owned()),
email_verified: c.email_verified(),
name,
}
}
#[cfg(test)]
mod tests {
use super::*;
use openidconnect::url::Url;
const FIXTURE_METADATA: &str = r#"{
"issuer": "https://idp.example",
"authorization_endpoint": "https://idp.example/o/oauth2/auth",
"token_endpoint": "https://idp.example/o/oauth2/token",
"jwks_uri": "https://idp.example/o/oauth2/jwks",
"response_types_supported": ["code"],
"subject_types_supported": ["public"],
"id_token_signing_alg_values_supported": ["RS256"]
}"#;
fn metadata() -> CoreProviderMetadata {
serde_json::from_str(FIXTURE_METADATA).expect("fixture parses as CoreProviderMetadata")
}
fn provider() -> OidcProvider<MemoryOidcFlowStore> {
OidcProvider::from_provider_metadata(
metadata(),
ClientId::new("test-client".into()),
Some(ClientSecret::new("test-secret".into())),
RedirectUrl::new("https://app.example/oauth/callback".into()).unwrap(),
MemoryOidcFlowStore::new(),
)
}
fn query_pairs(u: &Url) -> Vec<(String, String)> {
u.query_pairs()
.map(|(k, v)| (k.into_owned(), v.into_owned()))
.collect()
}
fn first(qs: &[(String, String)], k: &str) -> Option<String> {
qs.iter().find(|(kk, _)| kk == k).map(|(_, v)| v.clone())
}
#[tokio::test]
async fn memory_store_put_then_take_is_single_use() {
let s = MemoryOidcFlowStore::new();
let st = OidcFlowState {
csrf_token: CsrfToken::new("STATE".into()),
nonce: Nonce::new("NONCE".into()),
pkce_verifier: PkceCodeVerifier::new("VERIFIER".repeat(8)),
expires_at: 1_000,
};
s.put("k", st).await.unwrap();
assert_eq!(s.len(), 1);
let taken = s.take("k").await.unwrap().expect("entry");
assert_eq!(taken.csrf_token.secret(), "STATE");
assert!(s.take("k").await.unwrap().is_none());
assert!(s.is_empty());
}
#[tokio::test]
async fn memory_store_take_missing_returns_none() {
let s = MemoryOidcFlowStore::new();
assert!(s.take("nope").await.unwrap().is_none());
}
#[tokio::test]
async fn memory_store_gc_drops_expired_entries() {
let s = MemoryOidcFlowStore::new();
s.put(
"a",
OidcFlowState {
csrf_token: CsrfToken::new("a".into()),
nonce: Nonce::new("n".into()),
pkce_verifier: PkceCodeVerifier::new("v".repeat(64)),
expires_at: 100,
},
)
.await
.unwrap();
s.put(
"b",
OidcFlowState {
csrf_token: CsrfToken::new("b".into()),
nonce: Nonce::new("n".into()),
pkce_verifier: PkceCodeVerifier::new("v".repeat(64)),
expires_at: 500,
},
)
.await
.unwrap();
assert_eq!(s.len(), 2);
s.gc(200);
assert_eq!(s.len(), 1);
assert!(s.take("a").await.unwrap().is_none());
assert!(s.take("b").await.unwrap().is_some());
}
#[test]
fn from_provider_metadata_seeds_default_scopes_and_ttl() {
let p = provider();
let scope_strs: Vec<&str> = p.scopes().iter().map(|s| s.as_str()).collect();
assert_eq!(scope_strs, DEFAULT_SCOPES);
assert_eq!(p.flow_ttl_seconds(), DEFAULT_FLOW_TTL_SECONDS);
}
#[test]
fn with_scopes_replaces_the_set() {
let p = provider().with_scopes(["openid", "email"]);
let scope_strs: Vec<&str> = p.scopes().iter().map(|s| s.as_str()).collect();
assert_eq!(scope_strs, ["openid", "email"]);
}
#[test]
fn with_flow_ttl_overrides_default() {
let p = provider().with_flow_ttl_seconds(60);
assert_eq!(p.flow_ttl_seconds(), 60);
}
#[tokio::test]
async fn begin_url_has_oidc_params_pinned() {
let p = provider();
let begin = p.begin(1_000).await.unwrap();
let u = &begin.authorize_url;
assert_eq!(u.scheme(), "https");
assert_eq!(u.host_str(), Some("idp.example"));
assert_eq!(u.path(), "/o/oauth2/auth");
let qs = query_pairs(u);
assert_eq!(first(&qs, "response_type").as_deref(), Some("code"));
assert_eq!(first(&qs, "client_id").as_deref(), Some("test-client"));
assert_eq!(
first(&qs, "redirect_uri").as_deref(),
Some("https://app.example/oauth/callback")
);
let scope = first(&qs, "scope").unwrap();
for s in DEFAULT_SCOPES {
assert!(scope.split(' ').any(|x| x == *s), "missing scope {s}");
}
assert_eq!(first(&qs, "code_challenge_method").as_deref(), Some("S256"));
assert!(
!first(&qs, "code_challenge").unwrap_or_default().is_empty(),
"PKCE challenge missing"
);
assert!(
!first(&qs, "state").unwrap_or_default().is_empty(),
"csrf state missing"
);
assert!(
!first(&qs, "nonce").unwrap_or_default().is_empty(),
"nonce missing"
);
assert_eq!(first(&qs, "state").as_deref(), Some(begin.csrf_state.secret().as_str()));
}
#[tokio::test]
async fn begin_stashes_flow_state_keyed_by_csrf_secret() {
let p = provider();
let begin = p.begin(1_000).await.unwrap();
assert_eq!(p.flows().len(), 1);
let taken = p
.flows()
.take(begin.csrf_state.secret())
.await
.unwrap()
.expect("stashed");
assert_eq!(taken.csrf_token.secret(), begin.csrf_state.secret());
assert_eq!(taken.expires_at(), 1_000 + DEFAULT_FLOW_TTL_SECONDS);
}
#[tokio::test]
async fn begin_generates_unique_state_per_call() {
let p = provider();
let a = p.begin(1_000).await.unwrap();
let b = p.begin(1_000).await.unwrap();
assert_ne!(a.csrf_state.secret(), b.csrf_state.secret());
assert_eq!(p.flows().len(), 2);
}
#[tokio::test]
async fn begin_respects_custom_scopes() {
let p = provider().with_scopes(["openid", "profile"]);
let begin = p.begin(1_000).await.unwrap();
let qs = query_pairs(&begin.authorize_url);
let scope = first(&qs, "scope").unwrap();
let parts: Vec<&str> = scope.split(' ').collect();
assert!(parts.contains(&"openid"));
assert!(parts.contains(&"profile"));
assert!(!parts.contains(&"email"));
}
fn dummy_http() -> reqwest::Client {
reqwest::ClientBuilder::new()
.redirect(reqwest::redirect::Policy::none())
.build()
.expect("reqwest client builds")
}
#[tokio::test]
async fn finish_with_unknown_state_returns_unknown_flow() {
let p = provider();
let err = p
.finish(
OidcCallback::new(
AuthorizationCode::new("CODE".into()),
CsrfToken::new("never-stashed".into()),
),
&dummy_http(),
1_000,
)
.await
.unwrap_err();
assert!(matches!(err, OidcError::UnknownFlow), "got {err:?}");
}
#[tokio::test]
async fn finish_consumes_flow_on_first_call_replay_is_unknown() {
let p = provider();
let begin = p.begin(1_000).await.unwrap();
let state = begin.csrf_state.clone();
let cb = || {
OidcCallback::new(
AuthorizationCode::new("CODE".into()),
state.clone(),
)
};
let first = p.finish(cb(), &dummy_http(), 1_000).await.unwrap_err();
assert!(
matches!(first, OidcError::Http(_)),
"expected Http error, got {first:?}"
);
let second = p.finish(cb(), &dummy_http(), 1_000).await.unwrap_err();
assert!(
matches!(second, OidcError::UnknownFlow),
"expected UnknownFlow on replay, got {second:?}"
);
}
#[tokio::test]
async fn finish_with_expired_flow_returns_flow_expired() {
let p = provider().with_flow_ttl_seconds(60);
let begin = p.begin(1_000).await.unwrap();
let err = p
.finish(
OidcCallback::new(AuthorizationCode::new("CODE".into()), begin.csrf_state),
&dummy_http(),
1_000 + 60, )
.await
.unwrap_err();
assert!(matches!(err, OidcError::FlowExpired), "got {err:?}");
}
#[tokio::test]
async fn flow_state_debug_redacts_secrets() {
let st = OidcFlowState {
csrf_token: CsrfToken::new("SUPER-SECRET-STATE".into()),
nonce: Nonce::new("SUPER-SECRET-NONCE".into()),
pkce_verifier: PkceCodeVerifier::new("v".repeat(64)),
expires_at: 100,
};
let dbg = format!("{st:?}");
assert!(!dbg.contains("SUPER-SECRET-STATE"));
assert!(!dbg.contains("SUPER-SECRET-NONCE"));
assert!(dbg.contains("expires_at: 100"));
}
#[test]
fn verified_id_token_round_trips_through_eq() {
let v1 = VerifiedIdToken {
issuer: "https://idp".into(),
subject: "u-1".into(),
email: Some("a@b.co".into()),
email_verified: Some(true),
name: Some("Alice".into()),
};
let v2 = v1.clone();
assert_eq!(v1, v2);
}
}