use std::convert::Infallible;
use std::future::Future;
use std::marker::PhantomData;
use std::time::{Duration, SystemTime, UNIX_EPOCH};
use axum::extract::FromRequestParts;
use axum::response::{IntoResponse, Response};
use tower_sessions::Session as TowerSession;
use crate::auth::AuthUser;
const ABSOLUTE_AUTH_AT_KEY: &str = "__arcature_absolute_auth_at";
const CREDENTIAL_STAMP_KEY: &str = "__arcature_credential_stamp";
fn credential_stamp(credential: &[u8]) -> String {
use sha2::{Digest, Sha256};
let digest = Sha256::digest(credential);
let mut out = String::with_capacity(digest.len() * 2);
for byte in digest {
use std::fmt::Write;
let _ = write!(out, "{byte:02x}");
}
out
}
fn now_unix_millis() -> i64 {
SystemTime::now()
.duration_since(UNIX_EPOCH)
.map(|d| d.as_millis() as i64)
.unwrap_or(0)
}
pub struct Auth<U: AuthUser>(pub U);
impl<U: AuthUser> Auth<U> {
#[must_use]
pub fn into_inner(self) -> U {
self.0
}
#[must_use]
pub fn user(&self) -> &U {
&self.0
}
}
impl<U, S> FromRequestParts<S> for Auth<U>
where
U: UserLoader<S>,
S: Send + Sync,
{
type Rejection = Response;
async fn from_request_parts(
parts: &mut axum::http::request::Parts,
state: &S,
) -> Result<Self, Self::Rejection> {
let user = load_user::<U, S>(parts, state).await?;
user.map(Auth).ok_or_else(|| {
(
axum::http::StatusCode::UNAUTHORIZED,
"Authentication required",
)
.into_response()
})
}
}
pub struct OptionalAuth<U: AuthUser>(pub Option<U>);
impl<U: AuthUser> OptionalAuth<U> {
#[must_use]
pub fn user(&self) -> Option<&U> {
self.0.as_ref()
}
#[must_use]
pub fn is_authenticated(&self) -> bool {
self.0.is_some()
}
}
impl<U, S> FromRequestParts<S> for OptionalAuth<U>
where
U: UserLoader<S>,
S: Send + Sync,
{
type Rejection = Infallible;
async fn from_request_parts(
parts: &mut axum::http::request::Parts,
state: &S,
) -> Result<Self, Self::Rejection> {
let user = load_user::<U, S>(parts, state).await.unwrap_or(None);
Ok(OptionalAuth(user))
}
}
pub type Current<U> = Auth<U>;
pub type OptionalCurrent<U> = OptionalAuth<U>;
pub struct AuthManager<U: AuthUser> {
session: TowerSession,
_marker: PhantomData<U>,
}
impl<U: AuthUser> AuthManager<U> {
#[must_use]
pub fn login(&self, user: &U) -> LoginBuilder<'_, U> {
LoginBuilder {
session: &self.session,
user_id: user.id().clone(),
credential_stamp: user.stored_credential().map(credential_stamp),
remember: false,
}
}
pub async fn rebind_credential(&self, user: &U) -> Result<(), AuthError> {
let Some(credential) = user.stored_credential() else {
return Ok(());
};
self.session
.insert(CREDENTIAL_STAMP_KEY, credential_stamp(credential))
.await
.map_err(|e| AuthError::Session(e.to_string()))?;
Ok(())
}
pub async fn logout(&self) -> Result<(), AuthError> {
self.session
.flush()
.await
.map_err(|e| AuthError::Session(e.to_string()))?;
Ok(())
}
pub async fn regenerate(&self) -> Result<(), AuthError> {
self.session
.cycle_id()
.await
.map_err(|e| AuthError::Session(e.to_string()))?;
Ok(())
}
}
impl<U, S> FromRequestParts<S> for AuthManager<U>
where
U: AuthUser,
S: Send + Sync,
{
type Rejection = Infallible;
async fn from_request_parts(
parts: &mut axum::http::request::Parts,
state: &S,
) -> Result<Self, Self::Rejection> {
let session = TowerSession::from_request_parts(parts, state)
.await
.map_err(|_| unreachable!("Session extraction is infallible"))?;
Ok(AuthManager {
session,
_marker: PhantomData,
})
}
}
pub struct LoginBuilder<'a, U: AuthUser> {
session: &'a TowerSession,
user_id: U::Id,
credential_stamp: Option<String>,
remember: bool,
}
impl<'a, U: AuthUser> LoginBuilder<'a, U> {
#[must_use]
pub fn remember(mut self, remember: bool) -> Self {
self.remember = remember;
self
}
}
impl<'a, U: AuthUser> std::future::IntoFuture for LoginBuilder<'a, U> {
type Output = Result<(), AuthError>;
type IntoFuture = std::pin::Pin<
std::boxed::Box<dyn std::future::Future<Output = Result<(), AuthError>> + Send + 'a>,
>;
fn into_future(self) -> Self::IntoFuture {
Box::pin(async move {
self.session
.cycle_id()
.await
.map_err(|e| AuthError::Session(e.to_string()))?;
self.session
.insert(U::SESSION_KEY, &self.user_id)
.await
.map_err(|e| AuthError::Session(e.to_string()))?;
self.session
.insert(ABSOLUTE_AUTH_AT_KEY, now_unix_millis())
.await
.map_err(|e| AuthError::Session(e.to_string()))?;
if let Some(stamp) = self.credential_stamp {
self.session
.insert(CREDENTIAL_STAMP_KEY, stamp)
.await
.map_err(|e| AuthError::Session(e.to_string()))?;
}
if self.remember {
self.session
.insert("remember_me", true)
.await
.map_err(|e| AuthError::Session(e.to_string()))?;
}
Ok(())
})
}
}
#[derive(Debug)]
pub enum AuthError {
Session(String),
}
impl std::fmt::Display for AuthError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::Session(msg) => write!(f, "session error: {msg}"),
}
}
}
impl std::error::Error for AuthError {}
pub trait UserLoader<S>: AuthUser + Sized {
type Error: std::error::Error + Send + Sync + 'static;
fn load_user(
id: &Self::Id,
state: &S,
) -> impl Future<Output = Result<Option<Self>, Self::Error>> + Send;
#[must_use]
fn absolute_max_age() -> Duration {
Duration::from_secs(60 * 60 * 24 * 30)
}
}
#[allow(clippy::result_large_err)]
async fn load_user<U, S>(
parts: &mut axum::http::request::Parts,
state: &S,
) -> Result<Option<U>, Response>
where
U: UserLoader<S>,
S: Send + Sync,
{
let session = TowerSession::from_request_parts(parts, state)
.await
.map_err(|_| {
(
axum::http::StatusCode::INTERNAL_SERVER_ERROR,
"session extraction failed",
)
.into_response()
})?;
let user_id: Option<U::Id> = session.get(U::SESSION_KEY).await.map_err(|_err| {
(
axum::http::StatusCode::INTERNAL_SERVER_ERROR,
"session read failed",
)
.into_response()
})?;
let user_id = match user_id {
Some(id) => id,
None => return Ok(None),
};
let absolute_max_millis: i64 =
i64::try_from(U::absolute_max_age().as_millis()).unwrap_or(i64::MAX);
let auth_at: Option<i64> = session.get(ABSOLUTE_AUTH_AT_KEY).await.map_err(|_err| {
(
axum::http::StatusCode::INTERNAL_SERVER_ERROR,
"session read failed",
)
.into_response()
})?;
match auth_at {
Some(auth_at) => {
if now_unix_millis().saturating_sub(auth_at) > absolute_max_millis {
session.flush().await.map_err(|_err| {
(
axum::http::StatusCode::INTERNAL_SERVER_ERROR,
"session flush failed",
)
.into_response()
})?;
return Ok(None);
}
}
None => {
session
.insert(ABSOLUTE_AUTH_AT_KEY, now_unix_millis())
.await
.map_err(|_err| {
(
axum::http::StatusCode::INTERNAL_SERVER_ERROR,
"session write failed",
)
.into_response()
})?;
}
}
let user = U::load_user(&user_id, state).await.map_err(|_err| {
(
axum::http::StatusCode::INTERNAL_SERVER_ERROR,
"user load failed",
)
.into_response()
})?;
let Some(user) = user else { return Ok(None) };
if let Some(credential) = user.stored_credential() {
let expected = credential_stamp(credential);
let bound: Option<String> = session.get(CREDENTIAL_STAMP_KEY).await.map_err(|_err| {
(
axum::http::StatusCode::INTERNAL_SERVER_ERROR,
"session read failed",
)
.into_response()
})?;
match bound {
Some(bound) if bound == expected => {}
Some(_) => {
session.flush().await.map_err(|_err| {
(
axum::http::StatusCode::INTERNAL_SERVER_ERROR,
"session flush failed",
)
.into_response()
})?;
return Ok(None);
}
None => {
session
.insert(CREDENTIAL_STAMP_KEY, expected)
.await
.map_err(|_err| {
(
axum::http::StatusCode::INTERNAL_SERVER_ERROR,
"session write failed",
)
.into_response()
})?;
}
}
}
Ok(Some(user))
}