use axum::body::Body;
use axum::extract::{RawPathParams, State};
use axum::http::{Method, Request};
use axum::middleware::Next;
use axum::response::Response;
use crate::error::{DbError, DbResult};
use crate::server::auth::Claims;
use crate::server::authorization::{AuthorizationService, PermissionAction};
use crate::server::handlers::AppState;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum AuthzMode {
Enforce,
Warn,
}
fn authz_mode() -> AuthzMode {
static MODE: std::sync::OnceLock<AuthzMode> = std::sync::OnceLock::new();
*MODE.get_or_init(|| match std::env::var("SOLIDB_DB_AUTHZ_MODE").as_deref() {
Ok("warn")
if std::env::var("SOLIDB_DB_AUTHZ_ALLOW_WARN")
.map(|v| v == "1" || v.eq_ignore_ascii_case("true"))
.unwrap_or(false) =>
{
tracing::warn!(
"SOLIDB_DB_AUTHZ_MODE=warn: per-database authorization failures are \
logged but NOT enforced"
);
AuthzMode::Warn
}
Ok("warn") => {
tracing::error!(
"SOLIDB_DB_AUTHZ_MODE=warn is ignored without SOLIDB_DB_AUTHZ_ALLOW_WARN=1; \
enforcing authorization"
);
AuthzMode::Enforce
}
_ => AuthzMode::Enforce,
})
}
pub async fn enforce(
claims: &Claims,
state: &AppState,
action: PermissionAction,
database: Option<&str>,
) -> DbResult<()> {
match AuthorizationService::check_permission(claims, state, action.clone(), database).await {
Ok(()) => Ok(()),
Err(e) => match authz_mode() {
AuthzMode::Enforce => Err(e),
AuthzMode::Warn => {
tracing::warn!(
target: "audit",
user = %claims.sub,
database = database.unwrap_or("<global>"),
action = ?action,
"authz dry-run: request would be denied ({})",
e
);
Ok(())
}
},
}
}
pub fn enforce_raw(
permissions: &std::collections::HashSet<crate::server::authorization::Permission>,
action: PermissionAction,
database: Option<&str>,
scoped_databases: Option<&[String]>,
subject: &str,
) -> bool {
match AuthorizationService::check_permission_raw(
permissions,
action.clone(),
database,
scoped_databases,
) {
Ok(()) => true,
Err(e) => match authz_mode() {
AuthzMode::Enforce => false,
AuthzMode::Warn => {
tracing::warn!(
target: "audit",
user = subject,
database = database.unwrap_or("<global>"),
action = ?action,
"authz dry-run: request would be denied ({})",
e
);
true
}
},
}
}
fn required_action(method: &Method, path: &str) -> PermissionAction {
if path.ends_with("/truncate") {
return PermissionAction::Admin;
}
if path.ends_with("/repl") {
return PermissionAction::Admin;
}
if is_script_or_service_mutation(method, path) {
return PermissionAction::Admin;
}
if *method == Method::DELETE && is_collection_drop(path) {
return PermissionAction::Admin;
}
if matches!(*method, Method::GET | Method::HEAD | Method::OPTIONS) {
return PermissionAction::Read;
}
const READ_SUFFIXES: [&str; 9] = [
"/cursor",
"/explain",
"/nl",
"/near",
"/within",
"/search",
"/aggregate",
"/sql",
"/_verify",
];
if READ_SUFFIXES.iter().any(|s| path.ends_with(s)) {
return PermissionAction::Read;
}
if path.ends_with("/query") {
return PermissionAction::Read;
}
PermissionAction::Write
}
fn is_script_or_service_mutation(method: &Method, path: &str) -> bool {
if matches!(*method, Method::GET | Method::HEAD | Method::OPTIONS) {
return false;
}
let Some(rest) = path
.strip_prefix("/_api/database/")
.and_then(|r| r.split_once('/'))
.map(|(_, rest)| rest)
else {
return false;
};
rest == "scripts"
|| rest.starts_with("scripts/")
|| rest == "services"
|| rest.starts_with("services/")
}
fn is_collection_drop(path: &str) -> bool {
let Some(rest) = path
.strip_prefix("/_api/database/")
.and_then(|r| r.split_once('/'))
.map(|(_, rest)| rest)
else {
return false;
};
match rest.split_once('/') {
Some(("collection", tail)) | Some(("columnar", tail)) => !tail.contains('/'),
_ => false,
}
}
fn method_default_action(method: &Method) -> PermissionAction {
if matches!(*method, Method::GET | Method::HEAD | Method::OPTIONS) {
PermissionAction::Read
} else {
PermissionAction::Write
}
}
pub async fn db_authz_middleware(
State(state): State<AppState>,
params: RawPathParams,
req: Request<Body>,
next: Next,
) -> Result<Response, DbError> {
let db = params
.iter()
.find(|(k, _)| *k == "db")
.map(|(_, v)| v.to_string());
let Some(db) = db else {
return Ok(next.run(req).await);
};
let claims = req
.extensions()
.get::<Claims>()
.cloned()
.ok_or_else(|| DbError::Forbidden("Missing authentication context".to_string()))?;
let action = match req.extensions().get::<axum::extract::MatchedPath>() {
Some(matched) => required_action(req.method(), matched.as_str()),
None => {
tracing::warn!(
target: "audit",
path = %req.uri().path(),
"authz: no matched route template; falling back to method defaults"
);
method_default_action(req.method())
}
};
enforce(&claims, &state, action, Some(&db)).await?;
Ok(next.run(req).await)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_required_action_defaults() {
let get = Method::GET;
let post = Method::POST;
let put = Method::PUT;
let delete = Method::DELETE;
assert_eq!(
required_action(&get, "/_api/database/{db}/document/{collection}/{key}"),
PermissionAction::Read
);
assert_eq!(
required_action(&post, "/_api/database/{db}/document/{collection}"),
PermissionAction::Write
);
assert_eq!(
required_action(&delete, "/_api/database/{db}/document/{collection}/{key}"),
PermissionAction::Write
);
assert_eq!(
required_action(&put, "/_api/database/{db}/collection/{name}/properties"),
PermissionAction::Write
);
}
#[test]
fn test_required_action_read_overrides() {
let post = Method::POST;
for path in [
"/_api/database/{db}/cursor",
"/_api/database/{db}/explain",
"/_api/database/{db}/nl",
"/_api/database/{db}/geo/{collection}/{field}/near",
"/_api/database/{db}/geo/{collection}/{field}/within",
"/_api/database/{db}/vector/{collection}/{index}/search",
"/_api/database/{db}/hybrid/{collection}/search",
"/_api/database/{db}/columnar/{collection}/aggregate",
"/_api/database/{db}/columnar/{collection}/query",
"/_api/database/{db}/transaction/{tx_id}/query",
"/_api/database/{db}/sql",
"/_api/database/{db}/document/{collection}/_verify",
] {
assert_eq!(
required_action(&post, path),
PermissionAction::Read,
"expected Read for {}",
path
);
}
}
#[test]
fn test_read_suffix_cannot_be_forged_by_a_document_key() {
for (method, template) in [
(
Method::PUT,
"/_api/database/{db}/document/{collection}/{key}",
),
(
Method::DELETE,
"/_api/database/{db}/document/{collection}/{key}",
),
(Method::POST, "/_api/database/{db}/document/{collection}"),
(
Method::DELETE,
"/_api/database/{db}/index/{collection}/{index_name}",
),
] {
assert_eq!(
required_action(&method, template),
PermissionAction::Write,
"{} {} must stay a write",
method,
template
);
}
}
#[test]
fn test_method_default_action_never_downgrades_writes() {
assert_eq!(method_default_action(&Method::GET), PermissionAction::Read);
assert_eq!(method_default_action(&Method::HEAD), PermissionAction::Read);
for method in [Method::POST, Method::PUT, Method::PATCH, Method::DELETE] {
assert_eq!(
method_default_action(&method),
PermissionAction::Write,
"{} must default to Write",
method
);
}
}
#[test]
fn test_required_action_admin_overrides() {
let put = Method::PUT;
let post = Method::POST;
let delete = Method::DELETE;
assert_eq!(
required_action(&put, "/_api/database/{db}/collection/{name}/truncate"),
PermissionAction::Admin
);
assert_eq!(
required_action(&delete, "/_api/database/{db}/collection/{name}"),
PermissionAction::Admin
);
assert_eq!(
required_action(&delete, "/_api/database/{db}/columnar/{collection}"),
PermissionAction::Admin
);
assert_eq!(
required_action(&post, "/_api/database/{db}/repl"),
PermissionAction::Admin
);
assert_eq!(
required_action(&post, "/_api/database/{db}/scripts"),
PermissionAction::Admin
);
assert_eq!(
required_action(&put, "/_api/database/{db}/services/{key}"),
PermissionAction::Admin
);
assert_eq!(
required_action(&delete, "/_api/database/{db}/scripts/{script_id}"),
PermissionAction::Admin
);
assert_eq!(
required_action(&delete, "/_api/database/{db}/collection/{name}/schema"),
PermissionAction::Write
);
assert_eq!(
required_action(
&delete,
"/_api/database/{db}/columnar/{collection}/index/{column}"
),
PermissionAction::Write
);
}
}