use std::collections::HashSet;
use std::ops::Bound;
use crate::dbs::capabilities::Capabilities;
use crate::err::{EngineError, Error};
use crate::exec::Error as ExecError;
use crate::exec::function::FunctionRegistry;
use crate::expr::visit::{MutVisitor, VisitMut};
use crate::expr::{BinaryOperator, Cond, Expr};
use crate::val::{Number, RecordIdKey, Value};
pub(crate) fn substitution_is_lossy(value: &Value) -> bool {
fn bound_is_lossy(bound: &Bound<Value>) -> bool {
match bound {
Bound::Included(v) | Bound::Excluded(v) => substitution_is_lossy(v),
Bound::Unbounded => false,
}
}
fn key_bound_is_lossy(bound: &Bound<RecordIdKey>) -> bool {
match bound {
Bound::Included(k) | Bound::Excluded(k) => key_is_lossy(k),
Bound::Unbounded => false,
}
}
fn key_is_lossy(key: &RecordIdKey) -> bool {
match key {
RecordIdKey::Number(_) | RecordIdKey::String(_) | RecordIdKey::Uuid(_) => false,
RecordIdKey::Array(a) => a.iter().any(substitution_is_lossy),
RecordIdKey::Object(o) => o.values().any(substitution_is_lossy),
RecordIdKey::Range(r) => key_bound_is_lossy(&r.start) || key_bound_is_lossy(&r.end),
}
}
match value {
Value::Closure(_) | Value::Set(_) => true,
Value::Array(a) => a.iter().any(substitution_is_lossy),
Value::Object(o) => o.values().any(substitution_is_lossy),
Value::Range(r) => bound_is_lossy(&r.start) || bound_is_lossy(&r.end),
Value::RecordId(rid) => key_is_lossy(&rid.key),
_ => false,
}
}
pub(crate) fn try_literal_to_value(
lit: &crate::expr::literal::Literal,
) -> Option<crate::val::Value> {
use crate::expr::literal::Literal;
use crate::val::Value;
match lit {
Literal::None => Some(Value::None),
Literal::Null => Some(Value::Null),
Literal::Bool(x) => Some(Value::Bool(*x)),
Literal::Float(x) => Some(Value::Number(Number::Float(*x))),
Literal::Integer(i) => Some(Value::Number(Number::Int(*i))),
Literal::Decimal(d) => Some(Value::Number(Number::Decimal(*d))),
Literal::String(s) => Some(Value::String(s.clone())),
Literal::Uuid(u) => Some(Value::Uuid(*u)),
Literal::Datetime(dt) => Some(Value::Datetime(*dt)),
Literal::Duration(d) => Some(Value::Duration(*d)),
Literal::RecordId(rid) => {
use crate::expr::RecordIdKeyLit;
let key = match &rid.key {
RecordIdKeyLit::Number(n) => crate::val::RecordIdKey::Number(*n),
RecordIdKeyLit::String(s) => crate::val::RecordIdKey::String(s.clone()),
RecordIdKeyLit::Uuid(u) => crate::val::RecordIdKey::Uuid(*u),
_ => return None,
};
Some(Value::RecordId(crate::val::RecordId::new(rid.table.clone(), key)))
}
Literal::Array(arr) => {
let values: Option<Vec<Value>> = arr.iter().map(try_expr_to_value).collect();
values.map(|v| Value::Array(v.into()))
}
Literal::Bytes(_)
| Literal::Regex(_)
| Literal::Geometry(_)
| Literal::File(_)
| Literal::Object(_)
| Literal::Set(_)
| Literal::UnboundedRange => None,
}
}
pub(crate) fn try_expr_to_value(expr: &Expr) -> Option<crate::val::Value> {
match expr {
Expr::Literal(lit) => try_literal_to_value(lit),
Expr::Binary {
left,
op:
op @ (BinaryOperator::Range
| BinaryOperator::RangeInclusive
| BinaryOperator::RangeSkip
| BinaryOperator::RangeSkipInclusive),
right,
} => {
let start = try_expr_to_value(left)?;
let end = try_expr_to_value(right)?;
let start = match op {
BinaryOperator::Range | BinaryOperator::RangeInclusive => Bound::Included(start),
_ => Bound::Excluded(start),
};
let end = match op {
BinaryOperator::Range | BinaryOperator::RangeSkip => Bound::Excluded(end),
_ => Bound::Included(end),
};
Some(crate::val::Value::Range(Box::new(crate::val::Range {
start,
end,
})))
}
_ => None,
}
}
pub(crate) fn fold_condition_expressions(
cond: &mut Cond,
registry: &FunctionRegistry,
capabilities: &Capabilities,
restricted_fields: Option<&HashSet<String>>,
) {
let mut folder = ExpressionFolder {
registry,
capabilities,
restricted_fields,
};
let _ = folder.visit_mut_expr(&mut cond.0);
}
pub(crate) fn source_field_reads_are_inert(what: &[Expr]) -> bool {
use crate::expr::literal::Literal;
what.iter().all(|e| match e {
Expr::Table(_) => true,
Expr::Constant(_) => true,
Expr::Literal(Literal::Object(_)) => true,
Expr::Literal(lit) => matches!(
lit,
Literal::None
| Literal::Null
| Literal::Bool(_)
| Literal::Float(_)
| Literal::Integer(_)
| Literal::Decimal(_)
| Literal::String(_)
| Literal::Bytes(_)
| Literal::Regex(_)
| Literal::Duration(_)
| Literal::Datetime(_)
| Literal::Uuid(_)
| Literal::Geometry(_)
| Literal::File(_)
),
_ => false,
})
}
struct ExpressionFolder<'a> {
registry: &'a FunctionRegistry,
capabilities: &'a Capabilities,
restricted_fields: Option<&'a HashSet<String>>,
}
impl MutVisitor for ExpressionFolder<'_> {
type Error = std::convert::Infallible;
fn visit_mut_expr(&mut self, expr: &mut Expr) -> Result<(), Self::Error> {
expr.visit_mut(self)?;
if let Some(folded) =
try_fold_to_literal(expr, self.registry, self.capabilities, self.restricted_fields)
{
*expr = folded;
}
Ok(())
}
fn visit_mut_select(
&mut self,
_: &mut crate::expr::SelectStatement,
) -> Result<(), Self::Error> {
Ok(())
}
}
fn evaluation_is_inert(expr: &Expr, restricted_fields: Option<&HashSet<String>>) -> bool {
use crate::expr::part::Part;
fn operator_is_inert(op: &BinaryOperator) -> bool {
match op {
BinaryOperator::Subtract
| BinaryOperator::Add
| BinaryOperator::Multiply
| BinaryOperator::Divide
| BinaryOperator::Remainder
| BinaryOperator::Power
| BinaryOperator::Equal
| BinaryOperator::ExactEqual
| BinaryOperator::NotEqual
| BinaryOperator::AllEqual
| BinaryOperator::AnyEqual
| BinaryOperator::Or
| BinaryOperator::And
| BinaryOperator::NullCoalescing
| BinaryOperator::TenaryCondition
| BinaryOperator::LessThan
| BinaryOperator::LessThanEqual
| BinaryOperator::MoreThan
| BinaryOperator::MoreThanEqual
| BinaryOperator::Contain
| BinaryOperator::NotContain
| BinaryOperator::Inside
| BinaryOperator::NotInside
| BinaryOperator::ContainAll
| BinaryOperator::ContainAny
| BinaryOperator::ContainNone
| BinaryOperator::AllInside
| BinaryOperator::AnyInside
| BinaryOperator::NoneInside
| BinaryOperator::Outside
| BinaryOperator::Intersects
| BinaryOperator::Range
| BinaryOperator::RangeInclusive
| BinaryOperator::RangeSkip
| BinaryOperator::RangeSkipInclusive => true,
BinaryOperator::Matches(_) | BinaryOperator::NearestNeighbor(_) => false,
}
}
match expr {
Expr::Literal(lit) => literal_is_inert(lit, restricted_fields),
Expr::Table(_) | Expr::Mock(_) | Expr::Constant(_) => true,
Expr::Param(_) => false,
Expr::Idiom(idiom) => match (restricted_fields, idiom.0.as_slice()) {
(Some(restricted), [Part::Field(name)]) => !restricted.contains(name.as_str()),
_ => false,
},
Expr::Prefix {
expr,
..
} => evaluation_is_inert(expr, restricted_fields),
Expr::Binary {
left,
op,
right,
} => {
operator_is_inert(op)
&& evaluation_is_inert(left, restricted_fields)
&& evaluation_is_inert(right, restricted_fields)
}
Expr::Postfix {
..
}
| Expr::FunctionCall(_)
| Expr::Closure(_)
| Expr::Block(_)
| Expr::Break
| Expr::Continue
| Expr::Return(_)
| Expr::Throw(_)
| Expr::IfElse(_)
| Expr::Select(_)
| Expr::Create(_)
| Expr::Update(_)
| Expr::Upsert(_)
| Expr::Delete(_)
| Expr::Relate(_)
| Expr::Insert(_)
| Expr::Define(_)
| Expr::Remove(_)
| Expr::Rebuild(_)
| Expr::Alter(_)
| Expr::Info(_)
| Expr::Foreach(_)
| Expr::Let(_)
| Expr::Sleep(_)
| Expr::Explain {
..
}
| Expr::Match(_) => false,
}
}
fn literal_is_inert(
lit: &crate::expr::literal::Literal,
restricted_fields: Option<&HashSet<String>>,
) -> bool {
use crate::expr::literal::Literal;
match lit {
Literal::None
| Literal::Null
| Literal::Bool(_)
| Literal::Float(_)
| Literal::Integer(_)
| Literal::Decimal(_)
| Literal::String(_)
| Literal::Uuid(_)
| Literal::Datetime(_)
| Literal::Duration(_)
| Literal::Bytes(_)
| Literal::Regex(_)
| Literal::Geometry(_)
| Literal::File(_)
| Literal::UnboundedRange => true,
Literal::Array(exprs) | Literal::Set(exprs) => {
exprs.iter().all(|e| evaluation_is_inert(e, restricted_fields))
}
Literal::Object(entries) => {
entries.iter().all(|e| evaluation_is_inert(&e.value, restricted_fields))
}
Literal::RecordId(rid) => record_id_key_is_inert(&rid.key, restricted_fields),
}
}
fn try_fold_to_literal(
expr: &Expr,
registry: &FunctionRegistry,
capabilities: &Capabilities,
restricted_fields: Option<&HashSet<String>>,
) -> Option<Expr> {
use crate::expr::Function;
use crate::val::{Datetime, Value};
match expr {
Expr::FunctionCall(fc)
if matches!(&fc.receiver, Function::Normal(name) if name == "time::now")
&& fc.arguments.is_empty()
&& capabilities.allows_function_name("time::now") =>
{
Some(Value::Datetime(Datetime::now()).into_literal())
}
Expr::FunctionCall(fc) => {
let Function::Normal(name) = &fc.receiver else {
return None;
};
if !capabilities.allows_function_name(name.as_str()) {
return None;
}
if let Some(target) = crate::exec::function::experimental_target(name.as_str())
&& !capabilities.allows_experimental(&target)
{
return None;
}
let func = registry.get(name.as_str())?;
if !func.is_pure() || func.is_async() || !func.is_deterministic() {
return None;
}
let args: Option<Vec<Value>> = fc.arguments.iter().map(try_expr_to_value).collect();
let args = args?;
let result = func.invoke(args).ok()?;
Some(result.into_literal())
}
Expr::Binary {
op: BinaryOperator::Inside,
left,
right,
} if is_empty_array_literal(right) && evaluation_is_inert(left, restricted_fields) => {
Some(Expr::Literal(crate::expr::literal::Literal::Bool(false)))
}
Expr::Binary {
left,
op,
right,
} => {
if let Some(short) = try_short_circuit_bool(left, op, right, restricted_fields) {
return Some(short);
}
let left_val = try_expr_to_value(left)?;
let right_val = try_expr_to_value(right)?;
let result = try_eval_binary(op, left_val, right_val)?;
Some(result.into_literal())
}
Expr::Idiom(idiom) => {
let mut parts = idiom.0.iter();
let Some(crate::expr::Part::Start(Expr::Literal(root))) = parts.next() else {
return None;
};
let mut value = literal_root_value(root)?;
let mut walked = false;
for part in parts {
value = apply_literal_part(value, part)?;
walked = true;
}
if !walked || substitution_is_lossy(&value) {
return None;
}
Some(value.into_literal())
}
_ => None,
}
}
fn literal_root_value(lit: &crate::expr::literal::Literal) -> Option<Value> {
use crate::expr::literal::Literal;
let expr_value = |expr: &Expr| match expr {
Expr::Literal(lit) => literal_root_value(lit),
other => try_expr_to_value(other),
};
match lit {
Literal::Object(entries) => {
let mut obj = crate::val::Object::default();
for entry in entries {
obj.insert(entry.key.clone(), expr_value(&entry.value)?);
}
Some(Value::Object(obj))
}
Literal::Array(items) => {
let values: Option<Vec<Value>> = items.iter().map(expr_value).collect();
values.map(|v| Value::Array(v.into()))
}
Literal::Geometry(geo) => Some(Value::Geometry(geo.clone())),
other => try_literal_to_value(other),
}
}
pub(crate) fn apply_literal_part(value: Value, part: &crate::expr::Part) -> Option<Value> {
use crate::expr::Part;
match part {
Part::Field(name) => match value {
Value::Object(obj) => Some(obj.get(name.as_str()).cloned().unwrap_or(Value::None)),
Value::Geometry(geo) => {
Some(geo.as_object().get(name.as_str()).cloned().unwrap_or(Value::None))
}
Value::RecordId(_) | Value::Array(_) => None,
_ => Some(Value::None),
},
Part::Value(index) => {
let index = try_expr_to_value(index)?;
crate::exec::parts::index::evaluate_index(&value, &index).ok()
}
Part::First => Some(match value {
Value::Array(arr) => arr.first().cloned().unwrap_or(Value::None),
Value::Set(set) => set.first().cloned().unwrap_or(Value::None),
other => other,
}),
Part::Last => Some(match value {
Value::Array(arr) => arr.last().cloned().unwrap_or(Value::None),
Value::Set(set) => set.last().cloned().unwrap_or(Value::None),
other => other,
}),
_ => None,
}
}
fn is_empty_array_literal(expr: &Expr) -> bool {
matches!(expr, Expr::Literal(crate::expr::literal::Literal::Array(arr)) if arr.is_empty())
}
fn is_boolean_valued(expr: &Expr) -> bool {
use crate::expr::literal::Literal;
match expr {
Expr::Literal(Literal::Bool(_)) => true,
Expr::Prefix {
op: crate::expr::PrefixOperator::Not,
..
} => true,
Expr::Binary {
op,
..
} => match op {
BinaryOperator::Equal
| BinaryOperator::ExactEqual
| BinaryOperator::NotEqual
| BinaryOperator::AllEqual
| BinaryOperator::AnyEqual
| BinaryOperator::LessThan
| BinaryOperator::LessThanEqual
| BinaryOperator::MoreThan
| BinaryOperator::MoreThanEqual
| BinaryOperator::Contain
| BinaryOperator::NotContain
| BinaryOperator::ContainAll
| BinaryOperator::ContainAny
| BinaryOperator::ContainNone
| BinaryOperator::Inside
| BinaryOperator::NotInside
| BinaryOperator::AllInside
| BinaryOperator::AnyInside
| BinaryOperator::NoneInside
| BinaryOperator::Outside
| BinaryOperator::Intersects
| BinaryOperator::Matches(_) => true,
BinaryOperator::And
| BinaryOperator::Or
| BinaryOperator::NullCoalescing
| BinaryOperator::TenaryCondition
| BinaryOperator::Subtract
| BinaryOperator::Add
| BinaryOperator::Multiply
| BinaryOperator::Divide
| BinaryOperator::Remainder
| BinaryOperator::Power
| BinaryOperator::Range
| BinaryOperator::RangeInclusive
| BinaryOperator::RangeSkip
| BinaryOperator::RangeSkipInclusive
| BinaryOperator::NearestNeighbor(_) => false,
},
_ => false,
}
}
fn try_short_circuit_bool(
left: &Expr,
op: &BinaryOperator,
right: &Expr,
restricted_fields: Option<&HashSet<String>>,
) -> Option<Expr> {
use crate::expr::literal::Literal;
let l_bool = match left {
Expr::Literal(Literal::Bool(b)) => Some(*b),
_ => None,
};
let r_bool = match right {
Expr::Literal(Literal::Bool(b)) => Some(*b),
_ => None,
};
match (op, l_bool, r_bool) {
(BinaryOperator::And, Some(false), _) => Some(Expr::Literal(Literal::Bool(false))),
(BinaryOperator::Or, Some(true), _) => Some(Expr::Literal(Literal::Bool(true))),
(BinaryOperator::And, _, Some(false))
if is_boolean_valued(left) && evaluation_is_inert(left, restricted_fields) =>
{
Some(Expr::Literal(Literal::Bool(false)))
}
(BinaryOperator::Or, _, Some(true))
if is_boolean_valued(left) && evaluation_is_inert(left, restricted_fields) =>
{
Some(Expr::Literal(Literal::Bool(true)))
}
(BinaryOperator::And, Some(true), _) => Some(right.clone()),
(BinaryOperator::Or, Some(false), _) => Some(right.clone()),
(BinaryOperator::And, _, Some(true)) if is_boolean_valued(left) => Some(left.clone()),
(BinaryOperator::Or, _, Some(false)) if is_boolean_valued(left) => Some(left.clone()),
_ => None,
}
}
fn try_eval_binary(
op: &BinaryOperator,
left: crate::val::Value,
right: crate::val::Value,
) -> Option<crate::val::Value> {
use crate::val::{TryAdd, TrySub};
match op {
BinaryOperator::Add => left.try_add(right).ok(),
BinaryOperator::Subtract => left.try_sub(right).ok(),
_ => None,
}
}
fn record_id_key_is_inert(
key: &crate::expr::RecordIdKeyLit,
restricted_fields: Option<&HashSet<String>>,
) -> bool {
use std::ops::Bound;
use crate::expr::RecordIdKeyLit;
fn bound_is_inert(
bound: &Bound<RecordIdKeyLit>,
restricted_fields: Option<&HashSet<String>>,
) -> bool {
match bound {
Bound::Included(k) | Bound::Excluded(k) => record_id_key_is_inert(k, restricted_fields),
Bound::Unbounded => true,
}
}
match key {
RecordIdKeyLit::Number(_) | RecordIdKeyLit::String(_) | RecordIdKeyLit::Uuid(_) => true,
RecordIdKeyLit::Array(exprs) => {
exprs.iter().all(|e| evaluation_is_inert(e, restricted_fields))
}
RecordIdKeyLit::Object(entries) => {
entries.iter().all(|e| evaluation_is_inert(&e.value, restricted_fields))
}
RecordIdKeyLit::Range(range) => {
bound_is_inert(&range.start, restricted_fields)
&& bound_is_inert(&range.end, restricted_fields)
}
RecordIdKeyLit::Generate(_) => false,
}
}
pub(crate) fn literal_to_value(
lit: crate::expr::literal::Literal,
) -> Result<crate::val::Value, Error> {
use crate::expr::literal::Literal;
use crate::val::{Range, Value};
if let Some(value) = try_literal_to_value(&lit) {
return Ok(value);
}
match lit {
Literal::UnboundedRange => Ok(Value::Range(Box::new(Range::unbounded()))),
Literal::Bytes(b) => Ok(Value::Bytes(b)),
Literal::Regex(r) => Ok(Value::Regex(r)),
Literal::Geometry(g) => Ok(Value::Geometry(g)),
Literal::File(f) => Ok(Value::File(f)),
other => Err(EngineError::Internal(format!(
"Literal should be handled upstream in physical_expr(): {:?}",
std::mem::discriminant(&other)
))
.into()),
}
}
pub(crate) fn key_lit_to_expr(lit: &crate::expr::RecordIdKeyLit) -> Result<Expr, Error> {
use crate::expr::RecordIdKeyLit;
match lit {
RecordIdKeyLit::Number(n) => Ok(Expr::Literal(crate::expr::literal::Literal::Integer(*n))),
RecordIdKeyLit::String(s) => {
Ok(Expr::Literal(crate::expr::literal::Literal::String(s.clone())))
}
RecordIdKeyLit::Uuid(u) => Ok(Expr::Literal(crate::expr::literal::Literal::Uuid(*u))),
RecordIdKeyLit::Array(exprs) => {
Ok(Expr::Literal(crate::expr::literal::Literal::Array(exprs.clone())))
}
RecordIdKeyLit::Object(entries) => {
Ok(Expr::Literal(crate::expr::literal::Literal::Object(entries.clone())))
}
RecordIdKeyLit::Generate(_) => Err(ExecError::Query {
message: "Generated keys (rand, ulid, uuid) cannot be used in graph range bounds"
.to_string(),
}
.into()),
RecordIdKeyLit::Range(_) => Err(ExecError::Query {
message: "Nested range keys cannot be used in graph range bounds".to_string(),
}
.into()),
}
}
#[cfg(test)]
mod tests {
use surrealdb_types::ToSql;
use super::*;
fn fold(snippet: &str) -> String {
let src = format!("SELECT * FROM t WHERE {snippet}");
let mut exprs = crate::syn::parse(&src).expect("parse").expressions;
let crate::expr::TopLevelExpr::Expr(Expr::Select(s)) = exprs.remove(0).into() else {
panic!("expected a SELECT");
};
let mut cond = s.cond.expect("WHERE");
let registry = FunctionRegistry::with_builtins();
fold_condition_expressions(&mut cond, ®istry, &Capabilities::all(), None);
cond.0.to_sql()
}
#[test]
fn an_idiom_over_a_literal_folds_to_its_value() {
assert_eq!(fold("a = { a: { b: 'x' } }.a.b"), "a = 'x'");
assert_eq!(fold("a = { a: ['x', 'y'] }.a[1]"), "a = 'y'");
assert_eq!(fold("a = ['x', 'y'][0]"), "a = 'x'");
assert_eq!(fold("a = ['x', 'y'][$]"), "a = 'y'");
assert_eq!(fold("a = ['x', 'y'][5]"), "a = NONE");
assert_eq!(fold("a = { v: 'alice' }.v[$]"), "a = 'alice'");
assert_eq!(fold("a = { v: 'alice' }.v[0]"), "a = NONE");
assert_eq!(fold("a = { a: 1 }.b"), "a = NONE");
assert_eq!(fold("a = [{ a: 'x' }][0].a"), "a = 'x'");
assert_eq!(fold("a = 'x'.a"), "a = NONE");
assert_eq!(fold("a = (1.0, 2.0).type"), "a = 'Point'");
}
#[test]
fn an_idiom_the_runtime_maps_or_fetches_is_left_whole() {
assert_eq!(fold("a = [{ a: 'x' }].a"), "a = [{ a: 'x' }].a");
assert_eq!(fold("a = t:1.a"), "a = (t:1).a");
assert_eq!(fold("a = ['x', 'y'][b]"), "a = ['x', 'y'][b]");
}
}