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}
126
127pub async fn valid_session(
134 pool: &sqlx::SqlitePool,
135 codec: &super::crypto::Codec,
136 http: &reqwest::Client,
137 locks: &RefreshLocks,
138 sub: &str,
139 ctx: &RefreshContext<'_>,
140 now: i64,
141) -> Result<OAuthSession> {
142 let session = super::store::get_session(pool, codec, sub)
143 .await?
144 .ok_or_else(|| anyhow::anyhow!("no session for {sub}"))?;
145 if !super::token::is_stale(session.expires_at, now) {
146 return Ok(session);
147 }
148
149 let _guard = locks.lock(sub).await;
150
151 let session = super::store::get_session(pool, codec, sub)
154 .await?
155 .ok_or_else(|| anyhow::anyhow!("session for {sub} disappeared while waiting to refresh"))?;
156 if !super::token::is_stale(session.expires_at, now) {
157 return Ok(session);
158 }
159
160 refresh_locked(pool, codec, http, &session, ctx, now).await
166}
167
168async fn refresh_locked(
170 pool: &sqlx::SqlitePool,
171 codec: &super::crypto::Codec,
172 http: &reqwest::Client,
173 session: &OAuthSession,
174 ctx: &RefreshContext<'_>,
175 now: i64,
176) -> Result<OAuthSession> {
177 let key = super::keys::SigningKey::from_jwk_json(&session.dpop_key_jwk, "session-dpop")?;
178
179 let assertion = match ctx.auth_method {
180 super::client_auth::AuthMethod::PrivateKeyJwt => {
181 let client_key = ctx
182 .client_key
183 .ok_or_else(|| anyhow::anyhow!("private_key_jwt refresh needs the client key"))?;
184 Some(super::client_auth::client_assertion(
185 client_key,
186 ctx.client_id,
187 &session.issuer,
188 now,
189 )?)
190 }
191 super::client_auth::AuthMethod::None => None,
192 };
193
194 let mut params = super::token::refresh_request_params(&session.refresh_token);
195 params.extend(super::client_auth::credential_params(
196 ctx.auth_method,
197 ctx.client_id,
198 assertion.as_deref(),
199 )?);
200 let borrowed: Vec<(&str, &str)> = params.iter().map(|(k, v)| (*k, v.as_str())).collect();
201
202 let outcome = super::request::send_with_dpop(
203 http,
204 pool,
205 &super::request::DpopRequest {
206 endpoint: super::dpop::Endpoint::AuthorizationServer,
207 url: ctx.token_endpoint,
208 key: &key,
209 access_token: None,
210 body: super::request::DpopBody::Form(&borrowed),
211 retry: super::request::Retry::Allowed,
214 },
215 )
216 .await?;
217
218 if outcome.is_success() {
219 let response = super::token::parse_token_response(&outcome.json()?)?;
220 let updated = apply_refresh(session, &response, now)?;
221 super::store::put_session(pool, codec, &updated).await?;
222 tracing::info!(
229 sub = %updated.sub,
230 expires_at = ?updated.expires_at,
231 "refreshed the OAuth session"
232 );
233 return Ok(updated);
234 }
235
236 match super::token::classify_refresh_failure(outcome.status, &outcome.body) {
237 super::token::RefreshFailure::Transient => {
238 bail!(
242 "refresh for {} failed transiently (status {}); the session is left intact",
243 session.sub,
244 outcome.status
245 )
246 }
247 super::token::RefreshFailure::SessionInvalid => {
248 if let Some(current) = super::store::get_session(pool, codec, &session.sub).await? {
253 if current.refresh_token != session.refresh_token {
254 return Ok(current);
255 }
256 }
257 super::store::delete_session(pool, &session.sub).await?;
258 bail!(
259 "refresh for {} was rejected as invalid_grant; the session has been \
260 removed and the user must log in again",
261 session.sub
262 )
263 }
264 }
265}
266
267#[cfg(test)]
268mod tests {
269 use super::*;
270
271 const DID: &str = "did:plc:ewvi7nxzyoun6zhxrhs64oiz";
272 const NOW: i64 = 1_700_000_000;
273
274 fn session() -> OAuthSession {
275 OAuthSession {
276 sub: DID.into(),
277 issuer: "https://auth.example.com".into(),
278 aud: "https://pds.example.com".into(),
279 dpop_key_jwk: r#"{"kty":"EC","d":"k"}"#.into(),
280 access_token: "old-access".into(),
281 refresh_token: "old-refresh".into(),
282 token_type: "DPoP".into(),
283 granted_scope: "atproto transition:generic".into(),
284 expires_at: Some(NOW + 60),
285 }
286 }
287
288 fn response() -> TokenResponse {
289 TokenResponse {
290 access_token: "new-access".into(),
291 refresh_token: Some("new-refresh".into()),
292 token_type: "DPoP".into(),
293 granted_scope: "atproto transition:generic".into(),
294 sub: DID.into(),
295 expires_in: Some(3600),
296 }
297 }
298
299 #[test]
300 fn a_refresh_replaces_both_tokens_and_the_expiry() {
301 let updated = apply_refresh(&session(), &response(), NOW).unwrap();
302 assert_eq!(updated.access_token, "new-access");
303 assert_eq!(updated.refresh_token, "new-refresh");
304 assert_eq!(updated.expires_at, Some(NOW + 3600));
305 }
306
307 #[test]
311 fn an_omitted_refresh_token_keeps_the_existing_one() {
312 let mut response = response();
313 response.refresh_token = None;
314 let updated = apply_refresh(&session(), &response, NOW).unwrap();
315 assert_eq!(updated.refresh_token, "old-refresh");
316 assert_eq!(
317 updated.access_token, "new-access",
318 "the access token still rotates"
319 );
320 }
321
322 #[test]
325 fn an_omitted_expiry_clears_rather_than_invents_one() {
326 let mut response = response();
327 response.expires_in = None;
328 assert_eq!(
329 apply_refresh(&session(), &response, NOW)
330 .unwrap()
331 .expires_at,
332 None
333 );
334 }
335
336 #[test]
339 fn a_refresh_for_a_different_subject_is_rejected() {
340 let mut response = response();
341 response.sub = "did:plc:aaaaaaaaaaaaaaaaaaaaaaaa".into();
342 assert!(apply_refresh(&session(), &response, NOW).is_err());
343 }
344
345 #[test]
348 fn the_granted_scope_is_taken_from_the_response() {
349 let mut response = response();
350 response.granted_scope = "atproto".into();
351 assert_eq!(
352 apply_refresh(&session(), &response, NOW)
353 .unwrap()
354 .granted_scope,
355 "atproto"
356 );
357 }
358
359 #[test]
362 fn a_refresh_preserves_the_session_key_and_audience() {
363 let updated = apply_refresh(&session(), &response(), NOW).unwrap();
364 assert_eq!(updated.dpop_key_jwk, session().dpop_key_jwk);
365 assert_eq!(updated.aud, session().aud);
366 assert_eq!(updated.issuer, session().issuer);
367 assert_eq!(updated.sub, session().sub);
368 }
369
370 #[tokio::test]
376 async fn the_same_subject_is_serialized() {
377 let locks = RefreshLocks::default();
378 let held = locks.lock(DID).await;
379
380 let second = locks.lock(DID);
381 tokio::pin!(second);
382 assert!(
383 futures_lite_poll_pending(&mut second),
384 "a second holder acquired the lock while the first held it"
385 );
386 drop(held);
387 let _ = second.await;
389 }
390
391 #[tokio::test]
394 async fn different_subjects_do_not_block_each_other() {
395 let locks = RefreshLocks::default();
396 let _a = locks.lock(DID).await;
397 let b = locks.lock("did:plc:aaaaaaaaaaaaaaaaaaaaaaaa");
398 tokio::pin!(b);
399 assert!(
400 !futures_lite_poll_pending(&mut b),
401 "an unrelated subject was blocked"
402 );
403 }
404
405 fn futures_lite_poll_pending<F: std::future::Future>(fut: &mut std::pin::Pin<&mut F>) -> bool {
407 use std::task::{Context, Poll, Waker};
408 let mut cx = Context::from_waker(Waker::noop());
409 matches!(fut.as_mut().poll(&mut cx), Poll::Pending)
410 }
411
412 #[test]
420 fn a_re_discovered_issuer_must_match_the_grants_own() {
421 same_issuer("https://pds.example.com", "https://pds.example.com")
422 .expect("the same issuer must pass");
423
424 let err = same_issuer("https://evil.example", "https://pds.example.com")
425 .expect_err("a different authorization server must be refused");
426 let rendered = format!("{err:#}");
427 assert!(
428 rendered.contains("evil.example") && rendered.contains("pds.example.com"),
429 "the error must name both, or an operator cannot tell what moved: {rendered}"
430 );
431 }
432
433 #[test]
437 fn the_issuer_comparison_is_exact() {
438 assert!(same_issuer("https://pds.example.com/", "https://pds.example.com").is_err());
439 assert!(same_issuer("https://PDS.example.com", "https://pds.example.com").is_err());
440 assert!(same_issuer("", "https://pds.example.com").is_err());
441 }
442}