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 sqlx::query(
426 r#"
427 INSERT INTO oauth_session (
428 sub, issuer, aud, dpop_key_jwk, access_token, refresh_token,
429 token_type, granted_scope, expires_at
430 ) VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9)
431 ON CONFLICT(sub) DO UPDATE SET
432 issuer = excluded.issuer,
433 aud = excluded.aud,
434 dpop_key_jwk = excluded.dpop_key_jwk,
435 access_token = excluded.access_token,
436 refresh_token = excluded.refresh_token,
437 token_type = excluded.token_type,
438 granted_scope = excluded.granted_scope,
439 expires_at = excluded.expires_at
440 "#,
441 )
442 .bind(&session.sub)
443 .bind(&session.issuer)
444 .bind(&session.aud)
445 .bind(codec.encrypt_bound(
446 &session.dpop_key_jwk,
447 &session_aad(
448 &session.sub,
449 "dpop_key_jwk",
450 &session.issuer,
451 &session.aud,
452 &session.token_type,
453 &session.granted_scope,
454 session.expires_at,
455 ),
456 ))
457 .bind(codec.encrypt_bound(
458 &session.access_token,
459 &session_aad(
460 &session.sub,
461 "access_token",
462 &session.issuer,
463 &session.aud,
464 &session.token_type,
465 &session.granted_scope,
466 session.expires_at,
467 ),
468 ))
469 .bind(codec.encrypt_bound(
470 &session.refresh_token,
471 &session_aad(
472 &session.sub,
473 "refresh_token",
474 &session.issuer,
475 &session.aud,
476 &session.token_type,
477 &session.granted_scope,
478 session.expires_at,
479 ),
480 ))
481 .bind(&session.token_type)
482 .bind(&session.granted_scope)
483 .bind(session.expires_at)
484 .execute(pool)
485 .await
486 .context("storing the OAuth session")?;
487 Ok(())
488}
489
490pub async fn get_session(
492 pool: &SqlitePool,
493 codec: &Codec,
494 sub: &str,
495) -> Result<Option<OAuthSession>> {
496 let row: Option<SessionRow> = sqlx::query_as(
497 r#"
498 SELECT sub, issuer, aud, dpop_key_jwk, access_token, refresh_token,
499 token_type, granted_scope, expires_at
500 FROM oauth_session WHERE sub = ?1
501 "#,
502 )
503 .bind(sub)
504 .fetch_optional(pool)
505 .await
506 .context("reading the OAuth session")?;
507
508 let Some(row) = row else { return Ok(None) };
509 let aad = |column: &str| {
510 session_aad(
511 &row.sub,
512 column,
513 &row.issuer,
514 &row.aud,
515 &row.token_type,
516 &row.granted_scope,
517 row.expires_at,
518 )
519 };
520 Ok(Some(OAuthSession {
521 dpop_key_jwk: codec
522 .decrypt_bound(&row.dpop_key_jwk, &aad("dpop_key_jwk"))
523 .context("decrypting the session DPoP key (or its bound context was altered)")?,
524 access_token: codec
525 .decrypt_bound(&row.access_token, &aad("access_token"))
526 .context("decrypting the stored access token (or its bound context was altered)")?,
527 refresh_token: codec
528 .decrypt_bound(&row.refresh_token, &aad("refresh_token"))
529 .context("decrypting the stored refresh token (or its bound context was altered)")?,
530 sub: row.sub,
531 issuer: row.issuer,
532 aud: row.aud,
533 token_type: row.token_type,
534 granted_scope: row.granted_scope,
535 expires_at: row.expires_at,
536 }))
537}
538
539pub async fn delete_session(pool: &SqlitePool, sub: &str) -> Result<bool> {
541 let result = sqlx::query("DELETE FROM oauth_session WHERE sub = ?1")
542 .bind(sub)
543 .execute(pool)
544 .await
545 .context("deleting the OAuth session")?;
546 Ok(result.rows_affected() > 0)
547}
548
549pub async fn sweep_expired_pending(pool: &SqlitePool, now: i64) -> Result<u64> {
557 let result = sqlx::query("DELETE FROM oauth_state WHERE expires_at <= ?1")
558 .bind(now)
559 .execute(pool)
560 .await
561 .context("sweeping expired pending logins")?;
562 Ok(result.rows_affected())
563}
564
565pub async fn sweep_stale_nonces(pool: &SqlitePool, cutoff: i64) -> Result<u64> {
574 let result = sqlx::query("DELETE FROM oauth_nonce WHERE updated_at <= ?1")
575 .bind(cutoff)
576 .execute(pool)
577 .await
578 .context("sweeping stale DPoP nonces")?;
579 Ok(result.rows_affected())
580}
581
582pub async fn get_nonce(pool: &SqlitePool, origin: &str) -> Result<Option<String>> {
584 sqlx::query_scalar("SELECT nonce FROM oauth_nonce WHERE origin = ?1")
585 .bind(origin)
586 .fetch_optional(pool)
587 .await
588 .context("reading the stored DPoP nonce")
589}
590
591pub async fn put_nonce(pool: &SqlitePool, origin: &str, nonce: &str, now: i64) -> Result<()> {
596 sqlx::query(
597 r#"
598 INSERT INTO oauth_nonce (origin, nonce, updated_at) VALUES (?1, ?2, ?3)
599 ON CONFLICT(origin) DO UPDATE SET
600 nonce = excluded.nonce, updated_at = excluded.updated_at
601 "#,
602 )
603 .bind(origin)
604 .bind(nonce)
605 .bind(now)
606 .execute(pool)
607 .await
608 .context("storing the DPoP nonce")?;
609 Ok(())
610}
611
612#[cfg(test)]
613mod tests {
614 use super::*;
615 use crate::store::init_url;
616
617 const KEY: &str = "aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa";
618 const DID: &str = "did:plc:ewvi7nxzyoun6zhxrhs64oiz";
619 const NOW: i64 = 1_700_000_000;
620
621 async fn db() -> (sqlx::SqlitePool, Codec) {
622 let pool = init_url("sqlite::memory:").await.unwrap();
623 init_schema(&pool).await.unwrap();
624 (pool, Codec::new(Some(KEY)).unwrap())
625 }
626
627 fn leak(sql: String) -> &'static str {
630 Box::leak(sql.into_boxed_str())
631 }
632
633 fn pending(state: &str) -> PendingAuth {
634 PendingAuth {
635 state: state.to_string(),
636 browser_binding_hash: "hash-of-cookie".into(),
637 pkce_verifier: "verifier-secret".into(),
638 dpop_key_jwk: r#"{"kty":"EC","d":"secret"}"#.into(),
639 issuer: "https://auth.example.com".into(),
640 pds_url: "https://pds.example.com".into(),
641 did: DID.into(),
642 auth_method: "private_key_jwt".into(),
643 auth_kid: Some("featherreader-oauth-1".into()),
644 redirect_uri: "https://feather-reader.com/oauth/callback".into(),
645 requested_scope: "atproto transition:generic".into(),
646 request_uri: "urn:ietf:params:oauth:request_uri:abc".into(),
647 app_return_to: Some("/reader".into()),
648 expires_at: NOW + 600,
649 }
650 }
651
652 #[tokio::test]
655 async fn a_pending_login_round_trips() -> anyhow::Result<()> {
656 let (pool, codec) = db().await;
657 let want = pending("state-1");
658 put_pending(&pool, &codec, &want).await?;
659
660 let got = take_pending(&pool, &codec, "state-1", NOW).await?.unwrap();
661 assert_eq!(got, want);
662 Ok(())
663 }
664
665 #[tokio::test]
669 async fn a_pending_login_can_only_be_taken_once() -> anyhow::Result<()> {
670 let (pool, codec) = db().await;
671 put_pending(&pool, &codec, &pending("state-1")).await?;
672
673 assert!(take_pending(&pool, &codec, "state-1", NOW).await?.is_some());
674 assert!(take_pending(&pool, &codec, "state-1", NOW).await?.is_none());
675 Ok(())
676 }
677
678 #[tokio::test]
680 async fn concurrent_takes_yield_exactly_one_winner() -> anyhow::Result<()> {
681 let (pool, codec) = db().await;
682 put_pending(&pool, &codec, &pending("race")).await?;
683
684 let (a, b, c, d) = tokio::join!(
687 take_pending(&pool, &codec, "race", NOW),
688 take_pending(&pool, &codec, "race", NOW),
689 take_pending(&pool, &codec, "race", NOW),
690 take_pending(&pool, &codec, "race", NOW),
691 );
692 let winners = [a?, b?, c?, d?].iter().filter(|r| r.is_some()).count();
693 assert_eq!(winners, 1, "more than one caller consumed the same state");
694 Ok(())
695 }
696
697 #[tokio::test]
700 async fn an_expired_pending_login_is_rejected_and_removed() -> anyhow::Result<()> {
701 let (pool, codec) = db().await;
702 put_pending(&pool, &codec, &pending("stale")).await?;
703
704 let after_expiry = NOW + 601;
705 assert!(take_pending(&pool, &codec, "stale", after_expiry)
706 .await?
707 .is_none());
708 assert!(take_pending(&pool, &codec, "stale", NOW).await?.is_none());
710 Ok(())
711 }
712
713 #[tokio::test]
714 async fn an_unknown_state_is_simply_absent() -> anyhow::Result<()> {
715 let (pool, codec) = db().await;
716 assert!(take_pending(&pool, &codec, "never-existed", NOW)
717 .await?
718 .is_none());
719 Ok(())
720 }
721
722 #[tokio::test]
726 async fn secret_columns_are_stored_encrypted() -> anyhow::Result<()> {
727 let (pool, codec) = db().await;
728 put_pending(&pool, &codec, &pending("state-1")).await?;
729
730 let (verifier, jwk): (String, String) =
731 sqlx::query_as("SELECT pkce_verifier, dpop_key_jwk FROM oauth_state WHERE state = ?")
732 .bind("state-1")
733 .fetch_one(&pool)
734 .await?;
735 for stored in [&verifier, &jwk] {
736 assert!(stored.starts_with("enc.v2.gcm."), "not bound: {stored}");
737 }
738 assert!(!verifier.contains("verifier-secret"));
739 assert!(!jwk.contains("secret"));
740 Ok(())
741 }
742
743 #[tokio::test]
747 async fn a_secret_moved_between_rows_does_not_decrypt() -> anyhow::Result<()> {
748 let (pool, codec) = db().await;
749 put_pending(&pool, &codec, &pending("victim")).await?;
750 let mut attacker = pending("attacker");
751 attacker.dpop_key_jwk = r#"{"kty":"EC","d":"attacker-key"}"#.into();
752 put_pending(&pool, &codec, &attacker).await?;
753
754 let stolen: String =
756 sqlx::query_scalar("SELECT dpop_key_jwk FROM oauth_state WHERE state = ?")
757 .bind("attacker")
758 .fetch_one(&pool)
759 .await?;
760 sqlx::query("UPDATE oauth_state SET dpop_key_jwk = ? WHERE state = ?")
761 .bind(&stolen)
762 .bind("victim")
763 .execute(&pool)
764 .await?;
765
766 assert!(
767 take_pending(&pool, &codec, "victim", NOW).await.is_err(),
768 "a grafted ciphertext decrypted in the wrong row"
769 );
770 Ok(())
771 }
772
773 #[tokio::test]
775 async fn a_secret_moved_between_columns_does_not_decrypt() -> anyhow::Result<()> {
776 let (pool, codec) = db().await;
777 put_pending(&pool, &codec, &pending("state-1")).await?;
778
779 let verifier: String =
780 sqlx::query_scalar("SELECT pkce_verifier FROM oauth_state WHERE state = ?")
781 .bind("state-1")
782 .fetch_one(&pool)
783 .await?;
784 sqlx::query("UPDATE oauth_state SET dpop_key_jwk = ? WHERE state = ?")
785 .bind(&verifier)
786 .bind("state-1")
787 .execute(&pool)
788 .await?;
789
790 assert!(take_pending(&pool, &codec, "state-1", NOW).await.is_err());
791 Ok(())
792 }
793
794 #[tokio::test]
804 async fn a_session_secret_moved_between_columns_does_not_decrypt() -> anyhow::Result<()> {
805 let (pool, codec) = db().await;
806 put_session(&pool, &codec, &session()).await?;
807 let access: String =
808 sqlx::query_scalar("SELECT access_token FROM oauth_session WHERE sub = ?")
809 .bind(DID)
810 .fetch_one(&pool)
811 .await?;
812 sqlx::query("UPDATE oauth_session SET refresh_token = ? WHERE sub = ?")
813 .bind(&access)
814 .bind(DID)
815 .execute(&pool)
816 .await?;
817 assert!(
818 get_session(&pool, &codec, DID).await.is_err(),
819 "the access token's ciphertext was accepted in the refresh_token column"
820 );
821 Ok(())
822 }
823
824 #[tokio::test]
831 async fn an_absent_expiry_and_a_zero_expiry_are_different_sessions() -> anyhow::Result<()> {
832 for (stored, flipped_to) in [(None, "0"), (Some(0), "NULL")] {
833 let (pool, codec) = db().await;
834 put_session(
835 &pool,
836 &codec,
837 &OAuthSession {
838 expires_at: stored,
839 ..session()
840 },
841 )
842 .await?;
843 sqlx::query(sqlx::AssertSqlSafe(format!(
847 "UPDATE oauth_session SET expires_at = {flipped_to} WHERE sub = ?"
848 )))
849 .bind(DID)
850 .execute(&pool)
851 .await?;
852 assert!(
853 get_session(&pool, &codec, DID).await.is_err(),
854 "expires_at {stored:?} → {flipped_to} still decrypted"
855 );
856 }
857 Ok(())
858 }
859
860 #[tokio::test]
869 async fn tampering_with_a_pending_logins_destinations_breaks_it() -> anyhow::Result<()> {
870 for column in [
871 "issuer",
872 "pds_url",
873 "did",
874 "redirect_uri",
875 "browser_binding_hash",
876 "auth_method",
877 "auth_kid",
886 "requested_scope",
887 "request_uri",
888 "app_return_to",
889 ] {
890 let (pool, codec) = db().await;
891 put_pending(&pool, &codec, &pending("state-1")).await?;
892 sqlx::query(leak(format!(
893 "UPDATE oauth_state SET {column} = ? WHERE state = ?"
894 )))
895 .bind("https://evil.example")
896 .bind("state-1")
897 .execute(&pool)
898 .await?;
899 assert!(
900 take_pending(&pool, &codec, "state-1", NOW).await.is_err(),
901 "tampering with `{column}` went undetected"
902 );
903 }
904 Ok(())
905 }
906
907 #[tokio::test]
911 async fn tampering_with_a_sessions_destinations_breaks_it() -> anyhow::Result<()> {
912 for column in [
913 "aud",
914 "issuer",
915 "token_type",
918 "granted_scope",
919 ] {
920 let (pool, codec) = db().await;
921 put_session(&pool, &codec, &session()).await?;
922 sqlx::query(leak(format!(
923 "UPDATE oauth_session SET {column} = ? WHERE sub = ?"
924 )))
925 .bind("https://evil.example")
926 .bind(DID)
927 .execute(&pool)
928 .await?;
929 assert!(
930 get_session(&pool, &codec, DID).await.is_err(),
931 "tampering with `{column}` went undetected"
932 );
933 }
934 Ok(())
935 }
936
937 #[tokio::test]
945 async fn stale_nonces_are_swept_and_fresh_ones_kept() -> anyhow::Result<()> {
946 let (pool, _codec) = db().await;
947 put_nonce(&pool, "https://old.example", "n1", NOW - 10_000).await?;
948 put_nonce(&pool, "https://new.example", "n2", NOW).await?;
949
950 assert_eq!(sweep_stale_nonces(&pool, NOW - 5_000).await?, 1);
951 assert_eq!(get_nonce(&pool, "https://old.example").await?, None);
952 assert_eq!(
953 get_nonce(&pool, "https://new.example").await?.as_deref(),
954 Some("n2"),
955 "a nonce still in use was swept"
956 );
957 Ok(())
958 }
959
960 #[tokio::test]
967 async fn swapping_an_absent_optional_column_for_an_empty_one_breaks_it() -> anyhow::Result<()> {
968 for (column, set_to_empty) in [
969 ("auth_kid", true),
970 ("auth_kid", false),
971 ("app_return_to", true),
972 ("app_return_to", false),
973 ] {
974 let (pool, codec) = db().await;
975 let mut auth = pending("state-1");
976 if set_to_empty {
978 if column == "auth_kid" {
980 auth.auth_kid = None;
981 } else {
982 auth.app_return_to = None;
983 }
984 } else {
985 if column == "auth_kid" {
987 auth.auth_kid = Some(String::new());
988 } else {
989 auth.app_return_to = Some(String::new());
990 }
991 }
992 put_pending(&pool, &codec, &auth).await?;
993
994 let sql = leak(format!(
995 "UPDATE oauth_state SET {column} = ? WHERE state = ?"
996 ));
997 let query = if set_to_empty {
998 sqlx::query(sql).bind(Some(String::new()))
999 } else {
1000 sqlx::query(sql).bind(Option::<String>::None)
1001 };
1002 query.bind("state-1").execute(&pool).await?;
1003
1004 assert!(
1005 take_pending(&pool, &codec, "state-1", NOW).await.is_err(),
1006 "`{column}`: {} went undetected",
1007 if set_to_empty {
1008 "NULL -> ''"
1009 } else {
1010 "'' -> NULL"
1011 }
1012 );
1013 }
1014 Ok(())
1015 }
1016
1017 #[tokio::test]
1026 async fn extending_a_pending_logins_expiry_breaks_it() -> anyhow::Result<()> {
1027 let (pool, codec) = db().await;
1028 put_pending(&pool, &codec, &pending("state-1")).await?;
1029 sqlx::query("UPDATE oauth_state SET expires_at = ? WHERE state = ?")
1030 .bind(NOW + 31_536_000)
1031 .bind("state-1")
1032 .execute(&pool)
1033 .await?;
1034 assert!(
1035 take_pending(&pool, &codec, "state-1", NOW).await.is_err(),
1036 "the expiry was extended without breaking the row"
1037 );
1038 Ok(())
1039 }
1040
1041 #[tokio::test]
1045 async fn clearing_a_sessions_expiry_breaks_it() -> anyhow::Result<()> {
1046 let (pool, codec) = db().await;
1047 put_session(&pool, &codec, &session()).await?;
1048 sqlx::query("UPDATE oauth_session SET expires_at = NULL WHERE sub = ?")
1049 .bind(DID)
1050 .execute(&pool)
1051 .await?;
1052 assert!(
1053 get_session(&pool, &codec, DID).await.is_err(),
1054 "the expiry was cleared without breaking the row"
1055 );
1056 Ok(())
1057 }
1058
1059 #[test]
1064 fn the_aad_encoding_is_unambiguous_across_field_boundaries() {
1065 assert_ne!(
1066 structured_aad("t", &["ab", "c"]),
1067 structured_aad("t", &["a", "bc"])
1068 );
1069 assert_ne!(
1070 structured_aad("t", &["a:b"]),
1071 structured_aad("t", &["a", "b"])
1072 );
1073 assert_ne!(
1074 structured_aad("t", &["a", ""]),
1075 structured_aad("t", &["", "a"])
1076 );
1077 assert_ne!(structured_aad("t1", &["a"]), structured_aad("t2", &["a"]));
1078 }
1079
1080 #[tokio::test]
1082 async fn a_pending_login_round_trips_with_its_optional_fields_absent() -> anyhow::Result<()> {
1083 let (pool, codec) = db().await;
1084 let mut want = pending("state-1");
1085 want.auth_kid = None;
1086 want.app_return_to = None;
1087 put_pending(&pool, &codec, &want).await?;
1088 assert_eq!(
1089 take_pending(&pool, &codec, "state-1", NOW).await?.unwrap(),
1090 want
1091 );
1092 Ok(())
1093 }
1094
1095 #[tokio::test]
1098 async fn a_pending_login_is_expired_at_exactly_its_expiry() -> anyhow::Result<()> {
1099 let (pool, codec) = db().await;
1100 put_pending(&pool, &codec, &pending("edge")).await?;
1101 assert!(take_pending(&pool, &codec, "edge", NOW + 600)
1102 .await?
1103 .is_none());
1104
1105 let (pool, codec) = db().await;
1106 put_pending(&pool, &codec, &pending("edge")).await?;
1107 assert!(take_pending(&pool, &codec, "edge", NOW + 599)
1108 .await?
1109 .is_some());
1110 Ok(())
1111 }
1112
1113 #[tokio::test]
1118 async fn expired_pending_logins_are_swept() -> anyhow::Result<()> {
1119 let (pool, codec) = db().await;
1120 put_pending(&pool, &codec, &pending("old")).await?;
1121 let mut fresh = pending("fresh");
1122 fresh.expires_at = NOW + 3600;
1123 put_pending(&pool, &codec, &fresh).await?;
1124
1125 assert_eq!(sweep_expired_pending(&pool, NOW + 700).await?, 1);
1126 assert!(take_pending(&pool, &codec, "old", NOW).await?.is_none());
1127 assert!(take_pending(&pool, &codec, "fresh", NOW).await?.is_some());
1128 Ok(())
1129 }
1130
1131 #[tokio::test]
1134 async fn an_unbound_ciphertext_is_refused() -> anyhow::Result<()> {
1135 let (pool, codec) = db().await;
1136 put_pending(&pool, &codec, &pending("state-1")).await?;
1137
1138 sqlx::query("UPDATE oauth_state SET pkce_verifier = ? WHERE state = ?")
1139 .bind(codec.encrypt("verifier-secret"))
1140 .bind("state-1")
1141 .execute(&pool)
1142 .await?;
1143
1144 assert!(take_pending(&pool, &codec, "state-1", NOW).await.is_err());
1145 Ok(())
1146 }
1147
1148 #[tokio::test]
1152 async fn an_unbound_session_ciphertext_is_refused() -> anyhow::Result<()> {
1153 for column in ["access_token", "refresh_token", "dpop_key_jwk"] {
1154 let (pool, codec) = db().await;
1155 put_session(&pool, &codec, &session()).await?;
1156 sqlx::query(leak(format!(
1157 "UPDATE oauth_session SET {column} = ? WHERE sub = ?"
1158 )))
1159 .bind(codec.encrypt("some-value"))
1160 .bind(DID)
1161 .execute(&pool)
1162 .await?;
1163 assert!(
1164 get_session(&pool, &codec, DID).await.is_err(),
1165 "an unbound value was accepted in `{column}`"
1166 );
1167 }
1168 Ok(())
1169 }
1170
1171 fn session() -> OAuthSession {
1174 OAuthSession {
1175 sub: DID.into(),
1176 issuer: "https://auth.example.com".into(),
1177 aud: "https://pds.example.com".into(),
1178 dpop_key_jwk: r#"{"kty":"EC","d":"session-key"}"#.into(),
1179 access_token: "access-abc".into(),
1180 refresh_token: "refresh-xyz".into(),
1181 token_type: "DPoP".into(),
1182 granted_scope: "atproto transition:generic".into(),
1183 expires_at: Some(NOW + 3600),
1184 }
1185 }
1186
1187 #[tokio::test]
1188 async fn a_session_round_trips() -> anyhow::Result<()> {
1189 let (pool, codec) = db().await;
1190 put_session(&pool, &codec, &session()).await?;
1191 assert_eq!(get_session(&pool, &codec, DID).await?.unwrap(), session());
1192 Ok(())
1193 }
1194
1195 #[tokio::test]
1198 async fn re_login_replaces_the_existing_session() -> anyhow::Result<()> {
1199 let (pool, codec) = db().await;
1200 put_session(&pool, &codec, &session()).await?;
1201
1202 let mut second = session();
1203 second.access_token = "access-second".into();
1204 second.refresh_token = "refresh-second".into();
1205 put_session(&pool, &codec, &second).await?;
1206
1207 let got = get_session(&pool, &codec, DID).await?.unwrap();
1208 assert_eq!(got.access_token, "access-second");
1209 assert_eq!(got.refresh_token, "refresh-second");
1210 Ok(())
1211 }
1212
1213 #[tokio::test]
1216 async fn a_session_without_an_expiry_round_trips() -> anyhow::Result<()> {
1217 let (pool, codec) = db().await;
1218 let mut s = session();
1219 s.expires_at = None;
1220 put_session(&pool, &codec, &s).await?;
1221 assert_eq!(
1222 get_session(&pool, &codec, DID).await?.unwrap().expires_at,
1223 None
1224 );
1225 Ok(())
1226 }
1227
1228 #[tokio::test]
1229 async fn session_tokens_are_bound_to_their_subject() -> anyhow::Result<()> {
1230 let (pool, codec) = db().await;
1231 put_session(&pool, &codec, &session()).await?;
1232
1233 let other = OAuthSession {
1234 sub: "did:plc:aaaaaaaaaaaaaaaaaaaaaaaa".into(),
1235 access_token: "access-other".into(),
1236 ..session()
1237 };
1238 put_session(&pool, &codec, &other).await?;
1239
1240 let stolen: String =
1241 sqlx::query_scalar("SELECT access_token FROM oauth_session WHERE sub = ?")
1242 .bind(&other.sub)
1243 .fetch_one(&pool)
1244 .await?;
1245 sqlx::query("UPDATE oauth_session SET access_token = ? WHERE sub = ?")
1246 .bind(&stolen)
1247 .bind(DID)
1248 .execute(&pool)
1249 .await?;
1250
1251 assert!(get_session(&pool, &codec, DID).await.is_err());
1252 Ok(())
1253 }
1254
1255 #[tokio::test]
1256 async fn a_deleted_session_is_gone() -> anyhow::Result<()> {
1257 let (pool, codec) = db().await;
1258 put_session(&pool, &codec, &session()).await?;
1259 assert!(delete_session(&pool, DID).await?);
1260 assert!(get_session(&pool, &codec, DID).await?.is_none());
1261 assert!(!delete_session(&pool, DID).await?);
1262 Ok(())
1263 }
1264
1265 #[tokio::test]
1270 async fn nonces_are_stored_and_replaced_per_origin() -> anyhow::Result<()> {
1271 let (pool, _) = db().await;
1272 assert_eq!(get_nonce(&pool, "https://a.example").await?, None);
1273
1274 put_nonce(&pool, "https://a.example", "n1", NOW).await?;
1275 put_nonce(&pool, "https://b.example", "n2", NOW).await?;
1276 assert_eq!(
1277 get_nonce(&pool, "https://a.example").await?.as_deref(),
1278 Some("n1")
1279 );
1280 assert_eq!(
1281 get_nonce(&pool, "https://b.example").await?.as_deref(),
1282 Some("n2")
1283 );
1284
1285 put_nonce(&pool, "https://a.example", "n3", NOW).await?;
1287 assert_eq!(
1288 get_nonce(&pool, "https://a.example").await?.as_deref(),
1289 Some("n3")
1290 );
1291 Ok(())
1292 }
1293}