use super::backend::{AuthBackend, AuthError, Role};
use axum::{
extract::{Request, State},
http::{header, StatusCode},
middleware::Next,
response::{IntoResponse, Response},
};
use std::sync::Arc;
use tracing::Instrument as _;
#[derive(Clone)]
pub struct AuthLayerState {
pub backend: Arc<dyn AuthBackend>,
pub required: Role,
}
pub async fn require_role_middleware(
State(state): State<AuthLayerState>,
mut req: Request,
next: Next,
) -> Response {
let principal = if state.backend.enabled() {
let token = req
.headers()
.get(header::AUTHORIZATION)
.and_then(|v| v.to_str().ok())
.and_then(|s| s.strip_prefix("Bearer "))
.map(str::trim);
let Some(token) = token else {
return unauthorized("missing Authorization: Bearer <token>");
};
let principal = match state.backend.authenticate(token).await {
Ok(p) => p,
Err(AuthError::UnknownToken | AuthError::MissingToken) => {
return unauthorized("unknown token");
}
Err(AuthError::Backend(e)) => {
return (
StatusCode::INTERNAL_SERVER_ERROR,
format!("auth backend: {e}"),
)
.into_response();
}
};
if principal.role < state.required {
return forbidden(&format!(
"token role `{:?}` below required `{:?}`",
principal.role, state.required
));
}
principal
} else {
super::backend::Principal {
tenant: pensieve_core::tenant::DEFAULT_TENANT,
role: Role::Admin,
subject: None,
allowed_databases: None,
allowed_realms: None,
}
};
let tenant = principal.tenant;
let route = req
.extensions()
.get::<axum::extract::MatchedPath>()
.map_or_else(|| req.uri().path().to_string(), |m| m.as_str().to_string());
let method = req.method().clone();
let subject = principal.subject.clone();
req.extensions_mut().insert(principal);
req.extensions_mut().insert(tenant);
if route == "/health" || route.starts_with("/metrics") || route.starts_with("/v1/explore/live")
{
return next.run(req).await;
}
let span = tracing::info_span!(
target: "pensieve_telemetry",
"request",
otel.name = %format!("{method} {route}"),
http.method = %method,
http.route = %route,
pensieve.tenant = %tenant,
pensieve.subject = tracing::field::Empty,
http.status = tracing::field::Empty,
otel.status_code = tracing::field::Empty,
);
if let Some(s) = &subject {
span.record("pensieve.subject", s.as_str());
}
let resp = next.run(req).instrument(span.clone()).await;
span.record("http.status", resp.status().as_u16());
span.record(
"otel.status_code",
if resp.status().is_server_error() { "ERROR" } else { "OK" },
);
resp
}
fn unauthorized(msg: &str) -> Response {
(
StatusCode::UNAUTHORIZED,
[(header::WWW_AUTHENTICATE, r#"Bearer realm="pensieve""#)],
msg.to_owned(),
)
.into_response()
}
fn forbidden(msg: &str) -> Response {
(StatusCode::FORBIDDEN, msg.to_owned()).into_response()
}