use std::sync::Arc;
use anyhow::Result;
use http::HeaderValue;
use reblessive::TreeStack;
use reblessive::tree::Stk;
use tracing::{debug, error, trace};
use super::response::ApiResponse;
use crate::api::X_SURREAL_REQUEST_ID;
use crate::api::err::ApiError;
use crate::api::request::ApiRequest;
use crate::catalog::providers::DatabaseProvider;
use crate::catalog::{ApiDefinition, MiddlewareDefinition, Permission};
use crate::ctx::{Context, FrozenContext};
use crate::dbs::Options;
use crate::doc::CursorDoc;
use crate::expr::{Expr, FlowResultExt as _};
use crate::fnc::args::{Any, FromArgs, FromPublic};
use crate::iam::{Action, AuthLimit};
use crate::syn::function_with_capabilities;
use crate::val::{Closure, Value};
pub async fn process_api_request(
ctx: &FrozenContext,
opt: &Options,
api: &ApiDefinition,
req: ApiRequest,
) -> Result<ApiResponse> {
let mut stack = TreeStack::new();
stack.enter(|stk| process_api_request_with_stack(stk, ctx, opt, api, req)).finish().await
}
pub async fn process_api_request_with_stack(
stk: &mut Stk,
ctx: &FrozenContext,
opt: &Options,
api: &ApiDefinition,
req: ApiRequest,
) -> Result<ApiResponse> {
let (ns_name, db_name) = opt.ns_db()?;
if !opt.auth.can_access_ns_db(ns_name, db_name) {
trace!(
request_id = %req.request_id,
"API request denied: selected namespace/database is outside the authenticated session scope"
);
return Ok(ApiResponse::from_error(ApiError::PermissionDenied, req.request_id.clone()));
}
let method_action = api.actions.iter().find(|x| x.methods.contains(&req.method));
let (action_expr, method_config) = match (method_action, &api.fallback) {
(Some(x), _) => (x.action.clone(), Some(&x.config)),
(None, Some(x)) => (x.clone(), None),
_ => {
trace!(
request_id = %req.request_id,
method = ?req.method,
"No matching handler or fallback for API request"
);
let res = ApiResponse::from_error(ApiError::NotFound, req.request_id.clone());
return Ok(res);
}
};
let (ns, db) = ctx.expect_ns_db_ids(opt).await?;
let global_entry = ctx.tx().get_db_config(ns, db, "api", None).await?;
let global = global_entry.as_ref().map(|v| v.try_as_api()).transpose()?;
if ctx.check_perms(opt, Action::Edit)? {
let permissions: Vec<&Permission> = method_config
.map(|config| &config.permissions)
.into_iter()
.chain(std::iter::once(&api.config.permissions))
.chain(global.as_ref().map(|config| &config.permissions))
.collect();
for permission in permissions {
match permission {
Permission::None => {
trace!(
request_id = %req.request_id,
"API request denied by PERMISSIONS NONE"
);
let res =
ApiResponse::from_error(ApiError::PermissionDenied, req.request_id.clone());
return Ok(res);
}
Permission::Full => (),
Permission::Specific(e) => {
let opt = &opt.new_for_permission_predicate();
if !stk
.run(|stk| e.compute(stk, ctx, opt, None))
.await
.catch_return()?
.is_truthy()
{
trace!(
request_id = %req.request_id,
"API request denied by PERMISSIONS WHERE clause"
);
let res = ApiResponse::from_error(
ApiError::PermissionDenied,
req.request_id.clone(),
);
return Ok(res);
}
}
}
}
}
let middleware: Vec<_> = global
.into_iter()
.flat_map(|cfg| cfg.middleware.iter().cloned())
.chain(api.config.middleware.iter().cloned())
.chain(method_config.into_iter().flat_map(|config| config.middleware.iter().cloned()))
.collect();
let final_action = create_final_action_closure(req.request_id.clone(), action_expr);
let middleware_len = middleware.len();
let next = middleware.iter().rev().enumerate().fold(final_action, |next, (idx, def)| {
let is_initial = idx == middleware_len.saturating_sub(1);
create_middleware_closure(req.request_id.clone(), def.clone(), next, is_initial)
});
let opt = AuthLimit::try_from(&api.auth_limit)?.limit_opt(opt);
let opt = opt.new_with_perms(false);
debug!(
request_id = %req.request_id,
middleware_count = middleware.len(),
"Executing API middleware chain"
);
let mut res: ApiResponse =
next.invoke(stk, ctx, &opt, None, vec![req.into()]).await?.try_into()?;
res.ensure_request_id_header();
Ok(res)
}
fn create_final_action_closure(request_id: String, action_expr: Expr) -> Closure {
Closure::Builtin(Arc::new(
move |stk: &mut Stk,
ctx: &FrozenContext,
opt: &Options,
doc: Option<&CursorDoc>,
args: Any| {
let (FromPublic(mut req),): (FromPublic<ApiRequest>,) =
match FromArgs::from_args("", args.0) {
Ok(v) => v,
Err(_e) => {
return Box::pin(std::future::ready(Err(
ApiError::FinalActionRequestParseFailure.into(),
)));
}
};
req.request_id.clone_from(&request_id);
if !request_id.is_empty() {
let _ = req.headers.insert(
X_SURREAL_REQUEST_ID,
HeaderValue::from_str(&request_id)
.unwrap_or_else(|_| HeaderValue::from_static("unknown")),
);
}
let mut ctx_isolated = Context::new_isolated(ctx);
ctx_isolated.add_value("request", Arc::new(req.into()));
let ctx_frozen = ctx_isolated.freeze();
let action_expr = action_expr.clone();
let request_id = request_id.clone();
Box::pin(stk.run(async move |stk| {
let res = action_expr.compute(stk, &ctx_frozen, opt, doc).await.catch_return();
let mut res = match res {
Ok(res) => ApiResponse::try_from(res)
.unwrap_or_else(|e| ApiResponse::from_error(e, request_id.clone())),
Err(e) => ApiResponse::from_error(e, request_id.clone()),
};
res.request_id.clone_from(&request_id);
res.ensure_request_id_header();
Ok(Value::from(res))
}))
},
))
}
fn create_middleware_closure(
request_id: String,
def: MiddlewareDefinition,
next: Closure,
is_initial: bool,
) -> Closure {
Closure::Builtin(Arc::new(
move |stk: &mut Stk,
ctx: &FrozenContext,
opt: &Options,
doc: Option<&CursorDoc>,
args: Any| {
let def = def.clone();
let (FromPublic(mut req),): (FromPublic<ApiRequest>,) =
match FromArgs::from_args("", args.0) {
Ok(v) => v,
Err(_e) => {
return Box::pin(std::future::ready(Err(
ApiError::MiddlewareRequestParseFailure {
middleware: def.name.to_string(),
}
.into(),
)));
}
};
let function: crate::expr::Function = match function_with_capabilities(
def.name.as_str(),
ctx.get_capabilities().as_ref(),
&ctx.config,
) {
Ok(f) => f.into(),
Err(_e) => {
return Box::pin(std::future::ready(Err(
ApiError::MiddlewareFunctionNotFound {
function: def.name.to_string(),
}
.into(),
)));
}
};
req.request_id.clone_from(&request_id);
if !request_id.is_empty() {
let _ = req.headers.insert(
X_SURREAL_REQUEST_ID,
HeaderValue::from_str(&request_id)
.unwrap_or_else(|_| HeaderValue::from_static("unknown")),
);
}
let mut fn_args = vec![Value::from(req), Value::Closure(Box::new(next.clone()))];
fn_args.extend(def.args);
let ctx = Context::new_isolated(ctx).freeze();
let opt = opt.clone();
let doc = doc.cloned();
let middleware_name = def.name;
let request_id = request_id.clone();
Box::pin(stk.run(async move |stk| {
let res =
function.compute(stk, &ctx, &opt, doc.as_ref(), fn_args).await.catch_return();
let mut res = match res {
Ok(res) => match ApiResponse::try_from(res) {
Ok(mut res) => {
res.request_id.clone_from(&request_id);
res
}
Err(e) => {
if is_initial {
error!(
request_id = %request_id,
middleware = %middleware_name,
error = %e,
"API middleware error; converting to response (ApiError exposed, internal errors masked)"
);
ApiResponse::from_error_secure(e, request_id.clone())
} else {
ApiResponse::from_error(e, request_id.clone())
}
}
},
Err(e) => {
if is_initial {
error!(
request_id = %request_id,
middleware = %middleware_name,
error = %e,
"API middleware error; converting to response (ApiError exposed, internal errors masked)"
);
ApiResponse::from_error_secure(e, request_id.clone())
} else {
ApiResponse::from_error(e, request_id.clone())
}
}
};
res.ensure_request_id_header();
Ok(Value::from(res))
}))
},
))
}