1use anyhow::{bail, Result};
19use std::collections::HashMap;
20use std::sync::{Arc, Mutex};
21use tokio::sync::{Mutex as AsyncMutex, OwnedMutexGuard};
22
23use super::store::OAuthSession;
24use super::token::TokenResponse;
25
26pub fn apply_refresh(
32 session: &OAuthSession,
33 response: &TokenResponse,
34 now: i64,
35) -> Result<OAuthSession> {
36 if response.sub != session.sub {
38 bail!(
39 "refresh returned subject {:?}, expected {:?}; refusing to rebind the session",
40 response.sub,
41 session.sub
42 );
43 }
44
45 Ok(OAuthSession {
46 access_token: response.access_token.clone(),
47 refresh_token: response
51 .refresh_token
52 .clone()
53 .unwrap_or_else(|| session.refresh_token.clone()),
54 token_type: response.token_type.clone(),
55 granted_scope: response.granted_scope.clone(),
58 expires_at: response.expires_in.map(|seconds| now + seconds),
59 sub: session.sub.clone(),
61 issuer: session.issuer.clone(),
62 aud: session.aud.clone(),
63 dpop_key_jwk: session.dpop_key_jwk.clone(),
64 })
65}
66
67#[derive(Default, Clone)]
79pub struct RefreshLocks {
80 locks: Arc<Mutex<HashMap<String, Arc<AsyncMutex<()>>>>>,
81}
82
83impl RefreshLocks {
84 pub async fn lock(&self, sub: &str) -> OwnedMutexGuard<()> {
86 let entry = {
87 let mut locks = self.locks.lock().unwrap_or_else(|p| p.into_inner());
90 Arc::clone(locks.entry(sub.to_string()).or_default())
91 };
92 entry.lock_owned().await
93 }
94}
95
96pub fn same_issuer(discovered: &str, expected: &str) -> Result<()> {
108 if discovered != expected {
109 anyhow::bail!(
110 "the PDS now names a different authorization server ({discovered:?}) than this \
111 grant was issued by ({expected:?}); refusing to send credentials to it"
112 );
113 }
114 Ok(())
115}
116
117pub struct RefreshContext<'a> {
119 pub token_endpoint: &'a str,
120 pub client_id: &'a str,
121 pub auth_method: super::client_auth::AuthMethod,
122 pub client_key: Option<&'a super::keys::SigningKey>,
125 pub revocation_endpoint: Option<&'a str>,
129}
130
131pub async fn valid_session(
138 pool: &sqlx::SqlitePool,
139 codec: &super::crypto::Codec,
140 http: &reqwest::Client,
141 locks: &RefreshLocks,
142 sub: &str,
143 ctx: &RefreshContext<'_>,
144 now: i64,
145) -> Result<OAuthSession> {
146 let session = super::store::get_session(pool, codec, sub)
147 .await?
148 .ok_or_else(|| anyhow::anyhow!("no session for {sub}"))?;
149 if !super::token::is_stale(session.expires_at, now) {
150 return Ok(session);
151 }
152
153 let _guard = locks.lock(sub).await;
154
155 let (session, version) = super::store::get_session_versioned(pool, codec, sub)
158 .await?
159 .ok_or_else(|| anyhow::anyhow!("session for {sub} disappeared while waiting to refresh"))?;
160 if !super::token::is_stale(session.expires_at, now) {
161 return Ok(session);
162 }
163
164 refresh_locked(pool, codec, http, &session, &version, ctx, now).await
170}
171
172async fn lost_the_row(
187 pool: &sqlx::SqlitePool,
188 codec: &super::crypto::Codec,
189 http: &reqwest::Client,
190 ctx: &RefreshContext<'_>,
191 obtained: &OAuthSession,
192 now: i64,
193) -> Result<OAuthSession> {
194 let current = super::store::get_session(pool, codec, &obtained.sub).await?;
195 let outcome = super::revoke::revoke_orphaned(
196 pool,
197 http,
198 &super::revoke::RevokeContext {
199 revocation_endpoint: ctx.revocation_endpoint,
200 client_id: ctx.client_id,
201 auth_method: ctx.auth_method,
202 client_key: ctx.client_key,
203 deadline: super::revoke::ORPHAN_REVOKE_DEADLINE,
204 },
205 obtained,
206 now,
207 )
208 .await;
209 if let super::revoke::Revocation::Failed(reason) = &outcome {
210 tracing::warn!(
211 sub = %obtained.sub,
212 %reason,
213 "a session changed or was signed out during its refresh: the refresh's \
214 unstored new tokens could not be revoked"
215 );
216 }
217 if let Some(current) = current {
218 return Ok(current);
219 }
220 bail!(
221 "no session for {}: it was signed out while being refreshed (the new tokens were \
222 not stored{})",
223 obtained.sub,
224 if outcome == super::revoke::Revocation::Revoked {
225 ", and were revoked"
226 } else {
227 ""
228 }
229 )
230}
231
232async fn refresh_locked(
235 pool: &sqlx::SqlitePool,
236 codec: &super::crypto::Codec,
237 http: &reqwest::Client,
238 session: &OAuthSession,
239 version: &super::store::SessionVersion,
240 ctx: &RefreshContext<'_>,
241 now: i64,
242) -> Result<OAuthSession> {
243 let key = super::keys::SigningKey::from_jwk_json(&session.dpop_key_jwk, "session-dpop")?;
244
245 let assertion = match ctx.auth_method {
246 super::client_auth::AuthMethod::PrivateKeyJwt => {
247 let client_key = ctx
248 .client_key
249 .ok_or_else(|| anyhow::anyhow!("private_key_jwt refresh needs the client key"))?;
250 Some(super::client_auth::client_assertion(
251 client_key,
252 ctx.client_id,
253 &session.issuer,
254 now,
255 )?)
256 }
257 super::client_auth::AuthMethod::None => None,
258 };
259
260 let mut params = super::token::refresh_request_params(&session.refresh_token);
261 params.extend(super::client_auth::credential_params(
262 ctx.auth_method,
263 ctx.client_id,
264 assertion.as_deref(),
265 )?);
266 let borrowed: Vec<(&str, &str)> = params.iter().map(|(k, v)| (*k, v.as_str())).collect();
267
268 let outcome = super::request::send_with_dpop(
269 http,
270 pool,
271 &super::request::DpopRequest {
272 endpoint: super::dpop::Endpoint::AuthorizationServer,
273 url: ctx.token_endpoint,
274 key: &key,
275 access_token: None,
276 body: super::request::DpopBody::Form(&borrowed),
277 retry: super::request::Retry::Allowed,
280 },
281 )
282 .await?;
283
284 if outcome.is_success() {
285 let response = super::token::parse_token_response(&outcome.json()?)?;
286 let updated = apply_refresh(session, &response, now)?;
287 if !super::store::update_session_if_unchanged(pool, codec, &updated, version).await? {
294 return lost_the_row(pool, codec, http, ctx, &updated, now).await;
295 }
296 tracing::info!(
303 sub = %updated.sub,
304 expires_at = ?updated.expires_at,
305 "refreshed the OAuth session"
306 );
307 return Ok(updated);
308 }
309
310 match super::token::classify_refresh_failure(outcome.status, &outcome.body) {
311 super::token::RefreshFailure::Transient => {
312 bail!(
316 "refresh for {} failed transiently (status {}); the session is left intact",
317 session.sub,
318 outcome.status
319 )
320 }
321 super::token::RefreshFailure::SessionInvalid => {
322 if let Some(current) = super::store::get_session(pool, codec, &session.sub).await? {
327 if current.refresh_token != session.refresh_token {
328 return Ok(current);
329 }
330 }
331 super::store::delete_session(pool, &session.sub).await?;
332 bail!(
333 "refresh for {} was rejected as invalid_grant; the session has been \
334 removed and the user must log in again",
335 session.sub
336 )
337 }
338 }
339}
340
341#[cfg(test)]
342mod tests {
343 use super::*;
344
345 const DID: &str = "did:plc:ewvi7nxzyoun6zhxrhs64oiz";
346 const NOW: i64 = 1_700_000_000;
347
348 fn session() -> OAuthSession {
349 OAuthSession {
350 sub: DID.into(),
351 issuer: "https://auth.example.com".into(),
352 aud: "https://pds.example.com".into(),
353 dpop_key_jwk: r#"{"kty":"EC","d":"k"}"#.into(),
354 access_token: "old-access".into(),
355 refresh_token: "old-refresh".into(),
356 token_type: "DPoP".into(),
357 granted_scope: "atproto transition:generic".into(),
358 expires_at: Some(NOW + 60),
359 }
360 }
361
362 fn response() -> TokenResponse {
363 TokenResponse {
364 access_token: "new-access".into(),
365 refresh_token: Some("new-refresh".into()),
366 token_type: "DPoP".into(),
367 granted_scope: "atproto transition:generic".into(),
368 sub: DID.into(),
369 expires_in: Some(3600),
370 }
371 }
372
373 #[test]
374 fn a_refresh_replaces_both_tokens_and_the_expiry() {
375 let updated = apply_refresh(&session(), &response(), NOW).unwrap();
376 assert_eq!(updated.access_token, "new-access");
377 assert_eq!(updated.refresh_token, "new-refresh");
378 assert_eq!(updated.expires_at, Some(NOW + 3600));
379 }
380
381 #[test]
385 fn an_omitted_refresh_token_keeps_the_existing_one() {
386 let mut response = response();
387 response.refresh_token = None;
388 let updated = apply_refresh(&session(), &response, NOW).unwrap();
389 assert_eq!(updated.refresh_token, "old-refresh");
390 assert_eq!(
391 updated.access_token, "new-access",
392 "the access token still rotates"
393 );
394 }
395
396 #[test]
399 fn an_omitted_expiry_clears_rather_than_invents_one() {
400 let mut response = response();
401 response.expires_in = None;
402 assert_eq!(
403 apply_refresh(&session(), &response, NOW)
404 .unwrap()
405 .expires_at,
406 None
407 );
408 }
409
410 #[test]
413 fn a_refresh_for_a_different_subject_is_rejected() {
414 let mut response = response();
415 response.sub = "did:plc:aaaaaaaaaaaaaaaaaaaaaaaa".into();
416 assert!(apply_refresh(&session(), &response, NOW).is_err());
417 }
418
419 #[test]
422 fn the_granted_scope_is_taken_from_the_response() {
423 let mut response = response();
424 response.granted_scope = "atproto".into();
425 assert_eq!(
426 apply_refresh(&session(), &response, NOW)
427 .unwrap()
428 .granted_scope,
429 "atproto"
430 );
431 }
432
433 #[test]
436 fn a_refresh_preserves_the_session_key_and_audience() {
437 let updated = apply_refresh(&session(), &response(), NOW).unwrap();
438 assert_eq!(updated.dpop_key_jwk, session().dpop_key_jwk);
439 assert_eq!(updated.aud, session().aud);
440 assert_eq!(updated.issuer, session().issuer);
441 assert_eq!(updated.sub, session().sub);
442 }
443
444 #[tokio::test]
450 async fn the_same_subject_is_serialized() {
451 let locks = RefreshLocks::default();
452 let held = locks.lock(DID).await;
453
454 let second = locks.lock(DID);
455 tokio::pin!(second);
456 assert!(
457 futures_lite_poll_pending(&mut second),
458 "a second holder acquired the lock while the first held it"
459 );
460 drop(held);
461 let _ = second.await;
463 }
464
465 #[tokio::test]
468 async fn different_subjects_do_not_block_each_other() {
469 let locks = RefreshLocks::default();
470 let _a = locks.lock(DID).await;
471 let b = locks.lock("did:plc:aaaaaaaaaaaaaaaaaaaaaaaa");
472 tokio::pin!(b);
473 assert!(
474 !futures_lite_poll_pending(&mut b),
475 "an unrelated subject was blocked"
476 );
477 }
478
479 fn futures_lite_poll_pending<F: std::future::Future>(fut: &mut std::pin::Pin<&mut F>) -> bool {
481 use std::task::{Context, Poll, Waker};
482 let mut cx = Context::from_waker(Waker::noop());
483 matches!(fut.as_mut().poll(&mut cx), Poll::Pending)
484 }
485
486 #[test]
494 fn a_re_discovered_issuer_must_match_the_grants_own() {
495 same_issuer("https://pds.example.com", "https://pds.example.com")
496 .expect("the same issuer must pass");
497
498 let err = same_issuer("https://evil.example", "https://pds.example.com")
499 .expect_err("a different authorization server must be refused");
500 let rendered = format!("{err:#}");
501 assert!(
502 rendered.contains("evil.example") && rendered.contains("pds.example.com"),
503 "the error must name both, or an operator cannot tell what moved: {rendered}"
504 );
505 }
506
507 #[test]
511 fn the_issuer_comparison_is_exact() {
512 assert!(same_issuer("https://pds.example.com/", "https://pds.example.com").is_err());
513 assert!(same_issuer("https://PDS.example.com", "https://pds.example.com").is_err());
514 assert!(same_issuer("", "https://pds.example.com").is_err());
515 }
516
517 const KEY: &str = "aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa";
520
521 async fn token_server() -> (String, std::sync::Arc<std::sync::Mutex<Vec<String>>>) {
524 let (addr, log) = crate::net::spawn_tls(|_| {
525 let mut r = std::collections::HashMap::new();
526 r.insert(
527 "/token".to_string(),
528 vec![crate::net::TestResponse::json(
529 200,
530 serde_json::json!({
531 "access_token": "rotated-access",
532 "refresh_token": "rotated-refresh",
533 "token_type": "DPoP",
534 "scope": "atproto",
535 "sub": DID,
536 "expires_in": 3600,
537 })
538 .to_string(),
539 )],
540 );
541 r.insert(
542 "/revoke".to_string(),
543 vec![crate::net::TestResponse::json(200, "{}")],
544 );
545 r
546 })
547 .await;
548 crate::net::test_host_override("as-e2e.test", addr);
549 (format!("https://as-e2e.test:{}", addr.port()), log)
550 }
551
552 async fn stored_stale() -> (
554 sqlx::SqlitePool,
555 crate::oauth::crypto::Codec,
556 OAuthSession,
557 crate::oauth::store::SessionVersion,
558 ) {
559 let pool = crate::store::init_url("sqlite::memory:").await.unwrap();
560 let codec = crate::oauth::crypto::Codec::new(Some(KEY)).unwrap();
561 let stale = OAuthSession {
562 dpop_key_jwk: crate::oauth::keys::SigningKey::generate("session-dpop")
563 .to_jwk_json()
564 .unwrap(),
565 expires_at: Some(NOW - 1),
566 ..session()
567 };
568 crate::oauth::store::put_session(&pool, &codec, &stale)
569 .await
570 .unwrap();
571 let (read, version) = crate::oauth::store::get_session_versioned(&pool, &codec, DID)
572 .await
573 .unwrap()
574 .unwrap();
575 (pool, codec, read, version)
576 }
577
578 fn ctx<'a>(token: &'a str, revoke: &'a str) -> RefreshContext<'a> {
579 RefreshContext {
580 token_endpoint: token,
581 client_id: "http://localhost",
582 auth_method: super::super::client_auth::AuthMethod::None,
583 client_key: None,
584 revocation_endpoint: Some(revoke),
585 }
586 }
587
588 #[tokio::test]
598 async fn a_refresh_does_not_resurrect_a_session_signed_out_mid_refresh() {
599 let (base, log) = token_server().await;
600 let (token, revoke) = (format!("{base}/token"), format!("{base}/revoke"));
601 let (pool, codec, session, version) = stored_stale().await;
602
603 crate::oauth::store::delete_session(&pool, DID)
605 .await
606 .unwrap();
607
608 let result = refresh_locked(
609 &pool,
610 &codec,
611 &reqwest::Client::new(),
612 &session,
613 &version,
614 &ctx(&token, &revoke),
615 NOW,
616 )
617 .await;
618
619 assert!(
620 crate::oauth::store::get_session(&pool, &codec, DID)
621 .await
622 .unwrap()
623 .is_none(),
624 "the refresh RESURRECTED a signed-out session"
625 );
626 let err = result.expect_err("a signed-out session must not be handed back");
627 assert!(
628 format!("{err:#}").contains("signed out while being refreshed"),
629 "{err:#}"
630 );
631 let seen = log.lock().unwrap().join("\n---\n");
632 assert!(
633 seen.contains("POST /revoke") && seen.contains("token=rotated-refresh"),
634 "the orphaned fresh token was left live:\n{seen}"
635 );
636 }
637
638 #[tokio::test]
645 async fn a_refresh_does_not_overwrite_a_concurrent_rotation() {
646 let (base, log) = token_server().await;
647 let (token, revoke) = (format!("{base}/token"), format!("{base}/revoke"));
648 let (pool, codec, session, version) = stored_stale().await;
649
650 let theirs = OAuthSession {
651 refresh_token: "their-refresh".into(),
652 access_token: "their-access".into(),
653 expires_at: Some(NOW + 3600),
654 ..session.clone()
655 };
656 crate::oauth::store::put_session(&pool, &codec, &theirs)
657 .await
658 .unwrap();
659
660 let got = refresh_locked(
661 &pool,
662 &codec,
663 &reqwest::Client::new(),
664 &session,
665 &version,
666 &ctx(&token, &revoke),
667 NOW,
668 )
669 .await
670 .expect("a concurrent rotation is not an error");
671
672 let stored = crate::oauth::store::get_session(&pool, &codec, DID)
673 .await
674 .unwrap()
675 .unwrap();
676 assert_eq!(
677 stored.refresh_token, "their-refresh",
678 "the other writer's rotation was overwritten"
679 );
680 assert_eq!(got.refresh_token, "their-refresh");
681 let revokes: Vec<String> = log
682 .lock()
683 .unwrap()
684 .iter()
685 .filter(|r| r.starts_with("POST /revoke"))
686 .cloned()
687 .collect();
688 assert!(
689 revokes.iter().any(|r| r.contains("token=rotated-refresh")),
690 "the refresh's own fresh token was left live, stored nowhere:\n{revokes:#?}"
691 );
692 assert!(
693 !revokes.iter().any(|r| r.contains("their-refresh")),
694 "revoked the OTHER writer's live token:\n{revokes:#?}"
695 );
696 }
697
698 #[tokio::test]
700 async fn an_uncontested_refresh_stores_the_rotated_tokens() {
701 let (base, _log) = token_server().await;
702 let (token, revoke) = (format!("{base}/token"), format!("{base}/revoke"));
703 let (pool, codec, session, version) = stored_stale().await;
704
705 let got = refresh_locked(
706 &pool,
707 &codec,
708 &reqwest::Client::new(),
709 &session,
710 &version,
711 &ctx(&token, &revoke),
712 NOW,
713 )
714 .await
715 .expect("refresh");
716 assert_eq!(got.refresh_token, "rotated-refresh");
717 assert_eq!(
718 crate::oauth::store::get_session(&pool, &codec, DID)
719 .await
720 .unwrap()
721 .unwrap()
722 .refresh_token,
723 "rotated-refresh"
724 );
725 }
726}