use anyhow::{bail, Result};
use std::collections::HashMap;
use std::sync::{Arc, Mutex};
use tokio::sync::{Mutex as AsyncMutex, OwnedMutexGuard};
use super::store::OAuthSession;
use super::token::TokenResponse;
pub fn apply_refresh(
session: &OAuthSession,
response: &TokenResponse,
now: i64,
) -> Result<OAuthSession> {
if response.sub != session.sub {
bail!(
"refresh returned subject {:?}, expected {:?}; refusing to rebind the session",
response.sub,
session.sub
);
}
Ok(OAuthSession {
access_token: response.access_token.clone(),
refresh_token: response
.refresh_token
.clone()
.unwrap_or_else(|| session.refresh_token.clone()),
token_type: response.token_type.clone(),
granted_scope: response.granted_scope.clone(),
expires_at: response.expires_in.map(|seconds| now + seconds),
sub: session.sub.clone(),
issuer: session.issuer.clone(),
aud: session.aud.clone(),
dpop_key_jwk: session.dpop_key_jwk.clone(),
})
}
#[derive(Default, Clone)]
pub struct RefreshLocks {
locks: Arc<Mutex<HashMap<String, Arc<AsyncMutex<()>>>>>,
}
impl RefreshLocks {
pub async fn lock(&self, sub: &str) -> OwnedMutexGuard<()> {
let entry = {
let mut locks = self.locks.lock().unwrap_or_else(|p| p.into_inner());
Arc::clone(locks.entry(sub.to_string()).or_default())
};
entry.lock_owned().await
}
}
pub fn same_issuer(discovered: &str, expected: &str) -> Result<()> {
if discovered != expected {
anyhow::bail!(
"the PDS now names a different authorization server ({discovered:?}) than this \
grant was issued by ({expected:?}); refusing to send credentials to it"
);
}
Ok(())
}
pub struct RefreshContext<'a> {
pub token_endpoint: &'a str,
pub client_id: &'a str,
pub auth_method: super::client_auth::AuthMethod,
pub client_key: Option<&'a super::keys::SigningKey>,
pub revocation_endpoint: Option<&'a str>,
}
pub async fn valid_session(
pool: &sqlx::SqlitePool,
codec: &super::crypto::Codec,
http: &reqwest::Client,
locks: &RefreshLocks,
sub: &str,
ctx: &RefreshContext<'_>,
now: i64,
) -> Result<OAuthSession> {
let session = super::store::get_session(pool, codec, sub)
.await?
.ok_or_else(|| anyhow::anyhow!("no session for {sub}"))?;
if !super::token::is_stale(session.expires_at, now) {
return Ok(session);
}
let _guard = locks.lock(sub).await;
let (session, version) = super::store::get_session_versioned(pool, codec, sub)
.await?
.ok_or_else(|| anyhow::anyhow!("session for {sub} disappeared while waiting to refresh"))?;
if !super::token::is_stale(session.expires_at, now) {
return Ok(session);
}
refresh_locked(pool, codec, http, &session, &version, ctx, now).await
}
async fn lost_the_row(
pool: &sqlx::SqlitePool,
codec: &super::crypto::Codec,
http: &reqwest::Client,
ctx: &RefreshContext<'_>,
obtained: &OAuthSession,
now: i64,
) -> Result<OAuthSession> {
let current = super::store::get_session(pool, codec, &obtained.sub).await?;
let outcome = super::revoke::revoke_orphaned(
pool,
http,
&super::revoke::RevokeContext {
revocation_endpoint: ctx.revocation_endpoint,
client_id: ctx.client_id,
auth_method: ctx.auth_method,
client_key: ctx.client_key,
deadline: super::revoke::ORPHAN_REVOKE_DEADLINE,
},
obtained,
now,
)
.await;
if let super::revoke::Revocation::Failed(reason) = &outcome {
tracing::warn!(
sub = %obtained.sub,
%reason,
"a session changed or was signed out during its refresh: the refresh's \
unstored new tokens could not be revoked"
);
}
if let Some(current) = current {
return Ok(current);
}
bail!(
"no session for {}: it was signed out while being refreshed (the new tokens were \
not stored{})",
obtained.sub,
if outcome == super::revoke::Revocation::Revoked {
", and were revoked"
} else {
""
}
)
}
async fn refresh_locked(
pool: &sqlx::SqlitePool,
codec: &super::crypto::Codec,
http: &reqwest::Client,
session: &OAuthSession,
version: &super::store::SessionVersion,
ctx: &RefreshContext<'_>,
now: i64,
) -> Result<OAuthSession> {
let key = super::keys::SigningKey::from_jwk_json(&session.dpop_key_jwk, "session-dpop")?;
let assertion = match ctx.auth_method {
super::client_auth::AuthMethod::PrivateKeyJwt => {
let client_key = ctx
.client_key
.ok_or_else(|| anyhow::anyhow!("private_key_jwt refresh needs the client key"))?;
Some(super::client_auth::client_assertion(
client_key,
ctx.client_id,
&session.issuer,
now,
)?)
}
super::client_auth::AuthMethod::None => None,
};
let mut params = super::token::refresh_request_params(&session.refresh_token);
params.extend(super::client_auth::credential_params(
ctx.auth_method,
ctx.client_id,
assertion.as_deref(),
)?);
let borrowed: Vec<(&str, &str)> = params.iter().map(|(k, v)| (*k, v.as_str())).collect();
let outcome = super::request::send_with_dpop(
http,
pool,
&super::request::DpopRequest {
endpoint: super::dpop::Endpoint::AuthorizationServer,
url: ctx.token_endpoint,
key: &key,
access_token: None,
body: super::request::DpopBody::Form(&borrowed),
retry: super::request::Retry::Allowed,
},
)
.await?;
if outcome.is_success() {
let response = super::token::parse_token_response(&outcome.json()?)?;
let updated = apply_refresh(session, &response, now)?;
if !super::store::update_session_if_unchanged(pool, codec, &updated, version).await? {
return lost_the_row(pool, codec, http, ctx, &updated, now).await;
}
tracing::info!(
sub = %updated.sub,
expires_at = ?updated.expires_at,
"refreshed the OAuth session"
);
return Ok(updated);
}
match super::token::classify_refresh_failure(outcome.status, &outcome.body) {
super::token::RefreshFailure::Transient => {
bail!(
"refresh for {} failed transiently (status {}); the session is left intact",
session.sub,
outcome.status
)
}
super::token::RefreshFailure::SessionInvalid => {
if let Some(current) = super::store::get_session(pool, codec, &session.sub).await? {
if current.refresh_token != session.refresh_token {
return Ok(current);
}
}
super::store::delete_session(pool, &session.sub).await?;
bail!(
"refresh for {} was rejected as invalid_grant; the session has been \
removed and the user must log in again",
session.sub
)
}
}
}
#[cfg(test)]
mod tests {
use super::*;
const DID: &str = "did:plc:ewvi7nxzyoun6zhxrhs64oiz";
const NOW: i64 = 1_700_000_000;
fn session() -> OAuthSession {
OAuthSession {
sub: DID.into(),
issuer: "https://auth.example.com".into(),
aud: "https://pds.example.com".into(),
dpop_key_jwk: r#"{"kty":"EC","d":"k"}"#.into(),
access_token: "old-access".into(),
refresh_token: "old-refresh".into(),
token_type: "DPoP".into(),
granted_scope: "atproto transition:generic".into(),
expires_at: Some(NOW + 60),
}
}
fn response() -> TokenResponse {
TokenResponse {
access_token: "new-access".into(),
refresh_token: Some("new-refresh".into()),
token_type: "DPoP".into(),
granted_scope: "atproto transition:generic".into(),
sub: DID.into(),
expires_in: Some(3600),
}
}
#[test]
fn a_refresh_replaces_both_tokens_and_the_expiry() {
let updated = apply_refresh(&session(), &response(), NOW).unwrap();
assert_eq!(updated.access_token, "new-access");
assert_eq!(updated.refresh_token, "new-refresh");
assert_eq!(updated.expires_at, Some(NOW + 3600));
}
#[test]
fn an_omitted_refresh_token_keeps_the_existing_one() {
let mut response = response();
response.refresh_token = None;
let updated = apply_refresh(&session(), &response, NOW).unwrap();
assert_eq!(updated.refresh_token, "old-refresh");
assert_eq!(
updated.access_token, "new-access",
"the access token still rotates"
);
}
#[test]
fn an_omitted_expiry_clears_rather_than_invents_one() {
let mut response = response();
response.expires_in = None;
assert_eq!(
apply_refresh(&session(), &response, NOW)
.unwrap()
.expires_at,
None
);
}
#[test]
fn a_refresh_for_a_different_subject_is_rejected() {
let mut response = response();
response.sub = "did:plc:aaaaaaaaaaaaaaaaaaaaaaaa".into();
assert!(apply_refresh(&session(), &response, NOW).is_err());
}
#[test]
fn the_granted_scope_is_taken_from_the_response() {
let mut response = response();
response.granted_scope = "atproto".into();
assert_eq!(
apply_refresh(&session(), &response, NOW)
.unwrap()
.granted_scope,
"atproto"
);
}
#[test]
fn a_refresh_preserves_the_session_key_and_audience() {
let updated = apply_refresh(&session(), &response(), NOW).unwrap();
assert_eq!(updated.dpop_key_jwk, session().dpop_key_jwk);
assert_eq!(updated.aud, session().aud);
assert_eq!(updated.issuer, session().issuer);
assert_eq!(updated.sub, session().sub);
}
#[tokio::test]
async fn the_same_subject_is_serialized() {
let locks = RefreshLocks::default();
let held = locks.lock(DID).await;
let second = locks.lock(DID);
tokio::pin!(second);
assert!(
futures_lite_poll_pending(&mut second),
"a second holder acquired the lock while the first held it"
);
drop(held);
let _ = second.await;
}
#[tokio::test]
async fn different_subjects_do_not_block_each_other() {
let locks = RefreshLocks::default();
let _a = locks.lock(DID).await;
let b = locks.lock("did:plc:aaaaaaaaaaaaaaaaaaaaaaaa");
tokio::pin!(b);
assert!(
!futures_lite_poll_pending(&mut b),
"an unrelated subject was blocked"
);
}
fn futures_lite_poll_pending<F: std::future::Future>(fut: &mut std::pin::Pin<&mut F>) -> bool {
use std::task::{Context, Poll, Waker};
let mut cx = Context::from_waker(Waker::noop());
matches!(fut.as_mut().poll(&mut cx), Poll::Pending)
}
#[test]
fn a_re_discovered_issuer_must_match_the_grants_own() {
same_issuer("https://pds.example.com", "https://pds.example.com")
.expect("the same issuer must pass");
let err = same_issuer("https://evil.example", "https://pds.example.com")
.expect_err("a different authorization server must be refused");
let rendered = format!("{err:#}");
assert!(
rendered.contains("evil.example") && rendered.contains("pds.example.com"),
"the error must name both, or an operator cannot tell what moved: {rendered}"
);
}
#[test]
fn the_issuer_comparison_is_exact() {
assert!(same_issuer("https://pds.example.com/", "https://pds.example.com").is_err());
assert!(same_issuer("https://PDS.example.com", "https://pds.example.com").is_err());
assert!(same_issuer("", "https://pds.example.com").is_err());
}
const KEY: &str = "aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa";
async fn token_server() -> (String, std::sync::Arc<std::sync::Mutex<Vec<String>>>) {
let (addr, log) = crate::net::spawn_tls(|_| {
let mut r = std::collections::HashMap::new();
r.insert(
"/token".to_string(),
vec![crate::net::TestResponse::json(
200,
serde_json::json!({
"access_token": "rotated-access",
"refresh_token": "rotated-refresh",
"token_type": "DPoP",
"scope": "atproto",
"sub": DID,
"expires_in": 3600,
})
.to_string(),
)],
);
r.insert(
"/revoke".to_string(),
vec![crate::net::TestResponse::json(200, "{}")],
);
r
})
.await;
crate::net::test_host_override("as-e2e.test", addr);
(format!("https://as-e2e.test:{}", addr.port()), log)
}
async fn stored_stale() -> (
sqlx::SqlitePool,
crate::oauth::crypto::Codec,
OAuthSession,
crate::oauth::store::SessionVersion,
) {
let pool = crate::store::init_url("sqlite::memory:").await.unwrap();
let codec = crate::oauth::crypto::Codec::new(Some(KEY)).unwrap();
let stale = OAuthSession {
dpop_key_jwk: crate::oauth::keys::SigningKey::generate("session-dpop")
.to_jwk_json()
.unwrap(),
expires_at: Some(NOW - 1),
..session()
};
crate::oauth::store::put_session(&pool, &codec, &stale)
.await
.unwrap();
let (read, version) = crate::oauth::store::get_session_versioned(&pool, &codec, DID)
.await
.unwrap()
.unwrap();
(pool, codec, read, version)
}
fn ctx<'a>(token: &'a str, revoke: &'a str) -> RefreshContext<'a> {
RefreshContext {
token_endpoint: token,
client_id: "http://localhost",
auth_method: super::super::client_auth::AuthMethod::None,
client_key: None,
revocation_endpoint: Some(revoke),
}
}
#[tokio::test]
async fn a_refresh_does_not_resurrect_a_session_signed_out_mid_refresh() {
let (base, log) = token_server().await;
let (token, revoke) = (format!("{base}/token"), format!("{base}/revoke"));
let (pool, codec, session, version) = stored_stale().await;
crate::oauth::store::delete_session(&pool, DID)
.await
.unwrap();
let result = refresh_locked(
&pool,
&codec,
&reqwest::Client::new(),
&session,
&version,
&ctx(&token, &revoke),
NOW,
)
.await;
assert!(
crate::oauth::store::get_session(&pool, &codec, DID)
.await
.unwrap()
.is_none(),
"the refresh RESURRECTED a signed-out session"
);
let err = result.expect_err("a signed-out session must not be handed back");
assert!(
format!("{err:#}").contains("signed out while being refreshed"),
"{err:#}"
);
let seen = log.lock().unwrap().join("\n---\n");
assert!(
seen.contains("POST /revoke") && seen.contains("token=rotated-refresh"),
"the orphaned fresh token was left live:\n{seen}"
);
}
#[tokio::test]
async fn a_refresh_does_not_overwrite_a_concurrent_rotation() {
let (base, log) = token_server().await;
let (token, revoke) = (format!("{base}/token"), format!("{base}/revoke"));
let (pool, codec, session, version) = stored_stale().await;
let theirs = OAuthSession {
refresh_token: "their-refresh".into(),
access_token: "their-access".into(),
expires_at: Some(NOW + 3600),
..session.clone()
};
crate::oauth::store::put_session(&pool, &codec, &theirs)
.await
.unwrap();
let got = refresh_locked(
&pool,
&codec,
&reqwest::Client::new(),
&session,
&version,
&ctx(&token, &revoke),
NOW,
)
.await
.expect("a concurrent rotation is not an error");
let stored = crate::oauth::store::get_session(&pool, &codec, DID)
.await
.unwrap()
.unwrap();
assert_eq!(
stored.refresh_token, "their-refresh",
"the other writer's rotation was overwritten"
);
assert_eq!(got.refresh_token, "their-refresh");
let revokes: Vec<String> = log
.lock()
.unwrap()
.iter()
.filter(|r| r.starts_with("POST /revoke"))
.cloned()
.collect();
assert!(
revokes.iter().any(|r| r.contains("token=rotated-refresh")),
"the refresh's own fresh token was left live, stored nowhere:\n{revokes:#?}"
);
assert!(
!revokes.iter().any(|r| r.contains("their-refresh")),
"revoked the OTHER writer's live token:\n{revokes:#?}"
);
}
#[tokio::test]
async fn an_uncontested_refresh_stores_the_rotated_tokens() {
let (base, _log) = token_server().await;
let (token, revoke) = (format!("{base}/token"), format!("{base}/revoke"));
let (pool, codec, session, version) = stored_stale().await;
let got = refresh_locked(
&pool,
&codec,
&reqwest::Client::new(),
&session,
&version,
&ctx(&token, &revoke),
NOW,
)
.await
.expect("refresh");
assert_eq!(got.refresh_token, "rotated-refresh");
assert_eq!(
crate::oauth::store::get_session(&pool, &codec, DID)
.await
.unwrap()
.unwrap()
.refresh_token,
"rotated-refresh"
);
}
}