1use anyhow::{Context as _, Result};
24use sqlx::SqlitePool;
25
26use super::crypto::Codec;
27
28const SCHEMA: &str = r#"
31-- One row per in-flight login. Short-lived and single-use; see `take_pending`.
32CREATE TABLE IF NOT EXISTS oauth_state (
33 state TEXT PRIMARY KEY NOT NULL,
34 -- SHA-256 of the cookie value set before the redirect. The callback must
35 -- present the cookie; without it a callback URL fired by any other browser
36 -- would complete the login and hand out the session.
37 browser_binding_hash TEXT NOT NULL,
38 pkce_verifier TEXT NOT NULL, -- AAD-bound
39 dpop_key_jwk TEXT NOT NULL, -- AAD-bound
40 issuer TEXT NOT NULL,
41 pds_url TEXT NOT NULL,
42 did TEXT NOT NULL,
43 -- The negotiated client-auth method is stored so the callback re-creates the
44 -- same client rather than re-negotiating against possibly-changed metadata.
45 auth_method TEXT NOT NULL,
46 auth_kid TEXT,
47 -- The EXACT redirect_uri sent in PAR; it must match byte-for-byte at the
48 -- token endpoint.
49 redirect_uri TEXT NOT NULL,
50 requested_scope TEXT NOT NULL,
51 request_uri TEXT NOT NULL,
52 app_return_to TEXT,
53 expires_at INTEGER NOT NULL
54);
55CREATE INDEX IF NOT EXISTS oauth_state_expires_at ON oauth_state(expires_at);
56
57-- One row per authenticated account.
58CREATE TABLE IF NOT EXISTS oauth_session (
59 sub TEXT PRIMARY KEY NOT NULL,
60 issuer TEXT NOT NULL,
61 -- The PDS. Every XRPC request is built against this rather than re-derived,
62 -- so it belongs to the token set.
63 aud TEXT NOT NULL,
64 dpop_key_jwk TEXT NOT NULL, -- AAD-bound
65 access_token TEXT NOT NULL, -- AAD-bound
66 refresh_token TEXT NOT NULL, -- AAD-bound
67 token_type TEXT NOT NULL,
68 granted_scope TEXT NOT NULL,
69 -- NULL is legitimate: `expires_in` is optional in a token response.
70 expires_at INTEGER
71);
72
73-- Server-issued DPoP nonces, per origin. Persisted rather than used once,
74-- because a nonce is expected on every subsequent request to that origin.
75CREATE TABLE IF NOT EXISTS oauth_nonce (
76 origin TEXT PRIMARY KEY NOT NULL,
77 nonce TEXT NOT NULL,
78 updated_at INTEGER NOT NULL
79);
80"#;
81
82pub async fn init_schema(pool: &SqlitePool) -> Result<()> {
84 sqlx::query(SCHEMA)
85 .execute(pool)
86 .await
87 .context("creating the OAuth tables")?;
88 Ok(())
89}
90
91fn structured_aad(table: &str, fields: &[&str]) -> Vec<u8> {
99 let mut out = Vec::new();
100 for field in std::iter::once(&table).chain(fields.iter()) {
101 out.extend_from_slice(&(field.len() as u64).to_be_bytes());
102 out.extend_from_slice(field.as_bytes());
103 }
104 out
105}
106
107struct StateBinding<'a> {
126 state: &'a str,
127 issuer: &'a str,
128 pds_url: &'a str,
129 did: &'a str,
130 redirect_uri: &'a str,
131 browser_binding_hash: &'a str,
132 auth_method: &'a str,
133 auth_kid: Option<&'a str>,
134 requested_scope: &'a str,
135 request_uri: &'a str,
136 app_return_to: Option<&'a str>,
137 expires_at: i64,
138}
139
140fn present_or_absent(value: Option<&str>) -> &'static str {
146 match value {
147 Some(_) => "present",
148 None => "absent",
149 }
150}
151
152fn state_aad(binding: &StateBinding<'_>, column: &str) -> Vec<u8> {
158 let expires_at = binding.expires_at.to_string();
159 structured_aad(
160 "oauth_state",
161 &[
162 binding.state,
163 column,
164 binding.issuer,
165 binding.pds_url,
166 binding.did,
167 binding.redirect_uri,
168 binding.browser_binding_hash,
169 binding.auth_method,
170 present_or_absent(binding.auth_kid),
177 binding.auth_kid.unwrap_or(""),
178 binding.requested_scope,
179 binding.request_uri,
180 present_or_absent(binding.app_return_to),
186 binding.app_return_to.unwrap_or(""),
187 &expires_at,
195 ],
196 )
197}
198
199fn session_aad(
205 sub: &str,
206 column: &str,
207 issuer: &str,
208 aud: &str,
209 token_type: &str,
210 granted_scope: &str,
211 expires_at: Option<i64>,
212) -> Vec<u8> {
213 let expires_at = expires_at.map_or_else(|| "none".to_string(), |secs| secs.to_string());
216 structured_aad(
217 "oauth_session",
218 &[
219 sub,
220 column,
221 issuer,
222 aud,
223 token_type,
226 granted_scope,
227 &expires_at,
230 ],
231 )
232}
233
234#[derive(Debug, Clone, PartialEq, Eq)]
236pub struct PendingAuth {
237 pub state: String,
238 pub browser_binding_hash: String,
239 pub pkce_verifier: String,
240 pub dpop_key_jwk: String,
241 pub issuer: String,
242 pub pds_url: String,
243 pub did: String,
244 pub auth_method: String,
245 pub auth_kid: Option<String>,
246 pub redirect_uri: String,
247 pub requested_scope: String,
248 pub request_uri: String,
249 pub app_return_to: Option<String>,
250 pub expires_at: i64,
251}
252
253#[derive(sqlx::FromRow)]
255struct PendingRow {
256 state: String,
257 browser_binding_hash: String,
258 pkce_verifier: String,
259 dpop_key_jwk: String,
260 issuer: String,
261 pds_url: String,
262 did: String,
263 auth_method: String,
264 auth_kid: Option<String>,
265 redirect_uri: String,
266 requested_scope: String,
267 request_uri: String,
268 app_return_to: Option<String>,
269 expires_at: i64,
270}
271
272pub async fn put_pending(pool: &SqlitePool, codec: &Codec, auth: &PendingAuth) -> Result<()> {
274 let binding = StateBinding {
275 state: &auth.state,
276 issuer: &auth.issuer,
277 pds_url: &auth.pds_url,
278 did: &auth.did,
279 redirect_uri: &auth.redirect_uri,
280 browser_binding_hash: &auth.browser_binding_hash,
281 auth_method: &auth.auth_method,
282 auth_kid: auth.auth_kid.as_deref(),
283 requested_scope: &auth.requested_scope,
284 request_uri: &auth.request_uri,
285 app_return_to: auth.app_return_to.as_deref(),
286 expires_at: auth.expires_at,
287 };
288
289 sqlx::query(
290 r#"
291 INSERT INTO oauth_state (
292 state, browser_binding_hash, pkce_verifier, dpop_key_jwk, issuer,
293 pds_url, did, auth_method, auth_kid, redirect_uri, requested_scope,
294 request_uri, app_return_to, expires_at
295 ) VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9, ?10, ?11, ?12, ?13, ?14)
296 "#,
297 )
298 .bind(&auth.state)
299 .bind(&auth.browser_binding_hash)
300 .bind(codec.encrypt_bound(&auth.pkce_verifier, &state_aad(&binding, "pkce_verifier")))
301 .bind(codec.encrypt_bound(&auth.dpop_key_jwk, &state_aad(&binding, "dpop_key_jwk")))
302 .bind(&auth.issuer)
303 .bind(&auth.pds_url)
304 .bind(&auth.did)
305 .bind(&auth.auth_method)
306 .bind(&auth.auth_kid)
307 .bind(&auth.redirect_uri)
308 .bind(&auth.requested_scope)
309 .bind(&auth.request_uri)
310 .bind(&auth.app_return_to)
311 .bind(auth.expires_at)
312 .execute(pool)
313 .await
314 .context("recording the pending login")?;
315 Ok(())
316}
317
318pub async fn take_pending(
328 pool: &SqlitePool,
329 codec: &Codec,
330 state: &str,
331 now: i64,
332) -> Result<Option<PendingAuth>> {
333 let row: Option<PendingRow> = sqlx::query_as(
334 r#"
335 DELETE FROM oauth_state WHERE state = ?1
336 RETURNING state, browser_binding_hash, pkce_verifier, dpop_key_jwk,
337 issuer, pds_url, did, auth_method, auth_kid, redirect_uri,
338 requested_scope, request_uri, app_return_to, expires_at
339 "#,
340 )
341 .bind(state)
342 .fetch_optional(pool)
343 .await
344 .context("consuming the pending login")?;
345
346 let Some(row) = row else { return Ok(None) };
347 if row.expires_at <= now {
349 return Ok(None);
351 }
352
353 let binding = StateBinding {
356 state: &row.state,
357 issuer: &row.issuer,
358 pds_url: &row.pds_url,
359 did: &row.did,
360 redirect_uri: &row.redirect_uri,
361 browser_binding_hash: &row.browser_binding_hash,
362 auth_method: &row.auth_method,
363 auth_kid: row.auth_kid.as_deref(),
364 requested_scope: &row.requested_scope,
365 request_uri: &row.request_uri,
366 app_return_to: row.app_return_to.as_deref(),
367 expires_at: row.expires_at,
368 };
369 let aad = |column: &str| state_aad(&binding, column);
370 Ok(Some(PendingAuth {
371 pkce_verifier: codec
372 .decrypt_bound(&row.pkce_verifier, &aad("pkce_verifier"))
373 .context("decrypting the stored PKCE verifier (or its bound context was altered)")?,
374 dpop_key_jwk: codec
375 .decrypt_bound(&row.dpop_key_jwk, &aad("dpop_key_jwk"))
376 .context("decrypting the stored DPoP key (or its bound context was altered)")?,
377 state: row.state,
378 browser_binding_hash: row.browser_binding_hash,
379 issuer: row.issuer,
380 pds_url: row.pds_url,
381 did: row.did,
382 auth_method: row.auth_method,
383 auth_kid: row.auth_kid,
384 redirect_uri: row.redirect_uri,
385 requested_scope: row.requested_scope,
386 request_uri: row.request_uri,
387 app_return_to: row.app_return_to,
388 expires_at: row.expires_at,
389 }))
390}
391
392#[derive(Debug, Clone, PartialEq, Eq)]
394pub struct OAuthSession {
395 pub sub: String,
396 pub issuer: String,
397 pub aud: String,
398 pub dpop_key_jwk: String,
399 pub access_token: String,
400 pub refresh_token: String,
401 pub token_type: String,
402 pub granted_scope: String,
403 pub expires_at: Option<i64>,
404}
405
406#[derive(sqlx::FromRow)]
407struct SessionRow {
408 sub: String,
409 issuer: String,
410 aud: String,
411 dpop_key_jwk: String,
412 access_token: String,
413 refresh_token: String,
414 token_type: String,
415 granted_scope: String,
416 expires_at: Option<i64>,
417}
418
419pub async fn put_session(pool: &SqlitePool, codec: &Codec, session: &OAuthSession) -> Result<()> {
425 let [dpop_key_jwk, access_token, refresh_token] = encrypt_secrets(codec, session);
426 sqlx::query(
427 r#"
428 INSERT INTO oauth_session (
429 sub, issuer, aud, dpop_key_jwk, access_token, refresh_token,
430 token_type, granted_scope, expires_at
431 ) VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9)
432 ON CONFLICT(sub) DO UPDATE SET
433 issuer = excluded.issuer,
434 aud = excluded.aud,
435 dpop_key_jwk = excluded.dpop_key_jwk,
436 access_token = excluded.access_token,
437 refresh_token = excluded.refresh_token,
438 token_type = excluded.token_type,
439 granted_scope = excluded.granted_scope,
440 expires_at = excluded.expires_at
441 "#,
442 )
443 .bind(&session.sub)
444 .bind(&session.issuer)
445 .bind(&session.aud)
446 .bind(dpop_key_jwk)
447 .bind(access_token)
448 .bind(refresh_token)
449 .bind(&session.token_type)
450 .bind(&session.granted_scope)
451 .bind(session.expires_at)
452 .execute(pool)
453 .await
454 .context("storing the OAuth session")?;
455 Ok(())
456}
457
458pub async fn get_session(
460 pool: &SqlitePool,
461 codec: &Codec,
462 sub: &str,
463) -> Result<Option<OAuthSession>> {
464 Ok(get_session_versioned(pool, codec, sub)
465 .await?
466 .map(|(session, _)| session))
467}
468
469#[derive(Debug, Clone, PartialEq, Eq)]
476pub struct SessionVersion {
477 dpop_key_jwk: String,
478 access_token: String,
479 refresh_token: String,
480}
481
482pub async fn get_session_versioned(
486 pool: &SqlitePool,
487 codec: &Codec,
488 sub: &str,
489) -> Result<Option<(OAuthSession, SessionVersion)>> {
490 let row: Option<SessionRow> = sqlx::query_as(
491 r#"
492 SELECT sub, issuer, aud, dpop_key_jwk, access_token, refresh_token,
493 token_type, granted_scope, expires_at
494 FROM oauth_session WHERE sub = ?1
495 "#,
496 )
497 .bind(sub)
498 .fetch_optional(pool)
499 .await
500 .context("reading the OAuth session")?;
501
502 let Some(row) = row else { return Ok(None) };
503 let aad = |column: &str| {
504 session_aad(
505 &row.sub,
506 column,
507 &row.issuer,
508 &row.aud,
509 &row.token_type,
510 &row.granted_scope,
511 row.expires_at,
512 )
513 };
514 let session = OAuthSession {
515 dpop_key_jwk: codec
516 .decrypt_bound(&row.dpop_key_jwk, &aad("dpop_key_jwk"))
517 .context("decrypting the session DPoP key (or its bound context was altered)")?,
518 access_token: codec
519 .decrypt_bound(&row.access_token, &aad("access_token"))
520 .context("decrypting the stored access token (or its bound context was altered)")?,
521 refresh_token: codec
522 .decrypt_bound(&row.refresh_token, &aad("refresh_token"))
523 .context("decrypting the stored refresh token (or its bound context was altered)")?,
524 sub: row.sub,
525 issuer: row.issuer,
526 aud: row.aud,
527 token_type: row.token_type,
528 granted_scope: row.granted_scope,
529 expires_at: row.expires_at,
530 };
531 let version = SessionVersion {
532 dpop_key_jwk: row.dpop_key_jwk,
533 access_token: row.access_token,
534 refresh_token: row.refresh_token,
535 };
536 Ok(Some((session, version)))
537}
538
539pub async fn update_session_if_unchanged(
549 pool: &SqlitePool,
550 codec: &Codec,
551 session: &OAuthSession,
552 version: &SessionVersion,
553) -> Result<bool> {
554 let [dpop_key_jwk, access_token, refresh_token] = encrypt_secrets(codec, session);
555 let result = sqlx::query(
557 r#"
558 UPDATE oauth_session SET
559 issuer = ?2,
560 aud = ?3,
561 dpop_key_jwk = ?4,
562 access_token = ?5,
563 refresh_token = ?6,
564 token_type = ?7,
565 granted_scope = ?8,
566 expires_at = ?9
567 WHERE sub = ?1
568 AND dpop_key_jwk = ?10 AND access_token = ?11 AND refresh_token = ?12
569 "#,
570 )
571 .bind(&session.sub)
572 .bind(&session.issuer)
573 .bind(&session.aud)
574 .bind(dpop_key_jwk)
575 .bind(access_token)
576 .bind(refresh_token)
577 .bind(&session.token_type)
578 .bind(&session.granted_scope)
579 .bind(session.expires_at)
580 .bind(&version.dpop_key_jwk)
581 .bind(&version.access_token)
582 .bind(&version.refresh_token)
583 .execute(pool)
584 .await
585 .context("updating the OAuth session (if unchanged)")?;
586 Ok(result.rows_affected() > 0)
587}
588
589fn encrypt_secrets(codec: &Codec, session: &OAuthSession) -> [String; 3] {
592 let bound = |column: &str, plaintext: &str| {
593 codec.encrypt_bound(
594 plaintext,
595 &session_aad(
596 &session.sub,
597 column,
598 &session.issuer,
599 &session.aud,
600 &session.token_type,
601 &session.granted_scope,
602 session.expires_at,
603 ),
604 )
605 };
606 [
607 bound("dpop_key_jwk", &session.dpop_key_jwk),
608 bound("access_token", &session.access_token),
609 bound("refresh_token", &session.refresh_token),
610 ]
611}
612
613pub async fn delete_session_if_unchanged(
622 pool: &SqlitePool,
623 sub: &str,
624 version: &SessionVersion,
625) -> Result<bool> {
626 let result = sqlx::query(
629 "DELETE FROM oauth_session WHERE sub = ?1 AND dpop_key_jwk = ?2 \
630 AND access_token = ?3 AND refresh_token = ?4",
631 )
632 .bind(sub)
633 .bind(&version.dpop_key_jwk)
634 .bind(&version.access_token)
635 .bind(&version.refresh_token)
636 .execute(pool)
637 .await
638 .context("deleting the OAuth session (if unchanged)")?;
639 Ok(result.rows_affected() > 0)
640}
641
642pub async fn list_session_subs(pool: &SqlitePool) -> Result<Vec<String>> {
649 sqlx::query_scalar("SELECT sub FROM oauth_session ORDER BY sub")
650 .fetch_all(pool)
651 .await
652 .context("listing the OAuth sessions")
653}
654
655pub async fn delete_session(pool: &SqlitePool, sub: &str) -> Result<bool> {
657 let result = sqlx::query("DELETE FROM oauth_session WHERE sub = ?1")
658 .bind(sub)
659 .execute(pool)
660 .await
661 .context("deleting the OAuth session")?;
662 Ok(result.rows_affected() > 0)
663}
664
665pub async fn sweep_expired_pending(pool: &SqlitePool, now: i64) -> Result<u64> {
673 let result = sqlx::query("DELETE FROM oauth_state WHERE expires_at <= ?1")
674 .bind(now)
675 .execute(pool)
676 .await
677 .context("sweeping expired pending logins")?;
678 Ok(result.rows_affected())
679}
680
681pub async fn sweep_stale_nonces(pool: &SqlitePool, cutoff: i64) -> Result<u64> {
690 let result = sqlx::query("DELETE FROM oauth_nonce WHERE updated_at <= ?1")
691 .bind(cutoff)
692 .execute(pool)
693 .await
694 .context("sweeping stale DPoP nonces")?;
695 Ok(result.rows_affected())
696}
697
698pub async fn get_nonce(pool: &SqlitePool, origin: &str) -> Result<Option<String>> {
700 sqlx::query_scalar("SELECT nonce FROM oauth_nonce WHERE origin = ?1")
701 .bind(origin)
702 .fetch_optional(pool)
703 .await
704 .context("reading the stored DPoP nonce")
705}
706
707pub async fn put_nonce(pool: &SqlitePool, origin: &str, nonce: &str, now: i64) -> Result<()> {
712 sqlx::query(
713 r#"
714 INSERT INTO oauth_nonce (origin, nonce, updated_at) VALUES (?1, ?2, ?3)
715 ON CONFLICT(origin) DO UPDATE SET
716 nonce = excluded.nonce, updated_at = excluded.updated_at
717 "#,
718 )
719 .bind(origin)
720 .bind(nonce)
721 .bind(now)
722 .execute(pool)
723 .await
724 .context("storing the DPoP nonce")?;
725 Ok(())
726}
727
728#[cfg(test)]
729mod tests {
730 use super::*;
731 use crate::store::init_url;
732
733 const KEY: &str = "aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa";
734 const DID: &str = "did:plc:ewvi7nxzyoun6zhxrhs64oiz";
735 const NOW: i64 = 1_700_000_000;
736
737 async fn db() -> (sqlx::SqlitePool, Codec) {
738 let pool = init_url("sqlite::memory:").await.unwrap();
739 init_schema(&pool).await.unwrap();
740 (pool, Codec::new(Some(KEY)).unwrap())
741 }
742
743 fn leak(sql: String) -> &'static str {
746 Box::leak(sql.into_boxed_str())
747 }
748
749 fn pending(state: &str) -> PendingAuth {
750 PendingAuth {
751 state: state.to_string(),
752 browser_binding_hash: "hash-of-cookie".into(),
753 pkce_verifier: "verifier-secret".into(),
754 dpop_key_jwk: r#"{"kty":"EC","d":"secret"}"#.into(),
755 issuer: "https://auth.example.com".into(),
756 pds_url: "https://pds.example.com".into(),
757 did: DID.into(),
758 auth_method: "private_key_jwt".into(),
759 auth_kid: Some("featherreader-oauth-1".into()),
760 redirect_uri: "https://feather-reader.com/oauth/callback".into(),
761 requested_scope: "atproto transition:generic".into(),
762 request_uri: "urn:ietf:params:oauth:request_uri:abc".into(),
763 app_return_to: Some("/reader".into()),
764 expires_at: NOW + 600,
765 }
766 }
767
768 #[tokio::test]
771 async fn a_pending_login_round_trips() -> anyhow::Result<()> {
772 let (pool, codec) = db().await;
773 let want = pending("state-1");
774 put_pending(&pool, &codec, &want).await?;
775
776 let got = take_pending(&pool, &codec, "state-1", NOW).await?.unwrap();
777 assert_eq!(got, want);
778 Ok(())
779 }
780
781 #[tokio::test]
785 async fn a_pending_login_can_only_be_taken_once() -> anyhow::Result<()> {
786 let (pool, codec) = db().await;
787 put_pending(&pool, &codec, &pending("state-1")).await?;
788
789 assert!(take_pending(&pool, &codec, "state-1", NOW).await?.is_some());
790 assert!(take_pending(&pool, &codec, "state-1", NOW).await?.is_none());
791 Ok(())
792 }
793
794 #[tokio::test]
796 async fn concurrent_takes_yield_exactly_one_winner() -> anyhow::Result<()> {
797 let (pool, codec) = db().await;
798 put_pending(&pool, &codec, &pending("race")).await?;
799
800 let (a, b, c, d) = tokio::join!(
803 take_pending(&pool, &codec, "race", NOW),
804 take_pending(&pool, &codec, "race", NOW),
805 take_pending(&pool, &codec, "race", NOW),
806 take_pending(&pool, &codec, "race", NOW),
807 );
808 let winners = [a?, b?, c?, d?].iter().filter(|r| r.is_some()).count();
809 assert_eq!(winners, 1, "more than one caller consumed the same state");
810 Ok(())
811 }
812
813 #[tokio::test]
816 async fn an_expired_pending_login_is_rejected_and_removed() -> anyhow::Result<()> {
817 let (pool, codec) = db().await;
818 put_pending(&pool, &codec, &pending("stale")).await?;
819
820 let after_expiry = NOW + 601;
821 assert!(take_pending(&pool, &codec, "stale", after_expiry)
822 .await?
823 .is_none());
824 assert!(take_pending(&pool, &codec, "stale", NOW).await?.is_none());
826 Ok(())
827 }
828
829 #[tokio::test]
830 async fn an_unknown_state_is_simply_absent() -> anyhow::Result<()> {
831 let (pool, codec) = db().await;
832 assert!(take_pending(&pool, &codec, "never-existed", NOW)
833 .await?
834 .is_none());
835 Ok(())
836 }
837
838 #[tokio::test]
842 async fn secret_columns_are_stored_encrypted() -> anyhow::Result<()> {
843 let (pool, codec) = db().await;
844 put_pending(&pool, &codec, &pending("state-1")).await?;
845
846 let (verifier, jwk): (String, String) =
847 sqlx::query_as("SELECT pkce_verifier, dpop_key_jwk FROM oauth_state WHERE state = ?")
848 .bind("state-1")
849 .fetch_one(&pool)
850 .await?;
851 for stored in [&verifier, &jwk] {
852 assert!(stored.starts_with("enc.v2.gcm."), "not bound: {stored}");
853 }
854 assert!(!verifier.contains("verifier-secret"));
855 assert!(!jwk.contains("secret"));
856 Ok(())
857 }
858
859 #[tokio::test]
863 async fn a_secret_moved_between_rows_does_not_decrypt() -> anyhow::Result<()> {
864 let (pool, codec) = db().await;
865 put_pending(&pool, &codec, &pending("victim")).await?;
866 let mut attacker = pending("attacker");
867 attacker.dpop_key_jwk = r#"{"kty":"EC","d":"attacker-key"}"#.into();
868 put_pending(&pool, &codec, &attacker).await?;
869
870 let stolen: String =
872 sqlx::query_scalar("SELECT dpop_key_jwk FROM oauth_state WHERE state = ?")
873 .bind("attacker")
874 .fetch_one(&pool)
875 .await?;
876 sqlx::query("UPDATE oauth_state SET dpop_key_jwk = ? WHERE state = ?")
877 .bind(&stolen)
878 .bind("victim")
879 .execute(&pool)
880 .await?;
881
882 assert!(
883 take_pending(&pool, &codec, "victim", NOW).await.is_err(),
884 "a grafted ciphertext decrypted in the wrong row"
885 );
886 Ok(())
887 }
888
889 #[tokio::test]
891 async fn a_secret_moved_between_columns_does_not_decrypt() -> anyhow::Result<()> {
892 let (pool, codec) = db().await;
893 put_pending(&pool, &codec, &pending("state-1")).await?;
894
895 let verifier: String =
896 sqlx::query_scalar("SELECT pkce_verifier FROM oauth_state WHERE state = ?")
897 .bind("state-1")
898 .fetch_one(&pool)
899 .await?;
900 sqlx::query("UPDATE oauth_state SET dpop_key_jwk = ? WHERE state = ?")
901 .bind(&verifier)
902 .bind("state-1")
903 .execute(&pool)
904 .await?;
905
906 assert!(take_pending(&pool, &codec, "state-1", NOW).await.is_err());
907 Ok(())
908 }
909
910 #[tokio::test]
920 async fn a_session_secret_moved_between_columns_does_not_decrypt() -> anyhow::Result<()> {
921 let (pool, codec) = db().await;
922 put_session(&pool, &codec, &session()).await?;
923 let access: String =
924 sqlx::query_scalar("SELECT access_token FROM oauth_session WHERE sub = ?")
925 .bind(DID)
926 .fetch_one(&pool)
927 .await?;
928 sqlx::query("UPDATE oauth_session SET refresh_token = ? WHERE sub = ?")
929 .bind(&access)
930 .bind(DID)
931 .execute(&pool)
932 .await?;
933 assert!(
934 get_session(&pool, &codec, DID).await.is_err(),
935 "the access token's ciphertext was accepted in the refresh_token column"
936 );
937 Ok(())
938 }
939
940 #[tokio::test]
947 async fn an_absent_expiry_and_a_zero_expiry_are_different_sessions() -> anyhow::Result<()> {
948 for (stored, flipped_to) in [(None, "0"), (Some(0), "NULL")] {
949 let (pool, codec) = db().await;
950 put_session(
951 &pool,
952 &codec,
953 &OAuthSession {
954 expires_at: stored,
955 ..session()
956 },
957 )
958 .await?;
959 sqlx::query(sqlx::AssertSqlSafe(format!(
963 "UPDATE oauth_session SET expires_at = {flipped_to} WHERE sub = ?"
964 )))
965 .bind(DID)
966 .execute(&pool)
967 .await?;
968 assert!(
969 get_session(&pool, &codec, DID).await.is_err(),
970 "expires_at {stored:?} → {flipped_to} still decrypted"
971 );
972 }
973 Ok(())
974 }
975
976 #[tokio::test]
985 async fn tampering_with_a_pending_logins_destinations_breaks_it() -> anyhow::Result<()> {
986 for column in [
987 "issuer",
988 "pds_url",
989 "did",
990 "redirect_uri",
991 "browser_binding_hash",
992 "auth_method",
993 "auth_kid",
1002 "requested_scope",
1003 "request_uri",
1004 "app_return_to",
1005 ] {
1006 let (pool, codec) = db().await;
1007 put_pending(&pool, &codec, &pending("state-1")).await?;
1008 sqlx::query(leak(format!(
1009 "UPDATE oauth_state SET {column} = ? WHERE state = ?"
1010 )))
1011 .bind("https://evil.example")
1012 .bind("state-1")
1013 .execute(&pool)
1014 .await?;
1015 assert!(
1016 take_pending(&pool, &codec, "state-1", NOW).await.is_err(),
1017 "tampering with `{column}` went undetected"
1018 );
1019 }
1020 Ok(())
1021 }
1022
1023 #[tokio::test]
1027 async fn tampering_with_a_sessions_destinations_breaks_it() -> anyhow::Result<()> {
1028 for column in [
1029 "aud",
1030 "issuer",
1031 "token_type",
1034 "granted_scope",
1035 ] {
1036 let (pool, codec) = db().await;
1037 put_session(&pool, &codec, &session()).await?;
1038 sqlx::query(leak(format!(
1039 "UPDATE oauth_session SET {column} = ? WHERE sub = ?"
1040 )))
1041 .bind("https://evil.example")
1042 .bind(DID)
1043 .execute(&pool)
1044 .await?;
1045 assert!(
1046 get_session(&pool, &codec, DID).await.is_err(),
1047 "tampering with `{column}` went undetected"
1048 );
1049 }
1050 Ok(())
1051 }
1052
1053 #[tokio::test]
1061 async fn stale_nonces_are_swept_and_fresh_ones_kept() -> anyhow::Result<()> {
1062 let (pool, _codec) = db().await;
1063 put_nonce(&pool, "https://old.example", "n1", NOW - 10_000).await?;
1064 put_nonce(&pool, "https://new.example", "n2", NOW).await?;
1065
1066 assert_eq!(sweep_stale_nonces(&pool, NOW - 5_000).await?, 1);
1067 assert_eq!(get_nonce(&pool, "https://old.example").await?, None);
1068 assert_eq!(
1069 get_nonce(&pool, "https://new.example").await?.as_deref(),
1070 Some("n2"),
1071 "a nonce still in use was swept"
1072 );
1073 Ok(())
1074 }
1075
1076 #[tokio::test]
1083 async fn swapping_an_absent_optional_column_for_an_empty_one_breaks_it() -> anyhow::Result<()> {
1084 for (column, set_to_empty) in [
1085 ("auth_kid", true),
1086 ("auth_kid", false),
1087 ("app_return_to", true),
1088 ("app_return_to", false),
1089 ] {
1090 let (pool, codec) = db().await;
1091 let mut auth = pending("state-1");
1092 if set_to_empty {
1094 if column == "auth_kid" {
1096 auth.auth_kid = None;
1097 } else {
1098 auth.app_return_to = None;
1099 }
1100 } else {
1101 if column == "auth_kid" {
1103 auth.auth_kid = Some(String::new());
1104 } else {
1105 auth.app_return_to = Some(String::new());
1106 }
1107 }
1108 put_pending(&pool, &codec, &auth).await?;
1109
1110 let sql = leak(format!(
1111 "UPDATE oauth_state SET {column} = ? WHERE state = ?"
1112 ));
1113 let query = if set_to_empty {
1114 sqlx::query(sql).bind(Some(String::new()))
1115 } else {
1116 sqlx::query(sql).bind(Option::<String>::None)
1117 };
1118 query.bind("state-1").execute(&pool).await?;
1119
1120 assert!(
1121 take_pending(&pool, &codec, "state-1", NOW).await.is_err(),
1122 "`{column}`: {} went undetected",
1123 if set_to_empty {
1124 "NULL -> ''"
1125 } else {
1126 "'' -> NULL"
1127 }
1128 );
1129 }
1130 Ok(())
1131 }
1132
1133 #[tokio::test]
1142 async fn extending_a_pending_logins_expiry_breaks_it() -> anyhow::Result<()> {
1143 let (pool, codec) = db().await;
1144 put_pending(&pool, &codec, &pending("state-1")).await?;
1145 sqlx::query("UPDATE oauth_state SET expires_at = ? WHERE state = ?")
1146 .bind(NOW + 31_536_000)
1147 .bind("state-1")
1148 .execute(&pool)
1149 .await?;
1150 assert!(
1151 take_pending(&pool, &codec, "state-1", NOW).await.is_err(),
1152 "the expiry was extended without breaking the row"
1153 );
1154 Ok(())
1155 }
1156
1157 #[tokio::test]
1161 async fn clearing_a_sessions_expiry_breaks_it() -> anyhow::Result<()> {
1162 let (pool, codec) = db().await;
1163 put_session(&pool, &codec, &session()).await?;
1164 sqlx::query("UPDATE oauth_session SET expires_at = NULL WHERE sub = ?")
1165 .bind(DID)
1166 .execute(&pool)
1167 .await?;
1168 assert!(
1169 get_session(&pool, &codec, DID).await.is_err(),
1170 "the expiry was cleared without breaking the row"
1171 );
1172 Ok(())
1173 }
1174
1175 #[test]
1180 fn the_aad_encoding_is_unambiguous_across_field_boundaries() {
1181 assert_ne!(
1182 structured_aad("t", &["ab", "c"]),
1183 structured_aad("t", &["a", "bc"])
1184 );
1185 assert_ne!(
1186 structured_aad("t", &["a:b"]),
1187 structured_aad("t", &["a", "b"])
1188 );
1189 assert_ne!(
1190 structured_aad("t", &["a", ""]),
1191 structured_aad("t", &["", "a"])
1192 );
1193 assert_ne!(structured_aad("t1", &["a"]), structured_aad("t2", &["a"]));
1194 }
1195
1196 #[tokio::test]
1198 async fn a_pending_login_round_trips_with_its_optional_fields_absent() -> anyhow::Result<()> {
1199 let (pool, codec) = db().await;
1200 let mut want = pending("state-1");
1201 want.auth_kid = None;
1202 want.app_return_to = None;
1203 put_pending(&pool, &codec, &want).await?;
1204 assert_eq!(
1205 take_pending(&pool, &codec, "state-1", NOW).await?.unwrap(),
1206 want
1207 );
1208 Ok(())
1209 }
1210
1211 #[tokio::test]
1214 async fn a_pending_login_is_expired_at_exactly_its_expiry() -> anyhow::Result<()> {
1215 let (pool, codec) = db().await;
1216 put_pending(&pool, &codec, &pending("edge")).await?;
1217 assert!(take_pending(&pool, &codec, "edge", NOW + 600)
1218 .await?
1219 .is_none());
1220
1221 let (pool, codec) = db().await;
1222 put_pending(&pool, &codec, &pending("edge")).await?;
1223 assert!(take_pending(&pool, &codec, "edge", NOW + 599)
1224 .await?
1225 .is_some());
1226 Ok(())
1227 }
1228
1229 #[tokio::test]
1234 async fn expired_pending_logins_are_swept() -> anyhow::Result<()> {
1235 let (pool, codec) = db().await;
1236 put_pending(&pool, &codec, &pending("old")).await?;
1237 let mut fresh = pending("fresh");
1238 fresh.expires_at = NOW + 3600;
1239 put_pending(&pool, &codec, &fresh).await?;
1240
1241 assert_eq!(sweep_expired_pending(&pool, NOW + 700).await?, 1);
1242 assert!(take_pending(&pool, &codec, "old", NOW).await?.is_none());
1243 assert!(take_pending(&pool, &codec, "fresh", NOW).await?.is_some());
1244 Ok(())
1245 }
1246
1247 #[tokio::test]
1250 async fn an_unbound_ciphertext_is_refused() -> anyhow::Result<()> {
1251 let (pool, codec) = db().await;
1252 put_pending(&pool, &codec, &pending("state-1")).await?;
1253
1254 sqlx::query("UPDATE oauth_state SET pkce_verifier = ? WHERE state = ?")
1255 .bind(codec.encrypt("verifier-secret"))
1256 .bind("state-1")
1257 .execute(&pool)
1258 .await?;
1259
1260 assert!(take_pending(&pool, &codec, "state-1", NOW).await.is_err());
1261 Ok(())
1262 }
1263
1264 #[tokio::test]
1268 async fn an_unbound_session_ciphertext_is_refused() -> anyhow::Result<()> {
1269 for column in ["access_token", "refresh_token", "dpop_key_jwk"] {
1270 let (pool, codec) = db().await;
1271 put_session(&pool, &codec, &session()).await?;
1272 sqlx::query(leak(format!(
1273 "UPDATE oauth_session SET {column} = ? WHERE sub = ?"
1274 )))
1275 .bind(codec.encrypt("some-value"))
1276 .bind(DID)
1277 .execute(&pool)
1278 .await?;
1279 assert!(
1280 get_session(&pool, &codec, DID).await.is_err(),
1281 "an unbound value was accepted in `{column}`"
1282 );
1283 }
1284 Ok(())
1285 }
1286
1287 fn session() -> OAuthSession {
1290 OAuthSession {
1291 sub: DID.into(),
1292 issuer: "https://auth.example.com".into(),
1293 aud: "https://pds.example.com".into(),
1294 dpop_key_jwk: r#"{"kty":"EC","d":"session-key"}"#.into(),
1295 access_token: "access-abc".into(),
1296 refresh_token: "refresh-xyz".into(),
1297 token_type: "DPoP".into(),
1298 granted_scope: "atproto transition:generic".into(),
1299 expires_at: Some(NOW + 3600),
1300 }
1301 }
1302
1303 #[tokio::test]
1304 async fn a_session_round_trips() -> anyhow::Result<()> {
1305 let (pool, codec) = db().await;
1306 put_session(&pool, &codec, &session()).await?;
1307 assert_eq!(get_session(&pool, &codec, DID).await?.unwrap(), session());
1308 Ok(())
1309 }
1310
1311 #[tokio::test]
1314 async fn re_login_replaces_the_existing_session() -> anyhow::Result<()> {
1315 let (pool, codec) = db().await;
1316 put_session(&pool, &codec, &session()).await?;
1317
1318 let mut second = session();
1319 second.access_token = "access-second".into();
1320 second.refresh_token = "refresh-second".into();
1321 put_session(&pool, &codec, &second).await?;
1322
1323 let got = get_session(&pool, &codec, DID).await?.unwrap();
1324 assert_eq!(got.access_token, "access-second");
1325 assert_eq!(got.refresh_token, "refresh-second");
1326 Ok(())
1327 }
1328
1329 #[tokio::test]
1332 async fn a_session_without_an_expiry_round_trips() -> anyhow::Result<()> {
1333 let (pool, codec) = db().await;
1334 let mut s = session();
1335 s.expires_at = None;
1336 put_session(&pool, &codec, &s).await?;
1337 assert_eq!(
1338 get_session(&pool, &codec, DID).await?.unwrap().expires_at,
1339 None
1340 );
1341 Ok(())
1342 }
1343
1344 #[tokio::test]
1345 async fn session_tokens_are_bound_to_their_subject() -> anyhow::Result<()> {
1346 let (pool, codec) = db().await;
1347 put_session(&pool, &codec, &session()).await?;
1348
1349 let other = OAuthSession {
1350 sub: "did:plc:aaaaaaaaaaaaaaaaaaaaaaaa".into(),
1351 access_token: "access-other".into(),
1352 ..session()
1353 };
1354 put_session(&pool, &codec, &other).await?;
1355
1356 let stolen: String =
1357 sqlx::query_scalar("SELECT access_token FROM oauth_session WHERE sub = ?")
1358 .bind(&other.sub)
1359 .fetch_one(&pool)
1360 .await?;
1361 sqlx::query("UPDATE oauth_session SET access_token = ? WHERE sub = ?")
1362 .bind(&stolen)
1363 .bind(DID)
1364 .execute(&pool)
1365 .await?;
1366
1367 assert!(get_session(&pool, &codec, DID).await.is_err());
1368 Ok(())
1369 }
1370
1371 #[tokio::test]
1375 async fn a_rewritten_session_is_not_deleted_by_a_stale_version() -> anyhow::Result<()> {
1376 let (pool, codec) = db().await;
1377 put_session(&pool, &codec, &session()).await?;
1378 let (_, stale) = get_session_versioned(&pool, &codec, DID).await?.unwrap();
1379
1380 let rotated = OAuthSession {
1381 refresh_token: "refresh-rotated".into(),
1382 ..session()
1383 };
1384 put_session(&pool, &codec, &rotated).await?;
1385 assert!(
1386 !delete_session_if_unchanged(&pool, DID, &stale).await?,
1387 "reported deleting a row it should have left"
1388 );
1389 assert_eq!(
1390 get_session(&pool, &codec, DID)
1391 .await?
1392 .unwrap()
1393 .refresh_token,
1394 "refresh-rotated",
1395 "the ROTATED token was deleted on the strength of a stale read"
1396 );
1397
1398 let (_, before) = get_session_versioned(&pool, &codec, DID).await?.unwrap();
1400 put_session(&pool, &codec, &rotated).await?;
1401 assert!(!delete_session_if_unchanged(&pool, DID, &before).await?);
1402
1403 let (_, current) = get_session_versioned(&pool, &codec, DID).await?.unwrap();
1405 assert!(delete_session_if_unchanged(&pool, DID, ¤t).await?);
1406 assert!(get_session(&pool, &codec, DID).await?.is_none());
1407 assert!(!delete_session_if_unchanged(&pool, DID, ¤t).await?);
1408 Ok(())
1409 }
1410
1411 #[tokio::test]
1412 async fn a_deleted_session_is_gone() -> anyhow::Result<()> {
1413 let (pool, codec) = db().await;
1414 put_session(&pool, &codec, &session()).await?;
1415 assert!(delete_session(&pool, DID).await?);
1416 assert!(get_session(&pool, &codec, DID).await?.is_none());
1417 assert!(!delete_session(&pool, DID).await?);
1418 Ok(())
1419 }
1420
1421 #[tokio::test]
1426 async fn every_session_subject_is_listed_including_an_unreadable_one() -> anyhow::Result<()> {
1427 let (pool, codec) = db().await;
1428 assert!(list_session_subs(&pool).await?.is_empty());
1429
1430 put_session(&pool, &codec, &session()).await?;
1431 let other = OAuthSession {
1432 sub: "did:plc:bbbbbbbbbbbbbbbbbbbbbbbb".into(),
1433 ..session()
1434 };
1435 put_session(&pool, &codec, &other).await?;
1436 sqlx::query(
1438 "INSERT INTO oauth_session (sub, issuer, aud, dpop_key_jwk, access_token, \
1439 refresh_token, token_type, granted_scope, expires_at) \
1440 VALUES (?, 'https://auth.example.com', 'https://pds.example.com', \
1441 'garbage', 'garbage', 'garbage', 'DPoP', 'atproto', NULL)",
1442 )
1443 .bind("did:plc:cccccccccccccccccccccccc")
1444 .execute(&pool)
1445 .await?;
1446 assert!(
1447 get_session(&pool, &codec, "did:plc:cccccccccccccccccccccccc")
1448 .await
1449 .is_err(),
1450 "precondition: the raw row must be unreadable"
1451 );
1452
1453 assert_eq!(
1454 list_session_subs(&pool).await?,
1455 vec![
1456 "did:plc:bbbbbbbbbbbbbbbbbbbbbbbb".to_string(),
1457 "did:plc:cccccccccccccccccccccccc".to_string(),
1458 DID.to_string(),
1459 ]
1460 );
1461 Ok(())
1462 }
1463
1464 #[tokio::test]
1469 async fn nonces_are_stored_and_replaced_per_origin() -> anyhow::Result<()> {
1470 let (pool, _) = db().await;
1471 assert_eq!(get_nonce(&pool, "https://a.example").await?, None);
1472
1473 put_nonce(&pool, "https://a.example", "n1", NOW).await?;
1474 put_nonce(&pool, "https://b.example", "n2", NOW).await?;
1475 assert_eq!(
1476 get_nonce(&pool, "https://a.example").await?.as_deref(),
1477 Some("n1")
1478 );
1479 assert_eq!(
1480 get_nonce(&pool, "https://b.example").await?.as_deref(),
1481 Some("n2")
1482 );
1483
1484 put_nonce(&pool, "https://a.example", "n3", NOW).await?;
1486 assert_eq!(
1487 get_nonce(&pool, "https://a.example").await?.as_deref(),
1488 Some("n3")
1489 );
1490 Ok(())
1491 }
1492}