use axum::extract::FromRequestParts;
use axum::http::HeaderMap;
use axum::http::request::Parts;
use crate::webadmin::AdminState;
use crate::webadmin::error::AdminError;
use crate::webadmin::pages::error::PageError;
use crate::webadmin::session::{
AdminRead, AdminWrite, Authenticated, AuthenticatedWrite, EnrolWrite, PendingMfa,
PendingMfaSubmit, SelfServiceWrite,
};
const HX_REQUEST: &str = "hx-request";
#[must_use]
pub fn is_htmx(headers: &HeaderMap) -> bool {
headers
.get(HX_REQUEST)
.and_then(|value| value.to_str().ok())
.is_some_and(|value| value.eq_ignore_ascii_case("true"))
}
pub struct PageSession {
pub auth: Authenticated,
pub hx: bool,
}
pub struct PageSessionWrite {
pub auth: Authenticated,
pub hx: bool,
}
impl FromRequestParts<AdminState> for PageSession {
type Rejection = PageError;
async fn from_request_parts(
parts: &mut Parts,
state: &AdminState,
) -> Result<Self, Self::Rejection> {
let hx = is_htmx(&parts.headers);
match Authenticated::from_request_parts(parts, state).await {
Ok(auth) => Ok(Self { auth, hx }),
Err(error) => Err(to_page_error(error, hx)),
}
}
}
impl FromRequestParts<AdminState> for PageSessionWrite {
type Rejection = PageError;
async fn from_request_parts(
parts: &mut Parts,
state: &AdminState,
) -> Result<Self, Self::Rejection> {
let hx = is_htmx(&parts.headers);
match AuthenticatedWrite::from_request_parts(parts, state).await {
Ok(AuthenticatedWrite(auth)) => Ok(Self { auth, hx }),
Err(error) => Err(to_page_error(error, hx)),
}
}
}
pub struct PageAdminWrite {
pub auth: Authenticated,
pub hx: bool,
}
pub struct PageSelfServiceWrite {
pub auth: Authenticated,
pub hx: bool,
}
pub struct PageAdminRead {
pub auth: Authenticated,
pub hx: bool,
}
pub trait PageAuth {
fn auth(&self) -> &Authenticated;
}
impl PageAuth for PageSession {
fn auth(&self) -> &Authenticated {
&self.auth
}
}
impl PageAuth for PageAdminRead {
fn auth(&self) -> &Authenticated {
&self.auth
}
}
impl FromRequestParts<AdminState> for PageAdminWrite {
type Rejection = PageError;
async fn from_request_parts(
parts: &mut Parts,
state: &AdminState,
) -> Result<Self, Self::Rejection> {
let hx = is_htmx(&parts.headers);
match AdminWrite::from_request_parts(parts, state).await {
Ok(AdminWrite(auth)) => Ok(Self { auth, hx }),
Err(error) => Err(to_page_error(error, hx)),
}
}
}
impl FromRequestParts<AdminState> for PageAdminRead {
type Rejection = PageError;
async fn from_request_parts(
parts: &mut Parts,
state: &AdminState,
) -> Result<Self, Self::Rejection> {
let hx = is_htmx(&parts.headers);
match AdminRead::from_request_parts(parts, state).await {
Ok(AdminRead(auth)) => Ok(Self { auth, hx }),
Err(error) => Err(to_page_error(error, hx)),
}
}
}
impl FromRequestParts<AdminState> for PageSelfServiceWrite {
type Rejection = PageError;
async fn from_request_parts(
parts: &mut Parts,
state: &AdminState,
) -> Result<Self, Self::Rejection> {
let hx = is_htmx(&parts.headers);
match SelfServiceWrite::from_request_parts(parts, state).await {
Ok(SelfServiceWrite(auth)) => Ok(Self { auth, hx }),
Err(error) => Err(to_page_error(error, hx)),
}
}
}
pub struct PageMfaPending {
pub pending: PendingMfa,
pub hx: bool,
}
pub struct PageMfaSubmit {
pub pending: PendingMfa,
pub hx: bool,
}
pub struct PageEnrolWrite {
pub enrol: EnrolWrite,
pub hx: bool,
}
impl FromRequestParts<AdminState> for PageMfaPending {
type Rejection = PageError;
async fn from_request_parts(
parts: &mut Parts,
state: &AdminState,
) -> Result<Self, Self::Rejection> {
let hx = is_htmx(&parts.headers);
match PendingMfa::from_request_parts(parts, state).await {
Ok(pending) => Ok(Self { pending, hx }),
Err(error) => Err(to_page_error(error, hx)),
}
}
}
impl FromRequestParts<AdminState> for PageMfaSubmit {
type Rejection = PageError;
async fn from_request_parts(
parts: &mut Parts,
state: &AdminState,
) -> Result<Self, Self::Rejection> {
let hx = is_htmx(&parts.headers);
match PendingMfaSubmit::from_request_parts(parts, state).await {
Ok(PendingMfaSubmit(pending)) => Ok(Self { pending, hx }),
Err(error) => Err(to_page_error(error, hx)),
}
}
}
impl FromRequestParts<AdminState> for PageEnrolWrite {
type Rejection = PageError;
async fn from_request_parts(
parts: &mut Parts,
state: &AdminState,
) -> Result<Self, Self::Rejection> {
let hx = is_htmx(&parts.headers);
match EnrolWrite::from_request_parts(parts, state).await {
Ok(enrol) => Ok(Self { enrol, hx }),
Err(error) => Err(to_page_error(error, hx)),
}
}
}
fn to_page_error(error: AdminError, hx: bool) -> PageError {
if error.status == axum::http::StatusCode::UNAUTHORIZED {
PageError::login_required(hx)
} else {
error.into()
}
}
#[cfg(test)]
mod tests {
use super::*;
use axum::http::StatusCode;
fn headers(pairs: &[(&str, &str)]) -> HeaderMap {
let mut map = HeaderMap::new();
for (name, value) in pairs {
map.insert(
axum::http::HeaderName::from_bytes(name.as_bytes()).unwrap(),
value.parse().unwrap(),
);
}
map
}
#[test]
fn is_htmx_reads_the_header_case_insensitively() {
assert!(is_htmx(&headers(&[("hx-request", "true")])));
assert!(is_htmx(&headers(&[("HX-Request", "True")])));
assert!(!is_htmx(&headers(&[("hx-request", "false")])));
assert!(!is_htmx(&headers(&[("hx-request", "yes")])));
assert!(!is_htmx(&HeaderMap::new()));
}
#[test]
fn an_unauthorized_rejection_becomes_a_sign_in_redirect() {
assert_eq!(
to_page_error(AdminError::session_expired(), false),
PageError::login_required(false)
);
assert_eq!(
to_page_error(AdminError::session_invalid(), true),
PageError::login_required(true)
);
}
#[test]
fn a_csrf_failure_keeps_its_own_status() {
let error = to_page_error(AdminError::csrf_failed("no token"), true);
assert_eq!(error.status(), StatusCode::FORBIDDEN);
assert_eq!(
error,
PageError::Rendered {
status: StatusCode::FORBIDDEN,
code: "csrf_failed",
message: "no token".to_string(),
}
);
}
}