use std::collections::HashMap;
use std::fmt::{self, Debug};
#[cfg(storage)]
use std::path::PathBuf;
use std::sync::Arc;
use std::sync::atomic::{AtomicBool, Ordering};
use std::time::Duration;
#[cfg(feature = "surrealism")]
use anyhow::Context as _;
use anyhow::{Result, bail};
use surrealdb_strand::Strand;
#[cfg(feature = "surrealism")]
use surrealism_runtime::package::{SurrealismPackage, UnpackOptions};
#[cfg(feature = "surrealism")]
use surrealism_runtime::runtime::Runtime;
#[cfg(feature = "http")]
use url::Url;
use uuid::Uuid;
use web_time::Instant;
use crate::buc::manager::BucketsManager;
#[cfg(feature = "surrealism")]
use crate::buc::store::ObjectKey;
use crate::buc::store::ObjectStore;
use crate::catalog::providers::{CatalogProvider, DatabaseProvider, NamespaceProvider};
use crate::catalog::{DatabaseDefinition, DatabaseId, NamespaceId};
use crate::cnf::dynamic::DynamicConfiguration;
use crate::cnf::{CommonConfig, PROTECTED_PARAM_NAMES};
use crate::ctx::cancel::CancelHandle;
use crate::ctx::canceller::Canceller;
use crate::ctx::reason::Reason;
#[cfg(feature = "surrealism")]
use crate::dbs::capabilities::ExperimentalTarget;
#[cfg(feature = "http")]
use crate::dbs::capabilities::NetTarget;
#[cfg(all(feature = "http", feature = "surrealism"))]
use crate::dbs::capabilities::Targets;
use crate::dbs::{
Capabilities, MessageBroker, NewPlannerStrategy, Options, Session, StatementCounters, Variables,
};
use crate::err::Error;
use crate::exec::function::FunctionRegistry;
use crate::expr::Base;
#[cfg(feature = "http")]
use crate::http::HttpClient;
use crate::iam::{Action, ResourceKind};
use crate::idx::planner::executor::QueryExecutor;
use crate::idx::planner::{IterationStage, QueryPlanner};
use crate::idx::trees::store::IndexStores;
use crate::kvs::Transaction;
use crate::kvs::cache::ds::DatastoreCache;
use crate::kvs::index::IndexBuilder;
use crate::kvs::sequences::Sequences;
use crate::kvs::slowlog::SlowLog;
use crate::mem::ALLOC;
use crate::sql::expression::convert_public_value_to_internal;
#[cfg(feature = "surrealism")]
use crate::surrealism::cache::{SurrealismCache, SurrealismCacheLookup, SurrealismCachedModule};
use crate::types::PublicVariables;
use crate::val::Value;
pub type FrozenContext = Arc<Context>;
pub struct Context {
parent: Option<FrozenContext>,
deadline: Option<(Instant, Duration)>,
slow_log: Option<SlowLog>,
cancelled: Arc<AtomicBool>,
cancel_token: Option<tokio_util::sync::CancellationToken>,
values: HashMap<Strand, Arc<Value>>,
query_planner: Option<Arc<QueryPlanner>>,
query_executor: Option<QueryExecutor>,
iteration_stage: Option<IterationStage>,
cache: Option<Arc<DatastoreCache>>,
index_stores: IndexStores,
index_builder: Option<IndexBuilder>,
sequences: Option<Sequences>,
capabilities: Arc<Capabilities>,
#[cfg(storage)]
temporary_directory: Option<Arc<PathBuf>>,
transaction: Option<Arc<Transaction>>,
isolated: bool,
buckets: Option<BucketsManager>,
#[cfg(feature = "surrealism")]
surrealism_cache: Option<Arc<SurrealismCache>>,
function_registry: Arc<FunctionRegistry>,
new_planner_strategy: NewPlannerStrategy,
redact_volatile_explain_attrs: bool,
statement_counters: Option<Arc<StatementCounters>>,
tenant_identity: Option<Arc<crate::observe::TenantIdentity>>,
matches_context: Option<Arc<crate::exec::function::MatchesContext>>,
knn_context: Option<Arc<crate::exec::function::KnnContext>>,
#[cfg(feature = "http")]
http_client: Arc<HttpClient>,
node_id: Uuid,
pub(crate) auth_enabled: bool,
dynamic_configuration: DynamicConfiguration,
live: bool,
broker: Option<Arc<dyn MessageBroker>>,
pub config: Arc<CommonConfig>,
}
impl Debug for Context {
fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
f.debug_struct("Context")
.field("parent", &self.parent)
.field("deadline", &self.deadline)
.field("cancelled", &self.cancelled)
.field("values", &self.values)
.finish()
}
}
impl Context {
pub(crate) fn background(parent: &Context) -> Self {
Self {
values: HashMap::default(),
parent: None,
deadline: None,
slow_log: None,
cancelled: Arc::new(AtomicBool::new(false)),
cancel_token: None,
query_planner: None,
query_executor: None,
iteration_stage: None,
capabilities: Arc::clone(&parent.capabilities),
index_stores: IndexStores::new(
parent.config.hnsw_cache_size,
parent.config.diskann_cache_size,
),
cache: None,
index_builder: None,
sequences: None,
#[cfg(storage)]
temporary_directory: None,
transaction: None,
isolated: false,
buckets: None,
#[cfg(feature = "surrealism")]
surrealism_cache: None,
function_registry: Arc::clone(&parent.function_registry),
new_planner_strategy: NewPlannerStrategy::default(),
redact_volatile_explain_attrs: false,
statement_counters: None,
matches_context: None,
knn_context: None,
config: Arc::clone(&parent.config),
#[cfg(feature = "http")]
http_client: Arc::clone(&parent.http_client),
tenant_identity: None,
node_id: parent.node_id,
auth_enabled: parent.auth_enabled,
dynamic_configuration: parent.dynamic_configuration.clone(),
live: parent.live,
broker: parent.broker.clone(),
}
}
pub(crate) fn new_child(parent: &FrozenContext) -> Self {
Self::new_child_with_capabilities(
parent,
Arc::clone(&parent.capabilities),
#[cfg(feature = "http")]
Arc::clone(&parent.http_client),
)
}
pub(crate) fn new_child_with_capabilities(
parent: &FrozenContext,
cap: Arc<Capabilities>,
#[cfg(feature = "http")] http_client: Arc<HttpClient>,
) -> Self {
Context {
values: HashMap::default(),
deadline: parent.deadline,
slow_log: parent.slow_log.clone(),
cancelled: Arc::new(AtomicBool::new(false)),
cancel_token: None,
query_planner: parent.query_planner.clone(),
query_executor: parent.query_executor.clone(),
iteration_stage: parent.iteration_stage.clone(),
capabilities: cap,
index_stores: parent.index_stores.clone(),
cache: parent.cache.clone(),
index_builder: parent.index_builder.clone(),
sequences: parent.sequences.clone(),
#[cfg(storage)]
temporary_directory: parent.temporary_directory.clone(),
transaction: parent.transaction.clone(),
isolated: false,
parent: Some(Arc::clone(parent)),
buckets: parent.buckets.clone(),
#[cfg(feature = "surrealism")]
surrealism_cache: parent.surrealism_cache.clone(),
function_registry: Arc::clone(&parent.function_registry),
new_planner_strategy: parent.new_planner_strategy,
redact_volatile_explain_attrs: parent.redact_volatile_explain_attrs,
statement_counters: parent.statement_counters.clone(),
matches_context: parent.matches_context.clone(),
knn_context: parent.knn_context.clone(),
config: Arc::clone(&parent.config),
#[cfg(feature = "http")]
http_client,
tenant_identity: parent.tenant_identity.clone(),
node_id: parent.node_id,
auth_enabled: parent.auth_enabled,
dynamic_configuration: parent.dynamic_configuration.clone(),
live: parent.live,
broker: parent.broker.clone(),
}
}
pub(crate) fn new_isolated(parent: &FrozenContext) -> Self {
Self {
values: HashMap::default(),
deadline: parent.deadline,
slow_log: parent.slow_log.clone(),
cancelled: Arc::new(AtomicBool::new(false)),
cancel_token: None,
query_planner: parent.query_planner.clone(),
query_executor: parent.query_executor.clone(),
iteration_stage: parent.iteration_stage.clone(),
capabilities: Arc::clone(&parent.capabilities),
index_stores: parent.index_stores.clone(),
cache: parent.cache.clone(),
index_builder: parent.index_builder.clone(),
sequences: parent.sequences.clone(),
#[cfg(storage)]
temporary_directory: parent.temporary_directory.clone(),
transaction: parent.transaction.clone(),
isolated: true,
parent: Some(Arc::clone(parent)),
buckets: parent.buckets.clone(),
#[cfg(feature = "surrealism")]
surrealism_cache: parent.surrealism_cache.clone(),
function_registry: Arc::clone(&parent.function_registry),
new_planner_strategy: parent.new_planner_strategy,
redact_volatile_explain_attrs: parent.redact_volatile_explain_attrs,
statement_counters: parent.statement_counters.clone(),
matches_context: parent.matches_context.clone(),
knn_context: parent.knn_context.clone(),
config: Arc::clone(&parent.config),
#[cfg(feature = "http")]
http_client: Arc::clone(&parent.http_client),
tenant_identity: parent.tenant_identity.clone(),
node_id: parent.node_id,
auth_enabled: parent.auth_enabled,
dynamic_configuration: parent.dynamic_configuration.clone(),
live: parent.live,
broker: parent.broker.clone(),
}
}
pub(crate) fn snapshot(from: &FrozenContext) -> Self {
Self {
values: from.collect_values(HashMap::default()),
deadline: from.deadline,
slow_log: from.slow_log.clone(),
cancelled: Arc::clone(&from.cancelled),
cancel_token: from.cancel_token.clone(),
query_planner: from.query_planner.clone(),
query_executor: from.query_executor.clone(),
iteration_stage: from.iteration_stage.clone(),
capabilities: Arc::clone(&from.capabilities),
index_stores: from.index_stores.clone(),
cache: from.cache.clone(),
index_builder: from.index_builder.clone(),
sequences: from.sequences.clone(),
#[cfg(storage)]
temporary_directory: from.temporary_directory.clone(),
transaction: from.transaction.clone(),
isolated: false,
parent: None, buckets: from.buckets.clone(),
#[cfg(feature = "surrealism")]
surrealism_cache: from.surrealism_cache.clone(),
function_registry: Arc::clone(&from.function_registry),
new_planner_strategy: from.new_planner_strategy,
redact_volatile_explain_attrs: from.redact_volatile_explain_attrs,
statement_counters: from.statement_counters.clone(),
matches_context: from.matches_context.clone(),
knn_context: from.knn_context.clone(),
config: Arc::clone(&from.config),
#[cfg(feature = "http")]
http_client: Arc::clone(&from.http_client),
tenant_identity: from.tenant_identity.clone(),
node_id: from.node_id,
auth_enabled: from.auth_enabled,
dynamic_configuration: from.dynamic_configuration.clone(),
live: from.live,
broker: from.broker.clone(),
}
}
pub(crate) fn new_concurrent(from: &FrozenContext) -> Self {
Self {
values: HashMap::default(),
deadline: None,
slow_log: from.slow_log.clone(),
cancelled: Arc::new(AtomicBool::new(false)),
cancel_token: None,
query_planner: from.query_planner.clone(),
query_executor: from.query_executor.clone(),
iteration_stage: from.iteration_stage.clone(),
capabilities: Arc::clone(&from.capabilities),
index_stores: from.index_stores.clone(),
cache: from.cache.clone(),
index_builder: None,
sequences: from.sequences.clone(),
#[cfg(storage)]
temporary_directory: from.temporary_directory.clone(),
transaction: None,
isolated: false,
parent: None,
buckets: from.buckets.clone(),
#[cfg(feature = "surrealism")]
surrealism_cache: from.surrealism_cache.clone(),
function_registry: Arc::clone(&from.function_registry),
new_planner_strategy: from.new_planner_strategy,
redact_volatile_explain_attrs: from.redact_volatile_explain_attrs,
statement_counters: from.statement_counters.clone(),
matches_context: from.matches_context.clone(),
knn_context: from.knn_context.clone(),
config: Arc::clone(&from.config),
#[cfg(feature = "http")]
http_client: Arc::clone(&from.http_client),
tenant_identity: from.tenant_identity.clone(),
node_id: from.node_id,
auth_enabled: from.auth_enabled,
dynamic_configuration: from.dynamic_configuration.clone(),
live: from.live,
broker: from.broker.clone(),
}
}
#[expect(clippy::too_many_arguments)]
pub(crate) fn from_ds(
node_id: Uuid,
auth_enabled: bool,
dynamic_configuration: DynamicConfiguration,
time_out: Option<Duration>,
slow_log: Option<SlowLog>,
capabilities: Arc<Capabilities>,
index_stores: IndexStores,
index_builder: IndexBuilder,
sequences: Sequences,
cache: Arc<DatastoreCache>,
function_registry: Arc<FunctionRegistry>,
#[cfg(feature = "http")] http_client: Arc<HttpClient>,
#[cfg(storage)] temporary_directory: Option<Arc<PathBuf>>,
buckets: BucketsManager,
config: Arc<CommonConfig>,
#[cfg(feature = "surrealism")] surrealism_cache: Arc<SurrealismCache>,
) -> Result<Context> {
let planner_strategy = *capabilities.planner_strategy();
let mut ctx = Self {
values: HashMap::default(),
parent: None,
deadline: None,
slow_log,
cancelled: Arc::new(AtomicBool::new(false)),
cancel_token: None,
query_planner: None,
query_executor: None,
iteration_stage: None,
capabilities,
index_stores,
cache: Some(cache),
index_builder: Some(index_builder),
sequences: Some(sequences),
#[cfg(storage)]
temporary_directory,
transaction: None,
isolated: false,
buckets: Some(buckets),
#[cfg(feature = "surrealism")]
surrealism_cache: Some(surrealism_cache),
function_registry,
new_planner_strategy: planner_strategy,
redact_volatile_explain_attrs: false,
statement_counters: None,
matches_context: None,
knn_context: None,
config,
#[cfg(feature = "http")]
http_client,
tenant_identity: None,
node_id,
auth_enabled,
dynamic_configuration,
live: false,
broker: None,
};
if let Some(timeout) = time_out {
ctx.add_timeout(timeout)?;
}
Ok(ctx)
}
#[cfg(test)]
pub(crate) fn new_test() -> Context {
Self {
values: HashMap::default(),
parent: None,
deadline: None,
slow_log: None,
cancelled: Arc::new(AtomicBool::new(false)),
cancel_token: None,
query_planner: None,
query_executor: None,
iteration_stage: None,
capabilities: Arc::new(Capabilities::default()),
index_stores: IndexStores::new(256 * 1024 * 1024, 256 * 1024 * 1024),
cache: None,
index_builder: None,
sequences: None,
#[cfg(storage)]
temporary_directory: None,
transaction: None,
isolated: false,
buckets: None,
#[cfg(feature = "surrealism")]
surrealism_cache: None,
function_registry: Arc::new(FunctionRegistry::with_builtins()),
new_planner_strategy: NewPlannerStrategy::default(),
redact_volatile_explain_attrs: false,
statement_counters: None,
matches_context: None,
knn_context: None,
config: Default::default(),
#[cfg(feature = "http")]
http_client: Arc::new(
HttpClient::new(
crate::dbs::capabilities::Targets::All,
crate::dbs::capabilities::Targets::None,
&Default::default(),
)
.expect("http client to be created"),
),
tenant_identity: None,
node_id: Uuid::nil(),
auth_enabled: true,
dynamic_configuration: DynamicConfiguration::default(),
live: false,
broker: None,
}
}
pub(crate) fn freeze(self) -> FrozenContext {
Arc::new(self)
}
pub(crate) fn unfreeze(ctx: FrozenContext) -> Result<Context> {
let Some(x) = Arc::into_inner(ctx) else {
fail!("Tried to unfreeze a Context with multiple references")
};
Ok(x)
}
#[inline]
pub fn node_id(&self) -> Uuid {
self.node_id
}
#[inline]
pub(crate) fn auth_enabled(&self) -> bool {
self.auth_enabled
}
pub(crate) fn dynamic_configuration(&self) -> &DynamicConfiguration {
&self.dynamic_configuration
}
pub(crate) fn realtime(&self) -> Result<()> {
if !self.live {
bail!(Error::RealtimeDisabled);
}
Ok(())
}
pub(crate) fn broker(&self) -> Option<&Arc<dyn MessageBroker>> {
self.broker.as_ref()
}
pub(crate) fn set_broker(&mut self, broker: Option<Arc<dyn MessageBroker>>) {
self.broker = broker;
}
pub fn is_allowed(
&self,
opt: &Options,
action: Action,
res: ResourceKind,
base: Base,
) -> Result<()> {
let res = match base {
Base::Root => res.on_root(),
Base::Ns => res.on_ns(opt.ns()?),
Base::Db => {
let (ns, db) = opt.ns_db()?;
res.on_db(ns, db)
}
};
if !self.auth_enabled && opt.auth.is_anon() {
return Ok(());
}
opt.auth.is_allowed(action, &res)
}
pub fn check_perms(&self, opt: &Options, action: Action) -> Result<bool> {
if !opt.perms {
return Ok(false);
}
if !self.auth_enabled && opt.auth.is_anon() {
return Ok(false);
}
match action {
Action::Edit => {
let allowed = opt.auth.has_editor_role();
let (ns, db) = opt.ns_db()?;
let db_in_actor_level =
opt.auth.is_root() || opt.auth.is_ns_check(ns) || opt.auth.is_db_check(ns, db);
Ok(!allowed || !db_in_actor_level)
}
Action::View => {
let allowed = opt.auth.has_viewer_role();
let (ns, db) = opt.ns_db()?;
let db_in_actor_level =
opt.auth.is_root() || opt.auth.is_ns_check(ns) || opt.auth.is_db_check(ns, db);
Ok(!allowed || !db_in_actor_level)
}
}
}
pub(crate) async fn get_ns_id(&self, opt: &Options) -> Result<NamespaceId> {
let ns = opt.ns()?;
let tx = self.tx();
let ns_def = tx.get_or_add_ns(Some(self), ns).await?;
Ok(ns_def.namespace_id)
}
pub(crate) async fn expect_ns_id(&self, opt: &Options) -> Result<NamespaceId> {
let ns = opt.ns()?;
let Some(ns_def) = self.tx().get_ns_by_name(ns, None).await? else {
return Err(Error::NsNotFound {
name: ns.to_string(),
}
.into());
};
Ok(ns_def.namespace_id)
}
pub(crate) async fn get_ns_db_ids(&self, opt: &Options) -> Result<(NamespaceId, DatabaseId)> {
let (ns, db) = opt.ns_db()?;
let db_def = self.tx().ensure_ns_db(Some(self), ns, db).await?;
Ok((db_def.namespace_id, db_def.database_id))
}
pub(crate) async fn try_ns_db_ids(
&self,
opt: &Options,
) -> Result<Option<(NamespaceId, DatabaseId)>> {
let (ns, db) = opt.ns_db()?;
let Some(db_def) = self.tx().get_db_by_name(ns, db, None).await? else {
return Ok(None);
};
Ok(Some((db_def.namespace_id, db_def.database_id)))
}
pub(crate) async fn expect_ns_db_ids(
&self,
opt: &Options,
) -> Result<(NamespaceId, DatabaseId)> {
let (ns, db) = opt.ns_db()?;
let Some(db_def) = self.tx().get_db_by_name(ns, db, None).await? else {
return Err(Error::DbNotFound {
name: db.to_string(),
}
.into());
};
Ok((db_def.namespace_id, db_def.database_id))
}
pub(crate) async fn get_db(&self, opt: &Options) -> Result<Arc<DatabaseDefinition>> {
let (ns, db) = opt.ns_db()?;
let db_def = self.tx().ensure_ns_db(Some(self), ns, db).await?;
Ok(db_def)
}
pub(crate) fn add_value<K>(&mut self, key: K, value: Arc<Value>)
where
K: Into<Strand>,
{
self.values.insert(key.into(), value);
}
pub(crate) fn add_values<T, K, V>(&mut self, iter: T)
where
T: IntoIterator<Item = (K, V)>,
K: Into<Strand>,
V: Into<Arc<Value>>,
{
self.values.extend(iter.into_iter().map(|(k, v)| (k.into(), v.into())))
}
pub(crate) fn add_cancel(&mut self) -> Canceller {
let cancelled = Arc::clone(&self.cancelled);
Canceller::new(cancelled)
}
pub(crate) fn set_cancellation(&mut self, handle: &CancelHandle) {
self.cancelled = handle.flag();
self.cancel_token = Some(handle.token());
}
pub(crate) fn cancel_token(&self) -> Option<tokio_util::sync::CancellationToken> {
if let Some(token) = &self.cancel_token {
return Some(token.clone());
}
self.parent.as_ref().and_then(|p| p.cancel_token())
}
pub(crate) fn add_deadline(&mut self, deadline: Instant, duration: Duration) {
match self.deadline {
Some((current, _)) if current < deadline => (),
_ => self.deadline = Some((deadline, duration)),
}
}
pub(crate) fn add_timeout(&mut self, timeout: Duration) -> Result<(), Error> {
match Instant::now().checked_add(timeout) {
Some(deadline) => {
self.add_deadline(deadline, timeout);
Ok(())
}
None => Err(Error::InvalidTimeout(timeout.as_secs())),
}
}
pub(crate) fn set_query_planner(&mut self, qp: QueryPlanner) {
self.query_planner = Some(Arc::new(qp));
}
pub(crate) fn set_query_executor(&mut self, qe: QueryExecutor) {
self.query_executor = Some(qe);
}
pub(crate) fn set_iteration_stage(&mut self, is: IterationStage) {
self.iteration_stage = Some(is);
}
pub(crate) fn set_transaction(&mut self, txn: Arc<Transaction>) {
self.transaction = Some(txn);
}
pub(crate) fn set_statement_counters(&mut self, counters: Option<Arc<StatementCounters>>) {
self.statement_counters = counters;
}
pub(crate) fn statement_counters(&self) -> Option<&Arc<StatementCounters>> {
self.statement_counters.as_ref()
}
pub(crate) fn tx(&self) -> Arc<Transaction> {
self.transaction
.clone()
.unwrap_or_else(|| unreachable!("The context was not associated with a transaction"))
}
pub(crate) fn try_tx(&self) -> Option<&Arc<Transaction>> {
self.transaction.as_ref()
}
pub(crate) fn timeout(&self) -> Option<Duration> {
self.deadline.map(|(v, _)| v.saturating_duration_since(Instant::now()))
}
pub(crate) fn slow_log(&self) -> Option<&SlowLog> {
self.slow_log.as_ref()
}
pub(crate) fn get_query_planner(&self) -> Option<&QueryPlanner> {
self.query_planner.as_ref().map(|qp| qp.as_ref())
}
pub(crate) fn get_query_executor(&self) -> Option<&QueryExecutor> {
self.query_executor.as_ref()
}
pub(crate) fn get_iteration_stage(&self) -> Option<&IterationStage> {
self.iteration_stage.as_ref()
}
pub(crate) fn get_index_stores(&self) -> &IndexStores {
&self.index_stores
}
pub(crate) fn get_index_builder(&self) -> Option<&IndexBuilder> {
self.index_builder.as_ref()
}
pub(crate) fn get_sequences(&self) -> Option<&Sequences> {
self.sequences.as_ref()
}
pub(crate) fn try_get_sequences(&self) -> Result<&Sequences> {
if let Some(sqs) = self.get_sequences() {
Ok(sqs)
} else {
bail!(Error::Internal("Sequences are not supported in this context.".to_string(),))
}
}
pub(crate) fn get_cache(&self) -> Option<Arc<DatastoreCache>> {
self.cache.clone()
}
pub(crate) fn done(&self, deep_check: bool) -> Result<Option<Reason>> {
if self.cancelled.load(Ordering::Relaxed) {
return Ok(Some(Reason::Canceled));
}
if deep_check {
if ALLOC.is_beyond_threshold() {
bail!(Error::QueryBeyondMemoryThreshold);
}
let now = Instant::now();
if let Some((deadline, timeout)) = self.deadline
&& deadline <= now
{
return Ok(Some(Reason::Timedout(timeout.into())));
}
}
if let Some(ctx) = &self.parent {
return ctx.done(deep_check);
}
Ok(None)
}
pub(crate) async fn is_done(&self, count: Option<usize>) -> Result<bool> {
let deep_check = if let Some(count) = count {
if count % 32 == 0 {
yield_now!();
}
match count {
1 | 2 | 4 | 8 | 16 | 32 => true,
_ => count % 64 == 0,
}
} else {
true
};
Ok(self.done(deep_check)?.is_some())
}
pub(crate) async fn is_timedout(&self) -> Result<Option<Duration>> {
yield_now!();
if let Some(Reason::Timedout(d)) = self.done(true)? {
Ok(Some(d.0))
} else {
Ok(None)
}
}
pub(crate) async fn expect_not_timedout(&self) -> Result<()> {
if let Some(d) = self.is_timedout().await? {
bail!(Error::QueryTimedout(d.into()))
} else {
Ok(())
}
}
#[cfg(storage)]
pub(crate) fn temporary_directory(&self) -> Option<&Arc<PathBuf>> {
self.temporary_directory.as_ref()
}
pub(crate) fn value(&self, key: &str) -> Option<&Value> {
match self.values.get(key) {
Some(v) => Some(v.as_ref()),
None if PROTECTED_PARAM_NAMES.contains(&key) || !self.isolated => match &self.parent {
Some(p) => p.value(key),
_ => None,
},
None => None,
}
}
pub(crate) fn collect_values(
&self,
map: HashMap<Strand, Arc<Value>>,
) -> HashMap<Strand, Arc<Value>> {
let mut map = if !self.isolated
&& let Some(p) = &self.parent
{
p.collect_values(map)
} else {
map
};
self.values.iter().for_each(|(k, v)| {
map.insert(k.clone(), Arc::clone(v));
});
map
}
#[cfg(feature = "scripting")]
pub(crate) fn cancellation(&self) -> crate::ctx::cancellation::Cancellation {
crate::ctx::cancellation::Cancellation::new(
self.deadline.map(|(deadline, _)| deadline),
std::iter::successors(Some(self), |ctx| ctx.parent.as_ref().map(|c| c.as_ref()))
.map(|ctx| Arc::clone(&ctx.cancelled))
.collect(),
)
}
pub(crate) fn attach_session(&mut self, session: &Session) -> Result<(), Error> {
self.live = session.live();
self.add_values(session.values());
if session.new_planner_strategy != NewPlannerStrategy::default() {
self.new_planner_strategy = session.new_planner_strategy;
}
if session.redact_volatile_explain_attrs {
self.redact_volatile_explain_attrs = true;
}
if !session.variables.is_empty() {
self.attach_variables(session.variables.clone().into())?;
}
self.tenant_identity =
Some(Arc::new(crate::observe::TenantIdentity::from_session(session)));
Ok(())
}
pub(crate) fn tenant_identity(&self) -> Option<&Arc<crate::observe::TenantIdentity>> {
self.tenant_identity.as_ref()
}
pub(crate) fn attach_variables(&mut self, vars: Variables) -> Result<(), Error> {
for (name, val) in vars {
if PROTECTED_PARAM_NAMES.contains(&name.as_str()) {
return Err(Error::InvalidParam {
name: name.into_string(),
});
}
self.add_value(name, Arc::new(val));
}
Ok(())
}
pub(crate) fn attach_public_variables(&mut self, vars: PublicVariables) -> Result<(), Error> {
for (name, val) in vars {
if PROTECTED_PARAM_NAMES.contains(&name.as_str()) {
return Err(Error::InvalidParam {
name,
});
}
self.add_value(name, Arc::new(convert_public_value_to_internal(val)));
}
Ok(())
}
pub(crate) fn get_capabilities(&self) -> Arc<Capabilities> {
Arc::clone(&self.capabilities)
}
pub(crate) fn function_registry(&self) -> &Arc<FunctionRegistry> {
&self.function_registry
}
pub(crate) fn set_matches_context(&mut self, ctx: crate::exec::function::MatchesContext) {
self.matches_context = Some(Arc::new(ctx));
}
pub(crate) fn get_matches_context(
&self,
) -> Option<&Arc<crate::exec::function::MatchesContext>> {
self.matches_context.as_ref()
}
pub(crate) fn set_knn_context(&mut self, ctx: Arc<crate::exec::function::KnnContext>) {
self.knn_context = Some(ctx);
}
pub(crate) fn get_knn_context(&self) -> Option<&Arc<crate::exec::function::KnnContext>> {
self.knn_context.as_ref()
}
pub(crate) fn new_planner_strategy(&self) -> &NewPlannerStrategy {
&self.new_planner_strategy
}
pub(crate) fn redact_volatile_explain_attrs(&self) -> bool {
self.redact_volatile_explain_attrs
}
#[cfg_attr(not(feature = "scripting"), expect(dead_code))]
pub(crate) fn check_allowed_scripting(&self) -> Result<()> {
if !self.capabilities.allows_scripting() {
warn!("Capabilities denied scripting attempt");
bail!(Error::ScriptingNotAllowed);
}
trace!("Capabilities allowed scripting");
Ok(())
}
pub(crate) fn check_allowed_function(&self, target: &str) -> Result<()> {
if !self.capabilities.allows_function_name(target) {
warn!("Capabilities denied function execution attempt, target: '{target}'");
bail!(Error::FunctionNotAllowed(target.to_string()));
}
trace!("Capabilities allowed function execution, target: '{target}'");
Ok(())
}
#[cfg(feature = "http")]
pub(crate) async fn check_allowed_net(&self, url: &Url) -> Result<()> {
let match_any_deny_net = |t| {
if self.capabilities.matches_any_deny_net(t) {
warn!("Capabilities denied outgoing network connection attempt, target: '{t}'");
bail!(Error::NetTargetNotAllowed(t.to_string()));
}
Ok(())
};
match url.host() {
Some(host) => {
let target = NetTarget::Host(host.to_owned(), url.port_or_known_default());
let host_allowed = self.capabilities.matches_any_allow_net(&target);
if !host_allowed {
warn!(
"Capabilities denied outgoing network connection attempt, target: '{target}'"
);
bail!(Error::NetTargetNotAllowed(target.to_string()));
}
match_any_deny_net(&target)?;
#[cfg(not(target_family = "wasm"))]
let targets = target.resolve().await?;
#[cfg(target_family = "wasm")]
let targets = target.resolve()?;
for t in &targets {
match_any_deny_net(t)?;
}
trace!("Capabilities allowed outgoing network connection, target: '{target}'");
Ok(())
}
_ => bail!(Error::InvalidUrl(url.to_string())),
}
}
pub(crate) fn get_buckets(&self) -> Option<&BucketsManager> {
self.buckets.as_ref()
}
pub(crate) async fn get_bucket_store(
&self,
ns: NamespaceId,
db: DatabaseId,
bu: &str,
) -> Result<Arc<dyn ObjectStore>> {
if let Some(buckets) = &self.buckets {
buckets.get_bucket_store(&self.tx(), ns, db, bu).await
} else {
bail!(Error::BucketUnavailable(bu.into()))
}
}
#[cfg(feature = "surrealism")]
pub(crate) fn get_surrealism_cache(&self) -> Option<Arc<SurrealismCache>> {
self.surrealism_cache.as_ref().map(Arc::clone)
}
#[cfg(feature = "surrealism")]
pub(crate) async fn get_surrealism_module(
&self,
lookup: SurrealismCacheLookup<'_>,
) -> Result<SurrealismCachedModule> {
if !self.get_capabilities().allows_experimental(&ExperimentalTarget::Surrealism) {
bail!(
"Failed to get surrealism runtime: Experimental capability `surrealism` is not enabled"
);
}
let Some(cache) = self.get_surrealism_cache() else {
bail!("Surrealism cache is not available");
};
let max_pool_size = self.config.surrealism_max_pool_size;
let max_memory = self.config.surrealism_max_memory;
let max_execution_time =
self.config.surrealism_max_execution_time.map(Duration::from_millis);
let max_kv_entries = self.config.surrealism_max_kv_entries;
let max_kv_value_bytes = self.config.surrealism_max_kv_value_bytes;
#[cfg(feature = "http")]
let config = Arc::clone(&self.config);
cache
.get_or_insert_with(&lookup, async || {
let SurrealismCacheLookup::File(ns, db, bucket, key) = lookup else {
bail!("silo lookups are not supported yet");
};
let bucket = self.get_bucket_store(*ns, *db, bucket).await?;
let key = ObjectKey::new(key);
let surli = bucket
.get(&key)
.await
.map_err(|e| anyhow::anyhow!("failed to get file: {}", e))?;
let Some(surli) = surli else {
bail!("file not found");
};
let safe_key = key.to_string().replace(['/', '\\'], "_");
let temp_prefix = format!("SURREAL_MODFS_{ns}_{db}_{safe_key}_");
let unpack_opts = UnpackOptions {
#[cfg(storage)]
temp_base: self.temporary_directory().map(|p| p.as_path()),
#[cfg(not(storage))]
temp_base: None,
temp_prefix: &temp_prefix,
max_fs_bytes: self.config.surrealism_max_fs_bytes,
};
let package =
SurrealismPackage::from_reader(std::io::Cursor::new(surli), &unpack_opts)?;
self.get_capabilities()
.validate_surrealism_capabilities(&package.config.capabilities)?;
let org = package.config.meta.organisation.clone();
let name = package.config.meta.name.clone();
#[cfg(feature = "http")]
let module_net_targets =
crate::surrealism::host::module_allow_net_targets(&package.config.capabilities);
let runtime = tokio::task::spawn_blocking(move || {
Runtime::new(
package,
max_pool_size,
max_memory,
max_execution_time,
max_kv_entries,
max_kv_value_bytes,
)
})
.await
.context("WASM compile task aborted")??;
let runtime = Arc::new(runtime);
let module_display_name: Arc<str> = format!("{org}::{name}").into();
#[cfg(feature = "http")]
let client = if module_net_targets.is_empty() {
Arc::new(
HttpClient::new(Targets::None, Targets::All, &config)
.context("Failed to create http client for WASM module")?,
)
} else {
let allow = Targets::Some(module_net_targets);
Arc::new(
HttpClient::new(
allow,
self.capabilities.denied_network_targets_ref().clone(),
&config,
)
.context("Failed to create http client for WASM module")?,
)
};
Ok(SurrealismCachedModule {
runtime,
module_display_name,
#[cfg(feature = "http")]
client,
})
})
.await
}
#[cfg(feature = "http")]
pub(crate) fn http_client(&self) -> Arc<HttpClient> {
Arc::clone(&self.http_client)
}
#[cfg(feature = "surrealism")]
pub(crate) async fn get_surrealism_runtime(
&self,
lookup: SurrealismCacheLookup<'_>,
) -> Result<Arc<Runtime>> {
Ok(self.get_surrealism_module(lookup).await?.runtime)
}
}
#[cfg(test)]
mod tests {
#[cfg(feature = "http")]
use std::str::FromStr;
use std::time::Duration;
#[cfg(feature = "http")]
use url::Url;
use crate::cnf::CommonConfig;
#[cfg(all(feature = "allocation-tracking", feature = "allocator"))]
use crate::cnf::MEMORY_THRESHOLD;
use crate::ctx::Context;
use crate::ctx::reason::Reason;
#[cfg(feature = "http")]
use crate::dbs::Capabilities;
use crate::dbs::Options;
#[cfg(feature = "http")]
use crate::dbs::capabilities::{NetTarget, Targets};
use crate::expr::Base;
use crate::iam::{Action, Auth, ResourceKind, Role};
#[test]
fn is_allowed_respects_context_auth_toggle_and_base() {
let config = CommonConfig::default();
{
let mut ctx = Context::new_test();
ctx.auth_enabled = false;
let empty = Options::new(&config);
ctx.is_allowed(&empty, Action::View, ResourceKind::Any, Base::Ns).unwrap_err();
ctx.is_allowed(&empty, Action::View, ResourceKind::Any, Base::Db).unwrap_err();
let db_only = Options::new(&config).with_db(Some("db".into()));
ctx.is_allowed(&db_only, Action::View, ResourceKind::Any, Base::Db).unwrap_err();
ctx.is_allowed(&empty, Action::View, ResourceKind::Any, Base::Root).unwrap();
let ns = Options::new(&config).with_ns(Some("ns".into()));
ctx.is_allowed(&ns, Action::View, ResourceKind::Any, Base::Ns).unwrap();
let ns_db = Options::new(&config).with_ns(Some("ns".into())).with_db(Some("db".into()));
ctx.is_allowed(&ns_db, Action::View, ResourceKind::Any, Base::Db).unwrap();
}
{
let mut ctx = Context::new_test();
ctx.auth_enabled = true;
let opts = Options::new(&config).with_auth(Auth::for_root(Role::Owner).into());
ctx.is_allowed(&opts, Action::View, ResourceKind::Any, Base::Ns).unwrap_err();
ctx.is_allowed(&opts, Action::View, ResourceKind::Any, Base::Db).unwrap_err();
let db_only = opts.clone().with_db(Some("db".into()));
ctx.is_allowed(&db_only, Action::View, ResourceKind::Any, Base::Db).unwrap_err();
ctx.is_allowed(&opts, Action::View, ResourceKind::Any, Base::Root).unwrap();
let ns = opts.with_ns(Some("ns".into()));
ctx.is_allowed(&ns, Action::View, ResourceKind::Any, Base::Ns).unwrap();
let ns_db = ns.with_db(Some("db".into()));
ctx.is_allowed(&ns_db, Action::View, ResourceKind::Any, Base::Db).unwrap();
}
}
#[cfg(feature = "http")]
#[tokio::test]
async fn test_context_check_allowed_net() {
let cap = Capabilities::all().without_network_targets(Targets::Some(
[NetTarget::from_str("127.0.0.1").unwrap()].into(),
));
let mut ctx = Context::new_test();
ctx.capabilities = cap.into();
let ctx = ctx.freeze();
let r = ctx.check_allowed_net(&Url::parse("http://localhost").unwrap()).await;
assert_eq!(
r.err().unwrap().to_string(),
"Access to network target '127.0.0.1/32' is not allowed"
);
}
#[tokio::test]
async fn test_context_cancellation_priority() {
let mut ctx = Context::new_test();
ctx.add_timeout(Duration::from_nanos(1)).unwrap();
tokio::time::sleep(Duration::from_millis(10)).await;
let canceller = ctx.add_cancel();
canceller.cancel();
let ctx = ctx.freeze();
let result = ctx.done(true);
assert!(result.is_ok());
assert_eq!(result.unwrap(), Some(Reason::Canceled));
}
#[tokio::test]
async fn test_context_deadline_detection() {
let mut ctx = Context::new_test();
ctx.add_timeout(Duration::from_nanos(1)).unwrap();
tokio::time::sleep(Duration::from_millis(10)).await;
let ctx = ctx.freeze();
let result = ctx.done(true);
assert!(result.is_ok());
assert!(matches!(result.unwrap(), Some(Reason::Timedout(_))));
}
#[tokio::test]
async fn test_context_no_deadline() {
let ctx = Context::new_test();
let ctx = ctx.freeze();
let result = ctx.done(true);
assert!(result.is_ok());
assert_eq!(result.unwrap(), None);
}
#[tokio::test]
async fn test_context_is_done_adaptive_backoff() {
let ctx = Context::new_test();
let ctx = ctx.freeze();
for count in [1, 2, 4, 8, 16, 32] {
let result = ctx.is_done(Some(count)).await;
assert!(result.is_ok());
assert_eq!(result.unwrap(), false, "Count {} should not be done", count);
}
for count in 33..64 {
let result = ctx.is_done(Some(count)).await;
assert!(result.is_ok());
assert_eq!(result.unwrap(), false, "Count {} should not be done", count);
}
let result = ctx.is_done(Some(64)).await;
assert!(result.is_ok());
assert_eq!(result.unwrap(), false);
}
#[tokio::test]
async fn test_context_is_done_with_none() {
let ctx = Context::new_test();
let ctx = ctx.freeze();
let result = ctx.is_done(None).await;
assert!(result.is_ok());
assert_eq!(result.unwrap(), false);
}
#[tokio::test]
async fn test_context_is_done_detects_cancellation() {
let mut ctx = Context::new_test();
let canceller = ctx.add_cancel();
canceller.cancel();
let ctx = ctx.freeze();
let result = ctx.is_done(None).await;
assert!(result.is_ok());
assert_eq!(result.unwrap(), true);
let result = ctx.is_done(Some(1)).await;
assert!(result.is_ok());
assert_eq!(result.unwrap(), true);
}
#[tokio::test]
async fn test_snapshot_preserves_external_cancellation() {
use crate::ctx::CancelHandle;
let handle = CancelHandle::new();
let mut ctx = Context::new_test();
ctx.set_cancellation(&handle);
let root = ctx.freeze();
let snap = Context::snapshot(&root).freeze();
assert_eq!(
snap.is_done(Some(1)).await.unwrap(),
false,
"snapshot reports done before cancel was tripped"
);
handle.trip();
assert_eq!(
snap.is_done(Some(1)).await.unwrap(),
true,
"snapshot did not observe external cancel after trip — legacy `is_done` on \
streaming-exec snapshot would silently miss WebSocket disconnect"
);
assert!(
snap.cancel_token().is_some(),
"snapshot lost the awaitable cancel token — bare-await sites reached via \
the snapshot (legacy SLEEP fallback, etc.) could not `select!` against cancel"
);
}
#[tokio::test]
async fn test_context_memory_threshold_priority_documentation() {
let ctx = Context::new_test();
let ctx = ctx.freeze();
let result = ctx.done(true);
assert!(result.is_ok());
assert_eq!(result.unwrap(), None);
}
#[tokio::test]
#[cfg(all(feature = "allocation-tracking", feature = "allocator"))]
#[serial_test::serial]
async fn test_context_memory_threshold_integration() {
use crate::err::Error;
use crate::str::ParseBytes;
unsafe {
std::env::set_var(
"SURREAL_MEMORY_THRESHOLD",
"1MB".parse_bytes::<u64>().unwrap().to_string(),
);
}
assert_eq!(*MEMORY_THRESHOLD, 1048576);
let _large_allocation: Vec<u8> = Vec::with_capacity(20 * 1024 * 1024);
tokio::time::sleep(Duration::from_millis(10)).await;
let ctx = Context::new_test();
let ctx = ctx.freeze();
let result = ctx.done(true);
match result {
Err(e) => {
match e.downcast_ref::<Error>() {
Some(Error::QueryBeyondMemoryThreshold) => {
println!("✓ Memory threshold violation detected as expected");
}
other => {
panic!("Expected QueryBeyondMemoryThreshold error, got: {:?}", other);
}
}
}
Ok(None) => {
println!(
"⚠ Memory threshold not enforced - MEMORY_THRESHOLD was already initialized"
);
println!(" This is expected when running as part of the full test suite.");
println!(
" To properly test memory threshold enforcement, run this test in isolation:"
);
println!(
" cargo test --package surrealdb-core --features allocation-tracking,allocator test_context_memory_threshold_integration"
);
panic!("MEMORY_THRESHOLD was already initialized")
}
Ok(Some(reason)) => {
panic!("Unexpected reason returned: {:?}", reason);
}
}
unsafe {
std::env::remove_var("SURREAL_MEMORY_THRESHOLD");
}
}
}