use std::sync::Arc;
use futures::stream;
use surrealdb_strand::Strand;
use surrealdb_types::{SqlFormat, ToSql};
use crate::err::Error;
use crate::exec::context::{ContextLevel, ExecutionContext};
use crate::exec::plan_or_compute::collect_stream;
use crate::exec::{
AccessMode, BoxFut, CardinalityHint, Error as ExecError, ExecOperator, FlowResult,
OperatorMetrics, OutputShape, ValueBatchStream, buffer_stream,
};
use crate::expr::{ControlFlow, Kind};
use crate::val::{Array, Value};
#[derive(Debug)]
pub struct LetPlan {
pub name: Strand,
pub kind: Option<Kind>,
pub(crate) metrics: Arc<OperatorMetrics>,
pub value: Arc<dyn ExecOperator>,
}
impl LetPlan {
pub(crate) fn new(name: Strand, kind: Option<Kind>, value: Arc<dyn ExecOperator>) -> Self {
Self {
name,
kind,
value,
metrics: Arc::new(OperatorMetrics::new()),
}
}
fn coerce(&self, value: Value) -> Result<Value, Error> {
match &self.kind {
Some(kind) => value.coerce_to_kind(kind).map_err(|e| {
ExecError::SetCoerce {
name: self.name.to_string(),
error: Box::new(e),
}
.into()
}),
None => Ok(value),
}
}
async fn compute_value(&self, input: &ExecutionContext) -> crate::expr::FlowResult<Value> {
let stream = buffer_stream(
self.value.execute(input)?,
self.value.access_mode(),
self.value.cardinality_hint(),
input.root().ctx.config.exec.operator_buffer_size,
);
let results = collect_stream(stream).await?;
Ok(if self.value.output_shape().is_scalar() {
results.into_iter().next().unwrap_or(Value::None)
} else {
Value::Array(Array(results))
})
}
}
impl ExecOperator for LetPlan {
fn name(&self) -> &'static str {
"Let"
}
fn attrs(&self) -> Vec<(String, String)> {
vec![("name".to_string(), format!("${}", self.name.as_str()))]
}
fn required_context(&self) -> ContextLevel {
self.value.required_context()
}
fn access_mode(&self) -> AccessMode {
self.value.access_mode()
}
fn cardinality_hint(&self) -> CardinalityHint {
CardinalityHint::AtMostOne
}
fn execute(&self, _ctx: &ExecutionContext) -> FlowResult<ValueBatchStream> {
Ok(Box::pin(stream::once(async { Ok(crate::exec::ValueBatch::new(vec![Value::None])) })))
}
fn mutates_context(&self) -> bool {
true
}
fn output_context<'a>(
&'a self,
input: &'a ExecutionContext,
) -> BoxFut<'a, crate::expr::FlowResult<ExecutionContext>> {
Box::pin(async move {
let computed_value = match self.compute_value(input).await {
Ok(v) => v,
Err(ControlFlow::Return(v)) => v,
Err(ctrl) => return Err(ctrl),
};
let coerced = self.coerce(computed_value).map_err(|e| ControlFlow::Err(e.into()))?;
Ok(input.with_param(self.name.clone(), coerced))
})
}
fn children(&self) -> Vec<&Arc<dyn ExecOperator>> {
vec![&self.value]
}
fn metrics(&self) -> Option<&OperatorMetrics> {
Some(&self.metrics)
}
fn output_shape(&self) -> OutputShape {
OutputShape::Scalar
}
}
impl ToSql for LetPlan {
fn fmt_sql(&self, f: &mut String, _fmt: SqlFormat) {
f.push_str("LET $");
f.push_str(self.name.as_str());
f.push_str(" = ");
if self.value.output_shape().is_scalar() {
f.push_str("<expr>");
} else {
f.push_str("(<query>)");
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::exec::operators::test_util::root_ctx;
use crate::exec::{OutputOrdering, OutputShape, ValueBatch};
#[derive(Debug, Clone, Copy)]
enum Signal {
Break,
Continue,
Return(i64),
Throw(&'static str),
}
impl Signal {
fn build(self) -> ControlFlow {
match self {
Signal::Break => ControlFlow::Break,
Signal::Continue => ControlFlow::Continue,
Signal::Return(v) => ControlFlow::Return(Value::from(v)),
Signal::Throw(msg) => {
ControlFlow::Err(anyhow::Error::new(crate::exec::Error::Thrown(msg.to_owned())))
}
}
}
}
#[derive(Debug)]
struct StubValue {
rows: Vec<Value>,
signal: Option<Signal>,
eager: bool,
scalar: bool,
access_mode: AccessMode,
required_context: ContextLevel,
}
impl StubValue {
fn rows(rows: Vec<Value>) -> Self {
Self {
rows,
signal: None,
eager: false,
scalar: true,
access_mode: AccessMode::ReadOnly,
required_context: ContextLevel::Root,
}
}
fn eager(signal: Signal) -> Self {
Self {
rows: Vec::new(),
signal: Some(signal),
eager: true,
scalar: true,
access_mode: AccessMode::ReadOnly,
required_context: ContextLevel::Root,
}
}
fn rows_then(rows: Vec<Value>, signal: Signal) -> Self {
Self {
rows,
signal: Some(signal),
eager: false,
scalar: true,
access_mode: AccessMode::ReadOnly,
required_context: ContextLevel::Root,
}
}
fn non_scalar(mut self) -> Self {
self.scalar = false;
self
}
fn metadata(mut self, access_mode: AccessMode, required_context: ContextLevel) -> Self {
self.access_mode = access_mode;
self.required_context = required_context;
self
}
fn into_operator(self) -> Arc<dyn ExecOperator> {
Arc::new(self)
}
}
impl ExecOperator for StubValue {
fn name(&self) -> &'static str {
"StubValue"
}
fn required_context(&self) -> ContextLevel {
self.required_context
}
fn access_mode(&self) -> AccessMode {
self.access_mode
}
fn cardinality_hint(&self) -> CardinalityHint {
CardinalityHint::Unbounded
}
fn output_ordering(&self) -> OutputOrdering {
OutputOrdering::Unordered
}
fn output_shape(&self) -> OutputShape {
if self.scalar {
OutputShape::Scalar
} else {
OutputShape::Rows
}
}
fn execute(&self, _ctx: &ExecutionContext) -> FlowResult<ValueBatchStream> {
if let Some(signal) = self.signal.filter(|_| self.eager) {
return Err(signal.build());
}
let mut items: Vec<FlowResult<ValueBatch>> = Vec::new();
if !self.rows.is_empty() {
items.push(Ok(ValueBatch::new(self.rows.clone())));
}
if let Some(signal) = self.signal {
items.push(Err(signal.build()));
}
Ok(Box::pin(stream::iter(items)))
}
}
fn let_plan(value: Arc<dyn ExecOperator>) -> LetPlan {
LetPlan::new(Strand::new("x"), None, value)
}
async fn bound_value(plan: &LetPlan) -> crate::expr::FlowResult<Value> {
let ctx = root_ctx();
let out = plan.output_context(&ctx).await?;
Ok(out.value("x").cloned().unwrap_or(Value::None))
}
async fn propagated_signal(plan: &LetPlan) -> ControlFlow {
match bound_value(plan).await {
Ok(v) => panic!("expected a control-flow signal, got the binding: {v:?}"),
Err(ctrl) => ctrl,
}
}
#[tokio::test]
async fn a_scalar_value_plan_binds_its_single_row() {
let plan = let_plan(StubValue::rows(vec![Value::from(7)]).into_operator());
assert_eq!(bound_value(&plan).await.unwrap(), Value::from(7));
}
#[tokio::test]
async fn an_empty_scalar_value_plan_binds_none() {
let plan = let_plan(StubValue::rows(Vec::new()).into_operator());
assert_eq!(bound_value(&plan).await.unwrap(), Value::None);
}
#[tokio::test]
async fn a_non_scalar_value_plan_binds_all_rows_as_an_array() {
let plan = let_plan(
StubValue::rows(vec![Value::from(1), Value::from(2)]).non_scalar().into_operator(),
);
assert_eq!(
bound_value(&plan).await.unwrap(),
Value::Array(Array(vec![Value::from(1), Value::from(2)]))
);
}
#[tokio::test]
async fn a_loop_signal_raised_by_the_value_plans_execute_travels_onwards() {
let plan = let_plan(StubValue::eager(Signal::Break).into_operator());
assert!(matches!(propagated_signal(&plan).await, ControlFlow::Break));
let plan = let_plan(StubValue::eager(Signal::Continue).into_operator());
assert!(matches!(propagated_signal(&plan).await, ControlFlow::Continue));
}
#[tokio::test]
async fn a_loop_signal_from_inside_the_value_stream_travels_onwards_too() {
let rows = vec![Value::from(1)];
let plan = let_plan(StubValue::rows_then(rows.clone(), Signal::Break).into_operator());
assert!(matches!(propagated_signal(&plan).await, ControlFlow::Break));
let plan = let_plan(StubValue::rows_then(rows, Signal::Continue).into_operator());
assert!(matches!(propagated_signal(&plan).await, ControlFlow::Continue));
}
#[tokio::test]
async fn a_return_raised_by_the_value_plans_execute_supplies_the_bound_value() {
let plan = let_plan(StubValue::eager(Signal::Return(9)).into_operator());
assert_eq!(bound_value(&plan).await.unwrap(), Value::from(9));
}
#[tokio::test]
async fn a_return_from_inside_the_value_stream_wins_over_the_rows_before_it() {
let plan = let_plan(
StubValue::rows_then(vec![Value::from(1), Value::from(2)], Signal::Return(9))
.into_operator(),
);
assert_eq!(bound_value(&plan).await.unwrap(), Value::from(9));
}
#[tokio::test]
async fn a_declared_type_coerces_the_bound_value() {
let plan = LetPlan::new(
Strand::new("x"),
Some(Kind::String),
StubValue::rows(vec![Value::from(7)]).into_operator(),
);
let err = match bound_value(&plan).await {
Err(ControlFlow::Err(e)) => e,
other => panic!("expected a coercion failure, got: {other:?}"),
};
assert!(
format!("{err}").contains("$x"),
"coercion failure should name the parameter, got: {err}"
);
}
#[tokio::test]
async fn the_binding_is_published_through_output_context_not_through_execute() {
let ctx = root_ctx();
let plan = let_plan(StubValue::rows(vec![Value::from(3i64)]).into_operator());
assert!(plan.mutates_context());
assert_eq!(bound_value(&plan).await.expect("binding should succeed"), Value::from(3i64));
assert!(ctx.value("x").is_none());
}
#[tokio::test]
async fn execute_emits_a_single_none_row_whatever_the_bound_value_is() {
let ctx = root_ctx();
let plan: Arc<dyn ExecOperator> =
Arc::new(let_plan(StubValue::rows(vec![Value::from(3i64)]).into_operator()));
assert_eq!(
crate::exec::operators::test_util::collect(&plan, &ctx).await,
vec![Value::None],
"LET is not an expression: its own output is always NONE"
);
}
#[tokio::test]
async fn output_shape_is_scalar_so_a_lone_let_statement_reduces_to_bare_none() {
let plan = let_plan(StubValue::rows(vec![Value::from(3i64)]).into_operator());
assert_eq!(plan.output_shape(), OutputShape::Scalar);
}
#[tokio::test]
async fn a_binding_shadows_an_outer_one_of_the_same_name_and_the_outer_stays_intact() {
let outer = root_ctx().with_param("x", Value::from(1i64));
let plan = let_plan(StubValue::rows(vec![Value::from(2i64)]).into_operator());
let inner = plan.output_context(&outer).await.expect("binding should succeed");
assert_eq!(inner.value("x").cloned(), Some(Value::from(2i64)));
assert_eq!(outer.value("x").cloned(), Some(Value::from(1i64)));
}
#[tokio::test]
async fn metadata_is_inherited_from_the_value_plan() {
let value = StubValue::rows(Vec::new())
.metadata(AccessMode::ReadWrite, ContextLevel::Database)
.into_operator();
let plan = let_plan(value);
assert_eq!(plan.access_mode(), AccessMode::ReadWrite);
assert_eq!(plan.required_context(), ContextLevel::Database);
}
#[tokio::test]
async fn a_value_that_does_not_satisfy_the_declared_type_fails_naming_the_parameter() {
let ctx = root_ctx();
let plan = LetPlan::new(
Strand::new("x"),
Some(Kind::Int),
StubValue::rows(vec![Value::from("not a number")]).into_operator(),
);
let ctrl = plan.output_context(&ctx).await.expect_err("a string cannot coerce to int");
let ControlFlow::Err(err) = ctrl else {
panic!("a coercion failure is an error, not a control-flow signal");
};
assert!(
matches!(
err.downcast_ref::<crate::err::Error>(),
Some(crate::err::Error::Exec(crate::exec::Error::SetCoerce { name, .. }))
if name == "x"
),
"expected SetCoerce naming $x, got {err:?}"
);
}
#[tokio::test]
async fn an_error_reaches_the_caller_downcastable_from_either_raise_site() {
for value in [
StubValue::eager(Signal::Throw("boom")).into_operator(),
StubValue::rows_then(vec![Value::from(1i64)], Signal::Throw("boom")).into_operator(),
] {
let plan = let_plan(value);
let ControlFlow::Err(err) = propagated_signal(&plan).await else {
panic!("expected an error");
};
assert!(
matches!(
err.downcast_ref::<crate::exec::Error>(),
Some(crate::exec::Error::Thrown(msg)) if msg == "boom"
),
"expected the original Thrown error, got {err:?}"
);
}
}
}