use std::collections::{HashMap, HashSet};
use std::sync::Arc;
use surrealdb_types::ToSql;
use super::planner::Planner;
use crate::err::Error;
use crate::exec::PhysicalExpr;
use crate::expr::part::Part;
use crate::expr::{Expr, Idiom};
#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash)]
#[allow(dead_code)] pub enum ComputePoint {
Filter = 0,
Aggregate = 1,
Sort = 2,
Project = 3,
}
#[derive(Debug, Clone)]
pub struct ExpressionInfo {
pub internal_name: String,
pub expr: Arc<dyn PhysicalExpr>,
pub compute_point: ComputePoint,
}
#[derive(Debug, Default, Clone)]
pub struct ExpressionRegistry {
expressions: HashMap<String, ExpressionInfo>,
counter: usize,
reserved_names: HashSet<String>,
used_internal_names: HashSet<String>,
}
impl ExpressionRegistry {
#[cfg(test)]
pub(crate) fn new() -> Self {
Self {
expressions: HashMap::new(),
counter: 0,
reserved_names: HashSet::new(),
used_internal_names: HashSet::new(),
}
}
#[cfg(test)]
pub(crate) fn with_reserved_names(names: Vec<String>) -> Self {
Self::with_reserved_and_protected_names(names, Vec::new())
}
pub fn with_reserved_and_protected_names(
reserved_names: Vec<String>,
protected_names: Vec<String>,
) -> Self {
Self {
expressions: HashMap::new(),
counter: 0,
reserved_names: reserved_names.into_iter().collect(),
used_internal_names: protected_names.into_iter().collect(),
}
}
pub async fn register(
&mut self,
expr: &Expr,
compute_point: ComputePoint,
alias: Option<String>,
planner: &Planner<'_>,
) -> Result<String, Error> {
let expr_sql = expr.to_sql();
let dedup_key = match &alias {
Some(a) => format!("{expr_sql}\0{a}"),
None => expr_sql,
};
if let Some(info) = self.expressions.get(&dedup_key) {
if compute_point < info.compute_point {
let mut updated = info.clone();
updated.compute_point = compute_point;
self.expressions.insert(dedup_key.clone(), updated);
}
return Ok(self.expressions[&dedup_key].internal_name.clone());
}
let physical_expr = planner.physical_expr(expr.clone()).await?;
let internal_name = self.choose_internal_name(&alias);
let info = ExpressionInfo {
internal_name: internal_name.clone(),
expr: physical_expr,
compute_point,
};
self.expressions.insert(dedup_key, info);
Ok(internal_name)
}
#[allow(clippy::needless_pass_by_value)] pub fn register_physical(
&mut self,
expr_key: String,
physical_expr: Arc<dyn PhysicalExpr>,
compute_point: ComputePoint,
alias: Option<String>,
) -> String {
let dedup_key = match &alias {
Some(a) => format!("{expr_key}\0{a}"),
None => expr_key,
};
if let Some(info) = self.expressions.get(&dedup_key) {
if compute_point < info.compute_point {
let mut updated = info.clone();
updated.compute_point = compute_point;
self.expressions.insert(dedup_key.clone(), updated);
}
return self.expressions[&dedup_key].internal_name.clone();
}
let internal_name = self.choose_internal_name(&alias);
let info = ExpressionInfo {
internal_name: internal_name.clone(),
expr: physical_expr,
compute_point,
};
self.expressions.insert(dedup_key, info);
internal_name
}
fn choose_internal_name(&mut self, alias: &Option<String>) -> String {
if let Some(name) = alias
&& !self.used_internal_names.contains(name)
{
self.used_internal_names.insert(name.clone());
return name.clone();
}
loop {
let name = format!("_e{}", self.counter);
self.counter += 1;
if !self.reserved_names.contains(&name) && !self.used_internal_names.contains(&name) {
self.used_internal_names.insert(name.clone());
return name;
}
}
}
pub fn get_expressions_for_point(
&self,
point: ComputePoint,
) -> Vec<(String, Arc<dyn PhysicalExpr>)> {
let mut exprs: Vec<_> = self
.expressions
.values()
.filter(|info| info.compute_point == point)
.map(|info| (info.internal_name.clone(), Arc::clone(&info.expr)))
.collect();
exprs.sort_by(|(a, _), (b, _)| a.cmp(b));
exprs
}
pub fn has_expressions_for_point(&self, point: ComputePoint) -> bool {
self.expressions.values().any(|info| info.compute_point == point)
}
pub fn contains_name(&self, name: &str) -> bool {
self.expressions.values().any(|info| info.internal_name == name)
}
}
use crate::expr::field::{Field, Fields};
pub fn resolve_order_by_alias(order_idiom: &Idiom, fields: &Fields) -> Option<(Expr, String)> {
match fields {
Fields::Value(selector) => {
if let Some(ref alias) = selector.alias
&& alias.len() == 1
&& let Some(Part::Field(alias_name)) = alias.first()
&& order_idiom.len() == 1
&& let Some(Part::Field(order_name)) = order_idiom.first()
&& alias_name == order_name
{
return Some((selector.expr.clone(), alias_name.as_str().to_owned()));
}
if order_idiom.len() == 1
&& let Some(Part::Field(name)) = order_idiom.first()
&& let Expr::Idiom(ref expr_idiom) = selector.expr
&& expr_idiom.len() == 1
&& let Some(Part::Field(expr_name)) = expr_idiom.first()
&& name == expr_name
{
return Some((selector.expr.clone(), name.as_str().to_owned()));
}
None
}
Fields::Select(field_list) => {
for field in field_list {
if let Field::Single(selector) = field {
if let Some(ref alias) = selector.alias
&& *order_idiom == *alias
{
return Some((selector.expr.clone(), idiom_to_flat_name(alias)));
}
if order_idiom.len() == 1
&& let Some(Part::Field(order_name)) = order_idiom.first()
&& selector.alias.is_none()
&& let Expr::Idiom(ref expr_idiom) = selector.expr
&& expr_idiom.len() == 1
&& let Some(Part::Field(name)) = expr_idiom.first()
&& name.as_str() == order_name.as_str()
{
return Some((selector.expr.clone(), order_name.as_str().to_owned()));
}
}
}
None
}
}
}
pub(crate) fn idiom_to_flat_name(idiom: &Idiom) -> String {
if idiom.len() == 1
&& let Some(Part::Field(name)) = idiom.first()
{
return name.as_str().to_owned();
}
idiom.to_sql()
}
#[cfg(test)]
mod tests {
use super::*;
use crate::exec::physical_expr::Literal;
use crate::val::{Number, Value};
fn literal_expr(n: i64) -> Arc<dyn PhysicalExpr> {
Arc::new(Literal(Value::Number(Number::Int(n))))
}
#[test]
fn test_registry_deduplication() {
let mut registry = ExpressionRegistry::new();
let name1 = registry.choose_internal_name(&None);
assert_eq!(name1, "_e0");
let name2 = registry.choose_internal_name(&None);
assert_eq!(name2, "_e1");
let name3 = registry.choose_internal_name(&Some("city_population".into()));
assert_eq!(name3, "city_population");
}
#[test]
fn test_reserved_names_do_not_block_aliases() {
let mut registry =
ExpressionRegistry::with_reserved_names(vec!["id".into(), "name".into()]);
let name1 = registry.choose_internal_name(&Some("id".into()));
assert_eq!(name1, "id");
let name2 = registry.choose_internal_name(&Some("name".into()));
assert_eq!(name2, "name");
let name3 = registry.choose_internal_name(&Some("city".into()));
assert_eq!(name3, "city");
}
#[test]
fn test_reserved_names_block_synthetic_only() {
let mut registry =
ExpressionRegistry::with_reserved_names(vec!["id".into(), "name".into()]);
let name1 = registry.choose_internal_name(&Some("id".into()));
assert_eq!(name1, "id");
let mut registry2 =
ExpressionRegistry::with_reserved_names(vec!["_e0".into(), "_e1".into()]);
let syn1 = registry2.choose_internal_name(&None);
assert_eq!(syn1, "_e2"); }
#[test]
fn test_register_physical_dedup_returns_existing_name() {
let mut registry = ExpressionRegistry::new();
let name1 = registry.register_physical(
"val * 2".into(),
literal_expr(42),
ComputePoint::Sort,
Some("doubled".into()),
);
assert_eq!(name1, "doubled");
let name2 = registry.register_physical(
"val * 2".into(),
literal_expr(99),
ComputePoint::Project,
Some("doubled".into()),
);
assert_eq!(name2, "doubled");
let sort_exprs = registry.get_expressions_for_point(ComputePoint::Sort);
assert_eq!(sort_exprs.len(), 1);
assert_eq!(sort_exprs[0].0, "doubled");
}
#[test]
fn test_register_physical_promotes_compute_point() {
let mut registry = ExpressionRegistry::new();
registry.register_physical(
"complex_expr".into(),
literal_expr(1),
ComputePoint::Project,
Some("total".into()),
);
assert!(registry.has_expressions_for_point(ComputePoint::Project));
assert!(!registry.has_expressions_for_point(ComputePoint::Sort));
registry.register_physical(
"complex_expr".into(),
literal_expr(1),
ComputePoint::Sort,
Some("total".into()),
);
assert!(registry.has_expressions_for_point(ComputePoint::Sort));
assert!(!registry.has_expressions_for_point(ComputePoint::Project));
}
#[test]
fn test_duplicate_alias_gets_synthetic_name() {
let mut registry = ExpressionRegistry::new();
let name1 = registry.register_physical(
"expr_a".into(),
literal_expr(1),
ComputePoint::Project,
Some("total".into()),
);
assert_eq!(name1, "total");
let name2 = registry.register_physical(
"expr_b".into(),
literal_expr(2),
ComputePoint::Project,
Some("total".into()),
);
assert_eq!(name2, "_e0"); }
#[test]
fn test_synthetic_name_skips_reserved() {
let mut registry = ExpressionRegistry::with_reserved_names(vec!["_e0".into()]);
let name = registry.choose_internal_name(&None);
assert_eq!(name, "_e1");
}
#[test]
fn test_synthetic_name_skips_multiple_reserved() {
let mut registry =
ExpressionRegistry::with_reserved_names(vec!["_e0".into(), "_e1".into(), "_e3".into()]);
let name1 = registry.choose_internal_name(&None);
assert_eq!(name1, "_e2");
let name2 = registry.choose_internal_name(&None);
assert_eq!(name2, "_e4");
}
}