use std::fs;
use std::path::PathBuf;
use std::sync::Arc;
use std::time::Duration;
use anyhow::Result;
use clap::Args;
use rand::Rng;
use surrealdb::opt::capabilities::Capabilities as SdkCapabilities;
use surrealdb_cnf::ConfigMap;
use surrealdb_core::channel::Receiver;
use surrealdb_core::kvs::{Datastore, TransactionBuilderFactory};
use surrealdb_core::options::EngineOptions;
use surrealdb_observe::ExecutionObserver;
use surrealdb_rpc::capabilities::{
ArbitraryQueryTarget, Capabilities, EvalQueryTarget, ExperimentalTarget, FuncTarget,
MethodTarget, NetTarget, NewPlannerStrategy, RouteTarget, Targets,
};
use surrealdb_types::Notification;
use tokio::time::{Instant, sleep, timeout};
use tokio_util::sync::CancellationToken;
use crate::cli::Config;
use crate::core::dbs::Session;
const TARGET: &str = "surreal::dbs";
const STARTUP_RETRY_LOG_INTERVAL: Duration = Duration::from_secs(5);
#[derive(Args, Debug)]
pub struct StartCommandDbsOptions {
#[arg(help = "Whether strict mode is enabled on this database instance")]
#[arg(env = "SURREAL_STRICT", short = 's', long = "strict", hide = true)]
strict_mode: Option<bool>,
#[arg(help = "The maximum duration that a set of statements can run for")]
#[arg(env = "SURREAL_QUERY_TIMEOUT", long)]
#[arg(value_parser = super::cli::validator::duration)]
query_timeout: Option<Duration>,
#[arg(help = "The maximum duration that any single transaction can run for")]
#[arg(env = "SURREAL_TRANSACTION_TIMEOUT", long)]
#[arg(value_parser = super::cli::validator::duration)]
transaction_timeout: Option<Duration>,
#[arg(
help = "The maximum duration that each built-in startup datastore operation can retry for"
)]
#[arg(env = "SURREAL_STARTUP_OPERATION_TIMEOUT", long, default_value = "60s")]
#[arg(value_parser = super::cli::validator::duration)]
startup_operation_timeout: Duration,
#[arg(help = "Whether to allow unauthenticated access", help_heading = "Authentication")]
#[arg(env = "SURREAL_UNAUTHENTICATED", long = "unauthenticated")]
#[arg(default_value_t = false)]
unauthenticated: bool,
#[command(flatten)]
#[command(next_help_heading = "Capabilities")]
capabilities: DbsCapabilities,
#[arg(help = "Sets the directory for storing temporary database files")]
#[arg(env = "SURREAL_TEMPORARY_DIRECTORY", long = "temporary-directory")]
#[arg(value_parser = super::cli::validator::dir_exists)]
temporary_directory: Option<PathBuf>,
#[arg(help = "Path to a SurrealQL file that will be imported when starting the server")]
#[arg(env = "SURREAL_IMPORT_FILE", long = "import-file")]
#[arg(value_parser = super::cli::validator::file_exists)]
import_file: Option<PathBuf>,
#[arg(help = "The minimum execution time in milliseconds to trigger slow query logging")]
#[arg(env = "SURREAL_SLOW_QUERY_LOG_THRESHOLD", long = "slow-log-threshold")]
#[arg(value_parser = super::cli::validator::duration)]
slow_log_threshold: Option<Duration>,
#[arg(help = "A comma-separated list of parameter names to include in slow query logs")]
#[arg(env = "SURREAL_SLOW_QUERY_LOG_PARAM_ALLOW", long = "slow-log-param-allow")]
#[arg(value_delimiter = ',', num_args = 1..)]
slow_log_param_allow: Vec<String>,
#[arg(help = "A comma-separated list of parameter names to omit from slow query logs")]
#[arg(env = "SURREAL_SLOW_QUERY_LOG_PARAM_DENY", long = "slow-log-param-deny")]
#[arg(value_delimiter = ',', num_args = 1..)]
slow_log_param_deny: Vec<String>,
#[arg(help = "The default namespace for a new instance")]
#[arg(env = "SURREAL_DEFAULT_NAMESPACE", long = "default-namespace")]
default_namespace: Option<String>,
#[arg(help = "The default database for a new instance")]
#[arg(env = "SURREAL_DEFAULT_DATABASE", long = "default-database")]
default_database: Option<String>,
#[arg(help = "Whether to disable default namespace and database creation")]
#[arg(env = "SURREAL_NO_DEFAULTS", long = "no-defaults", conflicts_with_all = ["default_namespace", "default_database"])]
#[arg(default_value_t = false)]
no_defaults: bool,
#[arg(help = "Load Surrealism modules lazily on first use instead of eagerly at startup")]
#[arg(env = "SURREAL_LAZY_SURREALISM", long = "lazy-surrealism")]
#[arg(default_value_t = false)]
lazy_surrealism: bool,
}
#[derive(Args, Debug)]
pub struct DbsCapabilities {
#[arg(help = "Allow all capabilities except for those more specifically denied")]
#[arg(env = "SURREAL_CAPS_ALLOW_ALL", short = 'A', long, conflicts_with = "deny_all")]
allow_all: bool,
#[cfg(feature = "scripting")]
#[arg(help = "Allow execution of embedded scripting functions")]
#[arg(env = "SURREAL_CAPS_ALLOW_SCRIPT", long, conflicts_with_all = ["allow_all", "deny_scripting"])]
allow_scripting: bool,
#[arg(help = "Allow guest users to execute queries")]
#[arg(env = "SURREAL_CAPS_ALLOW_GUESTS", long, conflicts_with_all = ["allow_all", "deny_guests"])]
allow_guests: bool,
#[arg(
help = "Allow execution of all functions except for functions that are specifically denied. Alternatively, you can provide a comma-separated list of function names to allow",
long_help = r#"Allow execution of all functions except for functions that are specifically denied. Alternatively, you can provide a comma-separated list of function names to allow
Specifically denied functions and function families prevail over any other allowed function execution.
Function names must be in the form <family>[::<name>]. For example:
- 'http' or 'http::*' -> Include all functions in the 'http' family
- 'http::get' -> Include only the 'get' function in the 'http' family
"#
)]
#[arg(env = "SURREAL_CAPS_ALLOW_FUNC", long)]
#[arg(default_missing_value_os = "", num_args = 0..)]
#[arg(value_parser = super::cli::validator::func_targets)]
allow_funcs: Option<Targets<FuncTarget>>,
#[arg(hide = true)]
#[arg(env = "SURREAL_CAPS_ALLOW_EXPERIMENTAL", long)]
#[arg(default_missing_value_os = "", num_args = 0..)]
#[arg(value_parser = super::cli::validator::experimental_targets)]
allow_experimental: Option<Targets<ExperimentalTarget>>,
#[arg(
help = "Allow execution of arbitrary queries by certain user groups except when specifically denied.",
long_help = r#"Allow execution of arbitrary queries by certain user groups except when specifically denied. Alternatively, you can provide a comma-separated list of user groups to allow
Specifically denied user groups prevail over any other allowed user group.
User groups must be one of "guest", "record" or "system".
"#
)]
#[arg(env = "SURREAL_CAPS_ALLOW_ARBITRARY_QUERY", long)]
#[arg(default_missing_value_os = "", num_args = 0..)]
#[arg(value_parser = super::cli::validator::query_arbitrary_targets)]
allow_arbitrary_query: Option<Targets<ArbitraryQueryTarget>>,
#[arg(
help = "Allow the eval::* functions to be invoked by certain user groups. Denied for everyone by default.",
long_help = r#"Allow the eval::surql / eval::gql functions to be invoked by certain user groups. Provide a comma-separated list of user groups to allow. This is an additive gate on top of arbitrary-query: an eval invocation must also satisfy the arbitrary-query capability.
Specifically denied user groups prevail over any other allowed user group.
User groups must be one of "guest", "record" or "system".
Denied for everyone by default (even under --allow-all); must be explicitly enabled.
"#
)]
#[arg(env = "SURREAL_CAPS_ALLOW_EVAL_QUERY", long)]
#[arg(default_missing_value_os = "", num_args = 0..)]
#[arg(value_parser = super::cli::validator::eval_query_targets)]
allow_eval_query: Option<Targets<EvalQueryTarget>>,
#[arg(
help = "Allow all outbound network connections except for network targets that are specifically denied. Alternatively, you can provide a comma-separated list of network targets to allow",
long_help = r#"Allow all outbound network connections except for network targets that are specifically denied. Alternatively, you can provide a comma-separated list of network targets to allow
Specifically denied network targets prevail over any other allowed outbound network connections.
Targets must be in the form of <host>[:<port>], <ipv4|ipv6>[/<mask>]. For example:
- 'surrealdb.com', '127.0.0.1' or 'fd00::1' -> Match outbound connections to these hosts on any port
- 'surrealdb.com:80', '127.0.0.1:80' or 'fd00::1:80' -> Match outbound connections to these hosts on port 80
- '10.0.0.0/8' or 'fd00::/8' -> Match outbound connections to any host in these networks
"#
)]
#[arg(env = "SURREAL_CAPS_ALLOW_NET", long)]
#[arg(default_missing_value_os = "", num_args = 0..)]
#[arg(value_parser = super::cli::validator::net_targets)]
allow_net: Option<Targets<NetTarget>>,
#[arg(
help = "Allow all RPC methods to be called except for routes that are specifically denied. Alternatively, you can provide a comma-separated list of RPC methods to allow."
)]
#[arg(env = "SURREAL_CAPS_ALLOW_RPC", long)]
#[arg(default_missing_value_os = "", num_args = 0..)]
#[arg(default_value_os = "")] #[arg(value_parser = super::cli::validator::method_targets)]
allow_rpc: Option<Targets<MethodTarget>>,
#[arg(
help = "Allow all HTTP routes to be requested except for routes that are specifically denied. Alternatively, you can provide a comma-separated list of HTTP routes to allow."
)]
#[arg(env = "SURREAL_CAPS_ALLOW_HTTP", long)]
#[arg(default_missing_value_os = "", num_args = 0..)]
#[arg(default_value_os = "")] #[arg(value_parser = super::cli::validator::route_targets)]
allow_http: Option<Targets<RouteTarget>>,
#[arg(help = "Deny all capabilities except for those more specifically allowed")]
#[arg(env = "SURREAL_CAPS_DENY_ALL", short = 'D', long, conflicts_with = "allow_all")]
deny_all: bool,
#[cfg(feature = "scripting")]
#[arg(help = "Deny execution of embedded scripting functions")]
#[arg(env = "SURREAL_CAPS_DENY_SCRIPT", long, conflicts_with_all = ["deny_all", "allow_scripting"])]
deny_scripting: bool,
#[arg(help = "Deny guest users to execute queries")]
#[arg(env = "SURREAL_CAPS_DENY_GUESTS", long, conflicts_with_all = ["deny_all", "allow_guests"])]
deny_guests: bool,
#[arg(
help = "Deny execution of all functions except for functions that are specifically allowed. Alternatively, you can provide a comma-separated list of function names to deny",
long_help = r#"Deny execution of all functions except for functions that are specifically allowed. Alternatively, you can provide a comma-separated list of function names to deny.
Specifically allowed functions and function families prevail over a general denial of function execution.
Function names must be in the form <family>[::<name>]. For example:
- 'http' or 'http::*' -> Include all functions in the 'http' family
- 'http::get' -> Include only the 'get' function in the 'http' family
"#
)]
#[arg(env = "SURREAL_CAPS_DENY_FUNC", long)]
#[arg(default_missing_value_os = "", num_args = 0..)]
#[arg(value_parser = super::cli::validator::func_targets)]
deny_funcs: Option<Targets<FuncTarget>>,
#[arg(hide = true)]
#[arg(env = "SURREAL_CAPS_DENY_EXPERIMENTAL", long)]
#[arg(default_missing_value_os = "", num_args = 0..)]
#[arg(value_parser = super::cli::validator::experimental_targets)]
deny_experimental: Option<Targets<ExperimentalTarget>>,
#[arg(
help = "Deny execution of arbitrary queries by certain user groups except when specifically allowed.",
long_help = r#"Deny execution of arbitrary queries by certain user groups except when specifically allowed. Alternatively, you can provide a comma-separated list of user groups to deny
Specifically allowed user groups prevail over a general denial of user group.
User groups must be one of "guest", "record" or "system".
"#
)]
#[arg(env = "SURREAL_CAPS_DENY_ARBITRARY_QUERY", long)]
#[arg(default_missing_value_os = "", num_args = 0..)]
#[arg(value_parser = super::cli::validator::query_arbitrary_targets)]
deny_arbitrary_query: Option<Targets<ArbitraryQueryTarget>>,
#[arg(
help = "Deny the eval::* functions for certain user groups except when specifically allowed.",
long_help = r#"Deny the eval::surql / eval::gql functions for certain user groups. Provide a comma-separated list of user groups to deny.
Specifically denied user groups prevail over any allowed user group.
User groups must be one of "guest", "record" or "system".
"#
)]
#[arg(env = "SURREAL_CAPS_DENY_EVAL_QUERY", long)]
#[arg(default_missing_value_os = "", num_args = 0..)]
#[arg(value_parser = super::cli::validator::eval_query_targets)]
deny_eval_query: Option<Targets<EvalQueryTarget>>,
#[arg(
help = "Deny all outbound network connections except for network targets that are specifically allowed. Alternatively, you can provide a comma-separated list of network targets to deny",
long_help = r#"Deny all outbound network connections except for network targets that are specifically allowed. Alternatively, you can provide a comma-separated list of network targets to deny.
Specifically allowed network targets prevail over a general denial of outbound network connections.
Targets must be in the form of <host>[:<port>], <ipv4|ipv6>[/<mask>]. For example:
- 'surrealdb.com', '127.0.0.1' or 'fd00::1' -> Match outbound connections to these hosts on any port
- 'surrealdb.com:80', '127.0.0.1:80' or 'fd00::1:80' -> Match outbound connections to these hosts on port 80
- '10.0.0.0/8' or 'fd00::/8' -> Match outbound connections to any host in these networks
"#
)]
#[arg(env = "SURREAL_CAPS_DENY_NET", long)]
#[arg(default_missing_value_os = "", num_args = 0..)]
#[arg(value_parser = super::cli::validator::net_targets)]
deny_net: Option<Targets<NetTarget>>,
#[arg(
help = "Deny all RPC methods from being called except for methods that are specifically allowed. Alternatively, you can provide a comma-separated list of RPC methods to deny."
)]
#[arg(env = "SURREAL_CAPS_DENY_RPC", long)]
#[arg(default_missing_value_os = "", num_args = 0..)]
#[arg(value_parser = super::cli::validator::method_targets)]
deny_rpc: Option<Targets<MethodTarget>>,
#[arg(
help = "Deny all HTTP routes from being requested except for routes that are specifically allowed. Alternatively, you can provide a comma-separated list of HTTP routes to deny."
)]
#[arg(env = "SURREAL_CAPS_DENY_HTTP", long)]
#[arg(default_missing_value_os = "", num_args = 0..)]
#[arg(value_parser = super::cli::validator::route_targets)]
deny_http: Option<Targets<RouteTarget>>,
#[arg(
help = "Strategy for the streaming query planner: 'best-effort' (default), 'compute-only', or 'all-read-only'"
)]
#[arg(env = "SURREAL_PLANNER_STRATEGY", long = "planner-strategy")]
#[arg(default_value = "best-effort")]
planner_strategy: NewPlannerStrategy,
}
impl DbsCapabilities {
#[cfg(feature = "scripting")]
fn get_scripting(&self) -> bool {
self.allow_scripting || (self.allow_all && !self.deny_scripting)
}
#[cfg(not(feature = "scripting"))]
fn get_scripting(&self) -> bool {
false
}
fn get_allow_guests(&self) -> bool {
self.allow_guests || (self.allow_all && !self.deny_guests)
}
fn get_allow_funcs(&self) -> Targets<FuncTarget> {
if self.deny_all {
if let Some(targets) = &self.allow_funcs {
match targets {
Targets::None => {}
Targets::Some(_) => {
return targets.clone();
}
Targets::All => {
return Targets::All;
}
}
}
return Targets::None;
}
if let Some(Targets::All) = self.deny_funcs {
if let Some(targets) = &self.allow_funcs
&& let Targets::Some(_) = targets
{
return targets.clone();
}
return Targets::None;
}
if self.allow_all {
return Targets::All;
}
self.allow_funcs.clone().unwrap_or(Targets::All) }
fn get_allow_net(&self) -> Targets<NetTarget> {
if self.deny_all {
if let Some(targets) = &self.allow_net {
match targets {
Targets::None => {}
Targets::Some(_) => {
return targets.clone();
}
Targets::All => {
return Targets::All;
}
}
}
return Targets::None;
}
if let Some(Targets::All) = self.deny_net {
if let Some(targets) = &self.allow_net
&& let Targets::Some(_) = targets
{
return targets.clone();
}
return Targets::None;
}
if self.allow_all {
return Targets::All;
}
self.allow_net.clone().unwrap_or(Targets::None)
}
fn get_allow_rpc(&self) -> Targets<MethodTarget> {
if self.deny_all {
if let Some(targets) = &self.allow_rpc {
match targets {
Targets::None => {}
Targets::Some(_) => {
return targets.clone();
}
Targets::All => {
return Targets::All;
}
}
}
return Targets::None;
}
if let Some(Targets::All) = self.deny_rpc {
if let Some(targets) = self.allow_rpc.as_ref()
&& let Targets::Some(_) = targets
{
return targets.clone();
}
return Targets::None;
}
if self.allow_all {
return Targets::All;
}
self.allow_rpc.clone().unwrap_or(Targets::All) }
fn get_allow_http(&self) -> Targets<RouteTarget> {
if self.deny_all {
if let Some(targets) = self.allow_http.as_ref() {
match targets {
Targets::None => {}
Targets::Some(_) => {
return targets.clone();
}
Targets::All => {
return Targets::All;
}
}
}
return Targets::None;
}
if let Some(Targets::All) = self.deny_http {
if let Some(targets) = self.allow_http.as_ref()
&& let Targets::Some(_) = targets
{
return targets.clone();
}
return Targets::None;
}
if self.allow_all {
return Targets::All;
}
self.allow_http.clone().unwrap_or(Targets::All) }
fn get_allow_experimental(&self) -> Targets<ExperimentalTarget> {
if self.deny_all {
return self.allow_experimental.clone().unwrap_or(Targets::None);
}
if let Some(Targets::All) = self.deny_experimental {
match &self.allow_experimental {
Some(t @ Targets::Some(_)) => return t.clone(),
_ => return Targets::None,
}
}
self.allow_experimental.clone().unwrap_or(Targets::None) }
fn get_allow_arbitrary_query(&self) -> Targets<ArbitraryQueryTarget> {
if let Some(Targets::All) = self.deny_arbitrary_query {
match &self.allow_arbitrary_query {
Some(t @ Targets::Some(_)) => return t.clone(),
_ => return Targets::None,
}
}
if self.allow_all {
return Targets::All;
}
self.allow_arbitrary_query.clone().unwrap_or(Targets::All) }
fn get_allow_eval_query(&self) -> Targets<EvalQueryTarget> {
if let Some(Targets::All) = self.deny_eval_query {
match &self.allow_eval_query {
Some(t @ Targets::Some(_)) => return t.clone(),
_ => return Targets::None,
}
}
self.allow_eval_query.clone().unwrap_or(Targets::None)
}
fn get_deny_funcs(&self) -> Targets<FuncTarget> {
if let Some(targets) = &self.deny_funcs
&& let Targets::Some(_) = targets
{
return targets.clone();
}
Targets::None
}
fn get_deny_net(&self) -> Targets<NetTarget> {
if let Some(targets) = &self.deny_net
&& let Targets::Some(_) = targets
{
return targets.clone();
}
Targets::None
}
fn get_deny_all(&self) -> bool {
self.deny_all
}
fn get_deny_rpc(&self) -> Targets<MethodTarget> {
if let Some(targets) = &self.deny_rpc
&& let Targets::Some(_) = targets
{
return targets.clone();
}
Targets::None
}
fn get_deny_http(&self) -> Targets<RouteTarget> {
if let Some(targets) = self.deny_http.as_ref()
&& let Targets::Some(_) = targets
{
return targets.clone();
}
Targets::None
}
fn get_deny_experimental(&self) -> Targets<ExperimentalTarget> {
if let Some(t @ Targets::Some(_)) = &self.deny_experimental {
t.clone()
} else {
Targets::None
}
}
fn get_deny_arbitrary_query(&self) -> Targets<ArbitraryQueryTarget> {
if let Some(t @ Targets::Some(_)) = &self.deny_arbitrary_query {
t.clone()
} else {
Targets::None
}
}
fn get_deny_eval_query(&self) -> Targets<EvalQueryTarget> {
if let Some(t @ Targets::Some(_)) = &self.deny_eval_query {
t.clone()
} else {
Targets::None
}
}
pub fn into_cli_capabilities(self) -> Capabilities {
merge_capabilities(SdkCapabilities::all().into(), &self)
}
}
fn merge_capabilities(initial: Capabilities, caps: &DbsCapabilities) -> Capabilities {
initial
.with_scripting(caps.get_scripting())
.with_guest_access(caps.get_allow_guests())
.with_functions(caps.get_allow_funcs())
.without_functions(caps.get_deny_funcs())
.with_network_targets(caps.get_allow_net())
.without_network_targets(caps.get_deny_net())
.with_rpc_methods(caps.get_allow_rpc())
.without_rpc_methods(caps.get_deny_rpc())
.with_http_routes(caps.get_allow_http())
.without_http_routes(caps.get_deny_http())
.with_experimental(caps.get_allow_experimental())
.without_experimental(caps.get_deny_experimental())
.with_arbitrary_query(caps.get_allow_arbitrary_query())
.without_arbitrary_query(caps.get_deny_arbitrary_query())
.with_eval_query(caps.get_allow_eval_query())
.without_eval_query(caps.get_deny_eval_query())
.with_planner_strategy(caps.planner_strategy)
}
impl From<DbsCapabilities> for Capabilities {
fn from(caps: DbsCapabilities) -> Self {
merge_capabilities(Default::default(), &caps)
}
}
async fn retry_with_timeout<F, Fut, T, E>(
operation_name: &str,
timeout_duration: Duration,
f: F,
) -> Result<T, anyhow::Error>
where
F: Fn() -> Fut,
Fut: Future<Output = Result<T, E>>,
E: std::fmt::Display + std::fmt::Debug,
{
retry_with_timeout_check(operation_name, timeout_duration, f, |_| false).await
}
async fn retry_with_timeout_check<F, Fut, T, E, P>(
operation_name: &str,
timeout_duration: Duration,
f: F,
is_permanent: P,
) -> Result<T, anyhow::Error>
where
F: Fn() -> Fut,
Fut: Future<Output = Result<T, E>>,
E: std::fmt::Display + std::fmt::Debug,
P: Fn(&E) -> bool,
{
let start = Instant::now();
let mut attempt = 0;
let mut last_reported = start;
loop {
attempt += 1;
match timeout(timeout_duration.saturating_sub(start.elapsed()), f()).await {
Ok(Ok(result)) => {
if attempt > 1 {
info!(target: TARGET, operation = operation_name, attempts = attempt, "Operation succeeded after retry");
} else {
info!(target: TARGET, operation = operation_name, attempts = attempt, "Operation succeeded");
}
return Ok(result);
}
Ok(Err(e)) => {
if is_permanent(&e) {
error!(
target: TARGET,
operation = operation_name,
error = %e,
"Operation failed with a permanent error, not retrying"
);
return Err(anyhow::anyhow!("{e}"));
}
let elapsed = start.elapsed();
if elapsed >= timeout_duration {
return Err(anyhow::anyhow!(
"Operation '{}' failed after {} attempts over {:?}: {}",
operation_name,
attempt,
elapsed,
e
));
}
if attempt == 1 || last_reported.elapsed() >= STARTUP_RETRY_LOG_INTERVAL {
last_reported = Instant::now();
info!(
target: TARGET,
operation = operation_name,
attempt = attempt,
elapsed = ?elapsed,
budget = ?timeout_duration,
error = %e,
"Operation failed, still retrying within the startup budget"
);
} else {
debug!(
target: TARGET,
operation = operation_name,
attempt = attempt,
error = %e,
"Operation failed, retrying..."
);
}
let base_backoff = Duration::from_millis(100 * 2u64.pow((attempt - 1).min(5)));
let base_backoff = base_backoff.min(Duration::from_secs(5));
let jitter = rand::rng().random_range(0.5..=1.5);
let backoff = base_backoff.mul_f64(jitter);
sleep(backoff).await;
}
Err(_) => {
return Err(anyhow::anyhow!(
"Operation '{}' timed out after {} attempts over {:?}",
operation_name,
attempt,
start.elapsed()
));
}
}
}
}
fn is_permanent_storage_error(e: &anyhow::Error) -> bool {
let message = e.to_string();
message.contains("out-of-date")
|| message.contains("cannot be read by this one")
|| message.contains("invalid storage version")
}
#[instrument(level = "trace", target = "surreal::dbs", skip_all)]
#[allow(clippy::type_complexity)]
pub async fn init<C: TransactionBuilderFactory>(
composer: C,
opt: &Config,
canceller: CancellationToken,
observer: Arc<dyn ExecutionObserver>,
#[cfg_attr(not(storage), allow(unused_variables))] StartCommandDbsOptions {
strict_mode,
query_timeout,
transaction_timeout,
startup_operation_timeout,
unauthenticated,
capabilities,
temporary_directory,
import_file,
slow_log_threshold,
slow_log_param_allow,
slow_log_param_deny,
default_namespace,
default_database,
no_defaults,
#[cfg_attr(not(feature = "surrealism"), allow(unused_variables))]
lazy_surrealism,
}: StartCommandDbsOptions,
) -> Result<(Arc<Datastore>, Receiver<Notification>, C::RouterState, PendingStartup)> {
if let Some(true) = strict_mode {
warn!(
"Strict mode is no longer defined on the server level. Use `DEFINE DATABASE <db> STRICT` instead. Ignoring strict mode flag."
);
}
if let Some(v) = query_timeout {
debug!("Maximum query processing timeout is {v:?}");
}
if let Some(v) = transaction_timeout {
debug!("Maximum transaction processing timeout is {v:?}");
}
debug!("Startup operation timeout is {startup_operation_timeout:?}");
if unauthenticated {
warn!(
"❌🔒 IMPORTANT: Authentication is disabled. This is not recommended for production use. 🔒❌"
);
}
if capabilities.get_deny_all() {
warn!(
"You are denying all capabilities by default. Although this is recommended, beware that any new capabilities will also be denied."
);
}
if let Some(v) = slow_log_threshold {
debug!("Slow log threshold is {v:?}");
}
if !slow_log_param_allow.is_empty() {
debug!("Slow log param allow is {:?}", slow_log_param_allow);
}
if !slow_log_param_deny.is_empty() {
debug!("Slow log param deny is {:?}", slow_log_param_deny);
}
let capabilities = capabilities.into();
debug!("Server capabilities: {capabilities}");
let (send, recv) = surrealdb_core::channel::bounded(surrealdb_cnf::NOTIFICATIONS_CHANNEL_SIZE);
let config = ConfigMap::from_env();
let builder = Datastore::builder()
.with_engine_options(opt.engine)
.without_maintenance_tasks()
.with_config(config)
.with_runtime_worker_threads(*crate::cnf::RUNTIME_WORKER_THREADS)
.with_query_timeout(query_timeout)
.with_transaction_timeout(transaction_timeout)
.with_auth(!unauthenticated)
.with_capabilities(capabilities)
.with_notify(send)
.with_shutdown_cancel(canceller.clone())
.with_observer(observer);
#[cfg(storage)]
let builder = builder.with_temporary_directory(temporary_directory);
let builder = if let Some(slow_log_threshold) = slow_log_threshold {
builder.with_slow_log(slow_log_threshold, slow_log_param_allow, slow_log_param_deny)
} else {
builder
};
#[cfg(feature = "surrealism")]
let builder = builder.with_lazy_surrealism(lazy_surrealism);
let (dbs, router_state) =
builder.build_with_factory_path_and_router_state::<C>(&opt.path, composer).await?;
let pending_startup = PendingStartup {
import_file,
credentials: opt.user.clone().zip(opt.pass.clone()),
default_namespace,
default_database,
no_defaults,
engine: opt.engine,
canceller,
timeout: startup_operation_timeout,
};
Ok((dbs, recv, router_state, pending_startup))
}
fn spawn_first_node_maintenance_pass(
dbs: &Arc<Datastore>,
engine: &EngineOptions,
canceller: CancellationToken,
) {
let expire = !engine.node_membership_check_interval.is_zero();
let cleanup = !engine.node_membership_cleanup_interval.is_zero();
if !expire && !cleanup {
return;
}
fn report(canceller: &CancellationToken, step: &str, res: Result<()>) {
if let Err(err) = res
&& !canceller.is_cancelled()
{
warn!(target: TARGET, "Error {step}: {err}");
}
}
let dbs = Arc::clone(dbs);
tokio::spawn(async move {
if expire && !canceller.is_cancelled() {
report(&canceller, "expiring inactive nodes", dbs.expire_nodes().await);
}
if cleanup && !canceller.is_cancelled() {
report(&canceller, "removing archived nodes", dbs.remove_nodes().await);
}
});
}
pub(crate) struct PendingStartup {
import_file: Option<PathBuf>,
credentials: Option<(String, String)>,
default_namespace: Option<String>,
default_database: Option<String>,
no_defaults: bool,
engine: EngineOptions,
canceller: CancellationToken,
timeout: Duration,
}
impl PendingStartup {
pub(crate) fn has_import(&self) -> bool {
self.import_file.is_some()
}
}
pub(crate) async fn initialise_datastore(
dbs: &Arc<Datastore>,
pending: &PendingStartup,
) -> Result<()> {
dbs.wait_until_serve_ready().await?;
let (_, is_new) = retry_with_timeout_check(
"check_version",
pending.timeout,
|| async { dbs.check_version().await },
is_permanent_storage_error,
)
.await?;
if is_new && !pending.no_defaults {
let ns = pending.default_namespace.as_deref().unwrap_or("main");
let db = pending.default_database.as_deref().unwrap_or("main");
retry_with_timeout("initialise_defaults", pending.timeout, || async {
dbs.initialise_defaults(ns, db).await
})
.await?;
}
retry_with_timeout("Insert node", pending.timeout, || async { dbs.insert_node().await })
.await?;
dbs.start_maintenance_tasks();
spawn_first_node_maintenance_pass(dbs, &pending.engine, pending.canceller.clone());
Ok(())
}
pub(crate) async fn finish_startup(ds: &Datastore, pending: &PendingStartup) -> Result<()> {
if let Some(file) = &pending.import_file {
info!(target: TARGET, file = ?file, "Importing data from file");
let sql = fs::read_to_string(file)?;
let results = retry_with_timeout("startup", pending.timeout, || async {
ds.startup(&sql, &Session::owner()).await
})
.await?;
let failed = results.iter().filter(|r| r.result.is_err()).count();
if failed > 0 {
error!(
target: TARGET,
file = ?file,
failed,
total = results.len(),
"Startup import did not apply in full; the database is partially restored"
);
for result in results.iter() {
if let Err(err) = &result.result {
error!(target: TARGET, error = %err, "Startup import statement failed");
}
}
}
}
if let Some((user, pass)) = &pending.credentials {
info!(target: TARGET, user = %user, "Initialising credentials");
retry_with_timeout("initialise_credentials", pending.timeout, || async {
ds.initialise_credentials(user, pass).await
})
.await?;
}
Ok(())
}
#[cfg(test)]
mod tests {
use std::ffi::OsString;
use std::str::FromStr;
use std::sync::Arc;
use std::sync::atomic::{AtomicUsize, Ordering};
use clap::Parser;
use serial_test::serial;
#[cfg(feature = "storage-rocksdb")]
use surrealdb_datastore::key::schema::{NodeKey, NodeLiveQueryKey};
use surrealdb_types::ToSql;
use test_log::test;
use wiremock::matchers::{method, path};
use wiremock::{Mock, MockServer, ResponseTemplate};
use super::*;
const STARTUP_OPERATION_TIMEOUT_ENV: &str = "SURREAL_STARTUP_OPERATION_TIMEOUT";
struct EnvGuard {
key: &'static str,
old: Option<OsString>,
}
impl EnvGuard {
fn set(key: &'static str, value: &str) -> Self {
let old = std::env::var_os(key);
unsafe {
std::env::set_var(key, value);
}
Self {
key,
old,
}
}
fn remove(key: &'static str) -> Self {
let old = std::env::var_os(key);
unsafe {
std::env::remove_var(key);
}
Self {
key,
old,
}
}
}
impl Drop for EnvGuard {
fn drop(&mut self) {
unsafe {
match &self.old {
Some(value) => std::env::set_var(self.key, value),
None => std::env::remove_var(self.key),
}
}
}
}
#[derive(Parser, Debug)]
struct TestCli {
#[command(flatten)]
dbs: StartCommandDbsOptions,
}
#[test]
#[serial]
fn startup_operation_timeout_defaults_to_sixty_seconds() {
let _guard = EnvGuard::remove(STARTUP_OPERATION_TIMEOUT_ENV);
let cli = TestCli::try_parse_from(["surrealdb"]).unwrap();
assert_eq!(cli.dbs.startup_operation_timeout, Duration::from_secs(60));
}
#[test]
fn startup_operation_timeout_can_be_set_from_cli() {
let cli =
TestCli::try_parse_from(["surrealdb", "--startup-operation-timeout", "10m"]).unwrap();
assert_eq!(cli.dbs.startup_operation_timeout, Duration::from_secs(10 * 60));
}
#[test]
#[serial]
fn startup_operation_timeout_can_be_set_from_env() {
let _guard = EnvGuard::set(STARTUP_OPERATION_TIMEOUT_ENV, "75s");
let cli = TestCli::try_parse_from(["surrealdb"]).unwrap();
assert_eq!(cli.dbs.startup_operation_timeout, Duration::from_secs(75));
}
#[test(tokio::test(start_paused = true))]
async fn startup_retry_times_out_attempt_without_retrying() {
let attempts = Arc::new(AtomicUsize::new(0));
let err = retry_with_timeout("test operation", Duration::from_millis(10), || {
let attempts = Arc::clone(&attempts);
async move {
attempts.fetch_add(1, Ordering::SeqCst);
sleep(Duration::from_millis(50)).await;
Ok::<_, &'static str>(())
}
})
.await
.unwrap_err();
assert!(err.to_string().contains("timed out after 1 attempts"));
assert_eq!(attempts.load(Ordering::SeqCst), 1);
}
#[test(tokio::test(start_paused = true))]
async fn startup_retry_returns_permanent_errors_immediately() {
let attempts = Arc::new(AtomicUsize::new(0));
let err = retry_with_timeout_check(
"test operation",
Duration::from_millis(100),
|| {
let attempts = Arc::clone(&attempts);
async move {
attempts.fetch_add(1, Ordering::SeqCst);
Err::<(), _>("permanent")
}
},
|e| *e == "permanent",
)
.await
.unwrap_err();
assert_eq!(err.to_string(), "permanent");
assert_eq!(attempts.load(Ordering::SeqCst), 1);
}
#[test]
fn every_storage_state_error_is_classified_permanent() {
use surrealdb_datastore::DatastoreError;
for error in [
DatastoreError::InvalidStorageVersion,
DatastoreError::OutdatedStorageVersion {
expected: 3,
actual: 1,
},
DatastoreError::MigratedBeyondStorageVersion {
stored: "3.4.0".to_string(),
running: "3.3.0".to_string(),
migrations: "2".to_string(),
},
] {
let rendered = error.to_string();
assert!(
is_permanent_storage_error(&anyhow::Error::new(error)),
"a permanent storage-state error would be retried: {rendered}"
);
}
}
#[test]
fn a_pending_migration_is_not_classified_permanent() {
use surrealdb_datastore::DatastoreError;
assert!(!is_permanent_storage_error(&anyhow::Error::new(
DatastoreError::MigrationTimedOut {
migrations: "2".to_string(),
}
)));
}
#[test(tokio::test(flavor = "multi_thread"))]
async fn test_capabilities() {
let server1 = {
let s = MockServer::start().await;
let get = Mock::given(method("GET"))
.and(path("/"))
.respond_with(ResponseTemplate::new(200).set_body_string("SUCCESS"))
.expect(1);
let get2 = Mock::given(method("GET"))
.and(path("/test"))
.respond_with(ResponseTemplate::new(200).set_body_string("SUCCESS"))
.expect(1);
s.register(get).await;
s.register(get2).await;
s
};
let server2 = {
let s = MockServer::start().await;
let get = Mock::given(method("GET"))
.respond_with(ResponseTemplate::new(200).set_body_string("SUCCESS"))
.expect(1);
let head =
Mock::given(method("HEAD")).respond_with(ResponseTemplate::new(200)).expect(0);
s.register(get).await;
s.register(head).await;
s
};
let server3 = {
let s = MockServer::start().await;
let redirect_res = ResponseTemplate::new(301).append_header("Location", server1.uri());
let redirect = Mock::given(method("GET"))
.and(path("redirect"))
.respond_with(redirect_res)
.expect(1);
s.register(redirect).await;
s
};
let cases = vec![
(
Datastore::builder()
.with_capabilities(
Capabilities::default()
.with_functions(Targets::<FuncTarget>::All)
.with_network_targets(Targets::<NetTarget>::All),
)
.build_with_path("memory")
.await
.unwrap(),
Session::owner(),
format!("RETURN http::get('{}')", server1.uri()),
true,
"SUCCESS".to_string(),
),
(
Datastore::builder()
.with_capabilities(Capabilities::default().with_scripting(true))
.build_with_path("memory")
.await
.unwrap(),
Session::owner(),
"RETURN function() { return '1' }".to_string(),
true,
"1".to_string(),
),
(
Datastore::builder()
.with_capabilities(Capabilities::default().with_scripting(false))
.build_with_path("memory")
.await
.unwrap(),
Session::owner(),
"RETURN function() { return '1' }".to_string(),
false,
"Scripting functions are not allowed".to_string(),
),
(
Datastore::builder()
.with_capabilities(Capabilities::default().with_guest_access(true))
.with_auth(true)
.build_with_path("memory")
.await
.unwrap(),
Session::default(),
"RETURN 1".to_string(),
true,
"1".to_string(),
),
(
Datastore::builder()
.with_capabilities(Capabilities::default().with_guest_access(false))
.with_auth(true)
.build_with_path("memory")
.await
.unwrap(),
Session::default(),
"RETURN 1".to_string(),
false,
"Not enough permissions to perform this action".to_string(),
),
(
Datastore::builder()
.with_auth(false)
.with_capabilities(Capabilities::default().with_guest_access(false))
.build_with_path("memory")
.await
.unwrap(),
Session::default(),
"RETURN 1".to_string(),
true,
"1".to_string(),
),
(
Datastore::builder()
.with_auth(true)
.with_capabilities(Capabilities::default().with_guest_access(false))
.build_with_path("memory")
.await
.unwrap(),
Session::viewer(),
"RETURN 1".to_string(),
true,
"1".to_string(),
),
(
Datastore::builder()
.with_capabilities(
Capabilities::default().with_experimental(ExperimentalTarget::Files.into()),
)
.build_with_path("memory")
.await
.unwrap(),
Session::owner().with_ns("test").with_db("test"),
"DEFINE BUCKET test BACKEND \"memory\";".to_string(),
true,
"NONE".to_string(),
),
(
Datastore::builder()
.with_capabilities(
Capabilities::default()
.without_experimental(ExperimentalTarget::Files.into()),
)
.build_with_path("memory")
.await
.unwrap(),
Session::owner().with_ns("test").with_db("test"),
"DEFINE BUCKET test BACKEND \"memory\";".to_string(),
false,
"expected the experimental files feature to be enabled".to_string(),
),
(
Datastore::builder()
.with_capabilities(
Capabilities::default()
.with_functions(Targets::<FuncTarget>::Some(
[FuncTarget::from_str("string::*").unwrap()].into(),
))
.without_functions(Targets::<FuncTarget>::Some(
[FuncTarget::from_str("string::len").unwrap()].into(),
)),
)
.build_with_path("memory")
.await
.unwrap(),
Session::owner(),
"RETURN string::len('a')".to_string(),
false,
"Function 'string::len' is not allowed".to_string(),
),
(
Datastore::builder()
.with_capabilities(
Capabilities::default()
.with_functions(Targets::<FuncTarget>::Some(
[FuncTarget::from_str("string::*").unwrap()].into(),
))
.without_functions(Targets::<FuncTarget>::Some(
[FuncTarget::from_str("string::len").unwrap()].into(),
)),
)
.build_with_path("memory")
.await
.unwrap(),
Session::owner(),
"RETURN string::lowercase('A')".to_string(),
true,
"a".to_string(),
),
(
Datastore::builder()
.with_capabilities(
Capabilities::default()
.with_functions(Targets::<FuncTarget>::Some(
[FuncTarget::from_str("string::*").unwrap()].into(),
))
.without_functions(Targets::<FuncTarget>::Some(
[FuncTarget::from_str("string::len").unwrap()].into(),
)),
)
.build_with_path("memory")
.await
.unwrap(),
Session::owner(),
"RETURN time::now()".to_string(),
false,
"Function 'time::now' is not allowed".to_string(),
),
(
Datastore::builder()
.with_capabilities(
Capabilities::default()
.with_functions(Targets::<FuncTarget>::All)
.with_network_targets(Targets::<NetTarget>::Some(
[
NetTarget::from_str(&server1.address().to_string()).unwrap(),
NetTarget::from_str(&server2.address().to_string()).unwrap(),
]
.into(),
))
.without_network_targets(Targets::<NetTarget>::Some(
[NetTarget::from_str(&server1.address().to_string()).unwrap()]
.into(),
)),
)
.build_with_path("memory")
.await
.unwrap(),
Session::owner(),
format!("RETURN http::get('{}')", server1.uri()),
false,
format!("Access to network target '{}' is not allowed", server1.address()),
),
(
Datastore::builder()
.with_capabilities(
Capabilities::default()
.with_functions(Targets::<FuncTarget>::All)
.with_network_targets(Targets::<NetTarget>::Some(
[
NetTarget::from_str(&server1.address().to_string()).unwrap(),
NetTarget::from_str(&server2.address().to_string()).unwrap(),
]
.into(),
))
.without_network_targets(Targets::<NetTarget>::Some(
[NetTarget::from_str(&server1.address().to_string()).unwrap()]
.into(),
)),
)
.build_with_path("memory")
.await
.unwrap(),
Session::owner(),
"RETURN http::get('http://1.1.1.1')".to_string(),
false,
"Access to network target '1.1.1.1:80' is not allowed".to_string(),
),
(
Datastore::builder()
.with_capabilities(
Capabilities::default()
.with_functions(Targets::<FuncTarget>::All)
.with_network_targets(Targets::<NetTarget>::Some(
[
NetTarget::from_str(&server1.address().to_string()).unwrap(),
NetTarget::from_str(&server2.address().to_string()).unwrap(),
]
.into(),
))
.without_network_targets(Targets::<NetTarget>::Some(
[NetTarget::from_str(&server1.address().to_string()).unwrap()]
.into(),
)),
)
.build_with_path("memory")
.await
.unwrap(),
Session::owner(),
format!("RETURN http::get('{}')", server2.uri()),
true,
"SUCCESS".to_string(),
),
(
Datastore::builder()
.with_capabilities(
Capabilities::default()
.with_functions(Targets::<FuncTarget>::All)
.with_network_targets(Targets::<NetTarget>::Some(
[NetTarget::from_str(&server3.address().to_string()).unwrap()]
.into(),
))
.without_network_targets(Targets::<NetTarget>::Some(
[NetTarget::from_str(&server1.address().to_string()).unwrap()]
.into(),
)),
)
.build_with_path("memory")
.await
.unwrap(),
Session::owner(),
format!("RETURN http::get('{}/redirect')", server3.uri()),
false,
format!(
"There was an error processing a remote HTTP request: error following redirect for url ({}/redirect)",
server3.uri()
),
),
(
Datastore::builder()
.with_capabilities(
Capabilities::default()
.with_functions(Targets::<FuncTarget>::All)
.with_network_targets(Targets::<NetTarget>::All),
)
.build_with_path("memory")
.await
.unwrap(),
Session::owner(),
format!("RETURN http::get('http://localhost:{}/test')", server1.address().port()),
true,
"SUCCESS".to_string(),
),
(
Datastore::builder()
.with_capabilities(
Capabilities::default()
.with_functions(Targets::<FuncTarget>::All)
.with_network_targets(Targets::<NetTarget>::All)
.without_network_targets(Targets::<NetTarget>::Some(
[
NetTarget::from_str("127.0.0.1/0").unwrap(),
NetTarget::from_str("::/0").unwrap(),
]
.into(),
)),
)
.build_with_path("memory")
.await
.unwrap(),
Session::owner(),
format!("RETURN http::get('http://localhost:{}')", server1.address().port()),
false,
"is not allowed".to_string(),
),
(
Datastore::builder()
.with_capabilities(
Capabilities::default()
.with_functions(Targets::<FuncTarget>::All)
.with_network_targets(Targets::<NetTarget>::Some(
[NetTarget::from_str("github.com").unwrap()].into(),
))
.without_network_targets(Targets::<NetTarget>::Some(
[
NetTarget::from_str("0.0.0.0/8").unwrap(),
NetTarget::from_str("10.0.0.0/8").unwrap(),
NetTarget::from_str("10.18.0.0/16").unwrap(),
NetTarget::from_str("10.2.0.0/16").unwrap(),
NetTarget::from_str("100.64.0.0/10").unwrap(),
NetTarget::from_str("127.0.0.0/8").unwrap(),
NetTarget::from_str("169.254.0.0/16").unwrap(),
NetTarget::from_str("172.16.0.0/12").unwrap(),
NetTarget::from_str("172.20.0.0/16").unwrap(),
NetTarget::from_str("192.0.0.0/24").unwrap(),
NetTarget::from_str("192.168.0.0/16").unwrap(),
NetTarget::from_str("192.88.99.0/24").unwrap(),
NetTarget::from_str("198.18.0.0/15").unwrap(),
NetTarget::from_str("::1/128").unwrap(),
NetTarget::from_str("fc00::/7").unwrap(),
NetTarget::from_str("fc00::/8").unwrap(),
]
.into(),
)),
)
.build_with_path("memory")
.await
.unwrap(),
Session::owner(),
"RETURN http::get('https://github.com/surrealdb/surrealdb/issues/6293')"
.to_string(),
true,
"<!DOCTYPE html>".to_string(),
),
];
for (idx, (ds, sess, query, succeeds, contains)) in cases.into_iter().enumerate() {
info!("Test case {idx}: query={query}, succeeds={succeeds}");
let res = ds.execute(&query, &sess, None).await;
if !succeeds && res.is_err() {
let res = res.unwrap_err();
assert!(
res.to_string().contains(&contains),
"Unexpected error for test case {}: {:?}",
idx,
res.to_string()
);
continue;
}
let res = res.unwrap().remove(0).output();
let res = if succeeds {
assert!(res.is_ok(), "Unexpected error for test case {idx}: {res:?}");
res.unwrap().to_sql()
} else {
assert!(res.is_err(), "Unexpected success for test case {idx}: {res:?}");
res.unwrap_err().to_string()
};
assert!(
res.contains(&contains),
"Unexpected result for test case {idx}: expected to contain = `{contains}`, got `{res}`"
);
}
server1.verify().await;
server2.verify().await;
server3.verify().await;
}
#[test]
fn test_dbs_capabilities_target_all() {
let caps = DbsCapabilities {
allow_all: false,
#[cfg(feature = "scripting")]
allow_scripting: false,
allow_guests: false,
allow_funcs: None,
allow_experimental: Some(Targets::All),
allow_arbitrary_query: Some(Targets::All),
allow_eval_query: None,
allow_net: None,
allow_rpc: None,
allow_http: None,
deny_all: false,
#[cfg(feature = "scripting")]
deny_scripting: false,
deny_guests: false,
deny_funcs: None,
deny_experimental: None,
deny_arbitrary_query: None,
deny_eval_query: None,
deny_net: None,
deny_rpc: None,
deny_http: None,
planner_strategy: NewPlannerStrategy::default(),
};
assert_eq!(caps.get_allow_experimental(), Targets::All);
assert_eq!(caps.get_allow_arbitrary_query(), Targets::All);
}
#[test]
fn test_dbs_capabilities_eval_query_denied_by_default() {
let mut caps = DbsCapabilities {
allow_all: true,
#[cfg(feature = "scripting")]
allow_scripting: false,
allow_guests: false,
allow_funcs: None,
allow_experimental: None,
allow_arbitrary_query: None,
allow_eval_query: None,
allow_net: None,
allow_rpc: None,
allow_http: None,
deny_all: false,
#[cfg(feature = "scripting")]
deny_scripting: false,
deny_guests: false,
deny_funcs: None,
deny_experimental: None,
deny_arbitrary_query: None,
deny_eval_query: None,
deny_net: None,
deny_rpc: None,
deny_http: None,
planner_strategy: NewPlannerStrategy::default(),
};
assert_eq!(caps.get_allow_arbitrary_query(), Targets::All);
assert_eq!(caps.get_allow_eval_query(), Targets::None, "eval is not enabled by allow_all");
caps.allow_eval_query = Some(Targets::from(EvalQueryTarget::System));
assert_eq!(caps.get_allow_eval_query(), Targets::from(EvalQueryTarget::System));
}
#[test]
fn test_dbs_capabilities_deny_all_vs_allow_all() {
let caps = DbsCapabilities {
allow_all: false,
#[cfg(feature = "scripting")]
allow_scripting: false,
allow_guests: false,
allow_funcs: Some(Targets::All),
allow_experimental: None,
allow_arbitrary_query: None,
allow_eval_query: None,
allow_net: None,
allow_rpc: None,
allow_http: None,
deny_all: false,
#[cfg(feature = "scripting")]
deny_scripting: false,
deny_guests: false,
deny_funcs: Some(Targets::All),
deny_experimental: None,
deny_arbitrary_query: None,
deny_eval_query: None,
deny_net: None,
deny_rpc: None,
deny_http: None,
planner_strategy: NewPlannerStrategy::default(),
};
assert_eq!(
caps.get_allow_funcs(),
Targets::None,
"When deny_funcs=All and allow_funcs=All, should deny (return None)"
);
let caps = DbsCapabilities {
allow_all: false,
#[cfg(feature = "scripting")]
allow_scripting: false,
allow_guests: false,
allow_funcs: None,
allow_experimental: None,
allow_arbitrary_query: None,
allow_eval_query: None,
allow_net: None,
allow_rpc: None,
allow_http: Some(Targets::All),
deny_all: false,
#[cfg(feature = "scripting")]
deny_scripting: false,
deny_guests: false,
deny_funcs: None,
deny_experimental: None,
deny_arbitrary_query: None,
deny_eval_query: None,
deny_net: None,
deny_rpc: None,
deny_http: Some(Targets::All),
planner_strategy: NewPlannerStrategy::default(),
};
assert_eq!(
caps.get_allow_http(),
Targets::None,
"When deny_http=All and allow_http=All, should deny (return None)"
);
let caps = DbsCapabilities {
allow_all: false,
#[cfg(feature = "scripting")]
allow_scripting: false,
allow_guests: false,
allow_funcs: Some(Targets::None),
allow_experimental: None,
allow_arbitrary_query: None,
allow_eval_query: None,
allow_net: None,
allow_rpc: None,
allow_http: None,
deny_all: false,
#[cfg(feature = "scripting")]
deny_scripting: false,
deny_guests: false,
deny_funcs: Some(Targets::All),
deny_experimental: None,
deny_arbitrary_query: None,
deny_eval_query: None,
deny_net: None,
deny_rpc: None,
deny_http: None,
planner_strategy: NewPlannerStrategy::default(),
};
assert_eq!(
caps.get_allow_funcs(),
Targets::None,
"When deny_funcs=All and allow_funcs=None, should deny (return None)"
);
}
#[cfg(feature = "storage-rocksdb")]
async fn seed_archived_peer(path: &str) -> Result<(NodeKey, NodeLiveQueryKey)> {
use surrealdb_core::dbs::node::{Node, Timestamp};
use surrealdb_datastore::TransactionType;
use surrealdb_datastore::catalog::{DatabaseId, NamespaceId, NodeLiveQuery};
use uuid::Uuid;
let peer = Uuid::new_v4();
let node_key = NodeKey {
nd: peer,
};
let nlq_key = NodeLiveQueryKey {
nd: peer,
lq: Uuid::new_v4(),
};
let seed = Datastore::builder().without_maintenance_tasks().build_with_path(path).await?;
seed.check_version().await?;
let txn = seed.transaction(TransactionType::Write).await?;
txn.set_key(&node_key, &Node::new(peer, Timestamp::default(), true)).await?;
txn.set_key(
&nlq_key,
&NodeLiveQuery {
ns: NamespaceId(1),
db: DatabaseId(1),
tb: "test".to_string().into(),
},
)
.await?;
txn.commit().await?;
seed.shutdown().await?;
Ok((node_key, nlq_key))
}
#[cfg(feature = "storage-rocksdb")]
async fn seed_interrupted_bootstrap(path: &str) -> Result<()> {
use surrealdb_core::dbs::node::{Node, Timestamp};
use surrealdb_datastore::TransactionType;
use surrealdb_datastore::key::schema::BootstrapKey;
use surrealdb_datastore::version::MajorVersion;
use uuid::Uuid;
let seed = Datastore::builder().without_maintenance_tasks().build_with_path(path).await?;
let txn = seed.transaction(TransactionType::Write).await?;
txn.set_key(&BootstrapKey {}, &MajorVersion::latest()).await?;
let peer = Uuid::new_v4();
txn.set_key(
&NodeKey {
nd: peer,
},
&Node::new(peer, Timestamp::default(), false),
)
.await?;
txn.commit().await?;
seed.shutdown().await?;
Ok(())
}
#[cfg(any(feature = "storage-rocksdb", feature = "storage-mem"))]
fn node_maintenance_config(path: String, interval: Duration) -> Result<Config> {
use surrealdb_core::options::EngineOptions;
use crate::ntw::client_ip::ClientIp;
Ok(Config {
bind: "127.0.0.1:0".parse()?,
postgres_bind: None,
path,
client_ip: ClientIp::None,
user: None,
pass: None,
crt: None,
key: None,
engine: EngineOptions::default()
.with_node_membership_check_interval(interval)
.with_node_membership_cleanup_interval(interval),
no_identification_headers: false,
allow_origin: Vec::new(),
durable_session_ttl: None,
})
}
#[cfg(feature = "storage-rocksdb")]
async fn seeded_rows_present(
dbs: &Datastore,
node_key: &NodeKey,
nlq_key: &NodeLiveQueryKey,
) -> Result<(bool, bool)> {
use surrealdb_datastore::TransactionType;
let txn = dbs.transaction(TransactionType::Read).await?;
let node = txn.get_key(node_key, None).await?;
let nlq = txn.get_key(nlq_key, None).await?;
txn.cancel().await?;
Ok((node.is_some(), nlq.is_some()))
}
#[cfg(feature = "storage-rocksdb")]
#[test(tokio::test(flavor = "multi_thread"))]
#[serial]
async fn init_leaves_archived_nodes_for_the_maintenance_scheduler() -> Result<()> {
use surrealdb_core::CommunityComposer;
use surrealdb_observe::NoopObserver;
let dir = tempfile::tempdir()?;
let path = format!("rocksdb:{}", dir.path().join("store").display());
let (node_key, nlq_key) = seed_archived_peer(&path).await?;
let dbs_opts = TestCli::try_parse_from(["surrealdb"])?.dbs;
let config = node_maintenance_config(path, Duration::ZERO)?;
let canceller = CancellationToken::new();
let (dbs, _recv, _router_state, pending) =
init(CommunityComposer(), &config, canceller.clone(), Arc::new(NoopObserver), dbs_opts)
.await?;
initialise_datastore(&dbs, &pending).await?;
let (node, nlq) = seeded_rows_present(&dbs, &node_key, &nlq_key).await?;
assert!(node, "startup removed the archived node");
assert!(nlq, "startup removed the archived node's live query");
dbs.remove_nodes().await?;
let (node, nlq) = seeded_rows_present(&dbs, &node_key, &nlq_key).await?;
assert!(!node, "remove_nodes left the archived node behind");
assert!(!nlq, "remove_nodes left the archived node's live query behind");
canceller.cancel();
dbs.shutdown().await?;
Ok(())
}
#[cfg(feature = "storage-rocksdb")]
#[test(tokio::test(flavor = "multi_thread"))]
#[serial]
async fn init_spawns_the_first_cleanup_pass_after_registration() -> Result<()> {
use surrealdb_core::CommunityComposer;
use surrealdb_observe::NoopObserver;
let dir = tempfile::tempdir()?;
let path = format!("rocksdb:{}", dir.path().join("store").display());
let (node_key, nlq_key) = seed_archived_peer(&path).await?;
let dbs_opts = TestCli::try_parse_from(["surrealdb"])?.dbs;
let config = node_maintenance_config(path, Duration::from_secs(3600))?;
let canceller = CancellationToken::new();
let (dbs, _recv, _router_state, pending) =
init(CommunityComposer(), &config, canceller.clone(), Arc::new(NoopObserver), dbs_opts)
.await?;
initialise_datastore(&dbs, &pending).await?;
let deadline = Instant::now() + Duration::from_secs(15);
let (node, nlq) = loop {
let present = seeded_rows_present(&dbs, &node_key, &nlq_key).await?;
if present == (false, false) || Instant::now() >= deadline {
break present;
}
sleep(Duration::from_millis(100)).await;
};
assert!(!node, "the first cleanup pass left the archived node behind");
assert!(!nlq, "the first cleanup pass left the archived node's live query behind");
canceller.cancel();
dbs.shutdown().await?;
Ok(())
}
#[cfg(feature = "storage-mem")]
#[test(tokio::test(flavor = "multi_thread"))]
async fn init_defers_every_step_that_opens_a_transaction() -> Result<()> {
use surrealdb_core::CommunityComposer;
use surrealdb_observe::NoopObserver;
let dbs_opts = TestCli::try_parse_from(["surrealdb"])?.dbs;
let config = node_maintenance_config("memory".to_string(), Duration::from_secs(3600))?;
let canceller = CancellationToken::new();
let (dbs, _recv, _router_state, pending) =
init(CommunityComposer(), &config, canceller.clone(), Arc::new(NoopObserver), dbs_opts)
.await?;
assert!(
!dbs.maintenance_tasks_running(),
"init started the maintenance schedule, whose tasks all write"
);
assert!(
dbs.node_heartbeat_age().await?.is_none(),
"init registered this node's cluster membership row"
);
initialise_datastore(&dbs, &pending).await?;
assert!(
dbs.maintenance_tasks_running(),
"the deferred initialisation started no maintenance schedule"
);
assert!(
dbs.node_heartbeat_age().await?.is_some(),
"the deferred initialisation did not register this node"
);
canceller.cancel();
dbs.shutdown().await?;
Ok(())
}
#[cfg(feature = "storage-mem")]
#[serial]
#[test(tokio::test(flavor = "multi_thread"))]
async fn initialisation_waits_for_the_backend_to_become_serve_ready() -> Result<()> {
use std::sync::atomic::AtomicBool;
use common::future::BoxFut;
use surrealdb_core::CommunityComposer;
use surrealdb_core::kvs::TransactionBuilderParts;
use surrealdb_kvs::api::Transactable;
use surrealdb_kvs::err::Result as KvsResult;
use surrealdb_kvs::{Metrics, TransactionBuilder, TransactionType};
use surrealdb_observe::NoopObserver;
const SHORT_BUDGET: Duration = Duration::from_secs(1);
struct GatedBuilder {
inner: Box<dyn TransactionBuilder>,
gate: CancellationToken,
ready: Arc<AtomicBool>,
opened_early: Arc<AtomicBool>,
}
impl TransactionBuilder for GatedBuilder {
fn name(&self) -> &'static str {
self.inner.name()
}
fn new_transaction(
&self,
write: TransactionType,
) -> BoxFut<'_, KvsResult<(Box<dyn Transactable>, bool)>> {
if !self.ready.load(Ordering::SeqCst) {
self.opened_early.store(true, Ordering::SeqCst);
}
self.inner.new_transaction(write)
}
fn shutdown(&self) -> BoxFut<'_, KvsResult<()>> {
self.inner.shutdown()
}
fn register_metrics(&self) -> Option<Metrics> {
self.inner.register_metrics()
}
fn collect_u64_metric(&self, metric: &str) -> Option<u64> {
self.inner.collect_u64_metric(metric)
}
fn wait_until_serve_ready(&self) -> BoxFut<'_, KvsResult<()>> {
Box::pin(async move {
self.gate.cancelled().await;
self.ready.store(true, Ordering::SeqCst);
Ok(())
})
}
}
struct GatedComposer {
gate: CancellationToken,
ready: Arc<AtomicBool>,
opened_early: Arc<AtomicBool>,
}
impl TransactionBuilderFactory for GatedComposer {
type RouterState = ();
async fn new_transaction_builder(
&self,
path: &str,
canceller: CancellationToken,
config: ConfigMap,
) -> Result<TransactionBuilderParts<Self::RouterState>> {
let parts =
CommunityComposer().new_transaction_builder(path, canceller, config).await?;
Ok(TransactionBuilderParts::without_router_state(Box::new(GatedBuilder {
inner: parts.builder,
gate: self.gate.clone(),
ready: Arc::clone(&self.ready),
opened_early: Arc::clone(&self.opened_early),
})))
}
fn path_valid(&self, v: &str) -> Result<String> {
CommunityComposer().path_valid(v)
}
}
let gate = CancellationToken::new();
let ready = Arc::new(AtomicBool::new(false));
let opened_early = Arc::new(AtomicBool::new(false));
let composer = GatedComposer {
gate: gate.clone(),
ready: Arc::clone(&ready),
opened_early: Arc::clone(&opened_early),
};
let _budget = EnvGuard::set(STARTUP_OPERATION_TIMEOUT_ENV, "1s");
let dbs_opts = TestCli::try_parse_from(["surrealdb"])?.dbs;
let config = node_maintenance_config("memory".to_string(), Duration::from_secs(3600))?;
let canceller = CancellationToken::new();
let (dbs, _recv, _router_state, pending) =
init(composer, &config, canceller.clone(), Arc::new(NoopObserver), dbs_opts).await?;
assert!(!ready.load(Ordering::SeqCst), "init waited for the backend to become serve-ready");
assert_eq!(
pending.timeout, SHORT_BUDGET,
"the retry budget this test outlasts was not the one applied"
);
let mut deferred = tokio::spawn({
let dbs = Arc::clone(&dbs);
async move { initialise_datastore(&dbs, &pending).await }
});
assert!(
tokio::time::timeout(SHORT_BUDGET + SHORT_BUDGET / 2, &mut deferred).await.is_err(),
"the deferred initialisation ran without waiting for serve-readiness"
);
assert!(
!opened_early.load(Ordering::SeqCst),
"a transaction was opened before the backend reported serve-readiness"
);
gate.cancel();
deferred.await??;
assert!(
!opened_early.load(Ordering::SeqCst),
"a transaction was opened before the backend reported serve-readiness"
);
canceller.cancel();
dbs.shutdown().await?;
Ok(())
}
#[cfg(feature = "storage-mem")]
#[test(tokio::test(flavor = "multi_thread"))]
async fn init_starts_the_maintenance_schedule_after_the_version_gate() -> Result<()> {
use surrealdb_core::CommunityComposer;
use surrealdb_datastore::version::MajorVersion;
use surrealdb_observe::NoopObserver;
let dbs_opts = TestCli::try_parse_from(["surrealdb"])?.dbs;
let config = node_maintenance_config("memory".to_string(), Duration::from_secs(3600))?;
let canceller = CancellationToken::new();
let (dbs, _recv, _router_state, pending) =
init(CommunityComposer(), &config, canceller.clone(), Arc::new(NoopObserver), dbs_opts)
.await?;
initialise_datastore(&dbs, &pending).await?;
assert!(
dbs.maintenance_tasks_running(),
"startup left the datastore with no maintenance tasks"
);
assert_eq!(dbs.get_version().await?, (MajorVersion::latest(), false));
canceller.cancel();
dbs.shutdown().await?;
Ok(())
}
#[cfg(feature = "storage-rocksdb")]
#[test(tokio::test(flavor = "multi_thread"))]
#[serial]
async fn init_completes_a_bootstrap_a_previous_boot_left_unfinished() -> Result<()> {
use surrealdb_core::CommunityComposer;
use surrealdb_datastore::version::MajorVersion;
use surrealdb_observe::NoopObserver;
let dir = tempfile::tempdir()?;
let path = format!("rocksdb:{}", dir.path().join("store").display());
seed_interrupted_bootstrap(&path).await?;
let dbs_opts = TestCli::try_parse_from(["surrealdb"])?.dbs;
let config = node_maintenance_config(path, Duration::from_secs(3600))?;
let canceller = CancellationToken::new();
let (dbs, _recv, _router_state, pending) =
init(CommunityComposer(), &config, canceller.clone(), Arc::new(NoopObserver), dbs_opts)
.await?;
initialise_datastore(&dbs, &pending).await?;
assert_eq!(
dbs.get_version().await?,
(MajorVersion::latest(), false),
"the interrupted bootstrap was not completed at this build's version"
);
canceller.cancel();
dbs.shutdown().await?;
Ok(())
}
}