1use anyhow::{bail, Context as _, Result};
19use serde_json::Value;
20
21use super::identity::is_atproto_did;
22
23const ATPROTO_SCOPE: &str = "atproto";
25
26const MAX_EXPIRES_IN_SECS: i64 = 365 * 24 * 60 * 60;
32
33pub const MIN_REFRESH_MARGIN_SECS: i64 = 10;
35
36pub const REFRESH_JITTER_SECS: i64 = 30;
42
43pub fn token_request_params(
48 code: &str,
49 redirect_uri: &str,
50 code_verifier: &str,
51) -> Vec<(&'static str, String)> {
52 vec![
53 ("grant_type", "authorization_code".to_string()),
54 ("code", code.to_string()),
55 ("redirect_uri", redirect_uri.to_string()),
57 ("code_verifier", code_verifier.to_string()),
58 ]
59}
60
61pub fn refresh_request_params(refresh_token: &str) -> Vec<(&'static str, String)> {
66 vec![
67 ("grant_type", "refresh_token".to_string()),
68 ("refresh_token", refresh_token.to_string()),
69 ]
70}
71
72pub struct TokenResponse {
74 pub access_token: String,
75 pub refresh_token: Option<String>,
77 pub token_type: String,
78 pub granted_scope: String,
81 pub sub: String,
82 pub expires_in: Option<i64>,
84}
85
86pub fn parse_token_response(body: &Value) -> Result<TokenResponse> {
94 let field = |name: &str| body.get(name).and_then(Value::as_str);
95
96 if body.get("id_token").is_some() {
97 bail!("token response carries an `id_token`; this is not an OIDC client");
98 }
99
100 let token_type = field("token_type").context("token response has no `token_type`")?;
101 if token_type != "DPoP" {
102 bail!(
103 "token_type must be exactly \"DPoP\", got {token_type:?} — accepting \
104 anything else would discard the proof-of-possession binding"
105 );
106 }
107
108 let granted_scope = field("scope").context("token response has no `scope`")?;
109 if !granted_scope
110 .split_ascii_whitespace()
111 .any(|s| s == ATPROTO_SCOPE)
112 {
113 bail!("granted scope {granted_scope:?} does not include {ATPROTO_SCOPE:?}");
114 }
115
116 let sub = field("sub").context("token response has no `sub`")?;
117 if !is_atproto_did(sub) {
118 bail!("token response `sub` {sub:?} is not a well-formed atproto DID");
119 }
120
121 let access_token = field("access_token").context("token response has no `access_token`")?;
122 if access_token.is_empty() {
123 bail!("token response `access_token` is empty");
124 }
125
126 let expires_in = match body.get("expires_in") {
127 None | Some(Value::Null) => None,
128 Some(value) => {
129 let seconds = value
130 .as_i64()
131 .with_context(|| format!("`expires_in` is not an integer: {value}"))?;
132 if seconds <= 0 {
133 bail!("`expires_in` must be positive, got {seconds}");
134 }
135 if seconds > MAX_EXPIRES_IN_SECS {
142 bail!(
143 "`expires_in` of {seconds}s is beyond anything a session should claim \
144 (cap {MAX_EXPIRES_IN_SECS}s)"
145 );
146 }
147 Some(seconds)
148 }
149 };
150
151 Ok(TokenResponse {
152 access_token: access_token.to_string(),
153 refresh_token: field("refresh_token")
160 .filter(|t| !t.is_empty())
161 .map(str::to_string),
162 token_type: token_type.to_string(),
163 granted_scope: granted_scope.to_string(),
164 sub: sub.to_string(),
165 expires_in,
166 })
167}
168
169pub fn is_stale_with_margin(expires_at: Option<i64>, now: i64, margin: i64) -> bool {
171 let Some(expires_at) = expires_at else {
174 return false;
175 };
176 expires_at <= now + margin
177}
178
179pub fn refresh_margin() -> i64 {
181 let mut byte = [0u8; 4];
182 getrandom::fill(&mut byte).expect("OS CSPRNG unavailable");
183 let jitter = i64::from(u32::from_be_bytes(byte) % (REFRESH_JITTER_SECS as u32 + 1));
184 MIN_REFRESH_MARGIN_SECS + jitter
185}
186
187pub fn is_stale(expires_at: Option<i64>, now: i64) -> bool {
189 is_stale_with_margin(expires_at, now, refresh_margin())
190}
191
192#[derive(Debug, Clone, Copy, PartialEq, Eq)]
194pub enum RefreshFailure {
195 SessionInvalid,
197 Transient,
200}
201
202pub fn classify_refresh_failure(status: u16, body: &[u8]) -> RefreshFailure {
208 if status != 400 {
209 return RefreshFailure::Transient;
210 }
211 if !super::error_body_worth_parsing(body) {
218 return RefreshFailure::Transient;
219 }
220 let is_invalid_grant = serde_json::from_slice::<Value>(body)
221 .ok()
222 .as_ref()
223 .and_then(|v| v.get("error"))
224 .and_then(Value::as_str)
225 == Some("invalid_grant");
226 if is_invalid_grant {
227 RefreshFailure::SessionInvalid
228 } else {
229 RefreshFailure::Transient
230 }
231}
232
233#[cfg(test)]
234mod tests {
235
236 #[test]
246 fn an_oversized_400_body_is_transient_rather_than_invalidating() {
247 let small = br#"{"error":"invalid_grant"}"#;
248 assert_eq!(
249 classify_refresh_failure(400, small),
250 RefreshFailure::SessionInvalid,
251 "a real invalid_grant must still invalidate, or nobody is ever asked \
252 to log in again",
253 );
254
255 let mut huge = String::from(r#"{"error":"invalid_grant","pad":["#);
257 while huge.len() < super::super::MAX_ERROR_BODY + 1_024 {
258 huge.push_str("{},");
259 }
260 huge.push_str("{}]}");
261 assert!(huge.len() > super::super::MAX_ERROR_BODY);
262 assert_eq!(
263 classify_refresh_failure(400, huge.as_bytes()),
264 RefreshFailure::Transient,
265 "an oversized 400 body was parsed and acted on",
266 );
267 }
268 use super::*;
269 use serde_json::json;
270
271 const DID: &str = "did:plc:ewvi7nxzyoun6zhxrhs64oiz";
272 const NOW: i64 = 1_700_000_000;
273
274 fn token_body() -> serde_json::Value {
275 json!({
276 "access_token": "access-abc",
277 "refresh_token": "refresh-xyz",
278 "token_type": "DPoP",
279 "scope": "atproto transition:generic",
280 "sub": DID,
281 "expires_in": 3600
282 })
283 }
284
285 #[test]
288 fn the_code_exchange_sends_the_grant_code_redirect_and_verifier() {
289 let params = token_request_params("the-code", "https://x.example/cb", "the-verifier");
290 let get = |k: &str| {
291 params
292 .iter()
293 .find(|(n, _)| *n == k)
294 .map(|(_, v)| v.as_str())
295 };
296 assert_eq!(get("grant_type"), Some("authorization_code"));
297 assert_eq!(get("code"), Some("the-code"));
298 assert_eq!(get("redirect_uri"), Some("https://x.example/cb"));
299 assert_eq!(get("code_verifier"), Some("the-verifier"));
300 }
301
302 #[test]
303 fn the_refresh_sends_only_the_grant_and_token() {
304 let params = refresh_request_params("refresh-xyz");
305 assert_eq!(
306 params,
307 vec![
308 ("grant_type", "refresh_token".to_string()),
309 ("refresh_token", "refresh-xyz".to_string()),
310 ]
311 );
312 }
313
314 #[test]
317 fn the_refresh_does_not_send_a_scope() {
318 assert!(!refresh_request_params("t")
319 .iter()
320 .any(|(n, _)| *n == "scope"));
321 }
322
323 #[test]
326 fn a_well_formed_token_response_parses() {
327 let parsed = parse_token_response(&token_body()).unwrap();
328 assert_eq!(parsed.access_token, "access-abc");
329 assert_eq!(parsed.refresh_token.as_deref(), Some("refresh-xyz"));
330 assert_eq!(parsed.sub, DID);
331 assert_eq!(parsed.granted_scope, "atproto transition:generic");
332 assert_eq!(parsed.expires_in, Some(3600));
333 }
334
335 #[test]
339 fn only_a_dpop_token_type_is_accepted() {
340 for bad in ["Bearer", "bearer", "dpop", "DPoP ", "", "MAC"] {
341 let mut body = token_body();
342 body["token_type"] = json!(bad);
343 assert!(parse_token_response(&body).is_err(), "accepted {bad:?}");
344 }
345 }
346
347 #[test]
351 fn the_granted_scope_must_contain_atproto_as_a_whole_value() {
352 for bad in ["transition:generic", "atproto-ish", "notatproto", ""] {
353 let mut body = token_body();
354 body["scope"] = json!(bad);
355 assert!(
356 parse_token_response(&body).is_err(),
357 "accepted scope {bad:?}"
358 );
359 }
360 for good in [
361 "atproto",
362 "atproto transition:generic",
363 "transition:generic atproto",
364 ] {
365 let mut body = token_body();
366 body["scope"] = json!(good);
367 assert!(
368 parse_token_response(&body).is_ok(),
369 "rejected scope {good:?}"
370 );
371 }
372 }
373
374 #[test]
375 fn a_missing_scope_is_rejected() {
376 let mut body = token_body();
377 body.as_object_mut().unwrap().remove("scope");
378 assert!(parse_token_response(&body).is_err());
379 }
380
381 #[test]
384 fn the_subject_must_be_a_well_formed_atproto_did() {
385 for bad in [
386 "not-a-did",
387 "did:example:123",
388 "did:plc:tooshort",
389 "did:web:evil.com/path",
390 "",
391 ] {
392 let mut body = token_body();
393 body["sub"] = json!(bad);
394 assert!(parse_token_response(&body).is_err(), "accepted sub {bad:?}");
395 }
396 }
397
398 #[test]
401 fn an_id_token_is_rejected() {
402 let mut body = token_body();
403 body["id_token"] = json!("eyJ...");
404 assert!(parse_token_response(&body).is_err());
405 }
406
407 #[test]
411 fn a_response_without_expires_in_is_valid_and_has_no_expiry() {
412 let mut body = token_body();
413 body.as_object_mut().unwrap().remove("expires_in");
414 assert_eq!(parse_token_response(&body).unwrap().expires_in, None);
415 }
416
417 #[test]
426 fn an_absurd_expires_in_is_rejected() {
427 for absurd in [i64::MAX, MAX_EXPIRES_IN_SECS + 1] {
428 let body = json!({
429 "token_type": "DPoP",
430 "scope": "atproto",
431 "sub": "did:plc:ewvi7nxzyoun6zhxrhs64oiz",
432 "access_token": "at",
433 "expires_in": absurd,
434 });
435 assert!(
436 parse_token_response(&body).is_err(),
437 "accepted expires_in={absurd}, which overflows `now + seconds`"
438 );
439 }
440 let ok = json!({
442 "token_type": "DPoP",
443 "scope": "atproto",
444 "sub": "did:plc:ewvi7nxzyoun6zhxrhs64oiz",
445 "access_token": "at",
446 "expires_in": MAX_EXPIRES_IN_SECS,
447 });
448 assert_eq!(
449 parse_token_response(&ok).unwrap().expires_in,
450 Some(MAX_EXPIRES_IN_SECS)
451 );
452 }
453
454 #[test]
462 fn an_empty_refresh_token_is_treated_as_absent() {
463 let body = json!({
464 "token_type": "DPoP",
465 "scope": "atproto",
466 "sub": "did:plc:ewvi7nxzyoun6zhxrhs64oiz",
467 "access_token": "at",
468 "refresh_token": "",
469 });
470 assert_eq!(
471 parse_token_response(&body).unwrap().refresh_token,
472 None,
473 "an empty refresh token would overwrite a live one and log the user out"
474 );
475 }
476
477 #[test]
478 fn a_nonsensical_expires_in_is_rejected() {
479 for bad in [json!(0), json!(-1), json!("3600"), json!(3600.5)] {
480 let mut body = token_body();
481 body["expires_in"] = bad.clone();
482 assert!(parse_token_response(&body).is_err(), "accepted {bad}");
483 }
484 }
485
486 #[test]
489 fn a_response_without_a_refresh_token_is_valid() {
490 let mut body = token_body();
491 body.as_object_mut().unwrap().remove("refresh_token");
492 assert_eq!(parse_token_response(&body).unwrap().refresh_token, None);
493 }
494
495 #[test]
496 fn a_response_without_an_access_token_is_rejected() {
497 let mut body = token_body();
498 body.as_object_mut().unwrap().remove("access_token");
499 assert!(parse_token_response(&body).is_err());
500 }
501
502 #[test]
507 fn a_session_without_an_expiry_is_never_stale() {
508 assert!(!is_stale_with_margin(None, NOW, 30));
509 }
510
511 #[test]
512 fn staleness_is_measured_against_the_margin() {
513 assert!(!is_stale_with_margin(Some(NOW + 100), NOW, 30));
514 assert!(is_stale_with_margin(Some(NOW + 29), NOW, 30));
515 assert!(is_stale_with_margin(Some(NOW - 1), NOW, 30));
516 assert!(is_stale_with_margin(Some(NOW + 30), NOW, 30));
519 }
520
521 #[test]
524 fn the_refresh_margin_is_jittered_within_its_band() {
525 let mut seen = std::collections::HashSet::new();
526 for _ in 0..200 {
527 let margin = refresh_margin();
528 assert!(
529 (MIN_REFRESH_MARGIN_SECS..=MIN_REFRESH_MARGIN_SECS + REFRESH_JITTER_SECS)
530 .contains(&margin),
531 "margin {margin} outside its band"
532 );
533 seen.insert(margin);
534 }
535 assert!(seen.len() > 1, "the margin is not actually jittered");
536 }
537
538 #[test]
544 fn only_invalid_grant_invalidates_the_session() {
545 assert_eq!(
546 classify_refresh_failure(400, br#"{"error":"invalid_grant"}"#),
547 RefreshFailure::SessionInvalid
548 );
549 for (status, body) in [
550 (400u16, &br#"{"error":"invalid_request"}"#[..]),
551 (400, br#"{"error":"use_dpop_nonce"}"#),
552 (400, b"not json"),
553 (401, br#"{"error":"invalid_grant"}"#),
554 (500, br#"{"error":"invalid_grant"}"#),
555 (503, b""),
556 ] {
557 assert_eq!(
558 classify_refresh_failure(status, body),
559 RefreshFailure::Transient,
560 "status {status} wrongly invalidated the session"
561 );
562 }
563 }
564}