use std::sync::Arc;
use axum::{
extract::{FromRequestParts, Request, State},
http::{header, request::Parts, StatusCode},
middleware::Next,
response::{IntoResponse, Response},
Json,
};
use jsonwebtoken::{decode, Algorithm, DecodingKey, Validation};
use serde::{Deserialize, Serialize};
use uuid::Uuid;
use crate::org::TOKEN_TYPE_ACCESS;
#[derive(Debug, Clone)]
pub struct CompanyContext {
pub company_id: Uuid,
pub branch_id: Option<Uuid>,
pub user_id: String,
}
#[derive(Debug, Serialize, Deserialize)]
pub struct CompanyClaims {
pub sub: String,
pub exp: usize,
#[serde(default)]
pub company_id: Option<Uuid>,
#[serde(default)]
pub branch_id: Option<Uuid>,
#[serde(default)]
pub typ: Option<String>,
}
#[derive(Clone)]
pub struct CompanyVerifier {
key: Arc<DecodingKey>,
validation: Arc<Validation>,
}
impl CompanyVerifier {
pub fn hs256(secret: &[u8]) -> Self {
Self {
key: Arc::new(DecodingKey::from_secret(secret)),
validation: Arc::new(Validation::new(Algorithm::HS256)),
}
}
pub fn rs256(public_key_pem: &[u8]) -> Result<Self, jsonwebtoken::errors::Error> {
Ok(Self {
key: Arc::new(DecodingKey::from_rsa_pem(public_key_pem)?),
validation: Arc::new(Validation::new(Algorithm::RS256)),
})
}
pub fn verify(&self, token: &str) -> Option<CompanyContext> {
let data = decode::<CompanyClaims>(token, &self.key, &self.validation).ok()?;
let c = data.claims;
if let Some(typ) = &c.typ {
if typ != TOKEN_TYPE_ACCESS {
return None;
}
}
Some(CompanyContext {
company_id: c.company_id?,
branch_id: c.branch_id,
user_id: c.sub,
})
}
}
fn unauthorized(message: &str) -> Response {
(
StatusCode::UNAUTHORIZED,
Json(serde_json::json!({ "error": "unauthorized", "message": message })),
)
.into_response()
}
fn internal_error(message: &str) -> Response {
(
StatusCode::INTERNAL_SERVER_ERROR,
Json(serde_json::json!({ "error": "internal_error", "message": message })),
)
.into_response()
}
pub async fn company_auth(
State(verifier): State<CompanyVerifier>,
mut req: Request,
next: Next,
) -> Response {
let token = req
.headers()
.get(header::AUTHORIZATION)
.and_then(|h| h.to_str().ok())
.and_then(|raw| {
raw.strip_prefix("Bearer ")
.or_else(|| raw.strip_prefix("bearer "))
});
let Some(token) = token else {
return unauthorized("missing bearer token");
};
match verifier.verify(token) {
Some(ctx) => {
let company_id = ctx.company_id;
let pool = req.extensions().get::<backbone_orm::PgPool>().cloned();
req.extensions_mut().insert(ctx);
match pool {
Some(pool) => {
match backbone_orm::company_scope::with_request_scope(
&pool,
company_id,
next.run(req),
)
.await
{
Ok(resp) => resp,
Err(_) => internal_error("could not establish the request database scope"),
}
}
None => backbone_orm::with_company_scope(Some(company_id), next.run(req)).await,
}
}
None => unauthorized("invalid token or missing company_id claim"),
}
}
#[async_trait::async_trait]
impl<S: Send + Sync> FromRequestParts<S> for CompanyContext {
type Rejection = Response;
async fn from_request_parts(parts: &mut Parts, _state: &S) -> Result<Self, Self::Rejection> {
parts
.extensions
.get::<CompanyContext>()
.cloned()
.ok_or_else(|| unauthorized("unauthenticated"))
}
}