use serde::{Deserialize, Serialize};
use crate::client::CallbackParams;
use crate::error::IntegrationError;
use crate::session::OAuthSession;
#[cfg(feature = "axum")]
pub mod axum;
#[cfg(feature = "actix")]
pub mod actix;
#[cfg(feature = "tower")]
pub mod tower;
pub mod validator;
pub use validator::{
AccessTokenValidator, CnfClaim, InMemoryTokenValidator, JwtAccessTokenClaims,
JwtAccessTokenValidator, RegisteredToken,
};
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct OAuthCallbackQuery {
pub code: Option<String>,
pub state: Option<String>,
pub iss: Option<String>,
pub error: Option<String>,
pub error_description: Option<String>,
}
impl OAuthCallbackQuery {
#[must_use]
pub fn new(code: impl Into<String>, state: impl Into<String>) -> Self {
Self {
code: Some(code.into()),
state: Some(state.into()),
iss: None,
error: None,
error_description: None,
}
}
#[must_use]
pub fn new_error(error: impl Into<String>, description: Option<String>) -> Self {
Self {
code: None,
state: None,
iss: None,
error: Some(error.into()),
error_description: description,
}
}
#[must_use]
pub fn with_iss(mut self, iss: impl Into<String>) -> Self {
self.iss = Some(iss.into());
self
}
pub fn to_callback_params(&self) -> Result<CallbackParams, IntegrationError> {
if let Some(err) = &self.error {
return Err(IntegrationError::OAuthError {
error: err.clone(),
description: self.error_description.clone().unwrap_or_default(),
});
}
let code = self
.code
.as_ref()
.ok_or(IntegrationError::MissingCode)?
.clone();
let state = self
.state
.as_ref()
.ok_or(IntegrationError::MissingState)?
.clone();
let mut params = CallbackParams::new(code, state);
if let Some(iss) = &self.iss {
params = params.with_iss(iss.clone());
}
Ok(params)
}
}
use zeroize::{Zeroize, ZeroizeOnDrop};
#[derive(Clone, PartialEq, Eq, Serialize)]
pub struct AuthenticatedUser {
pub did: String,
#[serde(skip_serializing)]
pub access_token: String,
pub dpop_thumbprint: String,
pub scope: Option<String>,
}
impl<'de> Deserialize<'de> for AuthenticatedUser {
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
where
D: serde::Deserializer<'de>,
{
#[derive(Deserialize)]
struct RawUser {
did: String,
#[serde(default)]
access_token: Option<String>,
dpop_thumbprint: String,
scope: Option<String>,
}
let raw = RawUser::deserialize(deserializer)?;
let access_token = raw.access_token.ok_or_else(|| {
serde::de::Error::custom(
"AuthenticatedUser deserialization requires a non-empty access_token; serialized views omit the token by design, so re-authenticate instead",
)
})?;
if access_token.trim().is_empty() {
return Err(serde::de::Error::custom(
"AuthenticatedUser deserialized with an empty access_token; tokens must be supplied by the validator, not from untrusted data",
));
}
Ok(Self {
did: raw.did,
access_token,
dpop_thumbprint: raw.dpop_thumbprint,
scope: raw.scope,
})
}
}
impl std::fmt::Debug for AuthenticatedUser {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("AuthenticatedUser")
.field("did", &self.did)
.field("access_token", &"[REDACTED]")
.field("dpop_thumbprint", &self.dpop_thumbprint)
.field("scope", &self.scope)
.finish()
}
}
impl Zeroize for AuthenticatedUser {
fn zeroize(&mut self) {
self.access_token.zeroize();
}
}
impl Drop for AuthenticatedUser {
fn drop(&mut self) {
self.zeroize();
}
}
impl ZeroizeOnDrop for AuthenticatedUser {}
impl AuthenticatedUser {
#[must_use]
pub fn new(
did: impl Into<String>,
access_token: impl Into<String>,
dpop_thumbprint: impl Into<String>,
) -> Self {
Self {
did: did.into(),
access_token: access_token.into(),
dpop_thumbprint: dpop_thumbprint.into(),
scope: None,
}
}
#[must_use]
pub fn with_scope(mut self, scope: impl Into<String>) -> Self {
self.scope = Some(scope.into());
self
}
#[must_use]
pub fn did(&self) -> &str {
&self.did
}
#[must_use]
pub fn access_token(&self) -> &str {
&self.access_token
}
#[must_use]
pub fn dpop_thumbprint(&self) -> &str {
&self.dpop_thumbprint
}
#[must_use]
pub fn scope(&self) -> Option<&str> {
self.scope.as_deref()
}
#[must_use]
pub fn into_parts(self) -> (String, String, String, Option<String>) {
let mut this = self;
let did = std::mem::take(&mut this.did);
let access_token = std::mem::take(&mut this.access_token);
let dpop_thumbprint = std::mem::take(&mut this.dpop_thumbprint);
let scope = this.scope.take();
(did, access_token, dpop_thumbprint, scope)
}
#[must_use]
pub fn into_did(mut self) -> String {
std::mem::take(&mut self.did)
}
#[must_use]
pub fn into_access_token(mut self) -> String {
std::mem::take(&mut self.access_token)
}
#[must_use]
pub fn into_dpop_thumbprint(mut self) -> String {
std::mem::take(&mut self.dpop_thumbprint)
}
}
#[derive(Debug, Clone)]
pub struct OAuthSessionExtension {
pub user: AuthenticatedUser,
pub session: Option<OAuthSession>,
}
impl OAuthSessionExtension {
#[must_use]
pub fn new(user: AuthenticatedUser) -> Self {
Self {
user,
session: None,
}
}
#[must_use]
pub fn from_session(session: OAuthSession) -> Self {
let user = AuthenticatedUser {
did: session.sub.clone(),
access_token: session.access_token.clone(),
dpop_thumbprint: session.dpop_key.jwk_thumbprint(),
scope: session.scope.clone(),
};
Self {
user,
session: Some(session),
}
}
}
#[cfg(test)]
#[allow(clippy::unwrap_used, clippy::expect_used, clippy::panic, missing_docs)]
mod mutation_killer_tests {
use super::*;
#[test]
fn killer_authenticated_user_accessors() {
let user = AuthenticatedUser::new("did:plc:alice123", "at_token_123", "jkt_abc")
.with_scope("atproto transition:generic");
assert_eq!(user.did(), "did:plc:alice123");
assert_eq!(user.access_token(), "at_token_123");
assert_eq!(user.dpop_thumbprint(), "jkt_abc");
assert_eq!(user.scope(), Some("atproto transition:generic"));
let no_scope = AuthenticatedUser::new("did:plc:x", "at", "jkt");
assert_eq!(no_scope.scope(), None);
}
#[test]
fn killer_authenticated_user_owned_accessors() {
let user = AuthenticatedUser::new("did:plc:alice123", "at_token_123", "jkt_abc")
.with_scope("atproto");
let (did, access_token, jkt, scope) = user.into_parts();
assert_eq!(did, "did:plc:alice123");
assert_eq!(access_token, "at_token_123");
assert_eq!(jkt, "jkt_abc");
assert_eq!(scope.as_deref(), Some("atproto"));
let user2 = AuthenticatedUser::new("did:plc:bob", "at_b", "jkt_b");
assert_eq!(user2.into_did(), "did:plc:bob");
let user3 = AuthenticatedUser::new("did:plc:carol", "at_c", "jkt_c");
assert_eq!(user3.into_access_token(), "at_c");
let user4 = AuthenticatedUser::new("did:plc:dave", "at_d", "jkt_d");
assert_eq!(user4.into_dpop_thumbprint(), "jkt_d");
let user5 = AuthenticatedUser::new("did:plc:eve", "at_e", "jkt_e");
let (_, _, _, scope5) = user5.into_parts();
assert!(scope5.is_none(), "unset scope must remain None");
}
#[test]
fn killer_authenticated_user_zeroize_on_drop() {
let mut user = AuthenticatedUser::new("did:plc:alice", "secret_token_value", "jkt");
user.zeroize();
assert!(
user.access_token.as_bytes().iter().all(|&b| b == 0),
"zeroize must scrub the access token buffer"
);
let dropped = AuthenticatedUser::new("did:plc:bob", "another_secret", "jkt2");
drop(dropped);
}
#[test]
fn killer_oauth_callback_query_builder_paths() {
let q = OAuthCallbackQuery::new_error("access_denied", None).with_iss("https://auth");
assert_eq!(q.error.as_deref(), Some("access_denied"));
assert!(q.error_description.is_none());
assert_eq!(q.iss.as_deref(), Some("https://auth"));
assert!(q.code.is_none() && q.state.is_none());
let q2 = OAuthCallbackQuery::new_error("server_error", Some("desc".to_string()));
assert_eq!(q2.error_description.as_deref(), Some("desc"));
let q3 = OAuthCallbackQuery::new("code-1", "state-1");
assert!(q3.error.is_none() && q3.error_description.is_none());
let params = q3.to_callback_params().unwrap();
assert_eq!(params.code, "code-1");
assert_eq!(params.state, "state-1");
assert!(params.iss.is_none());
}
}