use anyhow::Result;
use super::client_auth::AuthMethod;
use super::store::OAuthSession;
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum Revocation {
Revoked,
NoSession,
Failed(String),
}
pub fn token_to_revoke(session: &OAuthSession) -> (&str, &'static str) {
if session.refresh_token.is_empty() {
(&session.access_token, "access_token")
} else {
(&session.refresh_token, "refresh_token")
}
}
pub fn revoke_params(
method: AuthMethod,
client_id: &str,
assertion: Option<&str>,
token: &str,
) -> Result<Vec<(&'static str, String)>> {
let mut params = vec![("token", token.to_string())];
params.extend(super::client_auth::credential_params(
method, client_id, assertion,
)?);
Ok(params)
}
const REVOKE_DEADLINE: std::time::Duration = std::time::Duration::from_secs(5);
pub struct RevokeContext<'a> {
pub revocation_endpoint: Option<&'a str>,
pub client_id: &'a str,
pub auth_method: AuthMethod,
pub client_key: Option<&'a super::keys::SigningKey>,
pub deadline: std::time::Duration,
}
pub async fn sign_out(
pool: &sqlx::SqlitePool,
codec: &super::crypto::Codec,
http: &reqwest::Client,
ctx: &RevokeContext<'_>,
sub: &str,
now: i64,
) -> Revocation {
let session = match super::store::get_session(pool, codec, sub).await {
Ok(Some(session)) => session,
Ok(None) => return Revocation::NoSession,
Err(err) => {
let _ = super::store::delete_session(pool, sub).await;
return Revocation::Failed(format!("reading the session: {err:#}"));
}
};
bounded_then_delete(
pool,
sub,
ctx.deadline,
revoke_tokens(pool, http, ctx, &session, now),
)
.await
}
async fn revoke_tokens(
pool: &sqlx::SqlitePool,
http: &reqwest::Client,
ctx: &RevokeContext<'_>,
session: &OAuthSession,
now: i64,
) -> Revocation {
match try_revoke(pool, http, ctx, session, now).await {
Ok(()) => Revocation::Revoked,
Err(err) => Revocation::Failed(format!("{err:#}")),
}
}
async fn try_revoke(
pool: &sqlx::SqlitePool,
http: &reqwest::Client,
ctx: &RevokeContext<'_>,
session: &OAuthSession,
now: i64,
) -> Result<()> {
let endpoint = ctx.revocation_endpoint.ok_or_else(|| {
anyhow::anyhow!("the authorization server advertises no revocation endpoint")
})?;
let (token, _hint) = token_to_revoke(session);
let assertion = match ctx.auth_method {
AuthMethod::PrivateKeyJwt => {
let key = ctx.client_key.ok_or_else(|| {
anyhow::anyhow!("private_key_jwt requires the client signing key")
})?;
Some(super::client_auth::client_assertion(
key,
ctx.client_id,
&session.issuer,
now,
)?)
}
AuthMethod::None => None,
};
let params = revoke_params(ctx.auth_method, ctx.client_id, assertion.as_deref(), token)?;
let form: Vec<(&str, &str)> = params.iter().map(|(k, v)| (*k, v.as_str())).collect();
let key = super::keys::SigningKey::from_jwk_json(&session.dpop_key_jwk, "session")?;
let outcome = super::request::send_with_dpop(
http,
pool,
&super::request::DpopRequest {
endpoint: super::dpop::Endpoint::AuthorizationServer,
url: endpoint,
key: &key,
access_token: None,
body: super::request::DpopBody::Form(&form),
retry: super::request::Retry::Allowed,
},
)
.await?;
if !(200..300).contains(&outcome.status) {
anyhow::bail!("the revocation endpoint returned status {}", outcome.status);
}
Ok(())
}
pub async fn sign_out_discovering(
runtime: &super::runtime::OauthRuntime,
http: &reqwest::Client,
pool: &sqlx::SqlitePool,
sub: &str,
now: i64,
) -> Revocation {
let session = match super::store::get_session(pool, &runtime.codec, sub).await {
Ok(Some(session)) => session,
Ok(None) => return Revocation::NoSession,
Err(err) => {
let _ = super::store::delete_session(pool, sub).await;
return Revocation::Failed(format!("reading the session: {err:#}"));
}
};
let endpoint = match tokio::time::timeout(
REVOKE_DEADLINE,
super::discovery::discover(
http,
&session.aud,
runtime.auth_method.as_str(),
Some(&session.issuer),
),
)
.await
{
Ok(Ok(server)) => server.revocation_endpoint,
Ok(Err(err)) => {
tracing::warn!(%err, %sub, "could not discover the revocation endpoint");
None
}
Err(_) => {
tracing::warn!(%sub, "discovering the revocation endpoint timed out");
None
}
};
sign_out(
pool,
&runtime.codec,
http,
&RevokeContext {
revocation_endpoint: endpoint.as_deref(),
client_id: &runtime.client_id,
auth_method: runtime.auth_method,
client_key: runtime.client_key.as_ref(),
deadline: REVOKE_DEADLINE,
},
sub,
now,
)
.await
}
async fn bounded_then_delete<F>(
pool: &sqlx::SqlitePool,
sub: &str,
deadline: std::time::Duration,
attempt: F,
) -> Revocation
where
F: std::future::Future<Output = Revocation>,
{
let outcome = match tokio::time::timeout(deadline, attempt).await {
Ok(outcome) => outcome,
Err(_) => Revocation::Failed(format!(
"revocation did not finish within {deadline:?}; signing out locally anyway"
)),
};
if let Err(err) = super::store::delete_session(pool, sub).await {
return Revocation::Failed(format!("deleting the local session: {err:#}"));
}
outcome
}
#[cfg(test)]
mod tests {
use super::*;
fn session(access: &str, refresh: &str) -> OAuthSession {
OAuthSession {
sub: "did:plc:ewvi7nxzyoun6zhxrhs64oiz".into(),
issuer: "https://pds.example.com".into(),
aud: "https://pds.example.com".into(),
dpop_key_jwk: r#"{"kty":"EC"}"#.into(),
access_token: access.into(),
refresh_token: refresh.into(),
token_type: "DPoP".into(),
granted_scope: "atproto".into(),
expires_at: Some(1_700_000_000),
}
}
#[test]
fn the_refresh_token_is_preferred_over_the_access_token() {
let session = session("access-abc", "refresh-xyz");
let (token, hint) = token_to_revoke(&session);
assert_eq!(
token, "refresh-xyz",
"revoked the access token, leaving the refresh token live"
);
assert_eq!(hint, "refresh_token");
}
#[test]
fn an_absent_refresh_token_falls_back_to_the_access_token() {
let session = session("access-abc", "");
let (token, hint) = token_to_revoke(&session);
assert_eq!(token, "access-abc");
assert_eq!(hint, "access_token");
}
#[test]
fn a_public_client_sends_the_token_and_its_client_id() {
let params = revoke_params(AuthMethod::None, "http://localhost", None, "refresh-xyz")
.expect("a public client needs no assertion");
assert!(params.contains(&("token", "refresh-xyz".to_string())));
assert!(params.contains(&("client_id", "http://localhost".to_string())));
assert!(
!params
.iter()
.any(|(k, _)| k.starts_with("client_assertion")),
"a public client must not send an assertion it never registered: {params:?}"
);
}
#[test]
fn a_confidential_client_carries_its_assertion() {
let params = revoke_params(
AuthMethod::PrivateKeyJwt,
"https://feather-reader.com/oauth/client-metadata.json",
Some("the.assertion.jwt"),
"refresh-xyz",
)
.expect("an assertion was supplied");
assert!(params.contains(&("client_assertion", "the.assertion.jwt".to_string())));
}
#[test]
fn a_confidential_client_without_an_assertion_is_an_error() {
let err = revoke_params(AuthMethod::PrivateKeyJwt, "https://client", None, "tok")
.expect_err("must not send an unauthenticated revocation");
assert!(format!("{err:#}").contains("requires a client assertion"));
}
const KEY: &str = "aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa";
const DID: &str = "did:plc:ewvi7nxzyoun6zhxrhs64oiz";
const NOW: i64 = 1_700_000_000;
async fn db() -> (sqlx::SqlitePool, super::super::crypto::Codec) {
let pool = crate::store::init_url("sqlite::memory:").await.unwrap();
super::super::store::init_schema(&pool).await.unwrap();
(pool, super::super::crypto::Codec::new(Some(KEY)).unwrap())
}
async fn stored(pool: &sqlx::SqlitePool, codec: &super::super::crypto::Codec) -> OAuthSession {
let key = super::super::keys::SigningKey::generate("session");
let session = OAuthSession {
dpop_key_jwk: key.to_jwk_json().unwrap(),
..session("access-abc", "refresh-xyz")
};
super::super::store::put_session(pool, codec, &session)
.await
.unwrap();
session
}
const TEST_DEADLINE: std::time::Duration = std::time::Duration::from_millis(250);
fn ctx(endpoint: Option<&str>) -> RevokeContext<'_> {
RevokeContext {
revocation_endpoint: endpoint,
client_id: "http://localhost",
auth_method: AuthMethod::None,
client_key: None,
deadline: TEST_DEADLINE,
}
}
#[tokio::test]
async fn signing_out_deletes_the_local_session_even_when_revocation_fails() {
let (pool, codec) = db().await;
stored(&pool, &codec).await;
let outcome = sign_out(
&pool,
&codec,
&reqwest::Client::new(),
&ctx(Some("http://127.0.0.1/oauth/revoke")),
DID,
NOW,
)
.await;
match &outcome {
Revocation::Failed(reason) => assert!(
reason.contains("forbidden (internal) address"),
"failed BEFORE reaching the network, so this proves nothing about a \
revocation failure: {reason}"
),
other => panic!("the loopback endpoint must not report success: {other:?}"),
}
assert!(
super::super::store::get_session(&pool, &codec, DID)
.await
.unwrap()
.is_none(),
"THE SESSION SURVIVED A FAILED REVOCATION — a signed-out user still has live credentials"
);
}
#[tokio::test]
async fn a_server_without_a_revocation_endpoint_still_signs_out_locally() {
let (pool, codec) = db().await;
stored(&pool, &codec).await;
let outcome = sign_out(&pool, &codec, &reqwest::Client::new(), &ctx(None), DID, NOW).await;
match &outcome {
Revocation::Failed(reason) => assert!(
reason.contains("no revocation endpoint"),
"failed for the wrong reason: {reason}"
),
other => panic!("expected a failure, got {other:?}"),
}
assert!(super::super::store::get_session(&pool, &codec, DID)
.await
.unwrap()
.is_none());
}
#[tokio::test]
async fn a_hanging_attempt_does_not_hold_the_sign_out_open() {
let (pool, codec) = db().await;
stored(&pool, &codec).await;
let started = std::time::Instant::now();
let outcome = super::bounded_then_delete(
&pool,
DID,
std::time::Duration::from_millis(50),
std::future::pending::<Revocation>(),
)
.await;
assert!(
started.elapsed() < std::time::Duration::from_secs(2),
"the bound did not fire"
);
match &outcome {
Revocation::Failed(reason) => assert!(
reason.contains("did not finish within"),
"failed for the wrong reason: {reason}"
),
other => panic!("a never-resolving attempt must time out, got {other:?}"),
}
assert!(
super::super::store::get_session(&pool, &codec, DID)
.await
.unwrap()
.is_none(),
"the session survived a timed-out revocation"
);
}
#[tokio::test]
async fn an_unreadable_session_is_still_signed_out() {
let (pool, codec) = db().await;
stored(&pool, &codec).await;
sqlx::query("UPDATE oauth_session SET issuer = ? WHERE sub = ?")
.bind("https://evil.example")
.bind(DID)
.execute(&pool)
.await
.unwrap();
assert!(
super::super::store::get_session(&pool, &codec, DID)
.await
.is_err(),
"precondition: the row must be unreadable"
);
let outcome = sign_out(
&pool,
&codec,
&reqwest::Client::new(),
&ctx(Some("https://pds.example.com/oauth/revoke")),
DID,
NOW,
)
.await;
assert!(matches!(outcome, Revocation::Failed(_)), "got {outcome:?}");
let still_there: i64 =
sqlx::query_scalar("SELECT COUNT(*) FROM oauth_session WHERE sub = ?")
.bind(DID)
.fetch_one(&pool)
.await
.unwrap();
assert_eq!(
still_there, 0,
"an unreadable row survived a sign-out, so the account stays wedged"
);
}
fn runtime() -> super::super::runtime::OauthRuntime {
super::super::runtime::OauthRuntime::new(&crate::config::Config {
repo_backend: crate::metrics::Backend::Rust,
public_url: "http://127.0.0.1:8080".into(),
oauth: crate::config::OauthConfig {
encryption_key: Some(KEY.to_string()),
..crate::config::OauthConfig::default()
},
..crate::config::Config::default()
})
.expect("the test runtime must build")
}
#[tokio::test]
async fn the_production_sign_out_deletes_the_session() {
let (pool, codec) = db().await;
stored(&pool, &codec).await;
let outcome =
sign_out_discovering(&runtime(), &reqwest::Client::new(), &pool, DID, NOW).await;
assert!(
matches!(outcome, Revocation::Failed(_)),
"an unreachable PDS must not report success: {outcome:?}"
);
assert!(
super::super::store::get_session(&pool, &codec, DID)
.await
.unwrap()
.is_none(),
"the production sign-out left the session behind"
);
}
#[tokio::test]
async fn the_production_sign_out_deletes_an_unreadable_session() {
let (pool, codec) = db().await;
stored(&pool, &codec).await;
sqlx::query("UPDATE oauth_session SET issuer = ? WHERE sub = ?")
.bind("https://evil.example")
.bind(DID)
.execute(&pool)
.await
.unwrap();
let outcome =
sign_out_discovering(&runtime(), &reqwest::Client::new(), &pool, DID, NOW).await;
assert!(matches!(outcome, Revocation::Failed(_)), "got {outcome:?}");
let rows: i64 = sqlx::query_scalar("SELECT COUNT(*) FROM oauth_session WHERE sub = ?")
.bind(DID)
.fetch_one(&pool)
.await
.unwrap();
assert_eq!(rows, 0, "account deletion would leave live tokens behind");
}
#[tokio::test]
async fn signing_out_without_a_session_is_idempotent() {
let (pool, codec) = db().await;
let outcome = sign_out(
&pool,
&codec,
&reqwest::Client::new(),
&ctx(Some("https://pds.example.com/oauth/revoke")),
DID,
NOW,
)
.await;
assert_eq!(outcome, Revocation::NoSession);
}
}