use anyhow::{bail, Context as _, Result};
use serde_json::Value;
use super::identity::is_atproto_did;
const ATPROTO_SCOPE: &str = "atproto";
const MAX_EXPIRES_IN_SECS: i64 = 365 * 24 * 60 * 60;
pub const MIN_REFRESH_MARGIN_SECS: i64 = 10;
pub const REFRESH_JITTER_SECS: i64 = 30;
pub fn token_request_params(
code: &str,
redirect_uri: &str,
code_verifier: &str,
) -> Vec<(&'static str, String)> {
vec![
("grant_type", "authorization_code".to_string()),
("code", code.to_string()),
("redirect_uri", redirect_uri.to_string()),
("code_verifier", code_verifier.to_string()),
]
}
pub fn refresh_request_params(refresh_token: &str) -> Vec<(&'static str, String)> {
vec![
("grant_type", "refresh_token".to_string()),
("refresh_token", refresh_token.to_string()),
]
}
pub struct TokenResponse {
pub access_token: String,
pub refresh_token: Option<String>,
pub token_type: String,
pub granted_scope: String,
pub sub: String,
pub expires_in: Option<i64>,
}
pub fn parse_token_response(body: &Value) -> Result<TokenResponse> {
let field = |name: &str| body.get(name).and_then(Value::as_str);
if body.get("id_token").is_some() {
bail!("token response carries an `id_token`; this is not an OIDC client");
}
let token_type = field("token_type").context("token response has no `token_type`")?;
if token_type != "DPoP" {
bail!(
"token_type must be exactly \"DPoP\", got {token_type:?} — accepting \
anything else would discard the proof-of-possession binding"
);
}
let granted_scope = field("scope").context("token response has no `scope`")?;
if !granted_scope
.split_ascii_whitespace()
.any(|s| s == ATPROTO_SCOPE)
{
bail!("granted scope {granted_scope:?} does not include {ATPROTO_SCOPE:?}");
}
let sub = field("sub").context("token response has no `sub`")?;
if !is_atproto_did(sub) {
bail!("token response `sub` {sub:?} is not a well-formed atproto DID");
}
let access_token = field("access_token").context("token response has no `access_token`")?;
if access_token.is_empty() {
bail!("token response `access_token` is empty");
}
let expires_in = match body.get("expires_in") {
None | Some(Value::Null) => None,
Some(value) => {
let seconds = value
.as_i64()
.with_context(|| format!("`expires_in` is not an integer: {value}"))?;
if seconds <= 0 {
bail!("`expires_in` must be positive, got {seconds}");
}
if seconds > MAX_EXPIRES_IN_SECS {
bail!(
"`expires_in` of {seconds}s is beyond anything a session should claim \
(cap {MAX_EXPIRES_IN_SECS}s)"
);
}
Some(seconds)
}
};
Ok(TokenResponse {
access_token: access_token.to_string(),
refresh_token: field("refresh_token")
.filter(|t| !t.is_empty())
.map(str::to_string),
token_type: token_type.to_string(),
granted_scope: granted_scope.to_string(),
sub: sub.to_string(),
expires_in,
})
}
pub fn is_stale_with_margin(expires_at: Option<i64>, now: i64, margin: i64) -> bool {
let Some(expires_at) = expires_at else {
return false;
};
expires_at <= now + margin
}
pub fn refresh_margin() -> i64 {
let mut byte = [0u8; 4];
getrandom::fill(&mut byte).expect("OS CSPRNG unavailable");
let jitter = i64::from(u32::from_be_bytes(byte) % (REFRESH_JITTER_SECS as u32 + 1));
MIN_REFRESH_MARGIN_SECS + jitter
}
pub fn is_stale(expires_at: Option<i64>, now: i64) -> bool {
is_stale_with_margin(expires_at, now, refresh_margin())
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum RefreshFailure {
SessionInvalid,
Transient,
}
pub fn classify_refresh_failure(status: u16, body: &[u8]) -> RefreshFailure {
if status != 400 {
return RefreshFailure::Transient;
}
if !super::error_body_worth_parsing(body) {
return RefreshFailure::Transient;
}
let is_invalid_grant = serde_json::from_slice::<Value>(body)
.ok()
.as_ref()
.and_then(|v| v.get("error"))
.and_then(Value::as_str)
== Some("invalid_grant");
if is_invalid_grant {
RefreshFailure::SessionInvalid
} else {
RefreshFailure::Transient
}
}
#[cfg(test)]
mod tests {
#[test]
fn an_oversized_400_body_is_transient_rather_than_invalidating() {
let small = br#"{"error":"invalid_grant"}"#;
assert_eq!(
classify_refresh_failure(400, small),
RefreshFailure::SessionInvalid,
"a real invalid_grant must still invalidate, or nobody is ever asked \
to log in again",
);
let mut huge = String::from(r#"{"error":"invalid_grant","pad":["#);
while huge.len() < super::super::MAX_ERROR_BODY + 1_024 {
huge.push_str("{},");
}
huge.push_str("{}]}");
assert!(huge.len() > super::super::MAX_ERROR_BODY);
assert_eq!(
classify_refresh_failure(400, huge.as_bytes()),
RefreshFailure::Transient,
"an oversized 400 body was parsed and acted on",
);
}
use super::*;
use serde_json::json;
const DID: &str = "did:plc:ewvi7nxzyoun6zhxrhs64oiz";
const NOW: i64 = 1_700_000_000;
fn token_body() -> serde_json::Value {
json!({
"access_token": "access-abc",
"refresh_token": "refresh-xyz",
"token_type": "DPoP",
"scope": "atproto transition:generic",
"sub": DID,
"expires_in": 3600
})
}
#[test]
fn the_code_exchange_sends_the_grant_code_redirect_and_verifier() {
let params = token_request_params("the-code", "https://x.example/cb", "the-verifier");
let get = |k: &str| {
params
.iter()
.find(|(n, _)| *n == k)
.map(|(_, v)| v.as_str())
};
assert_eq!(get("grant_type"), Some("authorization_code"));
assert_eq!(get("code"), Some("the-code"));
assert_eq!(get("redirect_uri"), Some("https://x.example/cb"));
assert_eq!(get("code_verifier"), Some("the-verifier"));
}
#[test]
fn the_refresh_sends_only_the_grant_and_token() {
let params = refresh_request_params("refresh-xyz");
assert_eq!(
params,
vec![
("grant_type", "refresh_token".to_string()),
("refresh_token", "refresh-xyz".to_string()),
]
);
}
#[test]
fn the_refresh_does_not_send_a_scope() {
assert!(!refresh_request_params("t")
.iter()
.any(|(n, _)| *n == "scope"));
}
#[test]
fn a_well_formed_token_response_parses() {
let parsed = parse_token_response(&token_body()).unwrap();
assert_eq!(parsed.access_token, "access-abc");
assert_eq!(parsed.refresh_token.as_deref(), Some("refresh-xyz"));
assert_eq!(parsed.sub, DID);
assert_eq!(parsed.granted_scope, "atproto transition:generic");
assert_eq!(parsed.expires_in, Some(3600));
}
#[test]
fn only_a_dpop_token_type_is_accepted() {
for bad in ["Bearer", "bearer", "dpop", "DPoP ", "", "MAC"] {
let mut body = token_body();
body["token_type"] = json!(bad);
assert!(parse_token_response(&body).is_err(), "accepted {bad:?}");
}
}
#[test]
fn the_granted_scope_must_contain_atproto_as_a_whole_value() {
for bad in ["transition:generic", "atproto-ish", "notatproto", ""] {
let mut body = token_body();
body["scope"] = json!(bad);
assert!(
parse_token_response(&body).is_err(),
"accepted scope {bad:?}"
);
}
for good in [
"atproto",
"atproto transition:generic",
"transition:generic atproto",
] {
let mut body = token_body();
body["scope"] = json!(good);
assert!(
parse_token_response(&body).is_ok(),
"rejected scope {good:?}"
);
}
}
#[test]
fn a_missing_scope_is_rejected() {
let mut body = token_body();
body.as_object_mut().unwrap().remove("scope");
assert!(parse_token_response(&body).is_err());
}
#[test]
fn the_subject_must_be_a_well_formed_atproto_did() {
for bad in [
"not-a-did",
"did:example:123",
"did:plc:tooshort",
"did:web:evil.com/path",
"",
] {
let mut body = token_body();
body["sub"] = json!(bad);
assert!(parse_token_response(&body).is_err(), "accepted sub {bad:?}");
}
}
#[test]
fn an_id_token_is_rejected() {
let mut body = token_body();
body["id_token"] = json!("eyJ...");
assert!(parse_token_response(&body).is_err());
}
#[test]
fn a_response_without_expires_in_is_valid_and_has_no_expiry() {
let mut body = token_body();
body.as_object_mut().unwrap().remove("expires_in");
assert_eq!(parse_token_response(&body).unwrap().expires_in, None);
}
#[test]
fn an_absurd_expires_in_is_rejected() {
for absurd in [i64::MAX, MAX_EXPIRES_IN_SECS + 1] {
let body = json!({
"token_type": "DPoP",
"scope": "atproto",
"sub": "did:plc:ewvi7nxzyoun6zhxrhs64oiz",
"access_token": "at",
"expires_in": absurd,
});
assert!(
parse_token_response(&body).is_err(),
"accepted expires_in={absurd}, which overflows `now + seconds`"
);
}
let ok = json!({
"token_type": "DPoP",
"scope": "atproto",
"sub": "did:plc:ewvi7nxzyoun6zhxrhs64oiz",
"access_token": "at",
"expires_in": MAX_EXPIRES_IN_SECS,
});
assert_eq!(
parse_token_response(&ok).unwrap().expires_in,
Some(MAX_EXPIRES_IN_SECS)
);
}
#[test]
fn an_empty_refresh_token_is_treated_as_absent() {
let body = json!({
"token_type": "DPoP",
"scope": "atproto",
"sub": "did:plc:ewvi7nxzyoun6zhxrhs64oiz",
"access_token": "at",
"refresh_token": "",
});
assert_eq!(
parse_token_response(&body).unwrap().refresh_token,
None,
"an empty refresh token would overwrite a live one and log the user out"
);
}
#[test]
fn a_nonsensical_expires_in_is_rejected() {
for bad in [json!(0), json!(-1), json!("3600"), json!(3600.5)] {
let mut body = token_body();
body["expires_in"] = bad.clone();
assert!(parse_token_response(&body).is_err(), "accepted {bad}");
}
}
#[test]
fn a_response_without_a_refresh_token_is_valid() {
let mut body = token_body();
body.as_object_mut().unwrap().remove("refresh_token");
assert_eq!(parse_token_response(&body).unwrap().refresh_token, None);
}
#[test]
fn a_response_without_an_access_token_is_rejected() {
let mut body = token_body();
body.as_object_mut().unwrap().remove("access_token");
assert!(parse_token_response(&body).is_err());
}
#[test]
fn a_session_without_an_expiry_is_never_stale() {
assert!(!is_stale_with_margin(None, NOW, 30));
}
#[test]
fn staleness_is_measured_against_the_margin() {
assert!(!is_stale_with_margin(Some(NOW + 100), NOW, 30));
assert!(is_stale_with_margin(Some(NOW + 29), NOW, 30));
assert!(is_stale_with_margin(Some(NOW - 1), NOW, 30));
assert!(is_stale_with_margin(Some(NOW + 30), NOW, 30));
}
#[test]
fn the_refresh_margin_is_jittered_within_its_band() {
let mut seen = std::collections::HashSet::new();
for _ in 0..200 {
let margin = refresh_margin();
assert!(
(MIN_REFRESH_MARGIN_SECS..=MIN_REFRESH_MARGIN_SECS + REFRESH_JITTER_SECS)
.contains(&margin),
"margin {margin} outside its band"
);
seen.insert(margin);
}
assert!(seen.len() > 1, "the margin is not actually jittered");
}
#[test]
fn only_invalid_grant_invalidates_the_session() {
assert_eq!(
classify_refresh_failure(400, br#"{"error":"invalid_grant"}"#),
RefreshFailure::SessionInvalid
);
for (status, body) in [
(400u16, &br#"{"error":"invalid_request"}"#[..]),
(400, br#"{"error":"use_dpop_nonce"}"#),
(400, b"not json"),
(401, br#"{"error":"invalid_grant"}"#),
(500, br#"{"error":"invalid_grant"}"#),
(503, b""),
] {
assert_eq!(
classify_refresh_failure(status, body),
RefreshFailure::Transient,
"status {status} wrongly invalidated the session"
);
}
}
}