#[cfg(feature = "oauth2")]
use std::collections::HashMap;
use std::future::Future;
use std::pin::Pin;
use std::sync::Arc;
use std::task::{Context, Poll};
#[cfg(feature = "oauth2")]
use std::time::Duration;
use axum::extract::FromRequestParts;
use axum::response::{IntoResponse, Response};
use http::StatusCode;
use http::request::Parts;
#[cfg(feature = "oauth2")]
use jsonwebtoken::jwk::JwkSet;
#[cfg(feature = "oauth2")]
use serde::Deserialize;
#[cfg(feature = "oauth2")]
use url::Url;
pub mod password;
pub use password::{
BreachCheck, PasswordConfig, PasswordFailure, PasswordPolicy, PasswordValidation,
validate_password,
};
pub mod remember;
pub use remember::{
DEFAULT_ROTATION_GRACE_SECS, RememberConfig, RememberCredential, RememberDecision,
RememberRecord, build_remember_clear_cookie, build_remember_cookie, constant_time_eq,
default_rotation_grace, evaluate_remember, format_remember_cookie_value,
generate_remember_credential, generate_token, hash_remember_token, parse_remember_cookie_value,
verify_remember_token,
};
const DEFAULT_BCRYPT_COST: u32 = 12;
pub async fn hash_password(password: &str) -> crate::AutumnResult<String> {
let password = password.to_string();
tokio::task::spawn_blocking(move || {
bcrypt::hash(password, DEFAULT_BCRYPT_COST)
.map_err(|e| crate::AutumnError::from(std::io::Error::other(e.to_string())))
})
.await
.map_err(|e| crate::AutumnError::from(std::io::Error::other(e.to_string())))?
}
pub async fn verify_password(password: &str, hash: &str) -> crate::AutumnResult<bool> {
let password = password.to_string();
let is_valid_format = hash.len() == 60 && hash.starts_with('$');
let hash_to_verify = if is_valid_format {
hash.to_string()
} else {
"$2b$12$KIXe8K4j1sH6/xH.x9d71uJ5Jk8t6O4m6Q110g4H8y1r6J6O6O6O6".to_string()
};
let result = tokio::task::spawn_blocking(move || bcrypt::verify(&password, &hash_to_verify))
.await
.map_err(|e| crate::AutumnError::from(std::io::Error::other(e.to_string())))?;
if !is_valid_format {
return Ok(false);
}
result.map_err(|e| crate::AutumnError::from(std::io::Error::other(e.to_string())))
}
#[doc(hidden)]
pub async fn __check_secured(
session: &crate::session::Session,
roles: &[&str],
) -> crate::AutumnResult<()> {
__check_secured_with_key(session, "user_id", roles).await
}
#[doc(hidden)]
pub async fn __check_secured_with_key(
session: &crate::session::Session,
auth_session_key: &str,
roles: &[&str],
) -> crate::AutumnResult<()> {
let Some(user_id) = session.get(auth_session_key).await else {
return Err(crate::AutumnError::unauthorized_msg(
"authentication required",
));
};
if crate::current::Current::actor().is_none() {
crate::current::Current::set_actor(user_id.clone());
}
crate::log::context::set_user_id(user_id);
if !roles.is_empty() {
let user_role = session.get("role").await.unwrap_or_default();
if !roles.iter().any(|&r| r == user_role) {
return Err(crate::AutumnError::forbidden_msg(
"insufficient permissions",
));
}
}
Ok(())
}
#[doc(hidden)]
#[allow(clippy::unused_async)]
pub async fn __check_secured_scopes(
granted: Option<&ApiTokenScopes>,
required_scopes: &[&str],
) -> crate::AutumnResult<()> {
if required_scopes.is_empty() {
return Ok(());
}
let granted: &[String] = granted.map_or(&[], |g| g.0.as_slice());
if required_scopes
.iter()
.all(|req| granted.iter().any(|g| g == req))
{
Ok(())
} else {
Err(crate::AutumnError::forbidden_msg("insufficient scope"))
}
}
pub struct Auth<T>(pub T);
impl<T, S> FromRequestParts<S> for Auth<T>
where
T: Clone + Send + Sync + 'static,
S: Send + Sync,
{
type Rejection = AuthRejection;
fn from_request_parts(
parts: &mut Parts,
_state: &S,
) -> impl Future<Output = Result<Self, Self::Rejection>> + Send {
let user = parts.extensions.get::<T>().cloned();
async move { user.map_or_else(|| Err(AuthRejection), |user| Ok(Self(user))) }
}
}
#[derive(Debug)]
pub struct AuthRejection;
impl IntoResponse for AuthRejection {
fn into_response(self) -> Response {
crate::AutumnError::unauthorized_msg("authentication required").into_response()
}
}
impl std::fmt::Display for AuthRejection {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str("authentication required")
}
}
#[derive(Clone)]
pub struct RequireAuth {
session_key: Arc<str>,
}
impl RequireAuth {
pub fn new(session_key: impl Into<String>) -> Self {
Self {
session_key: Arc::from(session_key.into()),
}
}
}
impl<S> tower::Layer<S> for RequireAuth {
type Service = RequireAuthService<S>;
fn layer(&self, inner: S) -> Self::Service {
RequireAuthService {
inner,
session_key: Arc::clone(&self.session_key),
}
}
}
#[derive(Clone)]
pub struct RequireAuthService<S> {
inner: S,
session_key: Arc<str>,
}
impl<S, ResBody> tower::Service<axum::extract::Request> for RequireAuthService<S>
where
S: tower::Service<axum::extract::Request, Response = Response<ResBody>>
+ Clone
+ Send
+ 'static,
S::Future: Send + 'static,
S::Error: Send + 'static,
ResBody: From<String> + Default + Send + 'static,
{
type Response = Response<ResBody>;
type Error = S::Error;
type Future = Pin<Box<dyn Future<Output = Result<Self::Response, Self::Error>> + Send>>;
fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
self.inner.poll_ready(cx)
}
fn call(&mut self, mut req: axum::extract::Request) -> Self::Future {
let session_key = Arc::clone(&self.session_key);
let mut inner = self.inner.clone();
std::mem::swap(&mut self.inner, &mut inner);
Box::pin(async move {
let session = req.extensions().get::<crate::session::Session>().cloned();
let user_id = if let Some(ref session) = session {
session.get(&session_key).await
} else {
None
};
if let Some(user_id) = user_id {
req.extensions_mut()
.insert(crate::security::RateLimitPrincipal(user_id.clone()));
if crate::current::Current::actor().is_none() {
crate::current::Current::set_actor(user_id.clone());
}
crate::log::context::set_user_id(user_id);
inner.call(req).await
} else {
let body = crate::error::problem_details_json_string(
StatusCode::UNAUTHORIZED,
"authentication required",
None,
None,
req.extensions()
.get::<crate::middleware::RequestId>()
.map(std::string::ToString::to_string),
Some(req.uri().path().to_owned()),
true,
);
let response = Response::builder()
.status(StatusCode::UNAUTHORIZED)
.header(http::header::CONTENT_TYPE, "application/problem+json")
.body(ResBody::from(body))
.unwrap_or_default();
Ok(response)
}
})
}
}
#[derive(Debug, Clone, serde::Deserialize)]
pub struct AuthConfig {
#[serde(default = "default_bcrypt_cost")]
pub bcrypt_cost: u32,
#[serde(default = "default_session_key")]
pub session_key: String,
#[cfg(feature = "oauth2")]
#[serde(default)]
pub oauth2: OAuth2Config,
#[cfg(feature = "oauth2")]
#[serde(default)]
pub oauth_linking_policy: OAuthLinkingPolicy,
#[cfg(feature = "webauthn")]
#[serde(default)]
pub webauthn: WebAuthnConfig,
#[serde(default)]
pub lockout: LockoutConfig,
#[serde(default)]
pub step_up: StepUpConfig,
#[serde(default)]
pub sessions: SessionTrackingConfig,
#[serde(default)]
pub password: PasswordConfig,
#[serde(default)]
pub remember: RememberConfig,
#[serde(default)]
pub magic_link: MagicLinkConfig,
}
#[derive(Debug, Clone, serde::Deserialize)]
pub struct LockoutConfig {
#[serde(default = "default_lockout_enabled")]
pub enabled: bool,
#[serde(default = "default_lockout_threshold")]
pub threshold: i32,
#[serde(default = "default_lockout_window_secs")]
pub window_secs: u64,
#[serde(default = "default_lockout_cooloff_secs")]
pub cooloff_secs: u64,
}
const fn default_lockout_enabled() -> bool {
true
}
const fn default_lockout_threshold() -> i32 {
10
}
const fn default_lockout_window_secs() -> u64 {
60
}
const fn default_lockout_cooloff_secs() -> u64 {
900
}
impl Default for LockoutConfig {
fn default() -> Self {
Self {
enabled: default_lockout_enabled(),
threshold: default_lockout_threshold(),
window_secs: default_lockout_window_secs(),
cooloff_secs: default_lockout_cooloff_secs(),
}
}
}
#[derive(Debug, Clone, serde::Deserialize)]
pub struct StepUpConfig {
#[serde(default = "default_step_up_max_age_secs")]
pub default_max_age_secs: u64,
}
const fn default_step_up_max_age_secs() -> u64 {
crate::step_up::DEFAULT_MAX_AGE_SECS
}
impl Default for StepUpConfig {
fn default() -> Self {
Self {
default_max_age_secs: crate::step_up::DEFAULT_MAX_AGE_SECS,
}
}
}
#[derive(Debug, Clone, Copy, serde::Deserialize)]
pub struct SessionTrackingConfig {
#[serde(default = "default_true_flag")]
pub revoke_on_credential_change: bool,
#[serde(default = "default_last_seen_update_secs")]
pub last_seen_update_secs: u64,
}
const fn default_true_flag() -> bool {
true
}
const fn default_last_seen_update_secs() -> u64 {
60
}
impl Default for SessionTrackingConfig {
fn default() -> Self {
Self {
revoke_on_credential_change: true,
last_seen_update_secs: default_last_seen_update_secs(),
}
}
}
#[derive(Debug, Clone, Copy, serde::Deserialize)]
pub struct MagicLinkConfig {
#[serde(default = "default_magic_link_ttl_minutes")]
pub ttl_minutes: u64,
#[serde(default = "default_magic_link_email_cooldown_secs")]
pub email_cooldown_secs: u64,
}
const fn default_magic_link_ttl_minutes() -> u64 {
15
}
const fn default_magic_link_email_cooldown_secs() -> u64 {
60
}
impl Default for MagicLinkConfig {
fn default() -> Self {
Self {
ttl_minutes: default_magic_link_ttl_minutes(),
email_cooldown_secs: default_magic_link_email_cooldown_secs(),
}
}
}
#[cfg(feature = "webauthn")]
#[derive(Debug, Clone, serde::Deserialize)]
pub struct WebAuthnConfig {
#[serde(default = "default_rp_id")]
pub rp_id: String,
#[serde(default = "default_rp_name")]
pub rp_name: String,
#[serde(default = "default_rp_origin")]
pub rp_origin: String,
}
#[cfg(feature = "webauthn")]
impl Default for WebAuthnConfig {
fn default() -> Self {
Self {
rp_id: default_rp_id(),
rp_name: default_rp_name(),
rp_origin: default_rp_origin(),
}
}
}
#[cfg(feature = "webauthn")]
const fn default_rp_id() -> String {
String::new()
}
#[cfg(feature = "webauthn")]
fn default_rp_name() -> String {
"My App".to_owned()
}
#[cfg(feature = "webauthn")]
const fn default_rp_origin() -> String {
String::new()
}
const fn default_bcrypt_cost() -> u32 {
DEFAULT_BCRYPT_COST
}
fn default_session_key() -> String {
"user_id".to_owned()
}
#[cfg(feature = "oauth2")]
const fn default_provider_scope() -> String {
String::new()
}
#[cfg(feature = "oauth2")]
const OAUTH_HTTP_TIMEOUT_SECS: u64 = 15;
#[cfg(feature = "oauth2")]
#[derive(Debug, Clone, Default, serde::Deserialize)]
pub struct OAuth2Config {
#[serde(flatten)]
pub providers: HashMap<String, OAuth2ProviderConfig>,
}
#[cfg(feature = "oauth2")]
#[derive(Debug, Clone, serde::Deserialize)]
pub struct OAuth2ProviderConfig {
#[serde(default)]
pub client_id: String,
#[serde(default)]
pub client_secret: String,
#[serde(default)]
pub authorize_url: String,
#[serde(default)]
pub token_url: String,
#[serde(default)]
pub userinfo_url: Option<String>,
#[serde(default)]
pub redirect_uri: String,
#[serde(default = "default_provider_scope")]
pub scope: String,
#[serde(default)]
pub issuer: Option<String>,
#[serde(default)]
pub jwks_url: Option<String>,
#[serde(default)]
pub discovery_url: Option<String>,
}
#[cfg(feature = "oauth2")]
#[derive(Debug, Clone, PartialEq, Eq, serde::Deserialize, Default)]
#[serde(rename_all = "snake_case")]
pub enum OAuthLinkingPolicy {
#[default]
CreateAccount,
RequireLocalSignupFirst,
}
#[cfg(feature = "oauth2")]
#[must_use]
pub fn provider_preset(name: &str) -> Option<OAuth2ProviderConfig> {
match name {
"google" => Some(OAuth2ProviderConfig {
client_id: String::new(),
client_secret: String::new(),
authorize_url: "https://accounts.google.com/o/oauth2/v2/auth".into(),
token_url: "https://oauth2.googleapis.com/token".into(),
userinfo_url: Some("https://openidconnect.googleapis.com/v1/userinfo".into()),
redirect_uri: String::new(),
scope: "openid profile email".into(),
issuer: Some("https://accounts.google.com".into()),
jwks_url: Some("https://www.googleapis.com/oauth2/v3/certs".into()),
discovery_url: Some("https://accounts.google.com".into()),
}),
"github" => Some(OAuth2ProviderConfig {
client_id: String::new(),
client_secret: String::new(),
authorize_url: "https://github.com/login/oauth/authorize".into(),
token_url: "https://github.com/login/oauth/access_token".into(),
userinfo_url: Some("https://api.github.com/user".into()),
redirect_uri: String::new(),
scope: "read:user user:email".into(),
issuer: None,
jwks_url: None,
discovery_url: None,
}),
"microsoft" => Some(OAuth2ProviderConfig {
client_id: String::new(),
client_secret: String::new(),
authorize_url: "https://login.microsoftonline.com/common/oauth2/v2.0/authorize".into(),
token_url: "https://login.microsoftonline.com/common/oauth2/v2.0/token".into(),
userinfo_url: None,
redirect_uri: String::new(),
scope: "openid profile email".into(),
issuer: Some("https://login.microsoftonline.com/common/v2.0".into()),
jwks_url: Some("https://login.microsoftonline.com/common/discovery/v2.0/keys".into()),
discovery_url: Some("https://login.microsoftonline.com/common/v2.0".into()),
}),
_ => None,
}
}
#[cfg(feature = "oauth2")]
#[derive(Debug, Clone, Deserialize)]
pub struct OAuth2Callback {
pub code: String,
pub state: String,
}
#[cfg(feature = "oauth2")]
#[derive(Debug, Clone)]
pub struct OidcIdentity {
pub subject: String,
pub email: Option<String>,
pub name: Option<String>,
pub preferred_username: Option<String>,
pub raw_claims: serde_json::Value,
}
#[cfg(feature = "oauth2")]
#[derive(Debug, Deserialize)]
struct OAuth2TokenResponse {
access_token: String,
#[allow(dead_code)]
token_type: Option<String>,
id_token: Option<String>,
}
#[cfg(feature = "oauth2")]
pub async fn oauth2_authorize_url(
session: &crate::session::Session,
provider_name: &str,
provider: &OAuth2ProviderConfig,
) -> crate::AutumnResult<String> {
use base64::Engine as _;
use sha2::Digest as _;
let state = uuid::Uuid::new_v4().to_string();
let nonce = uuid::Uuid::new_v4().to_string();
let mut verifier_bytes = [0u8; 32];
getrandom::getrandom(&mut verifier_bytes).map_err(|e| {
crate::AutumnError::service_unavailable_msg(format!("pkce rng failed: {e}"))
})?;
let code_verifier = base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(verifier_bytes);
let digest = sha2::Sha256::digest(code_verifier.as_bytes());
let code_challenge = base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(digest);
session
.insert(format!("oauth2:{provider_name}:state"), state.clone())
.await;
session
.insert(format!("oauth2:{provider_name}:nonce"), nonce.clone())
.await;
session
.insert(
format!("oauth2:{provider_name}:code_verifier"),
code_verifier,
)
.await;
let mut url = Url::parse(&provider.authorize_url)
.map_err(|e| crate::AutumnError::bad_request_msg(format!("invalid authorize_url: {e}")))?;
{
let mut q = url.query_pairs_mut();
q.append_pair("response_type", "code");
q.append_pair("client_id", &provider.client_id);
q.append_pair("redirect_uri", &provider.redirect_uri);
if !provider.scope.trim().is_empty() {
q.append_pair("scope", &provider.scope);
}
q.append_pair("state", &state);
q.append_pair("nonce", &nonce);
q.append_pair("code_challenge", &code_challenge);
q.append_pair("code_challenge_method", "S256");
}
Ok(url.into())
}
#[cfg(feature = "oauth2")]
pub async fn oauth2_finish_login(
session: &crate::session::Session,
provider_name: &str,
provider: &OAuth2ProviderConfig,
callback: &OAuth2Callback,
) -> crate::AutumnResult<OidcIdentity> {
validate_callback_state(session, provider_name, callback).await?;
let code_verifier = session
.remove(&format!("oauth2:{provider_name}:code_verifier"))
.await
.ok_or_else(|| {
crate::AutumnError::unauthorized_msg("oauth2 code_verifier missing from session")
})?;
let token = exchange_oauth2_token(provider, callback, code_verifier).await?;
let (claims, source) = load_identity_claims(provider, &token).await?;
validate_oidc_nonce(session, provider_name, &claims, source).await?;
let subject = extract_subject(&claims, source)?;
finalize_oauth2_session(session, provider_name, subject, claims).await
}
#[cfg(feature = "oauth2")]
async fn validate_callback_state(
session: &crate::session::Session,
provider_name: &str,
callback: &OAuth2Callback,
) -> crate::AutumnResult<()> {
let state_key = format!("oauth2:{provider_name}:state");
let expected_state = session.get(&state_key).await.ok_or_else(|| {
crate::AutumnError::unauthorized_msg("oauth2 state missing; restart login")
})?;
if subtle::ConstantTimeEq::ct_eq(expected_state.as_bytes(), callback.state.as_bytes())
.unwrap_u8()
!= 1
{
return Err(crate::AutumnError::unauthorized_msg(
"oauth2 state mismatch",
));
}
session.remove(&state_key).await;
Ok(())
}
#[cfg(feature = "oauth2")]
async fn exchange_oauth2_token(
provider: &OAuth2ProviderConfig,
callback: &OAuth2Callback,
code_verifier: String,
) -> crate::AutumnResult<OAuth2TokenResponse> {
let form_fields: Vec<(&str, String)> = vec![
("grant_type", "authorization_code".to_owned()),
("code", callback.code.clone()),
("redirect_uri", provider.redirect_uri.clone()),
("client_id", provider.client_id.clone()),
("client_secret", provider.client_secret.clone()),
("code_verifier", code_verifier),
];
let token_response = oauth_http_client()?
.post(&provider.token_url)
.header(reqwest::header::ACCEPT, "application/json")
.form(&form_fields)
.send()
.await
.map_err(|e| {
crate::AutumnError::service_unavailable_msg(format!("token request failed: {e}"))
})?
.error_for_status()
.map_err(|e| crate::AutumnError::unauthorized_msg(format!("token exchange failed: {e}")))?;
let token_content_type = token_response
.headers()
.get(reqwest::header::CONTENT_TYPE)
.and_then(|v| v.to_str().ok())
.map(str::to_owned);
let token_body = token_response.text().await.map_err(|e| {
crate::AutumnError::bad_request_msg(format!("invalid token response body: {e}"))
})?;
parse_oauth2_token_response(token_content_type.as_deref(), &token_body)
}
#[cfg(feature = "oauth2")]
async fn load_identity_claims(
provider: &OAuth2ProviderConfig,
token: &OAuth2TokenResponse,
) -> crate::AutumnResult<(serde_json::Value, IdentitySource)> {
if let Some(id_token) = token.id_token.as_deref() {
return Ok((
validate_and_decode_id_token(id_token, provider).await?,
IdentitySource::IdToken,
));
}
if let Some(userinfo_url) = &provider.userinfo_url {
let claims = oauth_http_client()?
.get(userinfo_url)
.header(
reqwest::header::USER_AGENT,
concat!("autumn-web/", env!("CARGO_PKG_VERSION")),
)
.bearer_auth(&token.access_token)
.send()
.await
.map_err(|e| {
crate::AutumnError::service_unavailable_msg(format!("userinfo request failed: {e}"))
})?
.error_for_status()
.map_err(|e| crate::AutumnError::unauthorized_msg(format!("userinfo failed: {e}")))?
.json()
.await
.map_err(|e| {
crate::AutumnError::bad_request_msg(format!("invalid userinfo payload: {e}"))
})?;
return Ok((claims, IdentitySource::UserInfo));
}
Err(crate::AutumnError::bad_request_msg(
"provider must return id_token or configure userinfo_url",
))
}
#[cfg(feature = "oauth2")]
async fn validate_oidc_nonce(
session: &crate::session::Session,
provider_name: &str,
claims: &serde_json::Value,
source: IdentitySource,
) -> crate::AutumnResult<()> {
let nonce_key = format!("oauth2:{provider_name}:nonce");
let stored_nonce = session.remove(&nonce_key).await;
if source == IdentitySource::IdToken {
let expected_nonce = stored_nonce.ok_or_else(|| {
crate::AutumnError::unauthorized_msg("oauth2 nonce missing from session")
})?;
let actual_nonce = claims
.get("nonce")
.and_then(serde_json::Value::as_str)
.ok_or_else(|| crate::AutumnError::unauthorized_msg("missing oidc nonce claim"))?;
if subtle::ConstantTimeEq::ct_eq(expected_nonce.as_bytes(), actual_nonce.as_bytes())
.unwrap_u8()
!= 1
{
return Err(crate::AutumnError::unauthorized_msg("oidc nonce mismatch"));
}
}
Ok(())
}
#[cfg(feature = "oauth2")]
async fn finalize_oauth2_session(
session: &crate::session::Session,
provider_name: &str,
subject: String,
claims: serde_json::Value,
) -> crate::AutumnResult<OidcIdentity> {
session.insert("auth_provider", provider_name).await;
session.rotate_id().await;
Ok(OidcIdentity {
subject,
email: claims
.get("email")
.and_then(serde_json::Value::as_str)
.map(str::to_owned),
name: claims
.get("name")
.and_then(serde_json::Value::as_str)
.map(str::to_owned),
preferred_username: claims
.get("preferred_username")
.and_then(serde_json::Value::as_str)
.map(str::to_owned),
raw_claims: claims,
})
}
#[cfg(feature = "oauth2")]
fn parse_oauth2_token_response(
content_type: Option<&str>,
body: &str,
) -> crate::AutumnResult<OAuth2TokenResponse> {
let looks_like_json = content_type.is_some_and(|v| v.contains("application/json"))
|| body.trim_start().starts_with('{');
if looks_like_json {
return serde_json::from_str(body).map_err(|e| {
crate::AutumnError::bad_request_msg(format!("invalid json token response: {e}"))
});
}
let mut access_token = None;
let mut token_type = None;
let mut id_token = None;
for (k, v) in url::form_urlencoded::parse(body.as_bytes()) {
match k.as_ref() {
"access_token" => access_token = Some(v.into_owned()),
"token_type" => token_type = Some(v.into_owned()),
"id_token" => id_token = Some(v.into_owned()),
_ => {}
}
}
let access_token = access_token.ok_or_else(|| {
crate::AutumnError::bad_request_msg("token response missing access_token")
})?;
Ok(OAuth2TokenResponse {
access_token,
token_type,
id_token,
})
}
#[cfg(feature = "oauth2")]
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum IdentitySource {
IdToken,
UserInfo,
}
#[cfg(feature = "oauth2")]
fn extract_subject(
claims: &serde_json::Value,
source: IdentitySource,
) -> crate::AutumnResult<String> {
if let Some(sub) = claims.get("sub").and_then(serde_json::Value::as_str) {
return Ok(sub.to_owned());
}
if source == IdentitySource::UserInfo {
if let Some(id) = claims.get("id").and_then(serde_json::Value::as_i64) {
return Ok(id.to_string());
}
if let Some(id) = claims.get("id").and_then(serde_json::Value::as_str) {
return Ok(id.to_owned());
}
return Err(crate::AutumnError::bad_request_msg(
"missing identity claim: expected sub or id from userinfo",
));
}
Err(crate::AutumnError::bad_request_msg("missing sub claim"))
}
#[cfg(feature = "oauth2")]
fn jwk_allowed_algorithms(
jwk: &jsonwebtoken::jwk::Jwk,
) -> crate::AutumnResult<Vec<jsonwebtoken::Algorithm>> {
use jsonwebtoken::Algorithm;
use jsonwebtoken::jwk::{AlgorithmParameters, EllipticCurve, KeyAlgorithm};
if let Some(key_alg) = jwk.common.key_algorithm {
let alg = match key_alg {
KeyAlgorithm::RS256 => Algorithm::RS256,
KeyAlgorithm::RS384 => Algorithm::RS384,
KeyAlgorithm::RS512 => Algorithm::RS512,
KeyAlgorithm::PS256 => Algorithm::PS256,
KeyAlgorithm::PS384 => Algorithm::PS384,
KeyAlgorithm::PS512 => Algorithm::PS512,
KeyAlgorithm::ES256 => Algorithm::ES256,
KeyAlgorithm::ES384 => Algorithm::ES384,
KeyAlgorithm::EdDSA => Algorithm::EdDSA,
other => {
return Err(crate::AutumnError::unauthorized_msg(format!(
"jwk algorithm {other} not allowed for id_token verification"
)));
}
};
return Ok(vec![alg]);
}
match &jwk.algorithm {
AlgorithmParameters::RSA(_) => Ok(vec![
Algorithm::RS256,
Algorithm::RS384,
Algorithm::RS512,
Algorithm::PS256,
Algorithm::PS384,
Algorithm::PS512,
]),
AlgorithmParameters::EllipticCurve(params) => match params.curve {
EllipticCurve::P256 => Ok(vec![Algorithm::ES256]),
EllipticCurve::P384 => Ok(vec![Algorithm::ES384]),
ref other => Err(crate::AutumnError::unauthorized_msg(format!(
"unsupported jwk curve {other:?} for id_token verification"
))),
},
AlgorithmParameters::OctetKeyPair(_) => Ok(vec![Algorithm::EdDSA]),
AlgorithmParameters::OctetKey(_) => Err(crate::AutumnError::unauthorized_msg(
"symmetric jwk not allowed for id_token verification",
)),
}
}
#[cfg(feature = "oauth2")]
async fn validate_and_decode_id_token(
token: &str,
provider: &OAuth2ProviderConfig,
) -> crate::AutumnResult<serde_json::Value> {
let issuer = provider
.issuer
.as_deref()
.ok_or_else(|| crate::AutumnError::bad_request_msg("provider.issuer required for oidc"))?;
let jwks_url = provider.jwks_url.as_deref().ok_or_else(|| {
crate::AutumnError::bad_request_msg("provider.jwks_url required for oidc")
})?;
let header = jsonwebtoken::decode_header(token).map_err(|e| {
crate::AutumnError::unauthorized_msg(format!("invalid id_token header: {e}"))
})?;
let kid = header
.kid
.as_deref()
.ok_or_else(|| crate::AutumnError::unauthorized_msg("id_token header missing kid"))?;
let alg = header.alg;
let jwks: JwkSet = oauth_http_client()?
.get(jwks_url)
.send()
.await
.map_err(|e| {
crate::AutumnError::service_unavailable_msg(format!("jwks request failed: {e}"))
})?
.error_for_status()
.map_err(|e| crate::AutumnError::unauthorized_msg(format!("jwks fetch failed: {e}")))?
.json()
.await
.map_err(|e| crate::AutumnError::bad_request_msg(format!("invalid jwks response: {e}")))?;
let jwk = jwks
.keys
.iter()
.find(|k| k.common.key_id.as_deref() == Some(kid))
.ok_or_else(|| crate::AutumnError::unauthorized_msg("no jwk matched id_token kid"))?;
let decoding_key = jsonwebtoken::DecodingKey::from_jwk(jwk)
.map_err(|e| crate::AutumnError::unauthorized_msg(format!("invalid jwk key: {e}")))?;
let allowed_algs = jwk_allowed_algorithms(jwk)?;
if !allowed_algs.contains(&alg) {
return Err(crate::AutumnError::unauthorized_msg(format!(
"id_token alg {alg:?} not permitted by matching jwk"
)));
}
let mut validation = jsonwebtoken::Validation::new(alg);
validation.algorithms = allowed_algs;
let mut issuers = vec![issuer.to_owned()];
let is_multi_tenant = issuer.contains("/common/")
|| issuer.contains("/organizations/")
|| issuer.contains("/consumers/");
if let (true, true, Some(payload_b64)) = (
issuer.contains("login.microsoftonline.com"),
is_multi_tenant,
token.split('.').nth(1),
) {
use base64::Engine as _;
let extract_microsoft_iss = || -> Option<String> {
let payload_bytes = base64::engine::general_purpose::URL_SAFE_NO_PAD
.decode(payload_b64)
.ok()?;
let claims: serde_json::Value = serde_json::from_slice(&payload_bytes).ok()?;
let unverified_iss = claims.get("iss")?.as_str()?;
if unverified_iss.starts_with("https://login.microsoftonline.com/")
&& unverified_iss.ends_with("/v2.0")
{
Some(unverified_iss.to_owned())
} else {
None
}
};
if let Some(unverified_iss) = extract_microsoft_iss() {
issuers.push(unverified_iss);
}
}
let issuer_refs: Vec<&str> = issuers.iter().map(String::as_str).collect();
validation.set_issuer(&issuer_refs);
validation.set_audience(std::slice::from_ref(&provider.client_id));
validation.required_spec_claims = ["exp", "iss", "aud", "sub"]
.into_iter()
.map(str::to_owned)
.collect();
validation.validate_exp = true;
validation.validate_nbf = true;
let claims = jsonwebtoken::decode::<serde_json::Value>(token, &decoding_key, &validation)
.map_err(|e| crate::AutumnError::unauthorized_msg(format!("invalid id_token: {e}")))?;
Ok(claims.claims)
}
#[cfg(feature = "oauth2")]
#[derive(Clone)]
pub struct HttpClient {
inner: reqwest::Client,
}
#[cfg(feature = "oauth2")]
pub struct HttpRequestBuilder {
client: reqwest::Client,
builder: reqwest::RequestBuilder,
}
#[cfg(feature = "oauth2")]
#[allow(
clippy::must_use_candidate,
clippy::missing_const_for_fn,
clippy::return_self_not_must_use,
clippy::missing_errors_doc,
clippy::redundant_closure_for_method_calls
)]
impl HttpClient {
#[must_use]
pub const fn new(inner: reqwest::Client) -> Self {
Self { inner }
}
#[must_use]
pub fn post(&self, url: &str) -> HttpRequestBuilder {
HttpRequestBuilder {
client: self.inner.clone(),
builder: self.inner.post(url),
}
}
#[must_use]
pub fn get(&self, url: &str) -> HttpRequestBuilder {
HttpRequestBuilder {
client: self.inner.clone(),
builder: self.inner.get(url),
}
}
}
#[cfg(feature = "oauth2")]
#[allow(
clippy::must_use_candidate,
clippy::missing_const_for_fn,
clippy::return_self_not_must_use,
clippy::missing_errors_doc,
clippy::redundant_closure_for_method_calls
)]
impl HttpRequestBuilder {
#[must_use]
pub fn header<K, V>(mut self, key: K, value: V) -> Self
where
reqwest::header::HeaderName: TryFrom<K>,
<reqwest::header::HeaderName as TryFrom<K>>::Error: Into<http::Error>,
reqwest::header::HeaderValue: TryFrom<V>,
<reqwest::header::HeaderValue as TryFrom<V>>::Error: Into<http::Error>,
{
self.builder = self.builder.header(key, value);
self
}
#[must_use]
pub fn bearer_auth<T>(mut self, token: T) -> Self
where
T: std::fmt::Display,
{
self.builder = self.builder.bearer_auth(token);
self
}
#[must_use]
pub fn form<T: serde::Serialize + ?Sized>(mut self, form: &T) -> Self {
self.builder = self.builder.form(form);
self
}
pub async fn send(self) -> Result<reqwest::Response, reqwest::Error> {
let req = self.builder.build()?;
let interceptors = crate::interceptor::ACTIVE_HTTP_INTERCEPTORS
.try_with(Clone::clone)
.unwrap_or_default();
run_http_chain(req, interceptors, self.client.clone(), 0).await
}
}
#[cfg(feature = "oauth2")]
fn run_http_chain(
req: reqwest::Request,
interceptors: Vec<Arc<dyn crate::interceptor::HttpInterceptor>>,
client: reqwest::Client,
idx: usize,
) -> std::pin::Pin<
Box<
dyn std::future::Future<Output = Result<reqwest::Response, reqwest::Error>>
+ Send
+ 'static,
>,
> {
Box::pin(async move {
if idx < interceptors.len() {
let interceptor = interceptors[idx].clone();
let next_interceptors = interceptors.clone();
let next_client = client.clone();
let next_fn = move |r: reqwest::Request| {
run_http_chain(r, next_interceptors.clone(), next_client.clone(), idx + 1)
};
let fut = interceptor.intercept(req, &next_fn);
fut.await
} else {
client.execute(req).await
}
})
}
#[cfg(feature = "oauth2")]
fn oauth_http_client() -> crate::AutumnResult<HttpClient> {
let client = reqwest::Client::builder()
.timeout(Duration::from_secs(OAUTH_HTTP_TIMEOUT_SECS))
.build()
.map_err(|e| {
crate::AutumnError::service_unavailable_msg(format!(
"failed to build oauth http client: {e}"
))
})?;
Ok(HttpClient::new(client))
}
impl Default for AuthConfig {
fn default() -> Self {
Self {
bcrypt_cost: default_bcrypt_cost(),
session_key: default_session_key(),
#[cfg(feature = "oauth2")]
oauth2: OAuth2Config::default(),
#[cfg(feature = "oauth2")]
oauth_linking_policy: OAuthLinkingPolicy::default(),
#[cfg(feature = "webauthn")]
webauthn: WebAuthnConfig::default(),
lockout: LockoutConfig::default(),
step_up: StepUpConfig::default(),
sessions: SessionTrackingConfig::default(),
password: PasswordConfig::default(),
remember: RememberConfig::default(),
magic_link: MagicLinkConfig::default(),
}
}
}
#[derive(Debug, Clone, PartialEq, Eq, Default)]
pub struct VerifiedToken {
pub principal_id: String,
pub scopes: Vec<String>,
pub name: String,
pub expires_at: Option<chrono::DateTime<chrono::Utc>>,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct TokenMetadata {
pub id: String,
pub name: String,
pub principal_id: String,
pub scopes: Vec<String>,
pub created_at: chrono::DateTime<chrono::Utc>,
pub expires_at: Option<chrono::DateTime<chrono::Utc>>,
pub last_used_at: Option<chrono::DateTime<chrono::Utc>>,
pub revoked_at: Option<chrono::DateTime<chrono::Utc>>,
}
#[derive(Debug, Clone, Default)]
pub struct IssueTokenSpec<'a> {
pub principal_id: &'a str,
pub name: &'a str,
pub scopes: &'a [String],
pub expires_at: Option<chrono::DateTime<chrono::Utc>>,
}
pub trait ApiTokenStore: Send + Sync + 'static {
fn issue<'a>(
&'a self,
principal_id: &'a str,
) -> Pin<Box<dyn Future<Output = crate::AutumnResult<String>> + Send + 'a>>;
fn verify<'a>(
&'a self,
raw_token: &'a str,
) -> Pin<Box<dyn Future<Output = crate::AutumnResult<Option<String>>> + Send + 'a>>;
fn revoke<'a>(
&'a self,
raw_token: &'a str,
) -> Pin<Box<dyn Future<Output = crate::AutumnResult<()>> + Send + 'a>>;
fn issue_scoped<'a>(
&'a self,
spec: IssueTokenSpec<'a>,
) -> Pin<Box<dyn Future<Output = crate::AutumnResult<String>> + Send + 'a>> {
Box::pin(async move { self.issue(spec.principal_id).await })
}
fn verify_scoped<'a>(
&'a self,
raw_token: &'a str,
) -> Pin<Box<dyn Future<Output = crate::AutumnResult<Option<VerifiedToken>>> + Send + 'a>> {
Box::pin(async move {
Ok(self
.verify(raw_token)
.await?
.map(|principal_id| VerifiedToken {
principal_id,
scopes: Vec::new(),
name: String::new(),
expires_at: None,
}))
})
}
fn list<'a>(
&'a self,
principal_id: &'a str,
) -> Pin<Box<dyn Future<Output = crate::AutumnResult<Vec<TokenMetadata>>> + Send + 'a>> {
let _ = principal_id;
Box::pin(async move { Ok(Vec::new()) })
}
fn rotate<'a>(
&'a self,
raw_token: &'a str,
) -> Pin<Box<dyn Future<Output = crate::AutumnResult<Option<String>>> + Send + 'a>> {
Box::pin(async move {
match self.verify_scoped(raw_token).await? {
Some(vt) => {
self.revoke(raw_token).await?;
let scopes = vt.scopes.clone();
let raw = self
.issue_scoped(IssueTokenSpec {
principal_id: &vt.principal_id,
name: &vt.name,
scopes: &scopes,
expires_at: vt.expires_at,
})
.await?;
Ok(Some(raw))
}
None => Ok(None),
}
})
}
}
#[must_use]
pub fn hash_api_token(raw: &str) -> String {
use sha2::Digest as _;
sha2::Sha256::digest(raw.as_bytes())
.iter()
.fold(String::with_capacity(64), |mut s, b| {
use std::fmt::Write as _;
let _ = write!(s, "{b:02x}");
s
})
}
#[must_use]
pub fn generate_raw_token() -> String {
let u1 = uuid::Uuid::new_v4();
let u2 = uuid::Uuid::new_v4();
format!("{}{}", u1.simple(), u2.simple())
}
#[derive(Debug, Clone)]
struct StoredToken {
id: u64,
principal_id: String,
name: String,
scopes: Vec<String>,
created_at: chrono::DateTime<chrono::Utc>,
expires_at: Option<chrono::DateTime<chrono::Utc>>,
last_used_at: Option<chrono::DateTime<chrono::Utc>>,
revoked_at: Option<chrono::DateTime<chrono::Utc>>,
}
#[derive(Clone)]
pub struct InMemoryApiTokenStore {
tokens: Arc<std::sync::RwLock<std::collections::HashMap<String, StoredToken>>>,
next_id: Arc<std::sync::atomic::AtomicU64>,
clock: Arc<dyn crate::time::ClockSource>,
}
impl Default for InMemoryApiTokenStore {
fn default() -> Self {
Self {
tokens: Arc::new(std::sync::RwLock::new(std::collections::HashMap::new())),
next_id: Arc::new(std::sync::atomic::AtomicU64::new(1)),
clock: Arc::new(crate::time::SystemClock),
}
}
}
impl InMemoryApiTokenStore {
#[must_use]
pub fn with_clock(mut self, clock: Arc<dyn crate::time::ClockSource>) -> Self {
self.clock = clock;
self
}
#[must_use]
pub fn with_token(self, raw_token: &str, principal_id: &str) -> Self {
self.with_scoped_token(raw_token, principal_id, &[])
}
#[must_use]
pub fn with_scoped_token(self, raw_token: &str, principal_id: &str, scopes: &[String]) -> Self {
self.store_raw_token(
raw_token,
&IssueTokenSpec {
principal_id,
scopes,
..Default::default()
},
);
self
}
pub fn from_env(var: &str, principal_id: &str) -> crate::AutumnResult<Self> {
let raw = std::env::var(var).unwrap_or_default();
if raw.trim().is_empty() {
return Err(crate::AutumnError::internal_server_error_msg(format!(
"InMemoryApiTokenStore::from_env: environment variable `{var}` is unset or empty"
)));
}
Ok(Self::default().with_token(&raw, principal_id))
}
fn insert_token(&self, spec: &IssueTokenSpec<'_>) -> String {
let raw = generate_raw_token();
self.store_raw_token(&raw, spec);
raw
}
fn store_raw_token(&self, raw_token: &str, spec: &IssueTokenSpec<'_>) {
if raw_token.trim().is_empty() {
tracing::warn!(
"InMemoryApiTokenStore: ignoring a blank (empty or whitespace-only) seed token; \
no credential was stored"
);
return;
}
let hash = hash_api_token(raw_token);
let id = self
.next_id
.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
let stored = StoredToken {
id,
principal_id: spec.principal_id.to_owned(),
name: spec.name.to_owned(),
scopes: spec.scopes.to_vec(),
created_at: self.clock.now(),
expires_at: spec.expires_at,
last_used_at: None,
revoked_at: None,
};
self.tokens
.write()
.expect("api token store lock poisoned")
.insert(hash, stored);
}
fn resolve_used(&self, raw_token: &str) -> Option<VerifiedToken> {
let hash = hash_api_token(raw_token);
let now = self.clock.now();
let mut guard = self.tokens.write().expect("api token store lock poisoned");
let stored = guard.get_mut(&hash)?;
if stored.revoked_at.is_some() || stored.expires_at.is_some_and(|exp| exp <= now) {
return None;
}
stored.last_used_at = Some(now);
let verified = VerifiedToken {
principal_id: stored.principal_id.clone(),
scopes: stored.scopes.clone(),
name: stored.name.clone(),
expires_at: stored.expires_at,
};
drop(guard);
Some(verified)
}
}
impl ApiTokenStore for InMemoryApiTokenStore {
fn issue<'a>(
&'a self,
principal_id: &'a str,
) -> Pin<Box<dyn Future<Output = crate::AutumnResult<String>> + Send + 'a>> {
Box::pin(async move {
Ok(self.insert_token(&IssueTokenSpec {
principal_id,
..Default::default()
}))
})
}
fn verify<'a>(
&'a self,
raw_token: &'a str,
) -> Pin<Box<dyn Future<Output = crate::AutumnResult<Option<String>>> + Send + 'a>> {
Box::pin(async move { Ok(self.resolve_used(raw_token).map(|vt| vt.principal_id)) })
}
fn revoke<'a>(
&'a self,
raw_token: &'a str,
) -> Pin<Box<dyn Future<Output = crate::AutumnResult<()>> + Send + 'a>> {
Box::pin(async move {
let hash = hash_api_token(raw_token);
let now = self.clock.now();
let mut guard = self.tokens.write().expect("api token store lock poisoned");
if let Some(stored) = guard.get_mut(&hash) {
stored.revoked_at.get_or_insert(now);
}
drop(guard);
Ok(())
})
}
fn issue_scoped<'a>(
&'a self,
spec: IssueTokenSpec<'a>,
) -> Pin<Box<dyn Future<Output = crate::AutumnResult<String>> + Send + 'a>> {
Box::pin(async move { Ok(self.insert_token(&spec)) })
}
fn verify_scoped<'a>(
&'a self,
raw_token: &'a str,
) -> Pin<Box<dyn Future<Output = crate::AutumnResult<Option<VerifiedToken>>> + Send + 'a>> {
Box::pin(async move { Ok(self.resolve_used(raw_token)) })
}
fn list<'a>(
&'a self,
principal_id: &'a str,
) -> Pin<Box<dyn Future<Output = crate::AutumnResult<Vec<TokenMetadata>>> + Send + 'a>> {
Box::pin(async move {
let mut out: Vec<TokenMetadata> = {
let guard = self.tokens.read().expect("api token store lock poisoned");
guard
.values()
.filter(|s| s.principal_id == principal_id)
.map(|s| TokenMetadata {
id: s.id.to_string(),
name: s.name.clone(),
principal_id: s.principal_id.clone(),
scopes: s.scopes.clone(),
created_at: s.created_at,
expires_at: s.expires_at,
last_used_at: s.last_used_at,
revoked_at: s.revoked_at,
})
.collect()
};
out.sort_by(|a, b| {
a.id.parse::<u64>()
.unwrap_or(0)
.cmp(&b.id.parse().unwrap_or(0))
});
Ok(out)
})
}
}
pub async fn issue_api_token(
store: &dyn ApiTokenStore,
principal_id: &str,
) -> crate::AutumnResult<String> {
store.issue(principal_id).await
}
pub async fn revoke_api_token(
store: &dyn ApiTokenStore,
raw_token: &str,
) -> crate::AutumnResult<()> {
store.revoke(raw_token).await
}
pub async fn issue_scoped_api_token(
store: &dyn ApiTokenStore,
spec: IssueTokenSpec<'_>,
) -> crate::AutumnResult<String> {
store.issue_scoped(spec).await
}
pub async fn list_api_tokens(
store: &dyn ApiTokenStore,
principal_id: &str,
) -> crate::AutumnResult<Vec<TokenMetadata>> {
store.list(principal_id).await
}
pub async fn rotate_api_token(
store: &dyn ApiTokenStore,
raw_token: &str,
) -> crate::AutumnResult<Option<String>> {
store.rotate(raw_token).await
}
#[derive(Clone)]
struct ApiTokenPrincipal(String);
#[derive(Debug, Clone, Default, PartialEq, Eq)]
pub struct ApiTokenScopes(pub Vec<String>);
#[derive(Debug, Clone)]
pub struct ApiToken(pub String);
impl<S> FromRequestParts<S> for ApiToken
where
S: Send + Sync,
{
type Rejection = AuthRejection;
fn from_request_parts(
parts: &mut Parts,
_state: &S,
) -> impl Future<Output = Result<Self, Self::Rejection>> + Send {
let principal = parts.extensions.get::<ApiTokenPrincipal>().cloned();
async move { principal.map(|p| Self(p.0)).ok_or(AuthRejection) }
}
}
#[derive(Clone)]
pub struct RequireApiToken {
store: Arc<dyn ApiTokenStore>,
}
impl RequireApiToken {
#[must_use]
pub fn new<S: ApiTokenStore + 'static>(store: Arc<S>) -> Self {
Self { store }
}
}
impl<S> tower::Layer<S> for RequireApiToken {
type Service = RequireApiTokenService<S>;
fn layer(&self, inner: S) -> Self::Service {
RequireApiTokenService {
inner,
store: Arc::clone(&self.store),
}
}
}
#[derive(Clone)]
pub struct RequireApiTokenService<S> {
inner: S,
store: Arc<dyn ApiTokenStore>,
}
impl<S, ResBody> tower::Service<axum::extract::Request> for RequireApiTokenService<S>
where
S: tower::Service<axum::extract::Request, Response = Response<ResBody>>
+ Clone
+ Send
+ 'static,
S::Future: Send + 'static,
S::Error: Send + 'static,
ResBody: From<String> + Default + Send + 'static,
{
type Response = Response<ResBody>;
type Error = S::Error;
type Future = Pin<Box<dyn Future<Output = Result<Self::Response, Self::Error>> + Send>>;
fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
self.inner.poll_ready(cx)
}
fn call(&mut self, mut req: axum::extract::Request) -> Self::Future {
let store = Arc::clone(&self.store);
let mut inner = self.inner.clone();
std::mem::swap(&mut self.inner, &mut inner);
Box::pin(async move {
let raw_token = req
.headers()
.get(http::header::AUTHORIZATION)
.and_then(|v| v.to_str().ok())
.and_then(parse_bearer_token)
.map(str::to_owned);
let Some(raw_token) = raw_token else {
let (request_id, instance) = api_token_problem_context(&req);
return Ok(api_token_unauthorized_response(request_id, instance));
};
match store.verify_scoped(&raw_token).await {
Ok(Some(verified)) => {
let VerifiedToken {
principal_id,
scopes,
..
} = verified;
req.extensions_mut()
.insert(crate::security::RateLimitPrincipal(principal_id.clone()));
req.extensions_mut().insert(ApiTokenScopes(scopes));
if crate::current::Current::actor().is_none() {
crate::current::Current::set_actor(principal_id.clone());
}
req.extensions_mut().insert(ApiTokenPrincipal(principal_id));
inner.call(req).await
}
Ok(None) => {
let (request_id, instance) = api_token_problem_context(&req);
Ok(api_token_unauthorized_response(request_id, instance))
}
Err(err) => {
let (request_id, instance) = api_token_problem_context(&req);
Ok(api_token_error_response(&err, request_id, instance))
}
}
})
}
}
fn parse_bearer_token(header: &str) -> Option<&str> {
let (scheme, token) = header.split_once(' ')?;
scheme.eq_ignore_ascii_case("Bearer").then_some(token)
}
fn api_token_unauthorized_response<ResBody: From<String> + Default>(
request_id: Option<String>,
instance: Option<String>,
) -> Response<ResBody> {
let body = crate::error::problem_details_json_string(
StatusCode::UNAUTHORIZED,
"authentication required",
None,
None,
request_id,
instance,
true,
);
Response::builder()
.status(StatusCode::UNAUTHORIZED)
.header(http::header::CONTENT_TYPE, "application/problem+json")
.body(ResBody::from(body))
.unwrap_or_default()
}
fn api_token_error_response<ResBody: From<String> + Default>(
err: &crate::AutumnError,
request_id: Option<String>,
instance: Option<String>,
) -> Response<ResBody> {
let status = err.status();
let message = err.to_string();
let body = crate::error::problem_details_json_string(
status,
message.clone(),
None,
None,
request_id,
instance,
true,
);
let mut response = Response::builder()
.status(status)
.header(http::header::CONTENT_TYPE, "application/problem+json")
.body(ResBody::from(body))
.unwrap_or_default();
response
.extensions_mut()
.insert(crate::middleware::AutumnErrorInfo {
status,
message,
details: None,
problem_type: None,
backtrace_string: None,
});
response
}
fn api_token_problem_context(req: &axum::extract::Request) -> (Option<String>, Option<String>) {
(
req.extensions()
.get::<crate::middleware::RequestId>()
.map(std::string::ToString::to_string),
Some(req.uri().path().to_owned()),
)
}
#[cfg(feature = "db")]
pub const API_TOKEN_MIGRATIONS: diesel_migrations::EmbeddedMigrations =
diesel_migrations::embed_migrations!("migrations");
#[cfg(feature = "db")]
mod db_store {
use std::future::Future;
use std::pin::Pin;
use std::sync::Arc;
use chrono::{DateTime, NaiveDateTime, Utc};
use diesel::OptionalExtension as _;
use diesel::prelude::*;
use diesel_async::AsyncPgConnection;
use diesel_async::RunQueryDsl;
use diesel_async::pooled_connection::deadpool::Pool;
use super::{
ApiTokenStore, IssueTokenSpec, TokenMetadata, VerifiedToken, generate_raw_token,
hash_api_token,
};
use crate::error::AutumnError;
use crate::time::{ClockSource, SystemClock};
diesel::table! {
api_tokens (id) {
id -> Int8,
token_hash -> Text,
principal_id -> Text,
created_at -> Timestamp,
revoked_at -> Nullable<Timestamp>,
name -> Text,
scopes -> Jsonb,
expires_at -> Nullable<Timestamp>,
last_used_at -> Nullable<Timestamp>,
}
}
#[derive(Insertable)]
#[diesel(table_name = api_tokens)]
struct NewApiToken<'a> {
token_hash: &'a str,
principal_id: &'a str,
name: &'a str,
scopes: serde_json::Value,
expires_at: Option<NaiveDateTime>,
}
#[derive(Queryable)]
struct TokenRow {
id: i64,
name: String,
principal_id: String,
scopes: serde_json::Value,
created_at: NaiveDateTime,
expires_at: Option<NaiveDateTime>,
last_used_at: Option<NaiveDateTime>,
revoked_at: Option<NaiveDateTime>,
}
const fn to_utc(naive: NaiveDateTime) -> DateTime<Utc> {
DateTime::from_naive_utc_and_offset(naive, Utc)
}
fn scopes_to_json(scopes: &[String]) -> serde_json::Value {
serde_json::Value::Array(
scopes
.iter()
.map(|s| serde_json::Value::String(s.clone()))
.collect(),
)
}
#[must_use]
pub fn scopes_from_json(value: &serde_json::Value) -> Vec<String> {
value
.as_array()
.map(|arr| {
arr.iter()
.filter_map(|v| v.as_str().map(str::to_owned))
.collect()
})
.unwrap_or_default()
}
#[derive(Clone)]
pub struct DbApiTokenStore {
pool: Pool<AsyncPgConnection>,
clock: Arc<dyn ClockSource>,
}
impl DbApiTokenStore {
#[must_use]
pub fn new(pool: Pool<AsyncPgConnection>) -> Self {
Self {
pool,
clock: Arc::new(SystemClock),
}
}
#[must_use]
pub fn with_clock(mut self, clock: Arc<dyn ClockSource>) -> Self {
self.clock = clock;
self
}
async fn conn(
&self,
) -> crate::AutumnResult<diesel_async::pooled_connection::deadpool::Object<AsyncPgConnection>>
{
self.pool
.get()
.await
.map_err(|e| AutumnError::internal_server_error_msg(e.to_string()))
}
}
impl ApiTokenStore for DbApiTokenStore {
fn issue<'a>(
&'a self,
principal_id: &'a str,
) -> Pin<Box<dyn Future<Output = crate::AutumnResult<String>> + Send + 'a>> {
self.issue_scoped(IssueTokenSpec {
principal_id,
..Default::default()
})
}
fn verify<'a>(
&'a self,
raw_token: &'a str,
) -> Pin<Box<dyn Future<Output = crate::AutumnResult<Option<String>>> + Send + 'a>>
{
Box::pin(async move {
Ok(self
.verify_scoped(raw_token)
.await?
.map(|vt| vt.principal_id))
})
}
fn revoke<'a>(
&'a self,
raw_token: &'a str,
) -> Pin<Box<dyn Future<Output = crate::AutumnResult<()>> + Send + 'a>> {
Box::pin(async move {
let hash = hash_api_token(raw_token);
let now = self.clock.now().naive_utc();
let mut conn = self.conn().await?;
diesel::update(api_tokens::table)
.filter(api_tokens::token_hash.eq(&hash))
.filter(api_tokens::revoked_at.is_null())
.set(api_tokens::revoked_at.eq(Some(now)))
.execute(&mut conn)
.await
.map_err(|e| AutumnError::internal_server_error_msg(e.to_string()))?;
Ok(())
})
}
fn issue_scoped<'a>(
&'a self,
spec: IssueTokenSpec<'a>,
) -> Pin<Box<dyn Future<Output = crate::AutumnResult<String>> + Send + 'a>> {
Box::pin(async move {
let raw = generate_raw_token();
let hash = hash_api_token(&raw);
let mut conn = self.conn().await?;
diesel::insert_into(api_tokens::table)
.values(NewApiToken {
token_hash: &hash,
principal_id: spec.principal_id,
name: spec.name,
scopes: scopes_to_json(spec.scopes),
expires_at: spec.expires_at.map(|dt| dt.naive_utc()),
})
.execute(&mut conn)
.await
.map_err(|e| AutumnError::internal_server_error_msg(e.to_string()))?;
Ok(raw)
})
}
fn rotate<'a>(
&'a self,
raw_token: &'a str,
) -> Pin<Box<dyn Future<Output = crate::AutumnResult<Option<String>>> + Send + 'a>>
{
Box::pin(async move {
#[derive(diesel::QueryableByName)]
struct CountRow {
#[diesel(sql_type = diesel::sql_types::BigInt)]
count: i64,
}
let old_hash = hash_api_token(raw_token);
let new_raw = generate_raw_token();
let new_hash = hash_api_token(&new_raw);
let now = self.clock.now().naive_utc();
let mut conn = self.conn().await?;
let row: CountRow = diesel::sql_query(
"WITH rotated AS ( \
UPDATE api_tokens \
SET revoked_at = $3 \
WHERE token_hash = $1 AND revoked_at IS NULL \
AND (expires_at IS NULL OR expires_at > $3) \
RETURNING principal_id, name, scopes, expires_at \
), \
inserted AS ( \
INSERT INTO api_tokens (token_hash, principal_id, name, scopes, expires_at) \
SELECT $2, principal_id, name, scopes, expires_at FROM rotated \
RETURNING 1 \
) \
SELECT COUNT(*)::bigint AS count FROM inserted",
)
.bind::<diesel::sql_types::Text, _>(&old_hash)
.bind::<diesel::sql_types::Text, _>(&new_hash)
.bind::<diesel::sql_types::Timestamp, _>(now)
.get_result(&mut conn)
.await
.map_err(|e| AutumnError::internal_server_error_msg(e.to_string()))?;
if row.count == 0 {
Ok(None)
} else {
Ok(Some(new_raw))
}
})
}
fn verify_scoped<'a>(
&'a self,
raw_token: &'a str,
) -> Pin<Box<dyn Future<Output = crate::AutumnResult<Option<VerifiedToken>>> + Send + 'a>>
{
Box::pin(async move {
let hash = hash_api_token(raw_token);
let now = self.clock.now().naive_utc();
let mut conn = self.conn().await?;
let row: Option<(
i64,
String,
String,
Option<NaiveDateTime>,
serde_json::Value,
)> = api_tokens::table
.filter(api_tokens::token_hash.eq(&hash))
.filter(api_tokens::revoked_at.is_null())
.filter(
api_tokens::expires_at
.is_null()
.or(api_tokens::expires_at.gt(now)),
)
.select((
api_tokens::id,
api_tokens::principal_id,
api_tokens::name,
api_tokens::expires_at,
api_tokens::scopes,
))
.first(&mut conn)
.await
.optional()
.map_err(|e| AutumnError::internal_server_error_msg(e.to_string()))?;
let Some((id, principal_id, name, expires_at_naive, scopes_json)) = row else {
return Ok(None);
};
let threshold = now - chrono::Duration::minutes(5);
let _ = diesel::update(
api_tokens::table.filter(api_tokens::id.eq(id)).filter(
api_tokens::last_used_at
.is_null()
.or(api_tokens::last_used_at.lt(threshold)),
),
)
.set(api_tokens::last_used_at.eq(Some(now)))
.execute(&mut conn)
.await;
Ok(Some(VerifiedToken {
principal_id,
scopes: scopes_from_json(&scopes_json),
name,
expires_at: expires_at_naive.map(to_utc),
}))
})
}
fn list<'a>(
&'a self,
principal_id: &'a str,
) -> Pin<Box<dyn Future<Output = crate::AutumnResult<Vec<TokenMetadata>>> + Send + 'a>>
{
Box::pin(async move {
let mut conn = self.conn().await?;
let rows: Vec<TokenRow> = api_tokens::table
.filter(api_tokens::principal_id.eq(principal_id))
.order(api_tokens::id.asc())
.select((
api_tokens::id,
api_tokens::name,
api_tokens::principal_id,
api_tokens::scopes,
api_tokens::created_at,
api_tokens::expires_at,
api_tokens::last_used_at,
api_tokens::revoked_at,
))
.load(&mut conn)
.await
.map_err(|e| AutumnError::internal_server_error_msg(e.to_string()))?;
Ok(rows
.into_iter()
.map(|r| TokenMetadata {
id: r.id.to_string(),
name: r.name,
principal_id: r.principal_id,
scopes: scopes_from_json(&r.scopes),
created_at: to_utc(r.created_at),
expires_at: r.expires_at.map(to_utc),
last_used_at: r.last_used_at.map(to_utc),
revoked_at: r.revoked_at.map(to_utc),
})
.collect())
})
}
}
}
#[cfg(feature = "db")]
pub use db_store::DbApiTokenStore;
#[cfg(feature = "db")]
#[doc(hidden)]
pub use db_store::scopes_from_json;
#[cfg(test)]
mod tests {
use super::*;
fn test_app_state(auth_session_key: &str) -> crate::state::AppState {
crate::state::AppState {
extensions: std::sync::Arc::new(std::sync::RwLock::new(
std::collections::HashMap::new(),
)),
#[cfg(feature = "db")]
pool: None,
#[cfg(feature = "db")]
replica_pool: None,
#[cfg(feature = "db")]
shards: None,
profile: None,
role: crate::config::ProcessRole::Combined,
started_at: std::time::Instant::now(),
health_detailed: false,
probes: crate::probe::ProbeState::ready_for_test(),
metrics: crate::middleware::MetricsCollector::new(),
log_levels: crate::actuator::LogLevels::new("info"),
task_registry: crate::actuator::TaskRegistry::new(),
job_registry: crate::actuator::JobRegistry::new(),
config_props: crate::actuator::ConfigProperties::default(),
metrics_source_registry: crate::actuator::MetricsSourceRegistry::new(),
health_indicator_registry: crate::actuator::HealthIndicatorRegistry::new(),
#[cfg(feature = "ws")]
channels: crate::channels::Channels::new(32),
#[cfg(feature = "presence")]
presence: crate::presence::Presence::new(crate::channels::Channels::new(32)),
#[cfg(feature = "ws")]
shutdown: tokio_util::sync::CancellationToken::new(),
policy_registry: crate::authorization::PolicyRegistry::default(),
forbidden_response: crate::authorization::ForbiddenResponse::default(),
auth_session_key: auth_session_key.to_owned(),
shared_cache: None,
clock: std::sync::Arc::new(crate::time::SystemClock),
app_id: crate::state::AppState::next_app_id(),
}
}
#[tokio::test]
async fn hash_and_verify_password() {
let hash = hash_password("test_password").await.unwrap();
assert!(hash.starts_with("$2b$"));
assert!(verify_password("test_password", &hash).await.unwrap());
assert!(!verify_password("wrong_password", &hash).await.unwrap());
}
#[tokio::test]
async fn verify_invalid_hash_returns_false() {
let result = verify_password("test", "not-a-valid-hash").await;
assert!(result.is_ok());
assert!(!result.unwrap());
}
#[tokio::test]
async fn verify_password_rejects_invalid_hash_format_safely() {
let result = verify_password("test", "short").await;
assert!(result.is_ok());
assert!(!result.unwrap());
let bad_prefix = "a".repeat(60);
let result = verify_password("test", &bad_prefix).await;
assert!(result.is_ok());
assert!(!result.unwrap());
let bad_length = "$2b$12$short";
let result = verify_password("test", bad_length).await;
assert!(result.is_ok());
assert!(!result.unwrap());
}
#[test]
fn auth_config_defaults() {
let config = AuthConfig::default();
assert_eq!(config.bcrypt_cost, 12);
assert_eq!(config.session_key, "user_id");
#[cfg(feature = "oauth2")]
assert!(config.oauth2.providers.is_empty());
}
#[test]
fn session_tracking_config_defaults_to_revoke_on_credential_change() {
let config = AuthConfig::default();
assert!(config.sessions.revoke_on_credential_change);
assert_eq!(config.sessions.last_seen_update_secs, 60);
}
#[test]
fn session_tracking_config_deserializes_from_toml() {
let cfg: crate::config::AutumnConfig = toml::from_str(
r"
[auth.sessions]
revoke_on_credential_change = false
last_seen_update_secs = 5
",
)
.expect("config must parse");
assert!(!cfg.auth.sessions.revoke_on_credential_change);
assert_eq!(cfg.auth.sessions.last_seen_update_secs, 5);
}
#[cfg(feature = "oauth2")]
#[test]
fn oauth2_config_deserializes_provider_tables() {
let cfg: crate::config::AutumnConfig = toml::from_str(
r#"
[auth.oauth2.github]
client_id = "cid"
client_secret = "secret"
authorize_url = "https://github.com/login/oauth/authorize"
token_url = "https://github.com/login/oauth/access_token"
redirect_uri = "http://localhost:3000/auth/github/callback"
"#,
)
.unwrap();
let provider = cfg.auth.oauth2.providers.get("github").unwrap();
assert_eq!(provider.client_id, "cid");
assert_eq!(provider.scope, "");
assert!(provider.issuer.is_none());
assert!(provider.jwks_url.is_none());
}
#[cfg(feature = "oauth2")]
#[tokio::test]
async fn oauth2_authorize_url_sets_state_and_nonce() {
let session = crate::session::Session::new_for_test("s1".into(), HashMap::new());
let provider = OAuth2ProviderConfig {
client_id: "cid".into(),
client_secret: "secret".into(),
authorize_url: "https://idp.example/authorize".into(),
token_url: "https://idp.example/token".into(),
userinfo_url: None,
redirect_uri: "http://localhost:3000/callback".into(),
scope: "openid profile".into(),
issuer: None,
jwks_url: None,
discovery_url: None,
};
let url = oauth2_authorize_url(&session, "github", &provider)
.await
.unwrap();
assert!(url.contains("response_type=code"));
assert!(session.get("oauth2:github:state").await.is_some());
assert!(session.get("oauth2:github:nonce").await.is_some());
}
#[cfg(feature = "oauth2")]
#[tokio::test]
async fn oauth2_authorize_url_omits_scope_when_empty() {
let session = crate::session::Session::new_for_test("s1".into(), HashMap::new());
let provider = OAuth2ProviderConfig {
client_id: "cid".into(),
client_secret: "secret".into(),
authorize_url: "https://idp.example/authorize".into(),
token_url: "https://idp.example/token".into(),
userinfo_url: None,
redirect_uri: "http://localhost:3000/callback".into(),
scope: String::new(),
issuer: None,
jwks_url: None,
discovery_url: None,
};
let url = oauth2_authorize_url(&session, "github", &provider)
.await
.unwrap();
assert!(!url.contains("scope="));
}
#[cfg(feature = "oauth2")]
#[tokio::test]
async fn validate_id_token_requires_oidc_metadata() {
let provider = OAuth2ProviderConfig {
client_id: "cid".into(),
client_secret: "secret".into(),
authorize_url: "https://idp.example/authorize".into(),
token_url: "https://idp.example/token".into(),
userinfo_url: None,
redirect_uri: "http://localhost:3000/callback".into(),
scope: "openid profile".into(),
issuer: None,
jwks_url: None,
discovery_url: None,
};
let err = validate_and_decode_id_token("bad.token.value", &provider)
.await
.unwrap_err();
assert_eq!(err.to_string(), "provider.issuer required for oidc");
}
#[cfg(feature = "oauth2")]
fn rsa_test_jwk_json(key_algorithm: Option<&str>) -> serde_json::Value {
let mut jwk = serde_json::json!({
"kty": "RSA",
"kid": "test-kid",
"n": "ofgWCuLjybRlzo0tZWJjNiuSfb4p4fAkd_wWJcyQoTbji9k0l8W26mPddxHmfHQp\
-Vaw-4qPCJrcS2mJPMEzP1Pt0Bm4d4QlL-yRT-SFd2lZS-pCgNMsD1W_YpRPEwOW\
vG6b32690r2jZ47soMZo9wGzjb_7OMg0LOL-bSf63kpaSHSXndS5z5rexMdbBYUs\
LA9e-KXBdQOS-UTo7WTBEMa2R2CapHg665xsmtdVMTBQY4uDZlxvb3qCo5ZwKh9k\
G4LT6_I5IhlJH7aGhyxXFvUK-DWNmoudF8NAco9_h9iaGNj8q2ethFkMLs91kzk2\
PAcDTW9gb54h4FRWyuXpoQ",
"e": "AQAB"
});
if let Some(alg) = key_algorithm {
jwk["alg"] = serde_json::json!(alg);
}
jwk
}
#[cfg(feature = "oauth2")]
fn rsa_test_jwk(key_algorithm: Option<&str>) -> jsonwebtoken::jwk::Jwk {
serde_json::from_value(rsa_test_jwk_json(key_algorithm)).unwrap()
}
#[cfg(feature = "oauth2")]
#[test]
fn jwk_allowed_algorithms_pins_declared_algorithm() {
let algs = jwk_allowed_algorithms(&rsa_test_jwk(Some("RS256"))).unwrap();
assert_eq!(algs, vec![jsonwebtoken::Algorithm::RS256]);
}
#[cfg(feature = "oauth2")]
#[test]
fn jwk_allowed_algorithms_rejects_symmetric_declared_algorithm() {
let err = jwk_allowed_algorithms(&rsa_test_jwk(Some("HS256"))).unwrap_err();
assert!(
err.to_string().contains("not allowed"),
"expected symmetric alg rejection, got: {err}"
);
}
#[cfg(feature = "oauth2")]
#[test]
fn jwk_allowed_algorithms_derives_asymmetric_set_from_key_type() {
let algs = jwk_allowed_algorithms(&rsa_test_jwk(None)).unwrap();
assert!(algs.contains(&jsonwebtoken::Algorithm::RS256));
assert!(algs.contains(&jsonwebtoken::Algorithm::PS256));
assert!(!algs.contains(&jsonwebtoken::Algorithm::HS256));
assert!(!algs.contains(&jsonwebtoken::Algorithm::HS384));
assert!(!algs.contains(&jsonwebtoken::Algorithm::HS512));
}
#[cfg(feature = "oauth2")]
#[test]
fn jwk_allowed_algorithms_rejects_symmetric_octet_key() {
let jwk: jsonwebtoken::jwk::Jwk = serde_json::from_value(serde_json::json!({
"kty": "oct",
"kid": "sym-kid",
"k": "c2VjcmV0"
}))
.unwrap();
let err = jwk_allowed_algorithms(&jwk).unwrap_err();
assert!(
err.to_string().contains("symmetric jwk not allowed"),
"expected symmetric jwk rejection, got: {err}"
);
}
#[cfg(feature = "oauth2")]
async fn spawn_jwks_stub(body: String) -> String {
use tokio::io::{AsyncReadExt as _, AsyncWriteExt as _};
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
tokio::spawn(async move {
while let Ok((mut stream, _)) = listener.accept().await {
let mut buf = [0u8; 4096];
let _ = stream.read(&mut buf).await;
let response = format!(
"HTTP/1.1 200 OK\r\ncontent-type: application/json\r\n\
content-length: {}\r\nconnection: close\r\n\r\n{body}",
body.len()
);
let _ = stream.write_all(response.as_bytes()).await;
let _ = stream.shutdown().await;
}
});
format!("http://{addr}/jwks")
}
#[cfg(feature = "oauth2")]
fn oidc_test_provider(jwks_url: String) -> OAuth2ProviderConfig {
OAuth2ProviderConfig {
client_id: "cid".into(),
client_secret: "secret".into(),
authorize_url: "https://idp.example/authorize".into(),
token_url: "https://idp.example/token".into(),
userinfo_url: None,
redirect_uri: "http://localhost:3000/callback".into(),
scope: "openid".into(),
issuer: Some("https://idp.example".into()),
jwks_url: Some(jwks_url),
discovery_url: None,
}
}
#[cfg(feature = "oauth2")]
fn unix_now_secs() -> u64 {
std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap()
.as_secs()
}
#[cfg(feature = "oauth2")]
#[tokio::test]
async fn validate_id_token_rejects_hs256_against_rsa_jwks_key() {
let jwks_body = serde_json::json!({ "keys": [rsa_test_jwk_json(None)] }).to_string();
let provider = oidc_test_provider(spawn_jwks_stub(jwks_body).await);
let mut header = jsonwebtoken::Header::new(jsonwebtoken::Algorithm::HS256);
header.kid = Some("test-kid".into());
let claims = serde_json::json!({
"sub": "attacker-controlled",
"iss": "https://idp.example",
"aud": "cid",
"exp": unix_now_secs() + 3600,
});
let token = jsonwebtoken::encode(
&header,
&claims,
&jsonwebtoken::EncodingKey::from_secret(b"guessed-public-key-material"),
)
.unwrap();
let err = validate_and_decode_id_token(&token, &provider)
.await
.unwrap_err();
assert!(
err.to_string().contains("not permitted by matching jwk"),
"expected algorithm pinning rejection, got: {err}"
);
}
#[cfg(feature = "oauth2")]
#[tokio::test]
async fn validate_id_token_rejects_alg_none() {
use base64::Engine as _;
let b64 = |v: &serde_json::Value| {
base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(v.to_string())
};
let header = b64(&serde_json::json!({"alg": "none", "kid": "test-kid", "typ": "JWT"}));
let payload = b64(&serde_json::json!({
"sub": "attacker-controlled",
"iss": "https://idp.example",
"aud": "cid",
"exp": unix_now_secs() + 3600,
}));
let token = format!("{header}.{payload}.");
let provider = oidc_test_provider("http://127.0.0.1:9/jwks".into());
let err = validate_and_decode_id_token(&token, &provider)
.await
.unwrap_err();
assert!(
err.to_string().contains("invalid id_token header"),
"expected header rejection for alg=none, got: {err}"
);
}
#[cfg(feature = "oauth2")]
#[test]
fn parse_oauth2_token_response_supports_form_encoded_payload() {
let token = parse_oauth2_token_response(
Some("application/x-www-form-urlencoded"),
"access_token=abc123&token_type=bearer&id_token=xyz789&extra_field=ignored",
)
.unwrap();
assert_eq!(token.access_token, "abc123");
assert_eq!(token.token_type.as_deref(), Some("bearer"));
assert_eq!(token.id_token.as_deref(), Some("xyz789"));
}
#[cfg(feature = "oauth2")]
#[test]
fn parse_oauth2_token_response_fails_without_access_token() {
let err = parse_oauth2_token_response(
Some("application/x-www-form-urlencoded"),
"token_type=bearer&id_token=xyz789",
)
.unwrap_err();
assert_eq!(err.to_string(), "token response missing access_token");
}
#[cfg(feature = "oauth2")]
#[test]
fn extract_subject_allows_userinfo_id_fallback() {
let claims = serde_json::json!({ "id": 42 });
let subject = extract_subject(&claims, IdentitySource::UserInfo).unwrap();
assert_eq!(subject, "42");
}
#[cfg(feature = "oauth2")]
#[tokio::test]
async fn validate_callback_state_preserves_state_on_mismatch() {
let session = crate::session::Session::new_for_test("s1".into(), HashMap::new());
session
.insert("oauth2:github:state".to_owned(), "real-state".to_owned())
.await;
let bad_callback = OAuth2Callback {
code: "c".into(),
state: "wrong-state".into(),
};
let err = validate_callback_state(&session, "github", &bad_callback)
.await
.unwrap_err();
assert!(err.to_string().contains("state mismatch"));
assert_eq!(
session.get("oauth2:github:state").await.as_deref(),
Some("real-state")
);
}
#[cfg(feature = "oauth2")]
#[tokio::test]
async fn validate_oidc_nonce_rejects_missing_nonce_for_id_token() {
let session = crate::session::Session::new_for_test("s1".into(), HashMap::new());
let claims = serde_json::json!({ "nonce": "any" });
let err = validate_oidc_nonce(&session, "github", &claims, IdentitySource::IdToken)
.await
.unwrap_err();
assert!(err.to_string().contains("nonce missing from session"));
}
#[cfg(feature = "oauth2")]
#[test]
fn extract_subject_requires_sub_for_id_token() {
let claims = serde_json::json!({ "id": "abc" });
let err = extract_subject(&claims, IdentitySource::IdToken).unwrap_err();
assert_eq!(err.to_string(), "missing sub claim");
}
#[test]
fn auth_rejection_is_401() {
let rejection = AuthRejection;
let response = rejection.into_response();
assert_eq!(response.status(), StatusCode::UNAUTHORIZED);
}
#[test]
fn auth_rejection_display() {
assert_eq!(AuthRejection.to_string(), "authentication required");
}
#[tokio::test]
async fn auth_extractor_returns_401_when_no_user() {
use axum::Router;
use axum::body::Body;
use axum::routing::get;
use tower::ServiceExt;
#[derive(Clone)]
struct TestUser {
name: String,
}
async fn handler(Auth(user): Auth<TestUser>) -> String {
user.name
}
let state = test_app_state("user_id");
let app = Router::new().route("/", get(handler)).with_state(state);
let response = app
.oneshot(
http::Request::builder()
.uri("/")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(response.status(), StatusCode::UNAUTHORIZED);
}
#[tokio::test]
async fn auth_extractor_returns_user_when_present() {
use axum::Router;
use axum::body::Body;
use axum::routing::get;
use tower::ServiceExt;
#[derive(Clone)]
struct TestUser {
name: String,
}
async fn handler(Auth(user): Auth<TestUser>) -> String {
user.name
}
let state = test_app_state("user_id");
let app = Router::new()
.route("/", get(handler))
.layer(axum::middleware::from_fn(
|mut req: axum::extract::Request, next: axum::middleware::Next| async move {
req.extensions_mut().insert(TestUser {
name: "alice".into(),
});
next.run(req).await
},
))
.with_state(state);
let response = app
.oneshot(
http::Request::builder()
.uri("/")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(response.status(), StatusCode::OK);
let body = axum::body::to_bytes(response.into_body(), usize::MAX)
.await
.unwrap();
assert_eq!(std::str::from_utf8(&body).unwrap(), "alice");
}
#[tokio::test]
async fn require_auth_rejects_unauthenticated() {
use axum::Router;
use axum::body::Body;
use axum::routing::get;
use tower::ServiceExt;
use crate::session::{MemoryStore, SessionConfig, SessionLayer};
let state = test_app_state("user_id");
let app = Router::new()
.route("/protected", get(|| async { "secret" }))
.layer(RequireAuth::new("user_id"))
.layer(SessionLayer::new(
MemoryStore::new(),
SessionConfig::default(),
))
.with_state(state);
let response = app
.oneshot(
http::Request::builder()
.uri("/protected")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(response.status(), StatusCode::UNAUTHORIZED);
}
#[tokio::test]
async fn check_secured_rejects_unauthenticated() {
let session =
crate::session::Session::new_for_test(String::new(), std::collections::HashMap::new());
let result = __check_secured(&session, &[]).await;
assert!(result.is_err());
let err = result.unwrap_err();
assert_eq!(err.status(), StatusCode::UNAUTHORIZED);
assert_eq!(err.to_string(), "authentication required");
}
#[tokio::test]
async fn check_secured_allows_authenticated() {
let data = std::collections::HashMap::from([("user_id".into(), "42".into())]);
let session = crate::session::Session::new_for_test("sess".into(), data);
let result = __check_secured(&session, &[]).await;
assert!(result.is_ok());
}
#[tokio::test]
async fn check_secured_rejects_wrong_role() {
let data = std::collections::HashMap::from([
("user_id".into(), "42".into()),
("role".into(), "viewer".into()),
]);
let session = crate::session::Session::new_for_test("sess".into(), data);
let result = __check_secured(&session, &["admin"]).await;
assert!(result.is_err());
let err = result.unwrap_err();
assert_eq!(err.status(), StatusCode::FORBIDDEN);
assert_eq!(err.to_string(), "insufficient permissions");
}
#[tokio::test]
async fn check_secured_allows_matching_role() {
let data = std::collections::HashMap::from([
("user_id".into(), "42".into()),
("role".into(), "admin".into()),
]);
let session = crate::session::Session::new_for_test("sess".into(), data);
let result = __check_secured(&session, &["admin"]).await;
assert!(result.is_ok());
}
#[tokio::test]
async fn check_secured_allows_any_of_multiple_roles() {
let data = std::collections::HashMap::from([
("user_id".into(), "42".into()),
("role".into(), "editor".into()),
]);
let session = crate::session::Session::new_for_test("sess".into(), data);
let result = __check_secured(&session, &["admin", "editor"]).await;
assert!(result.is_ok());
}
#[tokio::test]
async fn check_secured_seeds_actor_when_none_established() {
crate::current::scope_request(async {
assert_eq!(crate::current::Current::actor(), None);
let data = std::collections::HashMap::from([("user_id".into(), "42".into())]);
let session = crate::session::Session::new_for_test("sess".into(), data);
let result = __check_secured_with_key(&session, "user_id", &[]).await;
assert!(result.is_ok());
assert_eq!(crate::current::Current::actor(), Some("42".to_owned()));
})
.await;
}
#[tokio::test]
async fn check_secured_preserves_already_established_actor() {
crate::current::scope_request(async {
crate::current::Current::set_actor("token-principal".to_owned());
let data = std::collections::HashMap::from([("user_id".into(), "42".into())]);
let session = crate::session::Session::new_for_test("sess".into(), data);
let result = __check_secured_with_key(&session, "user_id", &[]).await;
assert!(result.is_ok());
assert_eq!(
crate::current::Current::actor(),
Some("token-principal".to_owned())
);
})
.await;
}
#[tokio::test]
async fn secured_macro_rejects_unauthenticated() {
use axum::Router;
use axum::body::Body;
use axum::routing::get;
use tower::ServiceExt;
use crate::session::{MemoryStore, SessionConfig, SessionLayer};
#[autumn_macros::secured]
async fn protected_handler() -> crate::AutumnResult<&'static str> {
Ok("secret")
}
let state = test_app_state("user_id");
let app = Router::new()
.route("/", get(protected_handler))
.layer(SessionLayer::new(
MemoryStore::new(),
SessionConfig::default(),
))
.with_state(state);
let response = app
.oneshot(
http::Request::builder()
.uri("/")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(response.status(), StatusCode::UNAUTHORIZED);
}
#[tokio::test]
async fn secured_macro_allows_authenticated() {
use axum::Router;
use axum::body::Body;
use axum::routing::get;
use http::header::COOKIE;
use tower::ServiceExt;
use crate::session::{MemoryStore, SessionConfig, SessionLayer, SessionStore};
#[autumn_macros::secured]
async fn protected_handler() -> crate::AutumnResult<&'static str> {
Ok("secret")
}
let store = MemoryStore::new();
store
.save(
"sess1",
std::collections::HashMap::from([("user_id".into(), "42".into())]),
)
.await
.unwrap();
let state = test_app_state("user_id");
let app = Router::new()
.route("/", get(protected_handler))
.layer(SessionLayer::new(store, SessionConfig::default()))
.with_state(state);
let response = app
.oneshot(
http::Request::builder()
.uri("/")
.header(COOKIE, "autumn.sid=sess1")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(response.status(), StatusCode::OK);
let body = axum::body::to_bytes(response.into_body(), usize::MAX)
.await
.unwrap();
assert_eq!(std::str::from_utf8(&body).unwrap(), "secret");
}
#[tokio::test]
async fn secured_macro_honors_configured_auth_session_key() {
use axum::Router;
use axum::body::Body;
use axum::routing::get;
use http::header::COOKIE;
use tower::ServiceExt;
use crate::session::{MemoryStore, SessionConfig, SessionLayer, SessionStore};
#[autumn_macros::secured]
async fn account_handler() -> crate::AutumnResult<&'static str> {
Ok("account")
}
let store = MemoryStore::new();
store
.save(
"sess1",
std::collections::HashMap::from([
("uid".into(), "42".into()),
("account_id".into(), "42".into()),
]),
)
.await
.unwrap();
let state = test_app_state("uid");
let app = Router::new()
.route("/account", get(account_handler))
.layer(SessionLayer::new(store, SessionConfig::default()))
.with_state(state);
let response = app
.oneshot(
http::Request::builder()
.uri("/account")
.header(COOKIE, "autumn.sid=sess1")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(response.status(), StatusCode::OK);
let body = axum::body::to_bytes(response.into_body(), usize::MAX)
.await
.unwrap();
assert_eq!(std::str::from_utf8(&body).unwrap(), "account");
}
#[tokio::test]
async fn secured_macro_with_role_rejects_wrong_role() {
use axum::Router;
use axum::body::Body;
use axum::routing::get;
use http::header::COOKIE;
use tower::ServiceExt;
use crate::session::{MemoryStore, SessionConfig, SessionLayer, SessionStore};
#[autumn_macros::secured("admin")]
async fn admin_only() -> crate::AutumnResult<&'static str> {
Ok("admin area")
}
let store = MemoryStore::new();
store
.save(
"sess1",
std::collections::HashMap::from([
("user_id".into(), "42".into()),
("role".into(), "viewer".into()),
]),
)
.await
.unwrap();
let state = test_app_state("user_id");
let app = Router::new()
.route("/", get(admin_only))
.layer(SessionLayer::new(store, SessionConfig::default()))
.with_state(state);
let response = app
.oneshot(
http::Request::builder()
.uri("/")
.header(COOKIE, "autumn.sid=sess1")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(response.status(), StatusCode::FORBIDDEN);
}
#[tokio::test]
async fn secured_macro_with_multiple_roles_allows_match() {
use axum::Router;
use axum::body::Body;
use axum::routing::get;
use http::header::COOKIE;
use tower::ServiceExt;
use crate::session::{MemoryStore, SessionConfig, SessionLayer, SessionStore};
#[autumn_macros::secured("admin", "editor")]
async fn content_handler() -> crate::AutumnResult<&'static str> {
Ok("content")
}
let store = MemoryStore::new();
store
.save(
"sess1",
std::collections::HashMap::from([
("user_id".into(), "42".into()),
("role".into(), "editor".into()),
]),
)
.await
.unwrap();
let state = test_app_state("user_id");
let app = Router::new()
.route("/", get(content_handler))
.layer(SessionLayer::new(store, SessionConfig::default()))
.with_state(state);
let response = app
.oneshot(
http::Request::builder()
.uri("/")
.header(COOKIE, "autumn.sid=sess1")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(response.status(), StatusCode::OK);
let body = axum::body::to_bytes(response.into_body(), usize::MAX)
.await
.unwrap();
assert_eq!(std::str::from_utf8(&body).unwrap(), "content");
}
#[tokio::test]
async fn require_auth_allows_authenticated() {
use axum::Router;
use axum::body::Body;
use axum::routing::get;
use http::header::COOKIE;
use tower::ServiceExt;
use crate::session::{MemoryStore, SessionConfig, SessionLayer, SessionStore};
let store = MemoryStore::new();
let mut session_data = std::collections::HashMap::new();
session_data.insert("user_id".into(), "42".into());
store.save("valid-session", session_data).await.unwrap();
let state = test_app_state("user_id");
let app = Router::new()
.route("/protected", get(|| async { "secret" }))
.layer(RequireAuth::new("user_id"))
.layer(SessionLayer::new(store, SessionConfig::default()))
.with_state(state);
let response = app
.oneshot(
http::Request::builder()
.uri("/protected")
.header(COOKIE, "autumn.sid=valid-session")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(response.status(), StatusCode::OK);
let body = axum::body::to_bytes(response.into_body(), usize::MAX)
.await
.unwrap();
assert_eq!(std::str::from_utf8(&body).unwrap(), "secret");
}
#[tokio::test]
async fn require_auth_sets_rate_limit_principal() {
use axum::Router;
use axum::body::Body;
use axum::routing::get;
use http::header::COOKIE;
use tower::ServiceExt;
use crate::security::RateLimitPrincipal;
use crate::session::{MemoryStore, SessionConfig, SessionLayer, SessionStore};
async fn handler(
axum::Extension(principal): axum::Extension<RateLimitPrincipal>,
) -> String {
principal.0
}
let store = MemoryStore::new();
let mut session_data = std::collections::HashMap::new();
session_data.insert("user_id".into(), "42".into());
store.save("valid-session", session_data).await.unwrap();
let state = test_app_state("user_id");
let app = Router::new()
.route("/protected", get(handler))
.layer(RequireAuth::new("user_id"))
.layer(SessionLayer::new(store, SessionConfig::default()))
.with_state(state);
let response = app
.oneshot(
http::Request::builder()
.uri("/protected")
.header(COOKIE, "autumn.sid=valid-session")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(response.status(), StatusCode::OK);
let body = axum::body::to_bytes(response.into_body(), usize::MAX)
.await
.unwrap();
assert_eq!(std::str::from_utf8(&body).unwrap(), "42");
}
#[tokio::test]
async fn require_auth_rate_limits_by_session_principal() {
use axum::Router;
use axum::body::Body;
use axum::routing::get;
use http::header::COOKIE;
use tower::ServiceExt;
use crate::security::{KeyStrategy, RateLimitConfig, RateLimitLayer};
use crate::session::{MemoryStore, SessionConfig, SessionLayer, SessionStore};
let session_store = MemoryStore::new();
let mut data_a = std::collections::HashMap::new();
data_a.insert("user_id".into(), "user-1".into());
session_store.save("sess-a", data_a).await.unwrap();
let mut data_b = std::collections::HashMap::new();
data_b.insert("user_id".into(), "user-2".into());
session_store.save("sess-b", data_b).await.unwrap();
let rl_config = RateLimitConfig {
enabled: true,
requests_per_second: 0.1,
burst: 1,
key_strategy: KeyStrategy::AuthenticatedPrincipal,
..Default::default()
};
let state = test_app_state("user_id");
let app = Router::new()
.route("/protected", get(|| async { "ok" }))
.layer(RateLimitLayer::from_config(&rl_config)) .layer(RequireAuth::new("user_id")) .layer(SessionLayer::new(session_store, SessionConfig::default()))
.with_state(state);
let r = app
.clone()
.oneshot(
http::Request::builder()
.uri("/protected")
.header(COOKIE, "autumn.sid=sess-a")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(r.status(), StatusCode::OK);
let r = app
.clone()
.oneshot(
http::Request::builder()
.uri("/protected")
.header(COOKIE, "autumn.sid=sess-a")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(r.status(), StatusCode::TOO_MANY_REQUESTS);
let r = app
.clone()
.oneshot(
http::Request::builder()
.uri("/protected")
.header(COOKIE, "autumn.sid=sess-b")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(r.status(), StatusCode::OK);
}
#[tokio::test]
async fn require_auth_poll_ready_propagates() {
use std::task::{Context, Poll};
use tower::{Layer, Service};
#[derive(Clone)]
struct MockService {
ready: bool,
poll_count: std::sync::Arc<std::sync::atomic::AtomicUsize>,
}
impl Service<axum::extract::Request> for MockService {
type Response = axum::response::Response;
type Error = std::convert::Infallible;
type Future = std::future::Ready<Result<Self::Response, Self::Error>>;
fn poll_ready(&mut self, _cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
self.poll_count
.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
if self.ready {
Poll::Ready(Ok(()))
} else {
Poll::Pending
}
}
fn call(&mut self, _req: axum::extract::Request) -> Self::Future {
std::future::ready(Ok(axum::response::Response::new(axum::body::Body::empty())))
}
}
let layer = RequireAuth::new("user_id");
let poll_count = std::sync::Arc::new(std::sync::atomic::AtomicUsize::new(0));
let mock_service = MockService {
ready: false,
poll_count: poll_count.clone(),
};
let mut service = layer.layer(mock_service);
let waker = futures::task::noop_waker();
let mut cx = Context::from_waker(&waker);
let poll = service.poll_ready(&mut cx);
assert!(poll.is_pending());
assert_eq!(poll_count.load(std::sync::atomic::Ordering::SeqCst), 1);
let mock_service_ready = MockService {
ready: true,
poll_count: poll_count.clone(),
};
let mut service_ready = layer.layer(mock_service_ready);
let poll_ready = service_ready.poll_ready(&mut cx);
assert!(poll_ready.is_ready());
assert_eq!(poll_count.load(std::sync::atomic::Ordering::SeqCst), 2);
}
#[tokio::test]
async fn auth_rejection_into_response() {
let rejection = AuthRejection;
let response = rejection.into_response();
assert_eq!(response.status(), axum::http::StatusCode::UNAUTHORIZED);
let body = axum::body::to_bytes(response.into_body(), usize::MAX)
.await
.unwrap();
let json: serde_json::Value = serde_json::from_slice(&body).unwrap();
assert_eq!(json["status"], 401);
assert_eq!(json["detail"], "authentication required");
assert_eq!(json["code"], "autumn.unauthorized");
}
#[test]
fn test_auth_config_defaults() {
let config = AuthConfig::default();
assert_eq!(config.bcrypt_cost, DEFAULT_BCRYPT_COST);
assert_eq!(config.session_key, "user_id");
}
#[tokio::test]
async fn test_hash_password() {
let test_input = uuid::Uuid::new_v4().to_string();
let hash = super::hash_password(&test_input)
.await
.expect("Failed to hash password");
assert!(hash.starts_with("$2b$"));
let is_valid = super::verify_password(&test_input, &hash)
.await
.expect("Failed to verify password");
assert!(is_valid, "Password should be verified successfully");
let is_invalid = super::verify_password(&uuid::Uuid::new_v4().to_string(), &hash)
.await
.expect("Failed to verify wrong password");
assert!(!is_invalid, "Wrong password should not be verified");
}
#[tokio::test]
async fn test_hash_password_empty() {
let test_input = String::new();
let hash = super::hash_password(&test_input)
.await
.expect("Failed to hash empty password");
assert!(hash.starts_with("$2b$"));
let is_valid = super::verify_password(&test_input, &hash)
.await
.expect("Failed to verify empty password");
assert!(is_valid, "Empty password should be verified successfully");
}
#[tokio::test]
async fn test_hash_password_long() {
let test_input = "a".repeat(100);
let hash = super::hash_password(&test_input)
.await
.expect("Failed to hash long password");
assert!(hash.starts_with("$2b$"));
let is_valid = super::verify_password(&test_input, &hash)
.await
.expect("Failed to verify long password");
assert!(is_valid, "Long password should be verified successfully");
}
#[tokio::test]
async fn test_hash_password_unicode() {
let test_input = format!("{}🚀my_secrët_passwörd🔑", uuid::Uuid::new_v4());
let hash = super::hash_password(&test_input)
.await
.expect("Failed to hash unicode password");
assert!(hash.starts_with("$2b$"));
let is_valid = super::verify_password(&test_input, &hash)
.await
.expect("Failed to verify unicode password");
assert!(is_valid, "Unicode password should be verified successfully");
}
#[tokio::test]
async fn test_verify_password_invalid_hash() {
let test_input = uuid::Uuid::new_v4().to_string();
let result = super::verify_password(&test_input, "invalid_hash_string").await;
assert!(result.is_err() || !result.unwrap());
let result2 = super::verify_password(&test_input, "$2b$04$").await;
assert!(result2.is_err() || !result2.unwrap());
}
}
#[cfg(feature = "oauth2")]
#[cfg(test)]
mod http_interceptor_task_local_tests {
use crate::interceptor::{ACTIVE_HTTP_INTERCEPTORS, HttpInterceptor, HttpInterceptorFuture};
use std::sync::{
Arc,
atomic::{AtomicBool, Ordering},
};
struct FlagInterceptor {
fired: Arc<AtomicBool>,
}
impl HttpInterceptor for FlagInterceptor {
fn intercept<'a>(
&'a self,
req: reqwest::Request,
next: &'a dyn Fn(reqwest::Request) -> HttpInterceptorFuture<'a>,
) -> HttpInterceptorFuture<'a> {
self.fired.store(true, Ordering::SeqCst);
next(req)
}
}
#[tokio::test]
async fn http_request_builder_send_fires_interceptor_inside_scope() {
let fired = Arc::new(AtomicBool::new(false));
let interceptor: Arc<dyn HttpInterceptor> = Arc::new(FlagInterceptor {
fired: Arc::clone(&fired),
});
let client = reqwest::Client::new();
let http_client = super::HttpClient::new(client);
ACTIVE_HTTP_INTERCEPTORS
.scope(vec![interceptor], async {
let _ = http_client
.get("http://127.0.0.1:54321/noreply")
.send()
.await;
})
.await;
assert!(
fired.load(Ordering::SeqCst),
"interceptor must fire when ACTIVE_HTTP_INTERCEPTORS scope is established"
);
}
#[tokio::test]
async fn http_request_builder_send_skips_interceptor_outside_scope() {
let fired = Arc::new(AtomicBool::new(false));
let _interceptor: Arc<dyn HttpInterceptor> = Arc::new(FlagInterceptor {
fired: Arc::clone(&fired),
});
let client = reqwest::Client::new();
let http_client = super::HttpClient::new(client);
let _ = http_client
.get("http://127.0.0.1:54321/noreply")
.send()
.await;
assert!(
!fired.load(Ordering::SeqCst),
"interceptor must NOT fire when ACTIVE_HTTP_INTERCEPTORS scope is absent"
);
}
}
#[cfg(test)]
mod api_token_tests {
use std::sync::Arc;
use http::StatusCode;
use super::{
ApiToken, ApiTokenStore, InMemoryApiTokenStore, RequireApiToken, hash_api_token,
issue_api_token, revoke_api_token,
};
struct FailingApiTokenStore;
impl ApiTokenStore for FailingApiTokenStore {
fn issue<'a>(
&'a self,
_principal_id: &'a str,
) -> std::pin::Pin<
Box<dyn std::future::Future<Output = crate::AutumnResult<String>> + Send + 'a>,
> {
Box::pin(async {
Err(crate::AutumnError::service_unavailable_msg(
"api token store unavailable",
))
})
}
fn verify<'a>(
&'a self,
_raw_token: &'a str,
) -> std::pin::Pin<
Box<dyn std::future::Future<Output = crate::AutumnResult<Option<String>>> + Send + 'a>,
> {
Box::pin(async {
Err(crate::AutumnError::service_unavailable_msg(
"api token store unavailable",
))
})
}
fn revoke<'a>(
&'a self,
_raw_token: &'a str,
) -> std::pin::Pin<Box<dyn std::future::Future<Output = crate::AutumnResult<()>> + Send + 'a>>
{
Box::pin(async {
Err(crate::AutumnError::service_unavailable_msg(
"api token store unavailable",
))
})
}
}
#[test]
fn hash_api_token_is_deterministic() {
let h1 = hash_api_token("abc123");
let h2 = hash_api_token("abc123");
assert_eq!(h1, h2);
}
#[test]
fn hash_api_token_produces_64_char_hex() {
let hash = hash_api_token("any_raw_token");
assert_eq!(hash.len(), 64, "SHA-256 hex must be 64 chars");
assert!(
hash.chars().all(|c| c.is_ascii_hexdigit()),
"hash must be lowercase hex digits"
);
}
#[test]
fn hash_api_token_differs_from_input() {
let raw = "my_raw_token";
assert_ne!(hash_api_token(raw), raw);
}
#[test]
fn hash_api_token_different_inputs_produce_different_hashes() {
assert_ne!(hash_api_token("token_a"), hash_api_token("token_b"));
}
#[tokio::test]
async fn in_memory_store_issue_returns_unique_tokens() {
let store = InMemoryApiTokenStore::default();
let t1 = store.issue("user:1").await.unwrap();
let t2 = store.issue("user:1").await.unwrap();
assert_ne!(t1, t2, "each issued token must be unique");
assert!(t1.len() >= 32, "token must have sufficient entropy");
}
#[tokio::test]
async fn in_memory_store_verify_returns_principal_for_valid_token() {
let store = InMemoryApiTokenStore::default();
let raw = store.issue("user:42").await.unwrap();
let principal = store.verify(&raw).await.unwrap();
assert_eq!(principal, Some("user:42".to_owned()));
}
#[tokio::test]
async fn in_memory_store_verify_returns_none_for_unknown_token() {
let store = InMemoryApiTokenStore::default();
let result = store.verify("not_a_real_token").await.unwrap();
assert_eq!(result, None);
}
#[tokio::test]
async fn in_memory_store_revoke_invalidates_token() {
let store = InMemoryApiTokenStore::default();
let raw = store.issue("user:7").await.unwrap();
assert_eq!(
store.verify(&raw).await.unwrap(),
Some("user:7".to_owned()),
"token must be valid before revoking"
);
store.revoke(&raw).await.unwrap();
assert_eq!(store.verify(&raw).await.unwrap(), None);
}
#[tokio::test]
async fn in_memory_store_raw_token_not_stored_verbatim() {
let store = InMemoryApiTokenStore::default();
let raw = store.issue("user:1").await.unwrap();
let tampered = format!("{raw}x");
assert_eq!(store.verify(&tampered).await.unwrap(), None);
}
#[tokio::test]
async fn issue_api_token_helper_issues_verifiable_token() {
let store = InMemoryApiTokenStore::default();
let raw = issue_api_token(&store, "user:5").await.unwrap();
assert_eq!(store.verify(&raw).await.unwrap(), Some("user:5".to_owned()));
}
#[tokio::test]
async fn revoke_api_token_helper_revokes_token() {
let store = InMemoryApiTokenStore::default();
let raw = store.issue("user:6").await.unwrap();
revoke_api_token(&store, &raw).await.unwrap();
assert_eq!(store.verify(&raw).await.unwrap(), None);
}
#[tokio::test]
async fn with_token_seeds_token_resolvable_through_same_path() {
let store = InMemoryApiTokenStore::default().with_token("known-dev-token", "user:dev");
assert_eq!(
store.verify("known-dev-token").await.unwrap(),
Some("user:dev".to_owned()),
);
assert_eq!(store.verify("known-dev-tokenx").await.unwrap(), None);
let verified = store
.verify_scoped("known-dev-token")
.await
.unwrap()
.unwrap();
assert_eq!(verified.principal_id, "user:dev");
assert!(verified.scopes.is_empty());
store.revoke("known-dev-token").await.unwrap();
assert_eq!(store.verify("known-dev-token").await.unwrap(), None);
}
#[tokio::test]
async fn with_scoped_token_seeds_scopes_via_verify_scoped() {
let granted = scopes(&["reports:read", "reports:write"]);
let store = InMemoryApiTokenStore::default().with_scoped_token(
"scoped-dev-token",
"svc:reports",
&granted,
);
let verified = store
.verify_scoped("scoped-dev-token")
.await
.unwrap()
.unwrap();
assert_eq!(verified.principal_id, "svc:reports");
assert_eq!(verified.scopes, granted);
}
#[tokio::test]
async fn blank_seed_token_stores_no_credential() {
for blank in ["", " ", "\t\n"] {
let store = InMemoryApiTokenStore::default().with_token(blank, "user:oops");
assert_eq!(store.verify(blank).await.unwrap(), None);
assert_eq!(store.verify("").await.unwrap(), None);
assert!(store.verify_scoped(blank).await.unwrap().is_none());
assert!(store.verify_scoped("").await.unwrap().is_none());
let scoped = InMemoryApiTokenStore::default().with_scoped_token(
blank,
"svc:oops",
&scopes(&["reports:read"]),
);
assert_eq!(scoped.verify(blank).await.unwrap(), None);
assert!(scoped.verify_scoped("").await.unwrap().is_none());
}
let store = InMemoryApiTokenStore::default().with_token("real-token", "user:ok");
assert_eq!(
store.verify("real-token").await.unwrap(),
Some("user:ok".to_owned()),
);
}
#[test]
fn from_env_errors_when_variable_unset() {
const VAR: &str = "AUTUMN_TEST_MCP_TOKEN_1970_UNSET";
assert!(std::env::var(VAR).is_err(), "test var must be unset");
assert!(InMemoryApiTokenStore::from_env(VAR, "user:mcp").is_err());
}
use super::{
__check_secured_scopes, ApiTokenScopes, IssueTokenSpec, issue_scoped_api_token,
list_api_tokens, rotate_api_token,
};
use crate::time::{FixedClock, TickingClock};
use chrono::{Duration as ChronoDuration, TimeZone as _, Utc};
fn scopes(s: &[&str]) -> Vec<String> {
s.iter().map(|x| (*x).to_owned()).collect()
}
#[derive(Default)]
struct LegacyOnlyStore(InMemoryApiTokenStore);
impl ApiTokenStore for LegacyOnlyStore {
fn issue<'a>(
&'a self,
principal_id: &'a str,
) -> std::pin::Pin<
Box<dyn std::future::Future<Output = crate::AutumnResult<String>> + Send + 'a>,
> {
self.0.issue(principal_id)
}
fn verify<'a>(
&'a self,
raw_token: &'a str,
) -> std::pin::Pin<
Box<dyn std::future::Future<Output = crate::AutumnResult<Option<String>>> + Send + 'a>,
> {
self.0.verify(raw_token)
}
fn revoke<'a>(
&'a self,
raw_token: &'a str,
) -> std::pin::Pin<Box<dyn std::future::Future<Output = crate::AutumnResult<()>> + Send + 'a>>
{
self.0.revoke(raw_token)
}
}
#[tokio::test]
async fn legacy_store_default_verify_scoped_yields_empty_scopes() {
let store = LegacyOnlyStore::default();
let raw = store.issue("user:1").await.unwrap();
let verified = store.verify_scoped(&raw).await.unwrap().unwrap();
assert_eq!(verified.principal_id, "user:1");
assert!(verified.scopes.is_empty());
}
#[tokio::test]
async fn issue_scoped_round_trips_name_and_scopes() {
let store = InMemoryApiTokenStore::default();
let granted = scopes(&["posts:read", "posts:write"]);
let raw = issue_scoped_api_token(
&store,
IssueTokenSpec {
principal_id: "service:ci",
name: "ci",
scopes: &granted,
expires_at: None,
},
)
.await
.unwrap();
let verified = store.verify_scoped(&raw).await.unwrap().unwrap();
assert_eq!(verified.principal_id, "service:ci");
assert_eq!(verified.scopes, granted);
let listed = list_api_tokens(&store, "service:ci").await.unwrap();
assert_eq!(listed.len(), 1);
assert_eq!(listed[0].name, "ci");
assert_eq!(listed[0].scopes, granted);
}
#[tokio::test]
async fn expired_token_verifies_as_none() {
let now = Utc.with_ymd_and_hms(2026, 1, 1, 0, 0, 0).unwrap();
let store = InMemoryApiTokenStore::default().with_clock(Arc::new(FixedClock::at(now)));
let granted = scopes(&["posts:read"]);
let raw = store
.issue_scoped(IssueTokenSpec {
principal_id: "service:ci",
name: "ci",
scopes: &granted,
expires_at: Some(now - ChronoDuration::seconds(1)),
})
.await
.unwrap();
assert_eq!(store.verify(&raw).await.unwrap(), None);
assert!(store.verify_scoped(&raw).await.unwrap().is_none());
}
#[tokio::test]
async fn unexpired_token_verifies_then_records_last_used_at() {
let start = Utc.with_ymd_and_hms(2026, 1, 1, 0, 0, 0).unwrap();
let clock = TickingClock::starting_at(start);
let store = InMemoryApiTokenStore::default().with_clock(Arc::new(clock));
let granted = scopes(&["posts:read"]);
let raw = store
.issue_scoped(IssueTokenSpec {
principal_id: "service:ci",
name: "ci",
scopes: &granted,
expires_at: Some(start + ChronoDuration::days(30)),
})
.await
.unwrap();
assert!(
list_api_tokens(&store, "service:ci").await.unwrap()[0]
.last_used_at
.is_none()
);
assert!(store.verify_scoped(&raw).await.unwrap().is_some());
assert!(
list_api_tokens(&store, "service:ci").await.unwrap()[0]
.last_used_at
.is_some()
);
}
#[tokio::test]
async fn list_metadata_carries_no_secret_and_reflects_revocation() {
let store = InMemoryApiTokenStore::default();
let granted = scopes(&["posts:read"]);
let raw = store
.issue_scoped(IssueTokenSpec {
principal_id: "service:ci",
name: "ci",
scopes: &granted,
expires_at: None,
})
.await
.unwrap();
let listed = list_api_tokens(&store, "service:ci").await.unwrap();
assert_eq!(listed.len(), 1);
assert!(listed[0].revoked_at.is_none());
store.revoke(&raw).await.unwrap();
assert_eq!(store.verify(&raw).await.unwrap(), None);
let listed = list_api_tokens(&store, "service:ci").await.unwrap();
assert!(listed[0].revoked_at.is_some());
}
#[tokio::test]
async fn rotate_revokes_old_and_preserves_scopes() {
let store = InMemoryApiTokenStore::default();
let granted = scopes(&["posts:read", "posts:write"]);
let old = store
.issue_scoped(IssueTokenSpec {
principal_id: "service:ci",
name: "ci",
scopes: &granted,
expires_at: None,
})
.await
.unwrap();
let new = rotate_api_token(&store, &old).await.unwrap().unwrap();
assert_ne!(new, old);
assert!(store.verify_scoped(&old).await.unwrap().is_none());
let verified = store.verify_scoped(&new).await.unwrap().unwrap();
assert_eq!(verified.scopes, granted);
assert_eq!(verified.principal_id, "service:ci");
assert!(rotate_api_token(&store, "nope").await.unwrap().is_none());
}
#[tokio::test]
async fn check_secured_scopes_is_default_deny_and_all_must_match() {
assert!(__check_secured_scopes(None, &[]).await.is_ok());
let err = __check_secured_scopes(None, &["posts:write"])
.await
.unwrap_err();
assert_eq!(err.status(), StatusCode::FORBIDDEN);
let granted = ApiTokenScopes(scopes(&["posts:read", "posts:write"]));
assert!(
__check_secured_scopes(Some(&granted), &["posts:write"])
.await
.is_ok()
);
assert!(
__check_secured_scopes(Some(&granted), &["posts:write", "posts:delete"])
.await
.is_err()
);
}
#[tokio::test]
async fn require_api_token_rejects_missing_authorization_header() {
use axum::body::Body;
use tower::ServiceExt;
let store = Arc::new(InMemoryApiTokenStore::default());
let app = axum::Router::new()
.route("/", axum::routing::get(|| async { "ok" }))
.layer(RequireApiToken::new(store));
let response = app
.oneshot(
http::Request::builder()
.uri("/")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(response.status(), StatusCode::UNAUTHORIZED);
}
#[tokio::test]
async fn require_api_token_rejects_non_bearer_scheme() {
use axum::body::Body;
use tower::ServiceExt;
let store = Arc::new(InMemoryApiTokenStore::default());
let app = axum::Router::new()
.route("/", axum::routing::get(|| async { "ok" }))
.layer(RequireApiToken::new(store));
let response = app
.oneshot(
http::Request::builder()
.uri("/")
.header(http::header::AUTHORIZATION, "Basic dXNlcjpwYXNz")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(response.status(), StatusCode::UNAUTHORIZED);
}
#[tokio::test]
async fn require_api_token_rejects_unknown_bearer_token() {
use axum::body::Body;
use tower::ServiceExt;
let store = Arc::new(InMemoryApiTokenStore::default());
let app = axum::Router::new()
.route("/", axum::routing::get(|| async { "ok" }))
.layer(RequireApiToken::new(store));
let response = app
.oneshot(
http::Request::builder()
.uri("/")
.header(http::header::AUTHORIZATION, "Bearer unknown_token_xyz")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(response.status(), StatusCode::UNAUTHORIZED);
}
#[tokio::test]
async fn require_api_token_propagates_store_verify_errors() {
use axum::body::Body;
use tower::ServiceExt;
let store = Arc::new(FailingApiTokenStore);
let app = axum::Router::new()
.route("/", axum::routing::get(|| async { "ok" }))
.layer(RequireApiToken::new(store));
let response = app
.oneshot(
http::Request::builder()
.uri("/")
.header(http::header::AUTHORIZATION, "Bearer valid_client_token")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(response.status(), StatusCode::SERVICE_UNAVAILABLE);
assert_eq!(
response
.headers()
.get(http::header::CONTENT_TYPE)
.map(|value| value.to_str().unwrap_or_default()),
Some("application/problem+json")
);
let body = axum::body::to_bytes(response.into_body(), usize::MAX)
.await
.unwrap();
let json: serde_json::Value = serde_json::from_slice(&body).unwrap();
assert_eq!(json["status"], 503);
assert_eq!(json["code"], "autumn.service_unavailable");
assert_eq!(json["detail"], "api token store unavailable");
}
#[tokio::test]
async fn require_api_token_allows_valid_bearer_token() {
use axum::body::Body;
use tower::ServiceExt;
let store = Arc::new(InMemoryApiTokenStore::default());
let raw = store.issue("user:1").await.unwrap();
let app = axum::Router::new()
.route("/", axum::routing::get(|| async { "ok" }))
.layer(RequireApiToken::new(Arc::clone(&store)));
let response = app
.oneshot(
http::Request::builder()
.uri("/")
.header(http::header::AUTHORIZATION, format!("Bearer {raw}"))
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(response.status(), StatusCode::OK);
}
#[tokio::test]
async fn require_api_token_accepts_case_insensitive_bearer_scheme() {
use axum::body::Body;
use tower::ServiceExt;
let store = Arc::new(InMemoryApiTokenStore::default());
let raw = store.issue("user:1").await.unwrap();
for scheme in ["bearer", "bEaReR"] {
let app = axum::Router::new()
.route("/", axum::routing::get(|| async { "ok" }))
.layer(RequireApiToken::new(Arc::clone(&store)));
let response = app
.oneshot(
http::Request::builder()
.uri("/")
.header(http::header::AUTHORIZATION, format!("{scheme} {raw}"))
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(response.status(), StatusCode::OK, "scheme {scheme}");
}
}
#[tokio::test]
async fn require_api_token_rejects_revoked_token() {
use axum::body::Body;
use tower::ServiceExt;
let store = Arc::new(InMemoryApiTokenStore::default());
let raw = store.issue("user:1").await.unwrap();
store.revoke(&raw).await.unwrap();
let app = axum::Router::new()
.route("/", axum::routing::get(|| async { "ok" }))
.layer(RequireApiToken::new(Arc::clone(&store)));
let response = app
.oneshot(
http::Request::builder()
.uri("/")
.header(http::header::AUTHORIZATION, format!("Bearer {raw}"))
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(response.status(), StatusCode::UNAUTHORIZED);
}
#[tokio::test]
async fn require_api_token_401_response_has_problem_details() {
use axum::body::Body;
use tower::ServiceExt;
let store = Arc::new(InMemoryApiTokenStore::default());
let app = axum::Router::new()
.route("/", axum::routing::get(|| async { "ok" }))
.layer(RequireApiToken::new(store));
let response = app
.oneshot(
http::Request::builder()
.uri("/")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(response.status(), StatusCode::UNAUTHORIZED);
assert_eq!(
response
.headers()
.get(http::header::CONTENT_TYPE)
.map(|v| v.to_str().unwrap_or_default()),
Some("application/problem+json")
);
let body = axum::body::to_bytes(response.into_body(), usize::MAX)
.await
.unwrap();
let json: serde_json::Value = serde_json::from_slice(&body).unwrap();
assert_eq!(json["status"], 401);
assert_eq!(json["code"], "autumn.unauthorized");
assert!(json["detail"].as_str().is_some());
}
#[tokio::test]
async fn require_api_token_401_problem_details_include_request_context() {
use crate::middleware::RequestIdLayer;
use axum::body::Body;
use tower::ServiceExt;
let store = Arc::new(InMemoryApiTokenStore::default());
let app = axum::Router::new()
.route("/api/private", axum::routing::get(|| async { "ok" }))
.layer(RequireApiToken::new(store))
.layer(RequestIdLayer);
let response = app
.oneshot(
http::Request::builder()
.uri("/api/private")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(response.status(), StatusCode::UNAUTHORIZED);
let request_id = response
.headers()
.get("x-request-id")
.and_then(|value| value.to_str().ok())
.expect("request id header should be present")
.to_owned();
let body = axum::body::to_bytes(response.into_body(), usize::MAX)
.await
.unwrap();
let json: serde_json::Value = serde_json::from_slice(&body).unwrap();
assert_eq!(json["request_id"], request_id);
assert_eq!(json["instance"], "/api/private");
}
#[tokio::test]
async fn api_token_extractor_yields_principal_id_to_handler() {
use axum::body::Body;
use tower::ServiceExt;
async fn handler(ApiToken(principal): ApiToken) -> String {
principal
}
let store = Arc::new(InMemoryApiTokenStore::default());
let raw = store.issue("user:99").await.unwrap();
let app = axum::Router::new()
.route("/", axum::routing::get(handler))
.layer(RequireApiToken::new(Arc::clone(&store)));
let response = app
.oneshot(
http::Request::builder()
.uri("/")
.header(http::header::AUTHORIZATION, format!("Bearer {raw}"))
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(response.status(), StatusCode::OK);
let body = axum::body::to_bytes(response.into_body(), usize::MAX)
.await
.unwrap();
assert_eq!(std::str::from_utf8(&body).unwrap(), "user:99");
}
#[tokio::test]
async fn api_token_extractor_rejects_when_no_principal_in_extensions() {
use axum::body::Body;
use tower::ServiceExt;
async fn handler(ApiToken(principal): ApiToken) -> String {
principal
}
let app = axum::Router::new().route("/", axum::routing::get(handler));
let response = app
.oneshot(
http::Request::builder()
.uri("/")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(response.status(), StatusCode::UNAUTHORIZED);
}
#[tokio::test]
async fn api_token_and_session_auth_compose_without_conflict() {
use axum::body::Body;
use tower::ServiceExt;
use crate::session::{MemoryStore, SessionConfig, SessionLayer, SessionStore};
async fn api_handler(ApiToken(principal): ApiToken) -> String {
principal
}
let store = Arc::new(InMemoryApiTokenStore::default());
let raw = store.issue("api_user").await.unwrap();
let session_store = MemoryStore::new();
session_store
.save(
"sess1",
std::collections::HashMap::from([("user_id".into(), "session_user".into())]),
)
.await
.unwrap();
let app = axum::Router::new()
.route(
"/api",
axum::routing::get(api_handler).layer(RequireApiToken::new(Arc::clone(&store))),
)
.layer(SessionLayer::new(session_store, SessionConfig::default()));
let response = app
.oneshot(
http::Request::builder()
.uri("/api")
.header(http::header::AUTHORIZATION, format!("Bearer {raw}"))
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(response.status(), StatusCode::OK);
let body = axum::body::to_bytes(response.into_body(), usize::MAX)
.await
.unwrap();
assert_eq!(std::str::from_utf8(&body).unwrap(), "api_user");
}
#[tokio::test]
async fn require_api_token_poll_ready_propagates_to_inner() {
use std::task::{Context, Poll};
use tower::{Layer, Service};
#[derive(Clone)]
struct MockService {
ready: bool,
}
impl tower::Service<axum::extract::Request> for MockService {
type Response = axum::response::Response;
type Error = std::convert::Infallible;
type Future = std::future::Ready<Result<Self::Response, Self::Error>>;
fn poll_ready(&mut self, _cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
if self.ready {
Poll::Ready(Ok(()))
} else {
Poll::Pending
}
}
fn call(&mut self, _req: axum::extract::Request) -> Self::Future {
std::future::ready(Ok(axum::response::Response::new(axum::body::Body::empty())))
}
}
let waker = futures::task::noop_waker();
let mut cx = Context::from_waker(&waker);
let store = Arc::new(InMemoryApiTokenStore::default());
let layer = RequireApiToken::new(store);
let mut svc = layer.layer(MockService { ready: false });
assert!(svc.poll_ready(&mut cx).is_pending());
let store2 = Arc::new(InMemoryApiTokenStore::default());
let layer2 = RequireApiToken::new(store2);
let mut svc2 = layer2.layer(MockService { ready: true });
assert!(svc2.poll_ready(&mut cx).is_ready());
}
#[tokio::test]
async fn require_api_token_rate_limits_by_principal() {
use axum::body::Body;
use tower::ServiceExt;
use crate::security::{KeyStrategy, RateLimitConfig, RateLimitLayer};
let rl_config = RateLimitConfig {
enabled: true,
requests_per_second: 0.1,
burst: 1,
key_strategy: KeyStrategy::AuthenticatedPrincipal,
..Default::default()
};
let store = Arc::new(InMemoryApiTokenStore::default());
let token_a1 = issue_api_token(&*store, "principal-1").await.unwrap();
let token_a2 = issue_api_token(&*store, "principal-1").await.unwrap(); let token_b = issue_api_token(&*store, "principal-2").await.unwrap();
let app = axum::Router::new()
.route("/", axum::routing::get(|| async { "ok" }))
.layer(RateLimitLayer::from_config(&rl_config)) .layer(RequireApiToken::new(Arc::clone(&store)));
let r = app
.clone()
.oneshot(
http::Request::builder()
.uri("/")
.header("authorization", format!("Bearer {token_a1}"))
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(r.status(), StatusCode::OK);
let r = app
.clone()
.oneshot(
http::Request::builder()
.uri("/")
.header("authorization", format!("Bearer {token_a2}"))
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(
r.status(),
StatusCode::TOO_MANY_REQUESTS,
"second token for the same principal must share the rate-limit bucket"
);
let r = app
.clone()
.oneshot(
http::Request::builder()
.uri("/")
.header("authorization", format!("Bearer {token_b}"))
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(r.status(), StatusCode::OK);
}
#[tokio::test]
async fn require_api_token_sets_rate_limit_principal() {
use axum::body::Body;
use tower::ServiceExt;
use crate::security::RateLimitPrincipal;
async fn handler(
axum::Extension(principal): axum::Extension<RateLimitPrincipal>,
) -> String {
principal.0
}
let store = Arc::new(InMemoryApiTokenStore::default());
let raw = issue_api_token(&*store, "agent:bot").await.unwrap();
let app = axum::Router::new()
.route("/", axum::routing::get(handler))
.layer(RequireApiToken::new(Arc::clone(&store)));
let response = app
.oneshot(
http::Request::builder()
.uri("/")
.header("authorization", format!("Bearer {raw}"))
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(response.status(), StatusCode::OK);
let body = axum::body::to_bytes(response.into_body(), usize::MAX)
.await
.unwrap();
assert_eq!(std::str::from_utf8(&body).unwrap(), "agent:bot");
}
}
#[cfg(feature = "oauth2")]
#[cfg(test)]
mod oauth2_unit_tests {
use std::collections::HashMap;
use super::{
AuthConfig, OAuth2ProviderConfig, OAuthLinkingPolicy, oauth2_authorize_url, provider_preset,
};
#[allow(dead_code)]
fn make_provider(authorize_url: &str) -> OAuth2ProviderConfig {
OAuth2ProviderConfig {
client_id: "cid".into(),
client_secret: "secret".into(),
authorize_url: authorize_url.into(),
token_url: "https://idp.example/token".into(),
userinfo_url: None,
redirect_uri: "http://localhost:3000/callback".into(),
scope: "openid profile".into(),
issuer: None,
jwks_url: None,
discovery_url: None,
}
}
#[test]
fn provider_preset_google_returns_oidc_config() {
let preset = provider_preset("google").expect("google preset must exist");
assert!(
!preset.authorize_url.is_empty(),
"google authorize_url must not be empty"
);
assert!(
!preset.token_url.is_empty(),
"google token_url must not be empty"
);
assert!(
preset.discovery_url.is_some(),
"google must have discovery_url for OIDC"
);
assert!(
preset.scope.contains("openid"),
"google preset scope must include openid: {}",
preset.scope
);
assert!(
preset.scope.contains("email"),
"google preset scope must include email: {}",
preset.scope
);
assert_eq!(
preset.client_id, "",
"client_id must be empty in preset (user fills in)"
);
assert_eq!(
preset.client_secret, "",
"client_secret must be empty in preset"
);
assert_eq!(
preset.redirect_uri, "",
"redirect_uri must be empty in preset"
);
}
#[cfg(feature = "oauth2")]
#[test]
fn provider_preset_github_returns_pure_oauth2_config() {
let preset = provider_preset("github").expect("github preset must exist");
assert!(
!preset.authorize_url.is_empty(),
"github authorize_url must not be empty"
);
assert!(
!preset.token_url.is_empty(),
"github token_url must not be empty"
);
assert!(
preset.userinfo_url.is_some(),
"github must have userinfo_url (it is not OIDC)"
);
assert!(
preset.discovery_url.is_none(),
"github must NOT have discovery_url (pure OAuth2)"
);
}
#[cfg(feature = "oauth2")]
#[test]
fn provider_preset_microsoft_returns_oidc_config() {
let preset = provider_preset("microsoft").expect("microsoft preset must exist");
assert!(
!preset.authorize_url.is_empty(),
"microsoft authorize_url must not be empty"
);
assert!(
preset.discovery_url.is_some(),
"microsoft must have discovery_url for OIDC"
);
assert!(
preset.scope.contains("openid"),
"microsoft preset scope must include openid: {}",
preset.scope
);
}
#[cfg(feature = "oauth2")]
#[test]
fn provider_preset_unknown_returns_none() {
assert!(
provider_preset("nonexistent_provider_xyz").is_none(),
"unknown provider must return None"
);
}
#[cfg(feature = "oauth2")]
#[tokio::test]
async fn oauth2_authorize_url_includes_pkce_code_challenge() {
let session = crate::session::Session::new_for_test("s1".into(), HashMap::new());
let provider = OAuth2ProviderConfig {
client_id: "cid".into(),
client_secret: "secret".into(),
authorize_url: "https://idp.example/authorize".into(),
token_url: "https://idp.example/token".into(),
userinfo_url: None,
redirect_uri: "http://localhost:3000/callback".into(),
scope: "openid profile".into(),
issuer: None,
jwks_url: None,
discovery_url: None,
};
let url = oauth2_authorize_url(&session, "testprovider", &provider)
.await
.unwrap();
assert!(
url.contains("code_challenge="),
"PKCE code_challenge must be present in URL: {url}"
);
assert!(
url.contains("code_challenge_method=S256"),
"PKCE method must be S256: {url}"
);
assert!(
session
.get("oauth2:testprovider:code_verifier")
.await
.is_some(),
"code_verifier must be stored in session for later exchange"
);
}
#[cfg(feature = "oauth2")]
#[test]
fn oauth2_provider_config_has_discovery_url_field() {
let provider = OAuth2ProviderConfig {
client_id: "cid".into(),
client_secret: "secret".into(),
authorize_url: "https://idp.example/authorize".into(),
token_url: "https://idp.example/token".into(),
userinfo_url: None,
redirect_uri: "http://localhost:3000/callback".into(),
scope: String::new(),
issuer: None,
jwks_url: None,
discovery_url: Some("https://idp.example".into()),
};
assert_eq!(
provider.discovery_url.as_deref(),
Some("https://idp.example"),
"discovery_url must be accessible as a field"
);
}
#[cfg(feature = "oauth2")]
#[test]
fn auth_config_has_oauth_linking_policy() {
let config = AuthConfig::default();
assert!(
matches!(
config.oauth_linking_policy,
OAuthLinkingPolicy::CreateAccount
),
"default linking policy must be CreateAccount"
);
}
}