use anyhow::{Context as _, Result};
use sqlx::SqlitePool;
use super::crypto::Codec;
const SCHEMA: &str = r#"
-- One row per in-flight login. Short-lived and single-use; see `take_pending`.
CREATE TABLE IF NOT EXISTS oauth_state (
state TEXT PRIMARY KEY NOT NULL,
-- SHA-256 of the cookie value set before the redirect. The callback must
-- present the cookie; without it a callback URL fired by any other browser
-- would complete the login and hand out the session.
browser_binding_hash TEXT NOT NULL,
pkce_verifier TEXT NOT NULL, -- AAD-bound
dpop_key_jwk TEXT NOT NULL, -- AAD-bound
issuer TEXT NOT NULL,
pds_url TEXT NOT NULL,
did TEXT NOT NULL,
-- The negotiated client-auth method is stored so the callback re-creates the
-- same client rather than re-negotiating against possibly-changed metadata.
auth_method TEXT NOT NULL,
auth_kid TEXT,
-- The EXACT redirect_uri sent in PAR; it must match byte-for-byte at the
-- token endpoint.
redirect_uri TEXT NOT NULL,
requested_scope TEXT NOT NULL,
request_uri TEXT NOT NULL,
app_return_to TEXT,
expires_at INTEGER NOT NULL
);
CREATE INDEX IF NOT EXISTS oauth_state_expires_at ON oauth_state(expires_at);
-- One row per authenticated account.
CREATE TABLE IF NOT EXISTS oauth_session (
sub TEXT PRIMARY KEY NOT NULL,
issuer TEXT NOT NULL,
-- The PDS. Every XRPC request is built against this rather than re-derived,
-- so it belongs to the token set.
aud TEXT NOT NULL,
dpop_key_jwk TEXT NOT NULL, -- AAD-bound
access_token TEXT NOT NULL, -- AAD-bound
refresh_token TEXT NOT NULL, -- AAD-bound
token_type TEXT NOT NULL,
granted_scope TEXT NOT NULL,
-- NULL is legitimate: `expires_in` is optional in a token response.
expires_at INTEGER
);
-- Server-issued DPoP nonces, per origin. Persisted rather than used once,
-- because a nonce is expected on every subsequent request to that origin.
CREATE TABLE IF NOT EXISTS oauth_nonce (
origin TEXT PRIMARY KEY NOT NULL,
nonce TEXT NOT NULL,
updated_at INTEGER NOT NULL
);
"#;
pub async fn init_schema(pool: &SqlitePool) -> Result<()> {
sqlx::query(SCHEMA)
.execute(pool)
.await
.context("creating the OAuth tables")?;
Ok(())
}
fn structured_aad(table: &str, fields: &[&str]) -> Vec<u8> {
let mut out = Vec::new();
for field in std::iter::once(&table).chain(fields.iter()) {
out.extend_from_slice(&(field.len() as u64).to_be_bytes());
out.extend_from_slice(field.as_bytes());
}
out
}
struct StateBinding<'a> {
state: &'a str,
issuer: &'a str,
pds_url: &'a str,
did: &'a str,
redirect_uri: &'a str,
browser_binding_hash: &'a str,
auth_method: &'a str,
auth_kid: Option<&'a str>,
requested_scope: &'a str,
request_uri: &'a str,
app_return_to: Option<&'a str>,
expires_at: i64,
}
fn present_or_absent(value: Option<&str>) -> &'static str {
match value {
Some(_) => "present",
None => "absent",
}
}
fn state_aad(binding: &StateBinding<'_>, column: &str) -> Vec<u8> {
let expires_at = binding.expires_at.to_string();
structured_aad(
"oauth_state",
&[
binding.state,
column,
binding.issuer,
binding.pds_url,
binding.did,
binding.redirect_uri,
binding.browser_binding_hash,
binding.auth_method,
present_or_absent(binding.auth_kid),
binding.auth_kid.unwrap_or(""),
binding.requested_scope,
binding.request_uri,
present_or_absent(binding.app_return_to),
binding.app_return_to.unwrap_or(""),
&expires_at,
],
)
}
fn session_aad(
sub: &str,
column: &str,
issuer: &str,
aud: &str,
token_type: &str,
granted_scope: &str,
expires_at: Option<i64>,
) -> Vec<u8> {
let expires_at = expires_at.map_or_else(|| "none".to_string(), |secs| secs.to_string());
structured_aad(
"oauth_session",
&[
sub,
column,
issuer,
aud,
token_type,
granted_scope,
&expires_at,
],
)
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct PendingAuth {
pub state: String,
pub browser_binding_hash: String,
pub pkce_verifier: String,
pub dpop_key_jwk: String,
pub issuer: String,
pub pds_url: String,
pub did: String,
pub auth_method: String,
pub auth_kid: Option<String>,
pub redirect_uri: String,
pub requested_scope: String,
pub request_uri: String,
pub app_return_to: Option<String>,
pub expires_at: i64,
}
#[derive(sqlx::FromRow)]
struct PendingRow {
state: String,
browser_binding_hash: String,
pkce_verifier: String,
dpop_key_jwk: String,
issuer: String,
pds_url: String,
did: String,
auth_method: String,
auth_kid: Option<String>,
redirect_uri: String,
requested_scope: String,
request_uri: String,
app_return_to: Option<String>,
expires_at: i64,
}
pub async fn put_pending(pool: &SqlitePool, codec: &Codec, auth: &PendingAuth) -> Result<()> {
let binding = StateBinding {
state: &auth.state,
issuer: &auth.issuer,
pds_url: &auth.pds_url,
did: &auth.did,
redirect_uri: &auth.redirect_uri,
browser_binding_hash: &auth.browser_binding_hash,
auth_method: &auth.auth_method,
auth_kid: auth.auth_kid.as_deref(),
requested_scope: &auth.requested_scope,
request_uri: &auth.request_uri,
app_return_to: auth.app_return_to.as_deref(),
expires_at: auth.expires_at,
};
sqlx::query(
r#"
INSERT INTO oauth_state (
state, browser_binding_hash, pkce_verifier, dpop_key_jwk, issuer,
pds_url, did, auth_method, auth_kid, redirect_uri, requested_scope,
request_uri, app_return_to, expires_at
) VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9, ?10, ?11, ?12, ?13, ?14)
"#,
)
.bind(&auth.state)
.bind(&auth.browser_binding_hash)
.bind(codec.encrypt_bound(&auth.pkce_verifier, &state_aad(&binding, "pkce_verifier")))
.bind(codec.encrypt_bound(&auth.dpop_key_jwk, &state_aad(&binding, "dpop_key_jwk")))
.bind(&auth.issuer)
.bind(&auth.pds_url)
.bind(&auth.did)
.bind(&auth.auth_method)
.bind(&auth.auth_kid)
.bind(&auth.redirect_uri)
.bind(&auth.requested_scope)
.bind(&auth.request_uri)
.bind(&auth.app_return_to)
.bind(auth.expires_at)
.execute(pool)
.await
.context("recording the pending login")?;
Ok(())
}
pub async fn take_pending(
pool: &SqlitePool,
codec: &Codec,
state: &str,
now: i64,
) -> Result<Option<PendingAuth>> {
let row: Option<PendingRow> = sqlx::query_as(
r#"
DELETE FROM oauth_state WHERE state = ?1
RETURNING state, browser_binding_hash, pkce_verifier, dpop_key_jwk,
issuer, pds_url, did, auth_method, auth_kid, redirect_uri,
requested_scope, request_uri, app_return_to, expires_at
"#,
)
.bind(state)
.fetch_optional(pool)
.await
.context("consuming the pending login")?;
let Some(row) = row else { return Ok(None) };
if row.expires_at <= now {
return Ok(None);
}
let binding = StateBinding {
state: &row.state,
issuer: &row.issuer,
pds_url: &row.pds_url,
did: &row.did,
redirect_uri: &row.redirect_uri,
browser_binding_hash: &row.browser_binding_hash,
auth_method: &row.auth_method,
auth_kid: row.auth_kid.as_deref(),
requested_scope: &row.requested_scope,
request_uri: &row.request_uri,
app_return_to: row.app_return_to.as_deref(),
expires_at: row.expires_at,
};
let aad = |column: &str| state_aad(&binding, column);
Ok(Some(PendingAuth {
pkce_verifier: codec
.decrypt_bound(&row.pkce_verifier, &aad("pkce_verifier"))
.context("decrypting the stored PKCE verifier (or its bound context was altered)")?,
dpop_key_jwk: codec
.decrypt_bound(&row.dpop_key_jwk, &aad("dpop_key_jwk"))
.context("decrypting the stored DPoP key (or its bound context was altered)")?,
state: row.state,
browser_binding_hash: row.browser_binding_hash,
issuer: row.issuer,
pds_url: row.pds_url,
did: row.did,
auth_method: row.auth_method,
auth_kid: row.auth_kid,
redirect_uri: row.redirect_uri,
requested_scope: row.requested_scope,
request_uri: row.request_uri,
app_return_to: row.app_return_to,
expires_at: row.expires_at,
}))
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct OAuthSession {
pub sub: String,
pub issuer: String,
pub aud: String,
pub dpop_key_jwk: String,
pub access_token: String,
pub refresh_token: String,
pub token_type: String,
pub granted_scope: String,
pub expires_at: Option<i64>,
}
#[derive(sqlx::FromRow)]
struct SessionRow {
sub: String,
issuer: String,
aud: String,
dpop_key_jwk: String,
access_token: String,
refresh_token: String,
token_type: String,
granted_scope: String,
expires_at: Option<i64>,
}
pub async fn put_session(pool: &SqlitePool, codec: &Codec, session: &OAuthSession) -> Result<()> {
let [dpop_key_jwk, access_token, refresh_token] = encrypt_secrets(codec, session);
sqlx::query(
r#"
INSERT INTO oauth_session (
sub, issuer, aud, dpop_key_jwk, access_token, refresh_token,
token_type, granted_scope, expires_at
) VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9)
ON CONFLICT(sub) DO UPDATE SET
issuer = excluded.issuer,
aud = excluded.aud,
dpop_key_jwk = excluded.dpop_key_jwk,
access_token = excluded.access_token,
refresh_token = excluded.refresh_token,
token_type = excluded.token_type,
granted_scope = excluded.granted_scope,
expires_at = excluded.expires_at
"#,
)
.bind(&session.sub)
.bind(&session.issuer)
.bind(&session.aud)
.bind(dpop_key_jwk)
.bind(access_token)
.bind(refresh_token)
.bind(&session.token_type)
.bind(&session.granted_scope)
.bind(session.expires_at)
.execute(pool)
.await
.context("storing the OAuth session")?;
Ok(())
}
pub async fn get_session(
pool: &SqlitePool,
codec: &Codec,
sub: &str,
) -> Result<Option<OAuthSession>> {
Ok(get_session_versioned(pool, codec, sub)
.await?
.map(|(session, _)| session))
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct SessionVersion {
dpop_key_jwk: String,
access_token: String,
refresh_token: String,
}
pub async fn get_session_versioned(
pool: &SqlitePool,
codec: &Codec,
sub: &str,
) -> Result<Option<(OAuthSession, SessionVersion)>> {
let row: Option<SessionRow> = sqlx::query_as(
r#"
SELECT sub, issuer, aud, dpop_key_jwk, access_token, refresh_token,
token_type, granted_scope, expires_at
FROM oauth_session WHERE sub = ?1
"#,
)
.bind(sub)
.fetch_optional(pool)
.await
.context("reading the OAuth session")?;
let Some(row) = row else { return Ok(None) };
let aad = |column: &str| {
session_aad(
&row.sub,
column,
&row.issuer,
&row.aud,
&row.token_type,
&row.granted_scope,
row.expires_at,
)
};
let session = OAuthSession {
dpop_key_jwk: codec
.decrypt_bound(&row.dpop_key_jwk, &aad("dpop_key_jwk"))
.context("decrypting the session DPoP key (or its bound context was altered)")?,
access_token: codec
.decrypt_bound(&row.access_token, &aad("access_token"))
.context("decrypting the stored access token (or its bound context was altered)")?,
refresh_token: codec
.decrypt_bound(&row.refresh_token, &aad("refresh_token"))
.context("decrypting the stored refresh token (or its bound context was altered)")?,
sub: row.sub,
issuer: row.issuer,
aud: row.aud,
token_type: row.token_type,
granted_scope: row.granted_scope,
expires_at: row.expires_at,
};
let version = SessionVersion {
dpop_key_jwk: row.dpop_key_jwk,
access_token: row.access_token,
refresh_token: row.refresh_token,
};
Ok(Some((session, version)))
}
pub async fn update_session_if_unchanged(
pool: &SqlitePool,
codec: &Codec,
session: &OAuthSession,
version: &SessionVersion,
) -> Result<bool> {
let [dpop_key_jwk, access_token, refresh_token] = encrypt_secrets(codec, session);
let result = sqlx::query(
r#"
UPDATE oauth_session SET
issuer = ?2,
aud = ?3,
dpop_key_jwk = ?4,
access_token = ?5,
refresh_token = ?6,
token_type = ?7,
granted_scope = ?8,
expires_at = ?9
WHERE sub = ?1
AND dpop_key_jwk = ?10 AND access_token = ?11 AND refresh_token = ?12
"#,
)
.bind(&session.sub)
.bind(&session.issuer)
.bind(&session.aud)
.bind(dpop_key_jwk)
.bind(access_token)
.bind(refresh_token)
.bind(&session.token_type)
.bind(&session.granted_scope)
.bind(session.expires_at)
.bind(&version.dpop_key_jwk)
.bind(&version.access_token)
.bind(&version.refresh_token)
.execute(pool)
.await
.context("updating the OAuth session (if unchanged)")?;
Ok(result.rows_affected() > 0)
}
fn encrypt_secrets(codec: &Codec, session: &OAuthSession) -> [String; 3] {
let bound = |column: &str, plaintext: &str| {
codec.encrypt_bound(
plaintext,
&session_aad(
&session.sub,
column,
&session.issuer,
&session.aud,
&session.token_type,
&session.granted_scope,
session.expires_at,
),
)
};
[
bound("dpop_key_jwk", &session.dpop_key_jwk),
bound("access_token", &session.access_token),
bound("refresh_token", &session.refresh_token),
]
}
pub async fn delete_session_if_unchanged(
pool: &SqlitePool,
sub: &str,
version: &SessionVersion,
) -> Result<bool> {
let result = sqlx::query(
"DELETE FROM oauth_session WHERE sub = ?1 AND dpop_key_jwk = ?2 \
AND access_token = ?3 AND refresh_token = ?4",
)
.bind(sub)
.bind(&version.dpop_key_jwk)
.bind(&version.access_token)
.bind(&version.refresh_token)
.execute(pool)
.await
.context("deleting the OAuth session (if unchanged)")?;
Ok(result.rows_affected() > 0)
}
pub async fn list_session_subs(pool: &SqlitePool) -> Result<Vec<String>> {
sqlx::query_scalar("SELECT sub FROM oauth_session ORDER BY sub")
.fetch_all(pool)
.await
.context("listing the OAuth sessions")
}
pub async fn delete_session(pool: &SqlitePool, sub: &str) -> Result<bool> {
let result = sqlx::query("DELETE FROM oauth_session WHERE sub = ?1")
.bind(sub)
.execute(pool)
.await
.context("deleting the OAuth session")?;
Ok(result.rows_affected() > 0)
}
pub async fn sweep_expired_pending(pool: &SqlitePool, now: i64) -> Result<u64> {
let result = sqlx::query("DELETE FROM oauth_state WHERE expires_at <= ?1")
.bind(now)
.execute(pool)
.await
.context("sweeping expired pending logins")?;
Ok(result.rows_affected())
}
pub async fn sweep_stale_nonces(pool: &SqlitePool, cutoff: i64) -> Result<u64> {
let result = sqlx::query("DELETE FROM oauth_nonce WHERE updated_at <= ?1")
.bind(cutoff)
.execute(pool)
.await
.context("sweeping stale DPoP nonces")?;
Ok(result.rows_affected())
}
pub async fn get_nonce(pool: &SqlitePool, origin: &str) -> Result<Option<String>> {
sqlx::query_scalar("SELECT nonce FROM oauth_nonce WHERE origin = ?1")
.bind(origin)
.fetch_optional(pool)
.await
.context("reading the stored DPoP nonce")
}
pub async fn put_nonce(pool: &SqlitePool, origin: &str, nonce: &str, now: i64) -> Result<()> {
sqlx::query(
r#"
INSERT INTO oauth_nonce (origin, nonce, updated_at) VALUES (?1, ?2, ?3)
ON CONFLICT(origin) DO UPDATE SET
nonce = excluded.nonce, updated_at = excluded.updated_at
"#,
)
.bind(origin)
.bind(nonce)
.bind(now)
.execute(pool)
.await
.context("storing the DPoP nonce")?;
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
use crate::store::init_url;
const KEY: &str = "aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa";
const DID: &str = "did:plc:ewvi7nxzyoun6zhxrhs64oiz";
const NOW: i64 = 1_700_000_000;
async fn db() -> (sqlx::SqlitePool, Codec) {
let pool = init_url("sqlite::memory:").await.unwrap();
init_schema(&pool).await.unwrap();
(pool, Codec::new(Some(KEY)).unwrap())
}
fn leak(sql: String) -> &'static str {
Box::leak(sql.into_boxed_str())
}
fn pending(state: &str) -> PendingAuth {
PendingAuth {
state: state.to_string(),
browser_binding_hash: "hash-of-cookie".into(),
pkce_verifier: "verifier-secret".into(),
dpop_key_jwk: r#"{"kty":"EC","d":"secret"}"#.into(),
issuer: "https://auth.example.com".into(),
pds_url: "https://pds.example.com".into(),
did: DID.into(),
auth_method: "private_key_jwt".into(),
auth_kid: Some("featherreader-oauth-1".into()),
redirect_uri: "https://feather-reader.com/oauth/callback".into(),
requested_scope: "atproto transition:generic".into(),
request_uri: "urn:ietf:params:oauth:request_uri:abc".into(),
app_return_to: Some("/reader".into()),
expires_at: NOW + 600,
}
}
#[tokio::test]
async fn a_pending_login_round_trips() -> anyhow::Result<()> {
let (pool, codec) = db().await;
let want = pending("state-1");
put_pending(&pool, &codec, &want).await?;
let got = take_pending(&pool, &codec, "state-1", NOW).await?.unwrap();
assert_eq!(got, want);
Ok(())
}
#[tokio::test]
async fn a_pending_login_can_only_be_taken_once() -> anyhow::Result<()> {
let (pool, codec) = db().await;
put_pending(&pool, &codec, &pending("state-1")).await?;
assert!(take_pending(&pool, &codec, "state-1", NOW).await?.is_some());
assert!(take_pending(&pool, &codec, "state-1", NOW).await?.is_none());
Ok(())
}
#[tokio::test]
async fn concurrent_takes_yield_exactly_one_winner() -> anyhow::Result<()> {
let (pool, codec) = db().await;
put_pending(&pool, &codec, &pending("race")).await?;
let (a, b, c, d) = tokio::join!(
take_pending(&pool, &codec, "race", NOW),
take_pending(&pool, &codec, "race", NOW),
take_pending(&pool, &codec, "race", NOW),
take_pending(&pool, &codec, "race", NOW),
);
let winners = [a?, b?, c?, d?].iter().filter(|r| r.is_some()).count();
assert_eq!(winners, 1, "more than one caller consumed the same state");
Ok(())
}
#[tokio::test]
async fn an_expired_pending_login_is_rejected_and_removed() -> anyhow::Result<()> {
let (pool, codec) = db().await;
put_pending(&pool, &codec, &pending("stale")).await?;
let after_expiry = NOW + 601;
assert!(take_pending(&pool, &codec, "stale", after_expiry)
.await?
.is_none());
assert!(take_pending(&pool, &codec, "stale", NOW).await?.is_none());
Ok(())
}
#[tokio::test]
async fn an_unknown_state_is_simply_absent() -> anyhow::Result<()> {
let (pool, codec) = db().await;
assert!(take_pending(&pool, &codec, "never-existed", NOW)
.await?
.is_none());
Ok(())
}
#[tokio::test]
async fn secret_columns_are_stored_encrypted() -> anyhow::Result<()> {
let (pool, codec) = db().await;
put_pending(&pool, &codec, &pending("state-1")).await?;
let (verifier, jwk): (String, String) =
sqlx::query_as("SELECT pkce_verifier, dpop_key_jwk FROM oauth_state WHERE state = ?")
.bind("state-1")
.fetch_one(&pool)
.await?;
for stored in [&verifier, &jwk] {
assert!(stored.starts_with("enc.v2.gcm."), "not bound: {stored}");
}
assert!(!verifier.contains("verifier-secret"));
assert!(!jwk.contains("secret"));
Ok(())
}
#[tokio::test]
async fn a_secret_moved_between_rows_does_not_decrypt() -> anyhow::Result<()> {
let (pool, codec) = db().await;
put_pending(&pool, &codec, &pending("victim")).await?;
let mut attacker = pending("attacker");
attacker.dpop_key_jwk = r#"{"kty":"EC","d":"attacker-key"}"#.into();
put_pending(&pool, &codec, &attacker).await?;
let stolen: String =
sqlx::query_scalar("SELECT dpop_key_jwk FROM oauth_state WHERE state = ?")
.bind("attacker")
.fetch_one(&pool)
.await?;
sqlx::query("UPDATE oauth_state SET dpop_key_jwk = ? WHERE state = ?")
.bind(&stolen)
.bind("victim")
.execute(&pool)
.await?;
assert!(
take_pending(&pool, &codec, "victim", NOW).await.is_err(),
"a grafted ciphertext decrypted in the wrong row"
);
Ok(())
}
#[tokio::test]
async fn a_secret_moved_between_columns_does_not_decrypt() -> anyhow::Result<()> {
let (pool, codec) = db().await;
put_pending(&pool, &codec, &pending("state-1")).await?;
let verifier: String =
sqlx::query_scalar("SELECT pkce_verifier FROM oauth_state WHERE state = ?")
.bind("state-1")
.fetch_one(&pool)
.await?;
sqlx::query("UPDATE oauth_state SET dpop_key_jwk = ? WHERE state = ?")
.bind(&verifier)
.bind("state-1")
.execute(&pool)
.await?;
assert!(take_pending(&pool, &codec, "state-1", NOW).await.is_err());
Ok(())
}
#[tokio::test]
async fn a_session_secret_moved_between_columns_does_not_decrypt() -> anyhow::Result<()> {
let (pool, codec) = db().await;
put_session(&pool, &codec, &session()).await?;
let access: String =
sqlx::query_scalar("SELECT access_token FROM oauth_session WHERE sub = ?")
.bind(DID)
.fetch_one(&pool)
.await?;
sqlx::query("UPDATE oauth_session SET refresh_token = ? WHERE sub = ?")
.bind(&access)
.bind(DID)
.execute(&pool)
.await?;
assert!(
get_session(&pool, &codec, DID).await.is_err(),
"the access token's ciphertext was accepted in the refresh_token column"
);
Ok(())
}
#[tokio::test]
async fn an_absent_expiry_and_a_zero_expiry_are_different_sessions() -> anyhow::Result<()> {
for (stored, flipped_to) in [(None, "0"), (Some(0), "NULL")] {
let (pool, codec) = db().await;
put_session(
&pool,
&codec,
&OAuthSession {
expires_at: stored,
..session()
},
)
.await?;
sqlx::query(sqlx::AssertSqlSafe(format!(
"UPDATE oauth_session SET expires_at = {flipped_to} WHERE sub = ?"
)))
.bind(DID)
.execute(&pool)
.await?;
assert!(
get_session(&pool, &codec, DID).await.is_err(),
"expires_at {stored:?} → {flipped_to} still decrypted"
);
}
Ok(())
}
#[tokio::test]
async fn tampering_with_a_pending_logins_destinations_breaks_it() -> anyhow::Result<()> {
for column in [
"issuer",
"pds_url",
"did",
"redirect_uri",
"browser_binding_hash",
"auth_method",
"auth_kid",
"requested_scope",
"request_uri",
"app_return_to",
] {
let (pool, codec) = db().await;
put_pending(&pool, &codec, &pending("state-1")).await?;
sqlx::query(leak(format!(
"UPDATE oauth_state SET {column} = ? WHERE state = ?"
)))
.bind("https://evil.example")
.bind("state-1")
.execute(&pool)
.await?;
assert!(
take_pending(&pool, &codec, "state-1", NOW).await.is_err(),
"tampering with `{column}` went undetected"
);
}
Ok(())
}
#[tokio::test]
async fn tampering_with_a_sessions_destinations_breaks_it() -> anyhow::Result<()> {
for column in [
"aud",
"issuer",
"token_type",
"granted_scope",
] {
let (pool, codec) = db().await;
put_session(&pool, &codec, &session()).await?;
sqlx::query(leak(format!(
"UPDATE oauth_session SET {column} = ? WHERE sub = ?"
)))
.bind("https://evil.example")
.bind(DID)
.execute(&pool)
.await?;
assert!(
get_session(&pool, &codec, DID).await.is_err(),
"tampering with `{column}` went undetected"
);
}
Ok(())
}
#[tokio::test]
async fn stale_nonces_are_swept_and_fresh_ones_kept() -> anyhow::Result<()> {
let (pool, _codec) = db().await;
put_nonce(&pool, "https://old.example", "n1", NOW - 10_000).await?;
put_nonce(&pool, "https://new.example", "n2", NOW).await?;
assert_eq!(sweep_stale_nonces(&pool, NOW - 5_000).await?, 1);
assert_eq!(get_nonce(&pool, "https://old.example").await?, None);
assert_eq!(
get_nonce(&pool, "https://new.example").await?.as_deref(),
Some("n2"),
"a nonce still in use was swept"
);
Ok(())
}
#[tokio::test]
async fn swapping_an_absent_optional_column_for_an_empty_one_breaks_it() -> anyhow::Result<()> {
for (column, set_to_empty) in [
("auth_kid", true),
("auth_kid", false),
("app_return_to", true),
("app_return_to", false),
] {
let (pool, codec) = db().await;
let mut auth = pending("state-1");
if set_to_empty {
if column == "auth_kid" {
auth.auth_kid = None;
} else {
auth.app_return_to = None;
}
} else {
if column == "auth_kid" {
auth.auth_kid = Some(String::new());
} else {
auth.app_return_to = Some(String::new());
}
}
put_pending(&pool, &codec, &auth).await?;
let sql = leak(format!(
"UPDATE oauth_state SET {column} = ? WHERE state = ?"
));
let query = if set_to_empty {
sqlx::query(sql).bind(Some(String::new()))
} else {
sqlx::query(sql).bind(Option::<String>::None)
};
query.bind("state-1").execute(&pool).await?;
assert!(
take_pending(&pool, &codec, "state-1", NOW).await.is_err(),
"`{column}`: {} went undetected",
if set_to_empty {
"NULL -> ''"
} else {
"'' -> NULL"
}
);
}
Ok(())
}
#[tokio::test]
async fn extending_a_pending_logins_expiry_breaks_it() -> anyhow::Result<()> {
let (pool, codec) = db().await;
put_pending(&pool, &codec, &pending("state-1")).await?;
sqlx::query("UPDATE oauth_state SET expires_at = ? WHERE state = ?")
.bind(NOW + 31_536_000)
.bind("state-1")
.execute(&pool)
.await?;
assert!(
take_pending(&pool, &codec, "state-1", NOW).await.is_err(),
"the expiry was extended without breaking the row"
);
Ok(())
}
#[tokio::test]
async fn clearing_a_sessions_expiry_breaks_it() -> anyhow::Result<()> {
let (pool, codec) = db().await;
put_session(&pool, &codec, &session()).await?;
sqlx::query("UPDATE oauth_session SET expires_at = NULL WHERE sub = ?")
.bind(DID)
.execute(&pool)
.await?;
assert!(
get_session(&pool, &codec, DID).await.is_err(),
"the expiry was cleared without breaking the row"
);
Ok(())
}
#[test]
fn the_aad_encoding_is_unambiguous_across_field_boundaries() {
assert_ne!(
structured_aad("t", &["ab", "c"]),
structured_aad("t", &["a", "bc"])
);
assert_ne!(
structured_aad("t", &["a:b"]),
structured_aad("t", &["a", "b"])
);
assert_ne!(
structured_aad("t", &["a", ""]),
structured_aad("t", &["", "a"])
);
assert_ne!(structured_aad("t1", &["a"]), structured_aad("t2", &["a"]));
}
#[tokio::test]
async fn a_pending_login_round_trips_with_its_optional_fields_absent() -> anyhow::Result<()> {
let (pool, codec) = db().await;
let mut want = pending("state-1");
want.auth_kid = None;
want.app_return_to = None;
put_pending(&pool, &codec, &want).await?;
assert_eq!(
take_pending(&pool, &codec, "state-1", NOW).await?.unwrap(),
want
);
Ok(())
}
#[tokio::test]
async fn a_pending_login_is_expired_at_exactly_its_expiry() -> anyhow::Result<()> {
let (pool, codec) = db().await;
put_pending(&pool, &codec, &pending("edge")).await?;
assert!(take_pending(&pool, &codec, "edge", NOW + 600)
.await?
.is_none());
let (pool, codec) = db().await;
put_pending(&pool, &codec, &pending("edge")).await?;
assert!(take_pending(&pool, &codec, "edge", NOW + 599)
.await?
.is_some());
Ok(())
}
#[tokio::test]
async fn expired_pending_logins_are_swept() -> anyhow::Result<()> {
let (pool, codec) = db().await;
put_pending(&pool, &codec, &pending("old")).await?;
let mut fresh = pending("fresh");
fresh.expires_at = NOW + 3600;
put_pending(&pool, &codec, &fresh).await?;
assert_eq!(sweep_expired_pending(&pool, NOW + 700).await?, 1);
assert!(take_pending(&pool, &codec, "old", NOW).await?.is_none());
assert!(take_pending(&pool, &codec, "fresh", NOW).await?.is_some());
Ok(())
}
#[tokio::test]
async fn an_unbound_ciphertext_is_refused() -> anyhow::Result<()> {
let (pool, codec) = db().await;
put_pending(&pool, &codec, &pending("state-1")).await?;
sqlx::query("UPDATE oauth_state SET pkce_verifier = ? WHERE state = ?")
.bind(codec.encrypt("verifier-secret"))
.bind("state-1")
.execute(&pool)
.await?;
assert!(take_pending(&pool, &codec, "state-1", NOW).await.is_err());
Ok(())
}
#[tokio::test]
async fn an_unbound_session_ciphertext_is_refused() -> anyhow::Result<()> {
for column in ["access_token", "refresh_token", "dpop_key_jwk"] {
let (pool, codec) = db().await;
put_session(&pool, &codec, &session()).await?;
sqlx::query(leak(format!(
"UPDATE oauth_session SET {column} = ? WHERE sub = ?"
)))
.bind(codec.encrypt("some-value"))
.bind(DID)
.execute(&pool)
.await?;
assert!(
get_session(&pool, &codec, DID).await.is_err(),
"an unbound value was accepted in `{column}`"
);
}
Ok(())
}
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":"session-key"}"#.into(),
access_token: "access-abc".into(),
refresh_token: "refresh-xyz".into(),
token_type: "DPoP".into(),
granted_scope: "atproto transition:generic".into(),
expires_at: Some(NOW + 3600),
}
}
#[tokio::test]
async fn a_session_round_trips() -> anyhow::Result<()> {
let (pool, codec) = db().await;
put_session(&pool, &codec, &session()).await?;
assert_eq!(get_session(&pool, &codec, DID).await?.unwrap(), session());
Ok(())
}
#[tokio::test]
async fn re_login_replaces_the_existing_session() -> anyhow::Result<()> {
let (pool, codec) = db().await;
put_session(&pool, &codec, &session()).await?;
let mut second = session();
second.access_token = "access-second".into();
second.refresh_token = "refresh-second".into();
put_session(&pool, &codec, &second).await?;
let got = get_session(&pool, &codec, DID).await?.unwrap();
assert_eq!(got.access_token, "access-second");
assert_eq!(got.refresh_token, "refresh-second");
Ok(())
}
#[tokio::test]
async fn a_session_without_an_expiry_round_trips() -> anyhow::Result<()> {
let (pool, codec) = db().await;
let mut s = session();
s.expires_at = None;
put_session(&pool, &codec, &s).await?;
assert_eq!(
get_session(&pool, &codec, DID).await?.unwrap().expires_at,
None
);
Ok(())
}
#[tokio::test]
async fn session_tokens_are_bound_to_their_subject() -> anyhow::Result<()> {
let (pool, codec) = db().await;
put_session(&pool, &codec, &session()).await?;
let other = OAuthSession {
sub: "did:plc:aaaaaaaaaaaaaaaaaaaaaaaa".into(),
access_token: "access-other".into(),
..session()
};
put_session(&pool, &codec, &other).await?;
let stolen: String =
sqlx::query_scalar("SELECT access_token FROM oauth_session WHERE sub = ?")
.bind(&other.sub)
.fetch_one(&pool)
.await?;
sqlx::query("UPDATE oauth_session SET access_token = ? WHERE sub = ?")
.bind(&stolen)
.bind(DID)
.execute(&pool)
.await?;
assert!(get_session(&pool, &codec, DID).await.is_err());
Ok(())
}
#[tokio::test]
async fn a_rewritten_session_is_not_deleted_by_a_stale_version() -> anyhow::Result<()> {
let (pool, codec) = db().await;
put_session(&pool, &codec, &session()).await?;
let (_, stale) = get_session_versioned(&pool, &codec, DID).await?.unwrap();
let rotated = OAuthSession {
refresh_token: "refresh-rotated".into(),
..session()
};
put_session(&pool, &codec, &rotated).await?;
assert!(
!delete_session_if_unchanged(&pool, DID, &stale).await?,
"reported deleting a row it should have left"
);
assert_eq!(
get_session(&pool, &codec, DID)
.await?
.unwrap()
.refresh_token,
"refresh-rotated",
"the ROTATED token was deleted on the strength of a stale read"
);
let (_, before) = get_session_versioned(&pool, &codec, DID).await?.unwrap();
put_session(&pool, &codec, &rotated).await?;
assert!(!delete_session_if_unchanged(&pool, DID, &before).await?);
let (_, current) = get_session_versioned(&pool, &codec, DID).await?.unwrap();
assert!(delete_session_if_unchanged(&pool, DID, ¤t).await?);
assert!(get_session(&pool, &codec, DID).await?.is_none());
assert!(!delete_session_if_unchanged(&pool, DID, ¤t).await?);
Ok(())
}
#[tokio::test]
async fn a_deleted_session_is_gone() -> anyhow::Result<()> {
let (pool, codec) = db().await;
put_session(&pool, &codec, &session()).await?;
assert!(delete_session(&pool, DID).await?);
assert!(get_session(&pool, &codec, DID).await?.is_none());
assert!(!delete_session(&pool, DID).await?);
Ok(())
}
#[tokio::test]
async fn every_session_subject_is_listed_including_an_unreadable_one() -> anyhow::Result<()> {
let (pool, codec) = db().await;
assert!(list_session_subs(&pool).await?.is_empty());
put_session(&pool, &codec, &session()).await?;
let other = OAuthSession {
sub: "did:plc:bbbbbbbbbbbbbbbbbbbbbbbb".into(),
..session()
};
put_session(&pool, &codec, &other).await?;
sqlx::query(
"INSERT INTO oauth_session (sub, issuer, aud, dpop_key_jwk, access_token, \
refresh_token, token_type, granted_scope, expires_at) \
VALUES (?, 'https://auth.example.com', 'https://pds.example.com', \
'garbage', 'garbage', 'garbage', 'DPoP', 'atproto', NULL)",
)
.bind("did:plc:cccccccccccccccccccccccc")
.execute(&pool)
.await?;
assert!(
get_session(&pool, &codec, "did:plc:cccccccccccccccccccccccc")
.await
.is_err(),
"precondition: the raw row must be unreadable"
);
assert_eq!(
list_session_subs(&pool).await?,
vec![
"did:plc:bbbbbbbbbbbbbbbbbbbbbbbb".to_string(),
"did:plc:cccccccccccccccccccccccc".to_string(),
DID.to_string(),
]
);
Ok(())
}
#[tokio::test]
async fn nonces_are_stored_and_replaced_per_origin() -> anyhow::Result<()> {
let (pool, _) = db().await;
assert_eq!(get_nonce(&pool, "https://a.example").await?, None);
put_nonce(&pool, "https://a.example", "n1", NOW).await?;
put_nonce(&pool, "https://b.example", "n2", NOW).await?;
assert_eq!(
get_nonce(&pool, "https://a.example").await?.as_deref(),
Some("n1")
);
assert_eq!(
get_nonce(&pool, "https://b.example").await?.as_deref(),
Some("n2")
);
put_nonce(&pool, "https://a.example", "n3", NOW).await?;
assert_eq!(
get_nonce(&pool, "https://a.example").await?.as_deref(),
Some("n3")
);
Ok(())
}
}