use std::sync::Arc;
use std::time::{Duration, SystemTime, UNIX_EPOCH};
use axum::{
extract::{FromRequestParts, Request, State},
http::{header, request::Parts, HeaderMap, HeaderValue, StatusCode},
middleware::Next,
response::{IntoResponse, Response},
Json,
};
use backbone_orm::audit_context::RequestAuditContext;
use jsonwebtoken::{decode, encode, Algorithm, DecodingKey, EncodingKey, Header, Validation};
use serde::{Deserialize, Serialize};
use uuid::Uuid;
pub const TOKEN_TYPE_ACCESS: &str = "access";
pub const TOKEN_TYPE_REFRESH: &str = "refresh";
#[derive(Debug, Clone)]
pub struct OrgContext {
pub acting_unit_id: Uuid,
pub entitled_units: Vec<Uuid>,
pub legacy_company_id: Option<Uuid>,
pub user_id: String,
}
#[derive(Debug, Serialize, Deserialize)]
pub struct OrgClaims {
pub sub: String,
pub exp: usize,
#[serde(default)]
pub org_unit_id: Option<Uuid>,
#[serde(default)]
pub entitled_units: Vec<Uuid>,
#[serde(default)]
pub company_id: Option<Uuid>,
#[serde(default)]
pub typ: Option<String>,
}
#[derive(Clone)]
pub struct OrgVerifier {
key: Arc<DecodingKey>,
validation: Arc<Validation>,
}
impl OrgVerifier {
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<OrgContext> {
let data = decode::<OrgClaims>(token, &self.key, &self.validation).ok()?;
let c = data.claims;
if let Some(typ) = &c.typ {
if typ != TOKEN_TYPE_ACCESS {
return None;
}
}
Some(OrgContext {
acting_unit_id: c.org_unit_id?,
entitled_units: c.entitled_units,
legacy_company_id: c.company_id,
user_id: c.sub,
})
}
}
#[derive(Clone)]
pub struct OrgIssuer {
key: Arc<EncodingKey>,
algorithm: Algorithm,
}
#[derive(Serialize)]
struct IssuedSessionClaims {
sub: String,
exp: usize,
iat: usize,
typ: &'static str,
org_unit_id: Uuid,
entitled_units: Vec<Uuid>,
#[serde(skip_serializing_if = "Option::is_none")]
company_id: Option<Uuid>,
}
impl OrgIssuer {
pub fn hs256(secret: &[u8]) -> Self {
Self {
key: Arc::new(EncodingKey::from_secret(secret)),
algorithm: Algorithm::HS256,
}
}
pub fn rs256(private_key_pem: &[u8]) -> Result<Self, jsonwebtoken::errors::Error> {
Ok(Self {
key: Arc::new(EncodingKey::from_rsa_pem(private_key_pem)?),
algorithm: Algorithm::RS256,
})
}
pub fn issue_access(
&self,
user_id: &str,
acting_unit: Uuid,
entitled_units: &[Uuid],
legacy_company: Option<Uuid>,
ttl: Duration,
) -> Result<String, jsonwebtoken::errors::Error> {
self.issue(user_id, acting_unit, entitled_units, legacy_company, ttl, TOKEN_TYPE_ACCESS)
}
pub fn issue_refresh(
&self,
user_id: &str,
acting_unit: Uuid,
entitled_units: &[Uuid],
legacy_company: Option<Uuid>,
ttl: Duration,
) -> Result<String, jsonwebtoken::errors::Error> {
self.issue(user_id, acting_unit, entitled_units, legacy_company, ttl, TOKEN_TYPE_REFRESH)
}
fn issue(
&self,
user_id: &str,
acting_unit: Uuid,
entitled_units: &[Uuid],
legacy_company: Option<Uuid>,
ttl: Duration,
typ: &'static str,
) -> Result<String, jsonwebtoken::errors::Error> {
let now = SystemTime::now().duration_since(UNIX_EPOCH).unwrap_or_default();
let claims = IssuedSessionClaims {
sub: user_id.to_string(),
exp: now.saturating_add(ttl).as_secs() as usize,
iat: now.as_secs() as usize,
typ,
org_unit_id: acting_unit,
entitled_units: entitled_units.to_vec(),
company_id: legacy_company,
};
encode(&Header::new(self.algorithm), &claims, &self.key)
}
}
fn unauthorized(message: &str) -> Response {
(
StatusCode::UNAUTHORIZED,
Json(serde_json::json!({ "error": "unauthorized", "message": message })),
)
.into_response()
}
fn forbidden(message: &str) -> Response {
(
StatusCode::FORBIDDEN,
Json(serde_json::json!({ "error": "forbidden", "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()
}
const CORRELATION_ID_HEADER: &str = "x-correlation-id";
const MAX_CORRELATION_ID_LEN: usize = 128;
const MAX_CLIENT_IP_LEN: usize = 64;
const MAX_USER_AGENT_LEN: usize = 256;
const MAX_RESOURCE_PATH_LEN: usize = 256;
fn normalize_fact(raw: &str, max: usize) -> String {
raw.chars().filter(|c| c.is_ascii_graphic()).take(max).collect()
}
fn correlation_id_of(headers: &HeaderMap) -> String {
headers
.get(CORRELATION_ID_HEADER)
.and_then(|v| v.to_str().ok())
.map(|raw| normalize_fact(raw, MAX_CORRELATION_ID_LEN))
.filter(|normalized| !normalized.is_empty())
.unwrap_or_else(|| Uuid::new_v4().to_string())
}
fn client_ip_of(headers: &HeaderMap) -> String {
headers
.get("x-forwarded-for")
.and_then(|v| v.to_str().ok())
.and_then(|v| v.split(',').next())
.map(|raw| normalize_fact(raw.trim(), MAX_CLIENT_IP_LEN))
.unwrap_or_default()
}
fn audit_context_of(ctx: &OrgContext, req: &Request) -> RequestAuditContext {
RequestAuditContext {
actor: ctx.user_id.clone(),
correlation_id: correlation_id_of(req.headers()),
client_ip: client_ip_of(req.headers()),
user_agent: req
.headers()
.get(header::USER_AGENT)
.and_then(|v| v.to_str().ok())
.map(|raw| normalize_fact(raw, MAX_USER_AGENT_LEN))
.unwrap_or_default(),
http_method: req.method().as_str().to_string(),
resource_path: normalize_fact(req.uri().path(), MAX_RESOURCE_PATH_LEN),
}
}
pub async fn org_auth(
State(verifier): State<OrgVerifier>,
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");
};
let Some(ctx) = verifier.verify(token) else {
return unauthorized("invalid token or missing org_unit_id claim");
};
let Some(pool) = req.extensions().get::<backbone_orm::PgPool>().cloned() else {
return internal_error(
"no tenant database on the request — mount the tenant router outside org_auth",
);
};
let scope = {
let mut conn = match pool.acquire().await {
Ok(conn) => conn,
Err(e) => {
tracing::error!(target: "backbone_auth::org", error = %e, "could not reach the tenant database to resolve the org scope");
return internal_error("could not reach the tenant database");
}
};
match backbone_orm::org_scope::resolve_org_scope(
&mut *conn,
ctx.acting_unit_id,
&ctx.entitled_units,
)
.await
{
Ok(scope) => scope,
Err(backbone_orm::org_scope::OrgScopeError::UnknownActingUnit(unit)) => {
return forbidden(&format!(
"org unit {unit} is not in this tenant's organization"
))
}
Err(backbone_orm::org_scope::OrgScopeError::MissingSpineHelpers) => {
return internal_error(
"the org spine is not migrated in this tenant's database",
)
}
Err(e) => {
tracing::error!(target: "backbone_auth::org", error = %e, "org scope resolution failed");
return internal_error("could not resolve the org scope");
}
}
};
let audit = audit_context_of(&ctx, &req);
let correlation_id = audit.correlation_id.clone();
req.extensions_mut().insert(ctx);
match backbone_orm::org_scope::with_org_request_scope_and_audit(
&pool,
scope,
audit,
next.run(req),
)
.await
{
Ok(mut resp) => {
if let Ok(value) = HeaderValue::from_str(&correlation_id) {
resp.headers_mut().insert(CORRELATION_ID_HEADER, value);
}
resp
}
Err(_) => internal_error("could not establish the request org scope"),
}
}
#[async_trait::async_trait]
impl<S: Send + Sync> FromRequestParts<S> for OrgContext {
type Rejection = Response;
async fn from_request_parts(parts: &mut Parts, _state: &S) -> Result<Self, Self::Rejection> {
parts
.extensions
.get::<OrgContext>()
.cloned()
.ok_or_else(|| unauthorized("unauthenticated"))
}
}