use std::sync::Arc;
use axum::extract::{Request, State};
use axum::http::header::AUTHORIZATION;
use axum::http::{HeaderMap, StatusCode};
use axum::middleware::Next;
use axum::response::{IntoResponse, Response};
use axum::Json;
use serde::Serialize;
use serde_json::json;
use crate::service::MailboxService;
pub fn extract_bearer(headers: &HeaderMap) -> Option<String> {
let value = headers.get(AUTHORIZATION)?;
let value = value.to_str().ok()?;
let token = value.strip_prefix("Bearer ").or_else(|| value.strip_prefix("bearer "))?;
Some(token.trim().to_string())
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize)]
#[serde(rename_all = "snake_case")]
pub enum Tier {
Authenticated,
Admin,
}
impl Tier {
fn rank(self) -> u8 {
match self {
Tier::Authenticated => 1,
Tier::Admin => 2,
}
}
}
#[derive(Clone)]
pub struct TierGuard {
pub service: Arc<MailboxService>,
pub required: Tier,
}
pub async fn require_tier(State(guard): State<TierGuard>, req: Request, next: Next) -> Response {
let Some(token) = extract_bearer(req.headers()) else {
return tier_insufficient(guard.required, StatusCode::UNAUTHORIZED);
};
let carried = match guard.service.authenticate(&token).await {
Ok(participant) => {
if participant.operator {
Tier::Admin
} else {
Tier::Authenticated
}
}
Err(_) => return auth_rejected("unknown or invalid bearer token"),
};
if carried.rank() >= guard.required.rank() {
next.run(req).await
} else {
tier_insufficient(guard.required, StatusCode::FORBIDDEN)
}
}
fn auth_rejected(reason: &str) -> Response {
(StatusCode::UNAUTHORIZED, Json(json!({ "ok": false, "error": "auth_rejected", "reason": reason }))).into_response()
}
fn tier_insufficient(required: Tier, status: StatusCode) -> Response {
(status, Json(json!({ "ok": false, "error": "tier_insufficient", "required": required }))).into_response()
}
#[cfg(test)]
mod tests {
use super::*;
use axum::http::Request as HttpRequest;
fn headers_with_bearer(token: &str) -> HeaderMap {
HttpRequest::builder().header(AUTHORIZATION, format!("Bearer {token}")).body(()).unwrap().headers().clone()
}
#[test]
fn extract_bearer_reads_the_token_out_of_the_header() {
let headers = headers_with_bearer("secret-123");
assert_eq!(extract_bearer(&headers).as_deref(), Some("secret-123"));
}
#[test]
fn extract_bearer_is_case_insensitive_on_the_scheme() {
let mut headers = HeaderMap::new();
headers.insert(AUTHORIZATION, "bearer secret-123".parse().unwrap());
assert_eq!(extract_bearer(&headers).as_deref(), Some("secret-123"));
}
#[test]
fn extract_bearer_is_none_without_the_header() {
assert_eq!(extract_bearer(&HeaderMap::new()), None);
}
#[test]
fn admin_rank_satisfies_an_authenticated_requirement() {
assert!(Tier::Admin.rank() >= Tier::Authenticated.rank());
}
#[test]
fn authenticated_rank_does_not_satisfy_an_admin_requirement() {
assert!(Tier::Authenticated.rank() < Tier::Admin.rank());
}
#[test]
fn tier_serializes_snake_case() {
assert_eq!(serde_json::to_value(Tier::Authenticated).unwrap(), json!("authenticated"));
assert_eq!(serde_json::to_value(Tier::Admin).unwrap(), json!("admin"));
}
}