use std::future::Future;
use crate::exec::{AccessMode, ExecutionContext};
use crate::expr::FlowResult;
use crate::val::Value;
pub(crate) const RECORD_DEREFERENCE: AccessMode = AccessMode::ReadOnly;
pub(crate) async fn evaluate_each<'a, T, Fut>(
ctx: &ExecutionContext,
mode: AccessMode,
items: &'a [T],
eval_one: impl Fn(&'a T) -> Fut,
) -> FlowResult<Vec<Value>>
where
Fut: Future<Output = FlowResult<Value>>,
{
evaluate_each_above(ctx.root().ctx.config.exec.fan_out_row_threshold, mode, items, eval_one)
.await
}
async fn evaluate_each_above<'a, T, Fut>(
threshold: usize,
mode: AccessMode,
items: &'a [T],
eval_one: impl Fn(&'a T) -> Fut,
) -> FlowResult<Vec<Value>>
where
Fut: Future<Output = FlowResult<Value>>,
{
if mode.is_read_write() || items.len() < threshold {
let mut results = Vec::with_capacity(items.len());
for item in items {
results.push(eval_one(item).await?);
}
return Ok(results);
}
futures::future::try_join_all(items.iter().map(eval_one)).await
}
#[cfg(test)]
mod tests {
use std::cell::RefCell;
use anyhow::anyhow;
use super::*;
use crate::expr::ControlFlow;
async fn traced(trace: &RefCell<Vec<String>>, item: &i64) -> FlowResult<Value> {
trace.borrow_mut().push(format!("start {item}"));
tokio::task::yield_now().await;
trace.borrow_mut().push(format!("end {item}"));
Ok(Value::from(*item))
}
const THRESHOLD: usize = 4;
async fn run(mode: AccessMode, items: &[i64]) -> (Vec<Value>, Vec<String>) {
let trace = RefCell::new(Vec::new());
let results = evaluate_each_above(THRESHOLD, mode, items, |item| traced(&trace, item))
.await
.expect("evaluation should succeed");
(results, trace.into_inner())
}
fn sequential_trace(items: &[i64]) -> Vec<String> {
items.iter().flat_map(|i| [format!("start {i}"), format!("end {i}")]).collect()
}
#[tokio::test]
async fn a_read_write_evaluation_runs_one_row_at_a_time_in_input_order() {
let rows = [1, 2, 3, 4, 5];
let (results, trace) = run(AccessMode::ReadWrite, &rows).await;
assert_eq!(
trace,
sequential_trace(&rows),
"a writing evaluation must not overlap its rows, however many there are"
);
assert_eq!(results, rows.map(Value::from));
}
#[tokio::test]
async fn a_read_only_evaluation_overlaps_its_rows() {
let rows = [1, 2, 3, 4, 5];
let (results, trace) = run(AccessMode::ReadOnly, &rows).await;
assert_eq!(
trace,
[
"start 1", "start 2", "start 3", "start 4", "start 5", "end 1", "end 2", "end 3",
"end 4", "end 5"
],
"a read-only evaluation should overlap the rows it is given"
);
assert_eq!(results, rows.map(Value::from));
}
#[tokio::test]
async fn too_few_rows_to_be_worth_overlapping_stay_sequential() {
for count in 0..THRESHOLD {
let rows: Vec<i64> = (1..=count as i64).collect();
let (results, trace) = run(AccessMode::ReadOnly, &rows).await;
assert_eq!(trace, sequential_trace(&rows), "{count} rows should not overlap");
assert_eq!(results, rows.iter().copied().map(Value::from).collect::<Vec<_>>());
}
}
#[cfg(feature = "kv-mem")]
#[tokio::test]
async fn the_threshold_comes_from_the_execution_config() {
use surrealdb_cnf::ConfigMap;
use crate::exec::operators::test_util::TestDb;
let rows = [1, 2, 3, 4, 5];
for (threshold, overlaps) in [("2", true), ("1000", false)] {
let db = TestDb::new_with_config(
"",
ConfigMap::empty().with_key_value("fan_out_row_threshold", threshold),
)
.await;
let ctx = db.exec_ctx().await;
let trace = RefCell::new(Vec::new());
evaluate_each(&ctx, AccessMode::ReadOnly, &rows, |item| traced(&trace, item))
.await
.expect("evaluation should succeed");
let overlapped = trace.into_inner() != sequential_trace(&rows);
assert_eq!(
overlapped,
overlaps,
"a threshold of {threshold} should {} five rows",
if overlaps {
"overlap"
} else {
"not overlap"
}
);
}
}
#[tokio::test]
async fn a_failing_row_stops_a_read_write_evaluation_where_it_failed() {
let trace = RefCell::new(Vec::new());
let recorded = &trace;
let err =
evaluate_each_above(THRESHOLD, AccessMode::ReadWrite, &[1, 2, 3], |item| async move {
if *item == 2 {
return Err(ControlFlow::from(anyhow!("row {item} failed")));
}
traced(recorded, item).await
})
.await
.expect_err("the failing row should abort the batch");
assert!(err.to_string().contains("row 2 failed"), "got {err}");
assert_eq!(trace.into_inner(), ["start 1", "end 1"]);
}
#[tokio::test]
async fn a_control_flow_signal_propagates_out_of_both_arms() {
for mode in [AccessMode::ReadOnly, AccessMode::ReadWrite] {
let signal =
evaluate_each_above(THRESHOLD, mode, &[1, 2, 3, 4, 5], |item| async move {
if *item == 5 {
return Err(ControlFlow::Break);
}
Ok(Value::from(*item))
})
.await
.expect_err("the signal should reach the caller");
assert!(matches!(signal, ControlFlow::Break), "{mode:?} swallowed the signal");
}
}
}
#[cfg(not(target_family = "wasm"))]
#[derive(Clone, Copy)]
pub(crate) enum SpawnedDocRoot {
RowValue,
Inherited,
}
#[cfg(not(target_family = "wasm"))]
const SPAWNED_CHUNK_CONCURRENCY: usize = 16;
#[cfg(not(target_family = "wasm"))]
const PROBE_WARM_ROWS: usize = 2;
#[cfg(not(target_family = "wasm"))]
const PROBE_TIMED_ROWS: usize = 6;
#[cfg(not(target_family = "wasm"))]
const SPAWN_PER_ROW_NANOS: u64 = 5_000;
#[cfg(not(target_family = "wasm"))]
const SPAWN_MIN_TOTAL_NANOS: u64 = 8_000_000;
#[cfg(not(target_family = "wasm"))]
pub(crate) async fn evaluate_rows_spawned(
expr: std::sync::Arc<dyn crate::exec::PhysicalExpr>,
ctx: &crate::exec::physical_expr::EvalContext<'_>,
values: &[Value],
doc_root: SpawnedDocRoot,
) -> crate::expr::FlowResult<Vec<Value>> {
use crate::exec::physical_expr::EvalContext;
let inherited_root: Option<std::sync::Arc<Value>> = match doc_root {
SpawnedDocRoot::Inherited => ctx.document_root.map(|v| std::sync::Arc::new(v.clone())),
SpawnedDocRoot::RowValue => None,
};
let local_params: Option<
std::sync::Arc<std::collections::HashMap<surrealdb_strand::Strand, Value>>,
> = ctx.local_params.map(|m| std::sync::Arc::new(m.clone()));
let probe = (PROBE_WARM_ROWS + PROBE_TIMED_ROWS).min(values.len());
let mut head = Vec::with_capacity(probe);
let mut started = web_time::Instant::now();
for (idx, value) in values.iter().take(probe).enumerate() {
if idx == PROBE_WARM_ROWS {
started = web_time::Instant::now();
}
let row_ctx = match doc_root {
SpawnedDocRoot::RowValue => ctx.with_value_and_doc(value),
SpawnedDocRoot::Inherited => ctx.with_value(value),
};
head.push(expr.evaluate(row_ctx).await?);
}
let timed = head.len().saturating_sub(PROBE_WARM_ROWS);
let per_row = started.elapsed().as_nanos() as u64 / timed.max(1) as u64;
let remaining_work = per_row.saturating_mul(values.len().saturating_sub(probe) as u64);
if timed == 0
|| per_row < SPAWN_PER_ROW_NANOS
|| remaining_work < SPAWN_MIN_TOTAL_NANOS
|| values.len() <= probe
{
let mut results = head;
results.reserve(values.len().saturating_sub(probe));
for value in values.iter().skip(probe) {
let row_ctx = match doc_root {
SpawnedDocRoot::RowValue => ctx.with_value_and_doc(value),
SpawnedDocRoot::Inherited => ctx.with_value(value),
};
results.push(expr.evaluate(row_ctx).await?);
}
return Ok(results);
}
let values = &values[probe..];
let chunk_size = values.len().div_ceil(SPAWNED_CHUNK_CONCURRENCY * 4).max(1);
let mut out: Vec<Option<Vec<Value>>> = Vec::new();
out.resize_with(values.len().div_ceil(chunk_size), || None);
let mut pending: std::collections::VecDeque<_> = values
.chunks(chunk_size)
.enumerate()
.map(|(idx, chunk)| {
let expr = std::sync::Arc::clone(&expr);
let exec_ctx = ctx.exec_ctx.clone();
let chunk = chunk.to_vec();
let inherited_root = inherited_root.clone();
let local_params = local_params.clone();
let skip_fetch_perms = ctx.skip_fetch_perms;
let computing_record = ctx.computing_record.clone();
let plan_depth = ctx.plan_depth;
async move {
let mut results = Vec::with_capacity(chunk.len());
for value in &chunk {
let mut row_ctx = EvalContext::from_exec_ctx(&exec_ctx);
row_ctx.local_params = local_params.as_deref();
row_ctx.skip_fetch_perms = skip_fetch_perms;
row_ctx.computing_record = computing_record.clone();
row_ctx.plan_depth = plan_depth;
let row_ctx = match doc_root {
SpawnedDocRoot::RowValue => row_ctx.with_value_and_doc(value),
SpawnedDocRoot::Inherited => {
row_ctx.document_root = inherited_root.as_deref();
row_ctx.with_value(value)
}
};
match expr.evaluate(row_ctx).await {
Ok(value) => results.push(value),
Err(err) => return (idx, Err(err)),
}
}
(idx, Ok(results))
}
})
.collect();
let mut set = tokio::task::JoinSet::new();
let mut first_err: Option<(usize, crate::expr::ControlFlow)> = None;
let mut join_err: Option<crate::expr::ControlFlow> = None;
loop {
if first_err.is_none() && join_err.is_none() {
while set.len() < SPAWNED_CHUNK_CONCURRENCY {
let Some(task) = pending.pop_front() else {
break;
};
set.spawn(task);
}
}
let Some(joined) = set.join_next().await else {
break;
};
match joined {
Ok((idx, Ok(results))) => out[idx] = Some(results),
Ok((idx, Err(err))) => {
if first_err.as_ref().is_none_or(|(lowest, _)| idx < *lowest) {
first_err = Some((idx, err));
}
}
Err(e) => {
join_err.get_or_insert_with(|| {
crate::expr::ControlFlow::Err(anyhow::anyhow!("row fan-out task failed: {e}"))
});
}
}
}
if let Some((_, err)) = first_err {
return Err(err);
}
if let Some(err) = join_err {
return Err(err);
}
let mut results = head;
results.reserve(values.len());
results.extend(out.into_iter().flat_map(|chunk| chunk.expect("every chunk task completed")));
Ok(results)
}
#[cfg(not(target_family = "wasm"))]
#[cfg(test)]
mod spawned_tests {
use std::sync::Arc;
use std::time::Duration;
use surrealdb_types::{SqlFormat, ToSql};
use super::*;
use crate::exec::operators::test_util::root_ctx;
use crate::exec::{BoxFut, ContextLevel, EvalContext, PhysicalExpr};
use crate::val::Number;
#[derive(Debug)]
struct SlowDoubler {
per_row_delay: Duration,
}
impl ToSql for SlowDoubler {
fn fmt_sql(&self, f: &mut String, _fmt: SqlFormat) {
f.push_str("SlowDoubler");
}
}
impl PhysicalExpr for SlowDoubler {
fn name(&self) -> &'static str {
"SlowDoubler"
}
fn as_any(&self) -> &dyn std::any::Any {
self
}
fn required_context(&self) -> ContextLevel {
ContextLevel::Root
}
fn evaluate<'a>(&'a self, ctx: EvalContext<'a>) -> BoxFut<'a, FlowResult<Value>> {
Box::pin(async move {
let n = match ctx.current_value {
Some(Value::Number(Number::Int(n))) => *n,
other => panic!("expected an integer row, got {other:?}"),
};
tokio::time::sleep(self.per_row_delay).await;
Ok(Value::from(n * 2 + 1))
})
}
fn access_mode(&self) -> AccessMode {
AccessMode::ReadOnly
}
}
#[tokio::test]
async fn batches_around_the_probe_boundary_keep_every_row() {
let probe = (PROBE_WARM_ROWS + PROBE_TIMED_ROWS) as i64;
for rows in [probe - 1, probe, probe + 1] {
let values: Vec<Value> = (0..rows).map(Value::from).collect();
let expr: Arc<dyn PhysicalExpr> = Arc::new(SlowDoubler {
per_row_delay: Duration::from_millis(10),
});
let exec_ctx = root_ctx();
let base = EvalContext::from_exec_ctx(&exec_ctx);
let results = evaluate_rows_spawned(expr, &base, &values, SpawnedDocRoot::RowValue)
.await
.expect("evaluation should succeed");
let expected: Vec<Value> = (0..rows).map(|n| Value::from(n * 2 + 1)).collect();
assert_eq!(results, expected, "{rows} rows must come back complete and in order");
}
}
#[derive(Debug)]
struct LocalParamReader;
impl ToSql for LocalParamReader {
fn fmt_sql(&self, f: &mut String, _fmt: SqlFormat) {
f.push_str("LocalParamReader");
}
}
impl PhysicalExpr for LocalParamReader {
fn name(&self) -> &'static str {
"LocalParamReader"
}
fn as_any(&self) -> &dyn std::any::Any {
self
}
fn required_context(&self) -> ContextLevel {
ContextLevel::Root
}
fn evaluate<'a>(&'a self, ctx: EvalContext<'a>) -> BoxFut<'a, FlowResult<Value>> {
Box::pin(async move {
tokio::time::sleep(Duration::from_millis(1)).await;
Ok(ctx
.local_params
.and_then(|params| params.get("x"))
.cloned()
.unwrap_or(Value::None))
})
}
fn access_mode(&self) -> AccessMode {
AccessMode::ReadOnly
}
}
#[tokio::test]
async fn spawned_rows_still_see_block_local_params() {
use std::collections::HashMap;
use surrealdb_strand::Strand;
const ROWS: i64 = 100;
let rows: Vec<Value> = (0..ROWS).map(Value::from).collect();
let expr: Arc<dyn PhysicalExpr> = Arc::new(LocalParamReader);
let exec_ctx = root_ctx();
let params: HashMap<Strand, Value> = HashMap::from([(Strand::from("x"), Value::from(42))]);
let mut base = EvalContext::from_exec_ctx(&exec_ctx);
base.local_params = Some(¶ms);
let results = evaluate_rows_spawned(expr, &base, &rows, SpawnedDocRoot::RowValue)
.await
.expect("evaluation should succeed");
assert_eq!(results.len(), rows.len());
for (idx, result) in results.iter().enumerate() {
assert_eq!(
*result,
Value::from(42),
"row {idx} must resolve $x from the restored local params"
);
}
}
#[derive(Debug)]
struct FailsTwice;
impl ToSql for FailsTwice {
fn fmt_sql(&self, f: &mut String, _fmt: SqlFormat) {
f.push_str("FailsTwice");
}
}
impl PhysicalExpr for FailsTwice {
fn name(&self) -> &'static str {
"FailsTwice"
}
fn as_any(&self) -> &dyn std::any::Any {
self
}
fn required_context(&self) -> ContextLevel {
ContextLevel::Root
}
fn evaluate<'a>(&'a self, ctx: EvalContext<'a>) -> BoxFut<'a, FlowResult<Value>> {
Box::pin(async move {
let n = match ctx.current_value {
Some(Value::Number(Number::Int(n))) => *n,
other => panic!("expected an integer row, got {other:?}"),
};
match n {
10 => {
tokio::time::sleep(Duration::from_millis(20)).await;
Err(crate::expr::ControlFlow::Err(anyhow::anyhow!("row 10 failed")))
}
90 => Err(crate::expr::ControlFlow::Err(anyhow::anyhow!("row 90 failed"))),
_ => {
tokio::time::sleep(Duration::from_millis(1)).await;
Ok(Value::from(n))
}
}
})
}
fn access_mode(&self) -> AccessMode {
AccessMode::ReadOnly
}
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn the_lowest_index_error_surfaces_regardless_of_completion_order() {
const ROWS: i64 = 100;
let rows: Vec<Value> = (0..ROWS).map(Value::from).collect();
let expr: Arc<dyn PhysicalExpr> = Arc::new(FailsTwice);
let exec_ctx = root_ctx();
let base = EvalContext::from_exec_ctx(&exec_ctx);
let err = evaluate_rows_spawned(expr, &base, &rows, SpawnedDocRoot::RowValue)
.await
.expect_err("two rows fail, so the batch must fail");
assert!(
err.to_string().contains("row 10 failed"),
"row 90's immediate failure completes first, but row 10's must surface: got {err}"
);
}
#[tokio::test]
async fn a_spawned_batch_keeps_every_row_including_the_timing_probe() {
const ROWS: i64 = 100;
let rows: Vec<Value> = (0..ROWS).map(Value::from).collect();
let expr: Arc<dyn PhysicalExpr> = Arc::new(SlowDoubler {
per_row_delay: Duration::from_millis(1),
});
let exec_ctx = root_ctx();
let base = EvalContext::from_exec_ctx(&exec_ctx);
let results = evaluate_rows_spawned(expr, &base, &rows, SpawnedDocRoot::RowValue)
.await
.expect("evaluation should succeed");
let expected: Vec<Value> = (0..ROWS).map(|n| Value::from(n * 2 + 1)).collect();
assert_eq!(
results.len(),
rows.len(),
"every row must produce a result, including the rows the timing probe ran"
);
assert_eq!(
results, expected,
"each result must correspond to the input row that produced it, in input order"
);
}
}