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,
}
}
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 = required_action(req.method(), req.uri().path());
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/db1/document/users/k1"),
PermissionAction::Read
);
assert_eq!(
required_action(&post, "/_api/database/db1/document/users"),
PermissionAction::Write
);
assert_eq!(
required_action(&delete, "/_api/database/db1/document/users/k1"),
PermissionAction::Write
);
assert_eq!(
required_action(&put, "/_api/database/db1/collection/users/properties"),
PermissionAction::Write
);
}
#[test]
fn test_required_action_read_overrides() {
let post = Method::POST;
for path in [
"/_api/database/db1/cursor",
"/_api/database/db1/explain",
"/_api/database/db1/nl",
"/_api/database/db1/geo/places/location/near",
"/_api/database/db1/geo/places/location/within",
"/_api/database/db1/vector/docs/embeddings/search",
"/_api/database/db1/hybrid/docs/search",
"/_api/database/db1/columnar/metrics/aggregate",
"/_api/database/db1/columnar/metrics/query",
"/_api/database/db1/transaction/tx1/query",
"/_api/database/db1/sql",
"/_api/database/db1/document/users/_verify",
] {
assert_eq!(
required_action(&post, path),
PermissionAction::Read,
"expected Read for {}",
path
);
}
}
#[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/db1/collection/users/truncate"),
PermissionAction::Admin
);
assert_eq!(
required_action(&delete, "/_api/database/db1/collection/users"),
PermissionAction::Admin
);
assert_eq!(
required_action(&delete, "/_api/database/db1/columnar/metrics"),
PermissionAction::Admin
);
assert_eq!(
required_action(&post, "/_api/database/db1/repl"),
PermissionAction::Admin
);
assert_eq!(
required_action(&post, "/_api/database/db1/scripts"),
PermissionAction::Admin
);
assert_eq!(
required_action(&put, "/_api/database/db1/services/users"),
PermissionAction::Admin
);
assert_eq!(
required_action(&delete, "/_api/database/db1/scripts/abc"),
PermissionAction::Admin
);
assert_eq!(
required_action(&delete, "/_api/database/db1/collection/users/schema"),
PermissionAction::Write
);
assert_eq!(
required_action(&delete, "/_api/database/db1/columnar/metrics/index/col1"),
PermissionAction::Write
);
}
}