use std::sync::Arc;
use std::time::{Duration, SystemTime};
use url::Url;
use crate::crypto::base64url_encode;
use crate::dpop::{compute_access_token_hash, extract_dpop_nonce, DPoPKey, DPoPNonceCache};
use crate::error::{AtprotoOAuthError, DPoPError, ParError, TokenError};
use crate::identity::{IdentityResolver, IdentityResolverBuilder};
use crate::par::{build_authorization_url, execute_par_request, ParParameters};
use crate::pkce::PkcePair;
use crate::session::OAuthSession;
use crate::ssrf::{read_bounded_body, SsrfFilter, MAX_OAUTH_RESPONSE_BYTES};
use crate::store::{OAuthStateStore, OAuthStore, DEFAULT_STATE_TTL};
use zeroize::{Zeroize, ZeroizeOnDrop};
#[derive(Clone, PartialEq, Eq)]
pub struct OAuthClientMetadata {
pub client_id: String,
pub redirect_uri: String,
pub scope: String,
pub client_name: Option<String>,
}
impl std::fmt::Debug for OAuthClientMetadata {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("OAuthClientMetadata")
.field("client_id", &self.client_id)
.field("redirect_uri", &self.redirect_uri)
.field("scope", &self.scope)
.field("client_name", &self.client_name)
.finish()
}
}
impl OAuthClientMetadata {
#[must_use]
pub fn new(client_id: impl Into<String>, redirect_uri: impl Into<String>) -> Self {
Self {
client_id: client_id.into(),
redirect_uri: redirect_uri.into(),
scope: "atproto".to_string(),
client_name: None,
}
}
#[must_use]
pub fn with_scope(mut self, scope: impl Into<String>) -> Self {
self.scope = scope.into();
self
}
#[must_use]
pub fn with_client_name(mut self, name: impl Into<String>) -> Self {
self.client_name = Some(name.into());
self
}
}
#[derive(Clone)]
pub struct StoredStateEntry {
pub state: String,
pub client_id: String,
pub code_verifier: String,
pub dpop_key: DPoPKey,
pub issuer: String,
pub did: Option<String>,
pub handle: Option<String>,
pub redirect_uri: String,
pub pds_endpoint: String,
pub token_endpoint: String,
pub scopes: String,
pub created_at: SystemTime,
pub expires_in_secs: u64,
}
impl std::fmt::Debug for StoredStateEntry {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("StoredStateEntry")
.field("state", &self.state)
.field("client_id", &self.client_id)
.field("code_verifier", &"[REDACTED]")
.field("dpop_key", &self.dpop_key)
.field("issuer", &self.issuer)
.field("did", &self.did)
.field("handle", &self.handle)
.field("redirect_uri", &self.redirect_uri)
.field("pds_endpoint", &self.pds_endpoint)
.field("token_endpoint", &self.token_endpoint)
.field("scopes", &self.scopes)
.field("created_at", &self.created_at)
.field("expires_in_secs", &self.expires_in_secs)
.finish()
}
}
impl Zeroize for StoredStateEntry {
fn zeroize(&mut self) {
self.code_verifier.zeroize();
}
}
impl Drop for StoredStateEntry {
fn drop(&mut self) {
self.zeroize();
}
}
impl ZeroizeOnDrop for StoredStateEntry {}
impl StoredStateEntry {
#[must_use]
pub fn is_expired(&self) -> bool {
let now = SystemTime::now();
let max_age = Duration::from_secs(self.expires_in_secs);
match now.duration_since(self.created_at) {
Ok(elapsed) => elapsed > max_age,
Err(_) => true,
}
}
}
#[derive(Debug, Clone)]
pub struct AuthorizationRequest {
pub authorization_url: Url,
pub state: String,
pub request_uri: String,
pub expires_in: u64,
pub stored_state: StoredStateEntry,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct CallbackParams {
pub code: String,
pub state: String,
pub iss: Option<String>,
}
impl CallbackParams {
#[must_use]
pub fn new(code: impl Into<String>, state: impl Into<String>) -> Self {
Self {
code: code.into(),
state: state.into(),
iss: None,
}
}
#[must_use]
pub fn with_iss(mut self, iss: impl Into<String>) -> Self {
self.iss = Some(iss.into());
self
}
}
#[derive(Clone, PartialEq, Eq, serde::Deserialize, serde::Serialize)]
pub struct TokenResponse {
pub access_token: String,
pub token_type: String,
#[serde(skip_serializing_if = "Option::is_none")]
pub expires_in: Option<u64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub refresh_token: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub scope: Option<String>,
pub sub: String,
}
impl std::fmt::Debug for TokenResponse {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("TokenResponse")
.field("access_token", &"[REDACTED]")
.field("token_type", &self.token_type)
.field("expires_in", &self.expires_in)
.field(
"refresh_token",
&self.refresh_token.as_ref().map(|_| "[REDACTED]"),
)
.field("scope", &self.scope)
.field("sub", &self.sub)
.finish()
}
}
impl zeroize::Zeroize for TokenResponse {
fn zeroize(&mut self) {
self.access_token.zeroize();
if let Some(ref mut rt) = self.refresh_token {
rt.zeroize();
}
}
}
impl Drop for TokenResponse {
fn drop(&mut self) {
self.zeroize();
}
}
impl zeroize::ZeroizeOnDrop for TokenResponse {}
impl TokenResponse {
#[must_use]
pub fn into_parts(
mut self,
) -> (
String, // sub
String, // access_token
Option<String>, // refresh_token
String, // token_type
Option<String>, // scope
Option<u64>, // expires_in
) {
let sub = std::mem::take(&mut self.sub);
let access_token = std::mem::take(&mut self.access_token);
let refresh_token = self.refresh_token.take();
let token_type = std::mem::take(&mut self.token_type);
let scope = self.scope.take();
let expires_in = self.expires_in;
(
sub,
access_token,
refresh_token,
token_type,
scope,
expires_in,
)
}
}
#[derive(Debug, Clone)]
pub struct AtprotoOAuthClientBuilder {
metadata: Option<OAuthClientMetadata>,
resolver: Option<IdentityResolver>,
nonce_cache: Option<DPoPNonceCache>,
ssrf_filter: SsrfFilter,
state_store: Option<Arc<OAuthStateStore>>,
state_ttl: Duration,
}
impl Default for AtprotoOAuthClientBuilder {
fn default() -> Self {
Self::new()
}
}
impl AtprotoOAuthClientBuilder {
#[must_use]
pub fn new() -> Self {
Self {
metadata: None,
resolver: None,
nonce_cache: None,
ssrf_filter: SsrfFilter::default(),
state_store: None,
state_ttl: DEFAULT_STATE_TTL,
}
}
#[must_use]
pub fn client_metadata(mut self, metadata: OAuthClientMetadata) -> Self {
self.metadata = Some(metadata);
self
}
#[must_use]
pub fn metadata(self, metadata: OAuthClientMetadata) -> Self {
self.client_metadata(metadata)
}
#[must_use]
pub fn identity_resolver(mut self, resolver: IdentityResolver) -> Self {
self.resolver = Some(resolver);
self
}
#[must_use]
pub fn nonce_cache(mut self, cache: DPoPNonceCache) -> Self {
self.nonce_cache = Some(cache);
self
}
#[must_use]
pub fn ssrf_filter(mut self, filter: SsrfFilter) -> Self {
self.ssrf_filter = filter;
self
}
#[must_use]
pub fn allow_insecure_localhost(mut self, allow: bool) -> Self {
self.ssrf_filter.allow_insecure_localhost = allow;
self
}
#[must_use]
pub fn state_store(mut self, store: Arc<OAuthStateStore>) -> Self {
self.state_store = Some(store);
self
}
#[must_use]
pub fn state_ttl(mut self, ttl: Duration) -> Self {
self.state_ttl = ttl;
self
}
pub fn build(self) -> Result<AtprotoOAuthClient, AtprotoOAuthError> {
if self.state_ttl.subsec_nanos() != 0 {
return Err(AtprotoOAuthError::Token(TokenError::InvalidStateTtl(
self.state_ttl,
)));
}
let metadata = self
.metadata
.ok_or(ParError::MissingField("client_metadata"))?;
let resolver = self.resolver.unwrap_or_else(|| {
IdentityResolverBuilder::new()
.ssrf_filter(self.ssrf_filter)
.build()
});
let nonce_cache = self.nonce_cache.unwrap_or_default();
let default_store_created = self.state_store.is_none();
let state_store = self
.state_store
.unwrap_or_else(|| Arc::new(OAuthStateStore::new(self.state_ttl)));
if default_store_created {
if let Ok(handle) = tokio::runtime::Handle::try_current() {
let pruner_store = Arc::clone(&state_store);
let prune_interval = self.state_ttl.max(Duration::from_secs(60));
handle.spawn(async move {
let mut interval = tokio::time::interval(prune_interval);
interval.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Skip);
loop {
interval.tick().await;
let pruned = pruner_store.prune_expired_sync();
if pruned > 0 {
tracing::trace!("default state store pruned {pruned} expired states");
}
}
});
}
}
Ok(AtprotoOAuthClient {
metadata,
resolver,
nonce_cache,
ssrf_filter: self.ssrf_filter,
state_store,
state_ttl: self.state_ttl,
refresh_single_flight: Arc::new(RefreshSingleFlight::new()),
})
}
}
#[derive(Debug, Default)]
struct RefreshSingleFlight {
locks: parking_lot::RwLock<std::collections::HashMap<String, Arc<tokio::sync::Mutex<()>>>>,
}
impl RefreshSingleFlight {
fn new() -> Self {
Self::default()
}
fn lock_for(&self, sub: &str) -> Arc<tokio::sync::Mutex<()>> {
if let Some(existing) = self.locks.read().get(sub) {
return Arc::clone(existing);
}
let mut guard = self.locks.write();
if let Some(existing) = guard.get(sub) {
return Arc::clone(existing);
}
let fresh = Arc::new(tokio::sync::Mutex::new(()));
guard.insert(sub.to_string(), Arc::clone(&fresh));
fresh
}
}
#[derive(Debug, Clone)]
pub struct AtprotoOAuthClient {
metadata: OAuthClientMetadata,
resolver: IdentityResolver,
nonce_cache: DPoPNonceCache,
ssrf_filter: SsrfFilter,
state_store: Arc<OAuthStateStore>,
state_ttl: Duration,
refresh_single_flight: Arc<RefreshSingleFlight>,
}
impl AtprotoOAuthClient {
#[must_use]
pub fn new(client_id: impl Into<String>, redirect_uri: impl Into<String>) -> Self {
let metadata = OAuthClientMetadata::new(client_id, redirect_uri);
let ssrf_filter = SsrfFilter::default();
let resolver = IdentityResolverBuilder::new()
.ssrf_filter(ssrf_filter)
.build();
Self {
metadata,
resolver,
nonce_cache: DPoPNonceCache::new(),
ssrf_filter,
state_store: Arc::new(OAuthStateStore::default()),
state_ttl: DEFAULT_STATE_TTL,
refresh_single_flight: Arc::new(RefreshSingleFlight::new()),
}
}
#[must_use]
pub fn builder() -> AtprotoOAuthClientBuilder {
AtprotoOAuthClientBuilder::new()
}
#[must_use]
pub const fn metadata(&self) -> &OAuthClientMetadata {
&self.metadata
}
#[must_use]
pub const fn resolver(&self) -> &IdentityResolver {
&self.resolver
}
#[must_use]
pub const fn nonce_cache(&self) -> &DPoPNonceCache {
&self.nonce_cache
}
#[must_use]
pub const fn ssrf_filter(&self) -> &SsrfFilter {
&self.ssrf_filter
}
#[must_use]
pub fn state_store(&self) -> &Arc<OAuthStateStore> {
&self.state_store
}
#[must_use]
pub const fn state_ttl(&self) -> Duration {
self.state_ttl
}
pub async fn authorize(
&self,
handle_or_did: &str,
) -> Result<AuthorizationRequest, AtprotoOAuthError> {
let (req, _) = self.initiate_login(handle_or_did).await?;
Ok(req)
}
pub async fn initiate_login(
&self,
handle_or_did: &str,
) -> Result<(AuthorizationRequest, StoredStateEntry), AtprotoOAuthError> {
self.initiate_login_with_scope(handle_or_did, &self.metadata.scope)
.await
}
pub async fn initiate_login_with_scope(
&self,
handle_or_did: &str,
scope: &str,
) -> Result<(AuthorizationRequest, StoredStateEntry), AtprotoOAuthError> {
let endpoints = self
.resolver
.discover_oauth_endpoints(handle_or_did)
.await?;
let pkce = PkcePair::generate();
let mut state_bytes = [0u8; 32];
rand::RngCore::fill_bytes(&mut rand::thread_rng(), &mut state_bytes);
let state = base64url_encode(&state_bytes);
let dpop_key = DPoPKey::generate();
let mut params = ParParameters::new(
&self.metadata.client_id,
&self.metadata.redirect_uri,
scope,
&state,
&pkce.challenge,
);
if let Some(ref handle) = endpoints.handle {
params = params.with_login_hint(handle);
} else {
params = params.with_login_hint(&endpoints.did);
}
let par_res = execute_par_request(
&self.ssrf_filter,
&endpoints.par_endpoint,
¶ms,
&dpop_key,
&self.nonce_cache,
)
.await?;
let auth_url = build_authorization_url(
&endpoints.authorization_endpoint,
&self.metadata.client_id,
&par_res.request_uri,
)?;
let stored_state = StoredStateEntry {
state: state.clone(),
client_id: self.metadata.client_id.clone(),
code_verifier: pkce.verifier,
dpop_key,
issuer: endpoints.auth_server_issuer.clone(),
did: Some(endpoints.did),
handle: endpoints.handle,
redirect_uri: self.metadata.redirect_uri.clone(),
pds_endpoint: endpoints.pds_endpoint,
token_endpoint: endpoints.token_endpoint,
scopes: scope.to_string(),
created_at: SystemTime::now(),
expires_in_secs: self.state_ttl.as_secs(),
};
self.state_store
.insert_state(state.clone(), stored_state.clone(), self.state_ttl)
.await?;
let auth_req = AuthorizationRequest {
authorization_url: auth_url,
state: state.clone(),
request_uri: par_res.request_uri,
expires_in: par_res.expires_in,
stored_state: stored_state.clone(),
};
Ok((auth_req, stored_state))
}
pub async fn exchange_code(
&self,
code: &str,
state_entry: &StoredStateEntry,
) -> Result<OAuthSession, AtprotoOAuthError> {
let form_pairs: Vec<(&str, &str)> = vec![
("grant_type", "authorization_code"),
("code", code),
("redirect_uri", state_entry.redirect_uri.as_str()),
("client_id", state_entry.client_id.as_str()),
("code_verifier", state_entry.code_verifier.as_str()),
];
let form_body =
serde_urlencoded::to_string(form_pairs).map_err(|e| TokenError::Http(e.to_string()))?;
let resp_json: TokenResponse = self
.send_dpop_token_request(
&state_entry.token_endpoint,
&state_entry.dpop_key,
form_body.into_bytes(),
)
.await?;
if !resp_json.token_type.eq_ignore_ascii_case("DPoP") {
return Err(TokenError::InvalidTokenType(resp_json.token_type.clone()).into());
}
if resp_json.sub.trim().is_empty() {
return Err(TokenError::MissingDid.into());
}
if let Some(ref expected_did) = state_entry.did {
if resp_json.sub != *expected_did {
return Err(TokenError::SubMismatch {
expected: expected_did.clone(),
actual: resp_json.sub.clone(),
}
.into());
}
}
let scope_str = resp_json
.scope
.as_deref()
.ok_or(TokenError::MissingScope)?
.to_string();
let has_atproto = scope_str.split_whitespace().any(|s| s == "atproto");
if !has_atproto {
return Err(TokenError::MissingAtprotoScope(scope_str.clone()).into());
}
let (sub, access_token, refresh_token, token_type, scope, expires_in) =
resp_json.into_parts();
OAuthSession::new(
sub,
access_token,
refresh_token,
token_type,
scope,
expires_in,
state_entry.dpop_key.clone(),
Some(state_entry.pds_endpoint.clone()),
Some(state_entry.issuer.clone()),
Some(state_entry.token_endpoint.clone()),
)
}
pub async fn handle_callback(
&self,
callback_params: &CallbackParams,
) -> Result<OAuthSession, AtprotoOAuthError> {
let state_entry = self
.state_store
.take_state(&callback_params.state)
.await?
.ok_or_else(|| {
TokenError::InvalidState(format!(
"State token '{}' not found, expired, or already consumed",
callback_params.state
))
})?;
self.handle_callback_with_entry(callback_params, &state_entry)
.await
}
pub async fn handle_callback_with_entry(
&self,
callback_params: &CallbackParams,
state_entry: &StoredStateEntry,
) -> Result<OAuthSession, AtprotoOAuthError> {
if callback_params.state != state_entry.state {
return Err(TokenError::InvalidState(format!(
"Callback state '{}' does not match expected state '{}'",
callback_params.state, state_entry.state
))
.into());
}
if state_entry.is_expired() {
return Err(TokenError::StateExpired.into());
}
let callback_iss = callback_params
.iss
.as_deref()
.ok_or(TokenError::MissingCallbackIssuer)?;
let norm_callback = callback_iss.trim().trim_end_matches('/');
let norm_expected = state_entry.issuer.trim().trim_end_matches('/');
if norm_callback != norm_expected {
return Err(TokenError::IssuerMismatch {
expected: state_entry.issuer.clone(),
actual: callback_iss.to_string(),
}
.into());
}
self.exchange_code(&callback_params.code, state_entry).await
}
pub async fn refresh_session(
&self,
session: &mut OAuthSession,
) -> Result<(), AtprotoOAuthError> {
let refresh_lock = self.refresh_single_flight.lock_for(session.sub());
let _refresh_guard = refresh_lock.lock().await;
let refresh_token = session
.refresh_token()
.ok_or(TokenError::MissingRefreshToken)?;
let token_endpoint = session
.token_endpoint()
.ok_or(TokenError::MissingField("token_endpoint"))?
.to_string();
let form_pairs: Vec<(&str, &str)> = vec![
("grant_type", "refresh_token"),
("refresh_token", refresh_token),
("client_id", self.metadata.client_id.as_str()),
];
let form_body =
serde_urlencoded::to_string(form_pairs).map_err(|e| TokenError::Http(e.to_string()))?;
let resp_json: TokenResponse = self
.send_dpop_token_request(&token_endpoint, session.dpop_key(), form_body.into_bytes())
.await?;
if !resp_json.token_type.eq_ignore_ascii_case("DPoP") {
return Err(TokenError::InvalidTokenType(resp_json.token_type.clone()).into());
}
if resp_json.sub.is_empty() {
return Err(TokenError::RequestFailed {
status: 200,
error: "invalid_request".to_string(),
description: Some(
"Refresh response is missing mandatory `sub` claim (review H4: fail-closed)"
.to_string(),
),
}
.into());
}
if resp_json.sub != session.sub() {
return Err(TokenError::SubMismatch {
expected: session.sub().to_string(),
actual: resp_json.sub.clone(),
}
.into());
}
let granted_scope = session.scope().unwrap_or("").to_string();
if let Some(ref new_scope) = resp_json.scope {
if !new_scope.split_whitespace().any(|s| s == "atproto") {
return Err(TokenError::MissingAtprotoScope(new_scope.clone()).into());
}
if !granted_scope.is_empty() {
let granted: std::collections::HashSet<&str> =
granted_scope.split_whitespace().collect();
let expanded: Vec<&str> = new_scope
.split_whitespace()
.filter(|s| !granted.contains(*s))
.collect();
if !expanded.is_empty() {
return Err(TokenError::ScopeExpansion {
granted: granted_scope,
requested: new_scope.clone(),
}
.into());
}
}
}
let (_sub, access_token, refresh_token, _token_type, scope, expires_in) =
resp_json.into_parts();
session.rotate_tokens_with_scope(access_token, refresh_token, expires_in, scope);
Ok(())
}
pub async fn refresh_token(
&self,
session: &OAuthSession,
) -> Result<OAuthSession, AtprotoOAuthError> {
let mut cloned = session.clone();
self.refresh_session(&mut cloned).await?;
Ok(cloned)
}
async fn send_dpop_token_request(
&self,
token_endpoint: &str,
dpop_key: &DPoPKey,
body_bytes: Vec<u8>,
) -> Result<TokenResponse, AtprotoOAuthError> {
send_dpop_token_request_inner(
&self.ssrf_filter,
&self.nonce_cache,
token_endpoint,
dpop_key,
body_bytes,
)
.await
}
pub async fn send_dpop_request(
&self,
dpop_key: &DPoPKey,
method: reqwest::Method,
url_str: &str,
access_token: Option<&str>,
body_bytes: Option<Vec<u8>>,
content_type: Option<&str>,
) -> Result<reqwest::Response, AtprotoOAuthError> {
let parsed_url = Url::parse(url_str)
.map_err(|e| TokenError::Http(format!("Invalid URL '{url_str}': {e}")))?;
let (client, _pinned_addr, host_header) = self
.ssrf_filter
.build_pinned_client(&parsed_url)
.await
.map_err(TokenError::from)?;
let server_origin = parsed_url.origin().ascii_serialization();
let ath = access_token.map(compute_access_token_hash);
let initial_nonce = self.nonce_cache.get_nonce(&server_origin);
let proof = dpop_key.create_proof(
method.as_str(),
url_str,
initial_nonce.as_deref(),
ath.as_deref(),
)?;
let mut req = client
.request(method.clone(), url_str)
.header(reqwest::header::HOST, host_header.clone())
.header("dpop", proof);
if let Some(token) = access_token {
req = req.header("authorization", format!("DPoP {token}"));
}
if let Some(ct) = content_type {
req = req.header("content-type", ct);
}
if let Some(ref bytes) = body_bytes {
req = req.body(bytes.clone());
}
let resp = req
.send()
.await
.map_err(|e| TokenError::Http(e.to_string()))?;
if let Some(new_nonce) = extract_dpop_nonce(
resp.headers()
.get("dpop-nonce")
.and_then(|h| h.to_str().ok()),
) {
self.nonce_cache.set_nonce(&server_origin, new_nonce);
}
let status = resp.status();
if status == reqwest::StatusCode::BAD_REQUEST || status == reqwest::StatusCode::UNAUTHORIZED
{
let resp_headers = resp.headers().clone();
if !is_rs_dpop_nonce_challenge(&resp_headers) {
let bytes = read_bounded_body(resp, MAX_OAUTH_RESPONSE_BYTES)
.await
.map_err(|e| TokenError::Http(e.to_string()))?;
let body_is_challenge =
is_use_dpop_nonce_error(serde_json::from_slice(&bytes).ok().as_ref());
if !body_is_challenge {
let mut builder = http::Response::builder().status(status);
for (name, value) in resp_headers.iter() {
builder = builder.header(name.clone(), value.clone());
}
let rebuilt = builder.body(bytes).map_err(|e| {
TokenError::Http(format!("Failed to rebuild response: {e}"))
})?;
return Ok(reqwest::Response::from(rebuilt));
}
}
{
let fresh_nonce = self.nonce_cache.get_nonce(&server_origin).ok_or_else(|| {
TokenError::RequestFailed {
status: status.as_u16(),
error: "use_dpop_nonce".to_string(),
description: Some("Missing DPoP-Nonce header".to_string()),
}
})?;
let retry_proof = dpop_key.create_proof(
method.as_str(),
url_str,
Some(&fresh_nonce),
ath.as_deref(),
)?;
let (retry_client, _retry_pinned_addr, retry_host_header) = self
.ssrf_filter
.build_pinned_client(&parsed_url)
.await
.map_err(TokenError::from)?;
let mut retry_req = retry_client
.request(method, url_str)
.header(reqwest::header::HOST, retry_host_header)
.header("dpop", retry_proof);
if let Some(token) = access_token {
retry_req = retry_req.header("authorization", format!("DPoP {token}"));
}
if let Some(ct) = content_type {
retry_req = retry_req.header("content-type", ct);
}
if let Some(bytes) = body_bytes {
retry_req = retry_req.body(bytes);
}
let retry_resp = retry_req
.send()
.await
.map_err(|e| TokenError::Http(e.to_string()))?;
if let Some(new_nonce) = extract_dpop_nonce(
retry_resp
.headers()
.get("dpop-nonce")
.and_then(|h| h.to_str().ok()),
) {
self.nonce_cache.set_nonce(&server_origin, new_nonce);
}
let retry_status = retry_resp.status();
if retry_status == reqwest::StatusCode::BAD_REQUEST
|| retry_status == reqwest::StatusCode::UNAUTHORIZED
{
let retry_headers = retry_resp.headers().clone();
if is_rs_dpop_nonce_challenge(&retry_headers) {
return Err(DPoPError::NonceRetryLimitExceeded.into());
}
let bytes = read_bounded_body(retry_resp, MAX_OAUTH_RESPONSE_BYTES)
.await
.map_err(|e| TokenError::Http(e.to_string()))?;
if is_use_dpop_nonce_error(serde_json::from_slice(&bytes).ok().as_ref()) {
return Err(DPoPError::NonceRetryLimitExceeded.into());
}
let mut builder = http::Response::builder().status(retry_status);
for (name, value) in retry_headers.iter() {
builder = builder.header(name.clone(), value.clone());
}
let rebuilt = builder.body(bytes).map_err(|e| {
TokenError::Http(format!("Failed to rebuild response: {e}"))
})?;
return Ok(reqwest::Response::from(rebuilt));
}
if retry_status.is_success() {
crate::dpop::require_dpop_nonce(retry_resp.headers())
.map_err(AtprotoOAuthError::from)?;
}
return Ok(retry_resp);
}
}
crate::dpop::require_dpop_nonce(resp.headers()).map_err(AtprotoOAuthError::from)?;
Ok(resp)
}
fn validate_xrpc_nsid(nsid: &str) -> Result<(), TokenError> {
if crate::kernels::nsid_bytes::is_valid_nsid(nsid) {
Ok(())
} else {
Err(TokenError::InvalidNsid(nsid.to_string()))
}
}
pub async fn send_xrpc_request(
&self,
session: &OAuthSession,
nsid: &str,
query_params: &[(&str, &str)],
) -> Result<reqwest::Response, AtprotoOAuthError> {
Self::validate_xrpc_nsid(nsid).map_err(AtprotoOAuthError::Token)?;
let pds_endpoint = session
.pds_endpoint()
.ok_or(TokenError::MissingField("pds_endpoint"))?;
let mut url = Url::parse(pds_endpoint)
.map_err(|e| TokenError::Http(format!("Invalid PDS endpoint: {e}")))?;
let trimmed_nsid = nsid.trim_start_matches('/');
let base_path = url.path().trim_end_matches('/');
url.set_path(&format!("{}/xrpc/{}", base_path, trimmed_nsid));
if !query_params.is_empty() {
url.query_pairs_mut()
.extend_pairs(query_params.iter().copied());
}
self.send_dpop_request(
session.dpop_key(),
reqwest::Method::GET,
url.as_str(),
Some(session.access_token()),
None,
None,
)
.await
}
}
fn parse_oauth_error_fields(json: Option<&serde_json::Value>) -> (String, Option<String>) {
let error_code = json
.and_then(|j| j.get("error"))
.and_then(|e| e.as_str())
.unwrap_or("token_request_failed")
.to_string();
let error_desc = json
.and_then(|j| j.get("error_description"))
.and_then(|d| d.as_str())
.map(ToString::to_string);
(error_code, error_desc)
}
fn is_use_dpop_nonce_error(json: Option<&serde_json::Value>) -> bool {
json.and_then(|j| j.get("error")).and_then(|e| e.as_str()) == Some("use_dpop_nonce")
}
fn is_rs_dpop_nonce_challenge(headers: &reqwest::header::HeaderMap) -> bool {
headers
.get_all(reqwest::header::WWW_AUTHENTICATE)
.iter()
.filter_map(|v| v.to_str().ok())
.any(|challenge| {
let challenge = challenge.trim();
let Some(space) = challenge.find(' ') else {
return false;
};
let (scheme, rest) = challenge.split_at(space);
if !scheme.eq_ignore_ascii_case("DPoP") {
return false;
}
let rest = rest.trim_start();
rest.split(',').any(|param| {
let param = param.trim();
let Some((key, value)) = param.split_once('=') else {
return false;
};
let key = key.trim();
if !key.eq_ignore_ascii_case("error") {
return false;
}
let value = value.trim().trim_matches('"');
value.eq_ignore_ascii_case("use_dpop_nonce")
})
})
}
async fn read_bounded_json(
resp: reqwest::Response,
) -> Result<Option<serde_json::Value>, TokenError> {
let bytes = read_bounded_body(resp, MAX_OAUTH_RESPONSE_BYTES)
.await
.map_err(|e| TokenError::Http(e.to_string()))?;
Ok(serde_json::from_slice(&bytes).ok())
}
async fn send_dpop_token_request_inner(
ssrf_filter: &SsrfFilter,
nonce_cache: &DPoPNonceCache,
token_endpoint: &str,
dpop_key: &DPoPKey,
body_bytes: Vec<u8>,
) -> Result<TokenResponse, AtprotoOAuthError> {
let parsed_url = Url::parse(token_endpoint).map_err(|e| {
TokenError::Http(format!(
"Invalid token endpoint URL '{token_endpoint}': {e}"
))
})?;
let (client, _pinned_addr, host_header) = ssrf_filter
.build_pinned_client(&parsed_url)
.await
.map_err(TokenError::from)?;
let server_origin = parsed_url.origin().ascii_serialization();
let initial_nonce = nonce_cache.get_nonce(&server_origin);
let proof = dpop_key.create_proof("POST", token_endpoint, initial_nonce.as_deref(), None)?;
let resp = client
.post(token_endpoint)
.header(reqwest::header::HOST, host_header.clone())
.header("content-type", "application/x-www-form-urlencoded")
.header("accept", "application/json")
.header("dpop", proof)
.body(body_bytes.clone())
.send()
.await
.map_err(|e| TokenError::Http(e.to_string()))?;
if let Some(new_nonce) = extract_dpop_nonce(
resp.headers()
.get("dpop-nonce")
.and_then(|h| h.to_str().ok()),
) {
nonce_cache.set_nonce(&server_origin, new_nonce);
}
let status = resp.status();
if status.is_redirection() {
return Err(TokenError::RequestFailed {
status: status.as_u16(),
error: "invalid_request".to_string(),
description: Some("Redirects are not permitted for token endpoints".to_string()),
}
.into());
}
if status == reqwest::StatusCode::BAD_REQUEST || status == reqwest::StatusCode::UNAUTHORIZED {
let json_err = read_bounded_json(resp).await?;
if is_use_dpop_nonce_error(json_err.as_ref()) {
let fresh_nonce =
nonce_cache
.get_nonce(&server_origin)
.ok_or_else(|| TokenError::RequestFailed {
status: status.as_u16(),
error: "use_dpop_nonce".to_string(),
description: Some(
"Missing DPoP-Nonce header in challenge response".to_string(),
),
})?;
let retry_proof =
dpop_key.create_proof("POST", token_endpoint, Some(&fresh_nonce), None)?;
let (retry_client, _retry_pinned_addr, retry_host_header) = ssrf_filter
.build_pinned_client(&parsed_url)
.await
.map_err(TokenError::from)?;
let retry_resp = retry_client
.post(token_endpoint)
.header(reqwest::header::HOST, retry_host_header)
.header("content-type", "application/x-www-form-urlencoded")
.header("accept", "application/json")
.header("dpop", retry_proof)
.body(body_bytes)
.send()
.await
.map_err(|e| TokenError::Http(e.to_string()))?;
if let Some(new_nonce) = extract_dpop_nonce(
retry_resp
.headers()
.get("dpop-nonce")
.and_then(|h| h.to_str().ok()),
) {
nonce_cache.set_nonce(&server_origin, new_nonce);
}
let retry_status = retry_resp.status();
if retry_status.is_redirection() {
return Err(TokenError::RequestFailed {
status: retry_status.as_u16(),
error: "invalid_request".to_string(),
description: Some(
"Redirects are not permitted for token endpoints".to_string(),
),
}
.into());
}
if retry_status.is_success() {
crate::dpop::require_dpop_nonce(retry_resp.headers()).map_err(TokenError::from)?;
let bytes = read_bounded_body(retry_resp, MAX_OAUTH_RESPONSE_BYTES)
.await
.map_err(|e| TokenError::Http(e.to_string()))?;
let res: TokenResponse =
serde_json::from_slice(&bytes).map_err(|e| TokenError::Json(e.to_string()))?;
return Ok(res);
}
let err_json = read_bounded_json(retry_resp).await?;
if is_use_dpop_nonce_error(err_json.as_ref()) {
return Err(DPoPError::NonceRetryLimitExceeded.into());
}
let (error_code, error_desc) = parse_oauth_error_fields(err_json.as_ref());
return Err(TokenError::RequestFailed {
status: retry_status.as_u16(),
error: error_code,
description: error_desc,
}
.into());
}
let (error_code, error_desc) = parse_oauth_error_fields(json_err.as_ref());
return Err(TokenError::RequestFailed {
status: status.as_u16(),
error: error_code,
description: error_desc,
}
.into());
}
if !status.is_success() {
let err_json = read_bounded_json(resp).await?;
let (error_code, error_desc) = parse_oauth_error_fields(err_json.as_ref());
return Err(TokenError::RequestFailed {
status: status.as_u16(),
error: error_code,
description: error_desc,
}
.into());
}
crate::dpop::require_dpop_nonce(resp.headers()).map_err(TokenError::from)?;
let bytes = read_bounded_body(resp, MAX_OAUTH_RESPONSE_BYTES)
.await
.map_err(|e| TokenError::Http(e.to_string()))?;
let res: TokenResponse =
serde_json::from_slice(&bytes).map_err(|e| TokenError::Json(e.to_string()))?;
Ok(res)
}
#[cfg(test)]
#[allow(clippy::unwrap_used, clippy::expect_used, clippy::panic, missing_docs)]
mod tests {
use super::*;
use reqwest::header::{HeaderMap, HeaderValue};
#[test]
fn test_rs_dpop_nonce_challenge_parsing() {
let mut h = HeaderMap::new();
h.insert(
reqwest::header::WWW_AUTHENTICATE,
HeaderValue::from_static("DPoP algs=\"ES256\", error=\"use_dpop_nonce\""),
);
assert!(is_rs_dpop_nonce_challenge(&h));
let mut h2 = HeaderMap::new();
h2.insert(
reqwest::header::WWW_AUTHENTICATE,
HeaderValue::from_static("dpop error=\"USE_DPOP_NONCE\""),
);
assert!(is_rs_dpop_nonce_challenge(&h2));
let mut h3 = HeaderMap::new();
h3.insert(
reqwest::header::WWW_AUTHENTICATE,
HeaderValue::from_static("Bearer realm=\"x\""),
);
h3.append(
reqwest::header::WWW_AUTHENTICATE,
HeaderValue::from_static("DPoP error=\"use_dpop_nonce\""),
);
assert!(is_rs_dpop_nonce_challenge(&h3));
let mut h4 = HeaderMap::new();
h4.insert(
reqwest::header::WWW_AUTHENTICATE,
HeaderValue::from_static("Bearer error=\"use_dpop_nonce\""),
);
assert!(!is_rs_dpop_nonce_challenge(&h4));
let mut h5 = HeaderMap::new();
h5.insert(
reqwest::header::WWW_AUTHENTICATE,
HeaderValue::from_static("DPoP error=\"invalid_token\""),
);
assert!(!is_rs_dpop_nonce_challenge(&h5));
let mut h6 = HeaderMap::new();
h6.insert(
reqwest::header::WWW_AUTHENTICATE,
HeaderValue::from_static("DPoP algs=\"ES256\""),
);
assert!(!is_rs_dpop_nonce_challenge(&h6));
assert!(!is_rs_dpop_nonce_challenge(&HeaderMap::new()));
}
#[test]
fn test_client_builder_and_metadata() {
let client = AtprotoOAuthClient::builder()
.client_metadata(
OAuthClientMetadata::new(
"https://app.example.com/client.json",
"https://app.example.com/callback",
)
.with_scope("atproto transition:generic")
.with_client_name("Example App"),
)
.allow_insecure_localhost(true)
.build()
.unwrap();
assert_eq!(
client.metadata().client_id,
"https://app.example.com/client.json"
);
assert_eq!(
client.metadata().redirect_uri,
"https://app.example.com/callback"
);
assert_eq!(client.metadata().scope, "atproto transition:generic");
assert_eq!(
client.metadata().client_name.as_deref(),
Some("Example App")
);
}
#[test]
fn test_stored_state_expiration() {
let entry = StoredStateEntry {
state: "state123".to_string(),
client_id: "client123".to_string(),
code_verifier: "verifier123".to_string(),
dpop_key: DPoPKey::generate(),
issuer: "https://auth.example.com".to_string(),
did: Some("did:plc:alice".to_string()),
handle: Some("alice.bsky.social".to_string()),
redirect_uri: "https://app.example.com/callback".to_string(),
pds_endpoint: "https://pds.example.com".to_string(),
token_endpoint: "https://auth.example.com/oauth/token".to_string(),
scopes: "atproto".to_string(),
created_at: SystemTime::now() - Duration::from_secs(400),
expires_in_secs: 300,
};
assert!(entry.is_expired());
}
#[tokio::test]
async fn test_handle_callback_with_expired_entry_rejected() {
let client = AtprotoOAuthClient::builder()
.client_metadata(OAuthClientMetadata::new(
"https://app.example.com/client.json",
"https://app.example.com/callback",
))
.allow_insecure_localhost(true)
.build()
.unwrap();
let expired_entry = StoredStateEntry {
state: "state_exp".to_string(),
client_id: "https://app.example.com/client.json".to_string(),
code_verifier: "verifier123".to_string(),
dpop_key: DPoPKey::generate(),
issuer: "https://auth.example.com".to_string(),
did: Some("did:plc:alice".to_string()),
handle: Some("alice.bsky.social".to_string()),
redirect_uri: "https://app.example.com/callback".to_string(),
pds_endpoint: "https://pds.example.com".to_string(),
token_endpoint: "https://auth.example.com/oauth/token".to_string(),
scopes: "atproto".to_string(),
created_at: SystemTime::now() - Duration::from_secs(400),
expires_in_secs: 300,
};
let cb = CallbackParams::new("code123", "state_exp").with_iss("https://auth.example.com");
let err = client.handle_callback_with_entry(&cb, &expired_entry).await;
assert!(matches!(
err,
Err(AtprotoOAuthError::Token(TokenError::StateExpired))
));
}
#[test]
fn test_callback_params() {
let cb = CallbackParams::new("code123", "state123").with_iss("https://auth.example.com");
assert_eq!(cb.code, "code123");
assert_eq!(cb.state, "state123");
assert_eq!(cb.iss.as_deref(), Some("https://auth.example.com"));
}
#[tokio::test]
async fn test_send_xrpc_request_missing_pds_endpoint() {
let client = AtprotoOAuthClient::builder()
.client_metadata(OAuthClientMetadata::new(
"https://app.example.com/client.json",
"https://app.example.com/callback",
))
.allow_insecure_localhost(true)
.build()
.unwrap();
let session = OAuthSession::new(
"did:plc:alice123",
"at_123",
None,
"DPoP",
None,
None,
DPoPKey::generate(),
None,
None,
None,
)
.unwrap();
let err = client
.send_xrpc_request(&session, "com.atproto.repo.describeRepo", &[])
.await;
assert!(matches!(
err,
Err(AtprotoOAuthError::Token(TokenError::MissingField(
"pds_endpoint"
)))
));
}
#[test]
fn test_validate_xrpc_nsid_accepts_valid_grammar() {
for nsid in [
"com.atproto.repo.describeRepo",
"/app.bsky.feed.getTimeline",
"a.b.c",
"one.2.three",
"one.two.three.four-and.FiVe",
"a-0.b-1.c",
"com.example.fooBar",
"com.example.fooBarV2",
"m.xn--masekowski-d0b.pl",
] {
assert!(
AtprotoOAuthClient::validate_xrpc_nsid(nsid).is_ok(),
"expected valid NSID: {nsid}"
);
}
}
#[test]
fn test_validate_xrpc_nsid_rejects_traversal_and_malformed() {
for nsid in [
"../../admin",
"com.atproto..describeRepo",
"com..atproto",
"com.atproto.repo.",
"com.atproto",
".one.two.three",
"1.0.0.127.record",
"0two.example.foo",
"3com.atproto.repo",
"com.atproto.-repo",
"com.atproto.repo.name-",
"com.atproto.repo.-name",
"a-0.b-1.c-3",
"a-0.b-1.c-o",
"com.example.foo.*",
"com.example.foo.blah*",
"com.atproto.re\\po",
"com.atproto.re%70o",
"",
] {
assert!(
matches!(
AtprotoOAuthClient::validate_xrpc_nsid(nsid),
Err(TokenError::InvalidNsid(_))
),
"expected InvalidNsid: {nsid:?}"
);
}
}
#[test]
fn test_validate_xrpc_nsid_length_limits() {
let seg63 = "o".repeat(63);
let seg64 = "o".repeat(64);
assert!(AtprotoOAuthClient::validate_xrpc_nsid(&format!("com.{seg63}.foo")).is_ok());
assert!(matches!(
AtprotoOAuthClient::validate_xrpc_nsid(&format!("com.{seg64}.foo")),
Err(TokenError::InvalidNsid(_))
));
let long_authority = format!("{}.{}.{}.{}", seg63, seg63, "o".repeat(62), "o".repeat(62));
let nsid_317 = format!("{long_authority}.{}", "f".repeat(63));
assert_eq!(nsid_317.len(), 317);
assert!(AtprotoOAuthClient::validate_xrpc_nsid(&nsid_317).is_ok());
let nsid_318 = format!("{nsid_317}x");
assert_eq!(nsid_318.len(), 318);
assert!(matches!(
AtprotoOAuthClient::validate_xrpc_nsid(&nsid_318),
Err(TokenError::InvalidNsid(_))
));
}
}