use crate::result::QueryError;
use powdb_storage::catalog::Catalog;
use powdb_storage::types::*;
use crate::executor::compiled::*;
use crate::executor::eval::date_unit_micros;
use super::*;
pub(crate) fn predicate_column_indices_json(expr: &Expr, columns: &[String]) -> Vec<usize> {
let mut indices = predicate_column_indices(expr, columns);
collect_json_path_base_indices(expr, columns, &mut indices);
indices.sort_unstable();
indices.dedup();
indices
}
fn collect_json_path_base_indices(expr: &Expr, columns: &[String], out: &mut Vec<usize>) {
match expr {
Expr::JsonPath { base, .. } => {
let name = match base.as_ref() {
Expr::Field(n) => n.clone(),
Expr::QualifiedField { qualifier, field } => format!("{qualifier}.{field}"),
other => {
collect_json_path_base_indices(other, columns, out);
return;
}
};
if let Some(idx) = columns.iter().position(|c| *c == name) {
out.push(idx);
}
}
Expr::BinaryOp(l, _, r) | Expr::Coalesce(l, r) => {
collect_json_path_base_indices(l, columns, out);
collect_json_path_base_indices(r, columns, out);
}
Expr::UnaryOp(_, i) | Expr::FunctionCall(_, i, _) | Expr::Cast(i, _) => {
collect_json_path_base_indices(i, columns, out);
}
Expr::ScalarFunc(_, args) => {
for a in args {
collect_json_path_base_indices(a, columns, out);
}
}
Expr::InList { expr, list, .. } => {
collect_json_path_base_indices(expr, columns, out);
for item in list {
collect_json_path_base_indices(item, columns, out);
}
}
Expr::InSubquery { expr, .. } => collect_json_path_base_indices(expr, columns, out),
Expr::Case { whens, else_expr } => {
for (c, r) in whens {
collect_json_path_base_indices(c, columns, out);
collect_json_path_base_indices(r, columns, out);
}
if let Some(e) = else_expr {
collect_json_path_base_indices(e, columns, out);
}
}
_ => {}
}
}
pub(crate) fn validate_json_path_types(
catalog: &Catalog,
plan: &PlanNode,
) -> Result<(), QueryError> {
let mut scope: Vec<(String, TypeId)> = Vec::new();
collect_scan_columns(catalog, plan, &mut scope);
let mut shadowed: std::collections::HashSet<String> = std::collections::HashSet::new();
collect_projected_names(plan, &mut shadowed);
check_plan_json_paths(plan, &scope, &shadowed)
}
fn collect_scan_columns(catalog: &Catalog, plan: &PlanNode, out: &mut Vec<(String, TypeId)>) {
match plan {
PlanNode::SeqScan { table }
| PlanNode::IndexScan { table, .. }
| PlanNode::RangeScan { table, .. } => {
if let Some(schema) = catalog.schema(table) {
for c in &schema.columns {
out.push((c.name.clone(), c.type_id));
}
}
}
PlanNode::AliasScan { table, alias } => {
if let Some(schema) = catalog.schema(table) {
for c in &schema.columns {
out.push((format!("{alias}.{}", c.name), c.type_id));
}
}
}
PlanNode::Filter { input, .. }
| PlanNode::Project { input, .. }
| PlanNode::Sort { input, .. }
| PlanNode::Limit { input, .. }
| PlanNode::Offset { input, .. }
| PlanNode::Aggregate { input, .. }
| PlanNode::Distinct { input }
| PlanNode::GroupBy { input, .. }
| PlanNode::Window { input, .. }
| PlanNode::Update { input, .. }
| PlanNode::Delete { input, .. }
| PlanNode::Explain { input } => collect_scan_columns(catalog, input, out),
PlanNode::NestedLoopJoin { left, right, .. } | PlanNode::Union { left, right, .. } => {
collect_scan_columns(catalog, left, out);
collect_scan_columns(catalog, right, out);
}
_ => {}
}
}
fn collect_projected_names(plan: &PlanNode, out: &mut std::collections::HashSet<String>) {
if let PlanNode::Project { fields, .. } = plan {
for f in fields {
if let Some(a) = &f.alias {
out.insert(a.clone());
} else {
match &f.expr {
Expr::Field(n) => {
out.insert(n.clone());
}
Expr::QualifiedField { qualifier, field } => {
out.insert(format!("{qualifier}.{field}"));
}
_ => {}
}
}
}
}
match plan {
PlanNode::Filter { input, .. }
| PlanNode::Project { input, .. }
| PlanNode::Sort { input, .. }
| PlanNode::Limit { input, .. }
| PlanNode::Offset { input, .. }
| PlanNode::Aggregate { input, .. }
| PlanNode::Distinct { input }
| PlanNode::GroupBy { input, .. }
| PlanNode::Window { input, .. }
| PlanNode::Update { input, .. }
| PlanNode::Delete { input, .. }
| PlanNode::Explain { input } => collect_projected_names(input, out),
PlanNode::NestedLoopJoin { left, right, .. } | PlanNode::Union { left, right, .. } => {
collect_projected_names(left, out);
collect_projected_names(right, out);
}
_ => {}
}
}
fn resolve_scan_type(name: &str, scope: &[(String, TypeId)]) -> Option<TypeId> {
let exact = resolve_scan_type_by(scope, |n| n == name);
if exact.is_some() || name.contains('.') {
return exact;
}
resolve_scan_type_by(
scope,
|n| matches!(n.split_once('.'), Some((_, field)) if field == name),
)
}
fn resolve_scan_type_by(
scope: &[(String, TypeId)],
matches_name: impl Fn(&str) -> bool,
) -> Option<TypeId> {
let mut found: Option<TypeId> = None;
for (n, t) in scope {
if matches_name(n) {
match found {
None => found = Some(*t),
Some(prev) if prev == *t => {}
Some(_) => return None, }
}
}
found
}
fn json_path_base_error(
base: &Expr,
scope: &[(String, TypeId)],
shadowed: &std::collections::HashSet<String>,
) -> Option<String> {
let name = match base {
Expr::Field(n) => n.clone(),
Expr::QualifiedField { qualifier, field } => format!("{qualifier}.{field}"),
_ => return None,
};
if shadowed.contains(&name) {
return None;
}
match resolve_scan_type(&name, scope) {
Some(TypeId::Json) | None => None,
Some(other) => Some(format!(
"'{}' is a {} column, not json: the '->' path operator requires a json column",
name,
type_id_to_name(other)
)),
}
}
fn check_expr_json_paths(
expr: &Expr,
scope: &[(String, TypeId)],
shadowed: &std::collections::HashSet<String>,
) -> Result<(), QueryError> {
match expr {
Expr::JsonPath { base, .. } => {
if let Some(msg) = json_path_base_error(base, scope, shadowed) {
return Err(QueryError::TypeError(msg));
}
check_expr_json_paths(base, scope, shadowed)
}
Expr::BinaryOp(l, _, r) | Expr::Coalesce(l, r) => {
check_expr_json_paths(l, scope, shadowed)?;
check_expr_json_paths(r, scope, shadowed)
}
Expr::UnaryOp(_, inner) | Expr::FunctionCall(_, inner, _) | Expr::Cast(inner, _) => {
check_expr_json_paths(inner, scope, shadowed)
}
Expr::ScalarFunc(_, args) => {
for a in args {
check_expr_json_paths(a, scope, shadowed)?;
}
Ok(())
}
Expr::Window {
args,
partition_by,
order_by,
..
} => {
for expr in args.iter().chain(partition_by) {
check_expr_json_paths(expr, scope, shadowed)?;
}
for key in order_by {
check_expr_json_paths(&key.expr, scope, shadowed)?;
}
Ok(())
}
Expr::InList { expr, list, .. } => {
check_expr_json_paths(expr, scope, shadowed)?;
for item in list {
check_expr_json_paths(item, scope, shadowed)?;
}
Ok(())
}
Expr::Case { whens, else_expr } => {
for (c, r) in whens {
check_expr_json_paths(c, scope, shadowed)?;
check_expr_json_paths(r, scope, shadowed)?;
}
if let Some(e) = else_expr {
check_expr_json_paths(e, scope, shadowed)?;
}
Ok(())
}
Expr::InSubquery { expr, .. } => check_expr_json_paths(expr, scope, shadowed),
_ => Ok(()),
}
}
fn check_plan_json_paths(
plan: &PlanNode,
scope: &[(String, TypeId)],
shadowed: &std::collections::HashSet<String>,
) -> Result<(), QueryError> {
match plan {
PlanNode::Filter { input, predicate } => {
check_expr_json_paths(predicate, scope, shadowed)?;
check_plan_json_paths(input, scope, shadowed)
}
PlanNode::Project { input, fields } => {
for f in fields {
check_expr_json_paths(&f.expr, scope, shadowed)?;
}
check_plan_json_paths(input, scope, shadowed)
}
PlanNode::GroupBy {
input,
keys,
aggregates,
having,
} => {
for key in keys {
check_expr_json_paths(&key.expr, scope, shadowed)?;
}
for aggregate in aggregates {
check_expr_json_paths(&aggregate.argument, scope, shadowed)?;
}
if let Some(h) = having {
check_expr_json_paths(h, scope, shadowed)?;
}
check_plan_json_paths(input, scope, shadowed)
}
PlanNode::NestedLoopJoin {
left, right, on, ..
} => {
if let Some(on) = on {
check_expr_json_paths(on, scope, shadowed)?;
}
check_plan_json_paths(left, scope, shadowed)?;
check_plan_json_paths(right, scope, shadowed)
}
PlanNode::Union { left, right, .. } => {
check_plan_json_paths(left, scope, shadowed)?;
check_plan_json_paths(right, scope, shadowed)
}
PlanNode::Sort { input, keys } => {
for key in keys {
check_expr_json_paths(&key.expr, scope, shadowed)?;
}
check_plan_json_paths(input, scope, shadowed)
}
PlanNode::Aggregate {
input, argument, ..
} => {
if let Some(argument) = argument {
check_expr_json_paths(argument, scope, shadowed)?;
}
check_plan_json_paths(input, scope, shadowed)
}
PlanNode::Window { input, windows } => {
for window in windows {
for expr in window.args.iter().chain(&window.partition_by) {
check_expr_json_paths(expr, scope, shadowed)?;
}
for key in &window.order_by {
check_expr_json_paths(&key.expr, scope, shadowed)?;
}
}
check_plan_json_paths(input, scope, shadowed)
}
PlanNode::Limit { input, .. }
| PlanNode::Offset { input, .. }
| PlanNode::Distinct { input }
| PlanNode::Update { input, .. }
| PlanNode::Delete { input, .. }
| PlanNode::Explain { input } => check_plan_json_paths(input, scope, shadowed),
_ => Ok(()),
}
}
pub(crate) fn validate_column_references(
catalog: &Catalog,
plan: &PlanNode,
) -> Result<(), QueryError> {
let mut scope: Vec<(String, TypeId)> = Vec::new();
collect_scan_columns(catalog, plan, &mut scope);
if scope.is_empty() {
return Ok(());
}
let mut rebound: std::collections::HashSet<String> = std::collections::HashSet::new();
let mut computed: std::collections::HashSet<String> = std::collections::HashSet::new();
collect_rebound_names(plan, &mut rebound, &mut computed);
let mut known: std::collections::HashSet<String> =
scope.iter().map(|(name, _)| name.clone()).collect();
for (name, _) in &scope {
if let Some((_, field)) = name.split_once('.') {
known.insert(field.to_string());
}
}
known.extend(rebound.iter().cloned());
let mut ambiguous = ambiguous_bare_names(plan, catalog);
for name in &computed {
ambiguous.remove(name);
}
let ctx = ColumnScope {
known,
rebound,
ambiguous,
scope,
};
check_plan_columns(plan, &ctx)
}
fn ambiguous_bare_names(plan: &PlanNode, catalog: &Catalog) -> std::collections::HashSet<String> {
let mut scope: Vec<(String, TypeId)> = Vec::new();
collect_join_scope_columns(catalog, plan, &mut scope);
let mut owners: std::collections::HashMap<&str, &str> = std::collections::HashMap::new();
let mut ambiguous = std::collections::HashSet::new();
for (name, _) in &scope {
let Some((qualifier, field)) = name.split_once('.') else {
continue;
};
match owners.get(field) {
Some(previous) if *previous != qualifier => {
ambiguous.insert(field.to_string());
}
Some(_) => {}
None => {
owners.insert(field, qualifier);
}
}
}
for (name, _) in &scope {
ambiguous.remove(name.as_str());
}
ambiguous
}
fn collect_join_scope_columns(catalog: &Catalog, plan: &PlanNode, out: &mut Vec<(String, TypeId)>) {
match plan {
PlanNode::Union { .. } => {}
PlanNode::SeqScan { .. }
| PlanNode::IndexScan { .. }
| PlanNode::RangeScan { .. }
| PlanNode::AliasScan { .. } => collect_scan_columns(catalog, plan, out),
PlanNode::NestedLoopJoin { left, right, .. } => {
collect_join_scope_columns(catalog, left, out);
collect_join_scope_columns(catalog, right, out);
}
PlanNode::Filter { input, .. }
| PlanNode::Project { input, .. }
| PlanNode::Sort { input, .. }
| PlanNode::Limit { input, .. }
| PlanNode::Offset { input, .. }
| PlanNode::Aggregate { input, .. }
| PlanNode::Distinct { input }
| PlanNode::GroupBy { input, .. }
| PlanNode::Window { input, .. }
| PlanNode::Update { input, .. }
| PlanNode::Delete { input, .. }
| PlanNode::Explain { input } => collect_join_scope_columns(catalog, input, out),
_ => {}
}
}
struct ColumnScope {
known: std::collections::HashSet<String>,
rebound: std::collections::HashSet<String>,
ambiguous: std::collections::HashSet<String>,
scope: Vec<(String, TypeId)>,
}
fn collect_rebound_names(
plan: &PlanNode,
out: &mut std::collections::HashSet<String>,
computed: &mut std::collections::HashSet<String>,
) {
if let PlanNode::Project { fields, .. } = plan {
for field in fields {
if let Some(alias) = &field.alias {
out.insert(alias.clone());
}
}
}
match plan {
PlanNode::GroupBy {
keys, aggregates, ..
} => {
for key in keys {
out.insert(key.output_name.clone());
computed.insert(key.output_name.clone());
}
for aggregate in aggregates {
out.insert(aggregate.output_name.clone());
computed.insert(aggregate.output_name.clone());
}
}
PlanNode::Window { windows, .. } => {
for window in windows {
out.insert(window.output_name.clone());
computed.insert(window.output_name.clone());
}
}
_ => {}
}
match plan {
PlanNode::Filter { input, .. }
| PlanNode::Project { input, .. }
| PlanNode::NestedProject { input, .. }
| PlanNode::Sort { input, .. }
| PlanNode::Limit { input, .. }
| PlanNode::Offset { input, .. }
| PlanNode::Aggregate { input, .. }
| PlanNode::Distinct { input }
| PlanNode::GroupBy { input, .. }
| PlanNode::Window { input, .. }
| PlanNode::Update { input, .. }
| PlanNode::Delete { input, .. }
| PlanNode::Explain { input } => collect_rebound_names(input, out, computed),
PlanNode::NestedLoopJoin { left, right, .. } | PlanNode::Union { left, right, .. } => {
collect_rebound_names(left, out, computed);
collect_rebound_names(right, out, computed);
}
_ => {}
}
}
fn column_is_known(name: &str, ctx: &ColumnScope) -> bool {
name == "*" || ctx.known.contains(name)
}
#[derive(PartialEq, Clone, Copy)]
enum TypeClass {
Numeric,
Text,
Bool,
Other,
}
fn column_class(type_id: TypeId) -> TypeClass {
match type_id {
TypeId::Int | TypeId::Float => TypeClass::Numeric,
TypeId::Str => TypeClass::Text,
TypeId::Bool => TypeClass::Bool,
_ => TypeClass::Other,
}
}
fn literal_class(literal: &Literal) -> TypeClass {
match literal {
Literal::Int(_) | Literal::Float(_) => TypeClass::Numeric,
Literal::String(_) => TypeClass::Text,
Literal::Bool(_) => TypeClass::Bool,
}
}
fn literal_type_name(literal: &Literal) -> &'static str {
match literal {
Literal::Int(_) => "int",
Literal::Float(_) => "float",
Literal::String(_) => "str",
Literal::Bool(_) => "bool",
}
}
fn comparable_column(expr: &Expr, ctx: &ColumnScope) -> Option<(String, TypeId)> {
let name = match expr {
Expr::Field(name) => name.clone(),
Expr::QualifiedField { qualifier, field } => format!("{qualifier}.{field}"),
_ => return None,
};
if ctx.rebound.contains(&name) {
return None;
}
let type_id = resolve_scan_type(&name, &ctx.scope)?;
Some((name, type_id))
}
fn comparison_type_error(left: &Expr, right: &Expr, ctx: &ColumnScope) -> Option<String> {
let (column, literal) = match (left, right) {
(column, Expr::Literal(literal)) => (column, literal),
(Expr::Literal(literal), column) => (column, literal),
_ => return None,
};
let (name, type_id) = comparable_column(column, ctx)?;
let column_class = column_class(type_id);
let literal_class = literal_class(literal);
if column_class == TypeClass::Other || column_class == literal_class {
return None;
}
Some(format!(
"type mismatch for column '{}': expected {:?}, got {}",
name,
type_id,
literal_type_name(literal)
))
}
enum ArithOperand {
Numeric,
NonNumeric(&'static str),
Unknown,
}
fn arith_operand(expr: &Expr, ctx: &ColumnScope) -> ArithOperand {
match expr {
Expr::Literal(Literal::Int(_) | Literal::Float(_)) => ArithOperand::Numeric,
Expr::Literal(literal) => ArithOperand::NonNumeric(literal_type_name(literal)),
Expr::Field(_) | Expr::QualifiedField { .. } => match comparable_column(expr, ctx) {
Some((_, TypeId::Int | TypeId::Float)) => ArithOperand::Numeric,
Some((_, other)) => ArithOperand::NonNumeric(type_id_to_name(other)),
None => ArithOperand::Unknown,
},
_ => ArithOperand::Unknown,
}
}
fn arith_op_symbol(op: BinOp) -> &'static str {
match op {
BinOp::Add => "+",
BinOp::Sub => "-",
BinOp::Mul => "*",
_ => "/",
}
}
fn arithmetic_type_error(
left: &Expr,
op: BinOp,
right: &Expr,
ctx: &ColumnScope,
) -> Option<QueryError> {
for operand in [left, right] {
if let ArithOperand::NonNumeric(type_name) = arith_operand(operand, ctx) {
return Some(QueryError::TypeError(format!(
"operator '{}' is not defined for {}; arithmetic requires int or float{}",
arith_op_symbol(op),
type_name,
if type_name == "datetime" {
" (use date_add / date_diff for timestamps)"
} else {
""
}
)));
}
}
let divisor_is_zero = match right {
Expr::Literal(Literal::Int(value)) => *value == 0,
Expr::Literal(Literal::Float(value)) => *value == 0.0,
_ => false,
};
if op == BinOp::Div && divisor_is_zero {
return Some(QueryError::Execution(
"cannot divide by zero: the divisor is the literal 0".to_string(),
));
}
None
}
fn date_add_overflow_error(args: &[Expr]) -> Option<QueryError> {
let (Some(Expr::Literal(Literal::Int(amount))), Some(Expr::Literal(Literal::String(unit)))) =
(args.get(1), args.get(2))
else {
return None;
};
let factor = date_unit_micros(unit)?;
amount.checked_mul(factor).is_none().then(|| {
QueryError::Execution(format!(
"cannot compute date_add: amount {amount} overflows the representable range in units of '{unit}'"
))
})
}
fn check_expr_columns(expr: &Expr, ctx: &ColumnScope) -> Result<(), QueryError> {
match expr {
Expr::Field(name) => {
if !column_is_known(name, ctx) {
return Err(QueryError::ColumnNotFound {
table: String::new(),
column: name.clone(),
});
}
if ctx.ambiguous.contains(name) {
return Err(QueryError::Execution(format!(
"cannot resolve column '{name}': more than one joined table exposes it, qualify it as <alias>.{name}"
)));
}
Ok(())
}
Expr::QualifiedField { qualifier, field } => {
if !column_is_known(&format!("{qualifier}.{field}"), ctx) {
return Err(QueryError::ColumnNotFound {
table: qualifier.clone(),
column: field.clone(),
});
}
Ok(())
}
Expr::BinaryOp(left, op, right) => {
if matches!(
op,
BinOp::Eq | BinOp::Neq | BinOp::Lt | BinOp::Gt | BinOp::Lte | BinOp::Gte
) {
if let Some(message) = comparison_type_error(left, right, ctx) {
return Err(QueryError::Execution(message));
}
}
if matches!(op, BinOp::Add | BinOp::Sub | BinOp::Mul | BinOp::Div) {
if let Some(error) = arithmetic_type_error(left, *op, right, ctx) {
return Err(error);
}
}
check_expr_columns(left, ctx)?;
check_expr_columns(right, ctx)
}
Expr::Coalesce(left, right) => {
check_expr_columns(left, ctx)?;
check_expr_columns(right, ctx)
}
Expr::UnaryOp(_, inner) | Expr::FunctionCall(_, inner, _) | Expr::Cast(inner, _) => {
check_expr_columns(inner, ctx)
}
Expr::JsonPath { base, .. } => check_expr_columns(base, ctx),
Expr::ScalarFunc(func, args) => {
if *func == ScalarFn::DateAdd {
if let Some(error) = date_add_overflow_error(args) {
return Err(error);
}
}
for arg in args {
check_expr_columns(arg, ctx)?;
}
Ok(())
}
Expr::InList { expr, list, .. } => {
check_expr_columns(expr, ctx)?;
for item in list {
check_expr_columns(item, ctx)?;
}
Ok(())
}
Expr::Case { whens, else_expr } => {
for (when, then) in whens {
check_expr_columns(when, ctx)?;
check_expr_columns(then, ctx)?;
}
if let Some(else_expr) = else_expr {
check_expr_columns(else_expr, ctx)?;
}
Ok(())
}
Expr::Window {
args,
partition_by,
order_by,
..
} => {
for expr in args.iter().chain(partition_by) {
check_expr_columns(expr, ctx)?;
}
for key in order_by {
check_expr_columns(&key.expr, ctx)?;
}
Ok(())
}
Expr::InSubquery { expr, .. } => check_expr_columns(expr, ctx),
_ => Ok(()),
}
}
fn check_plan_columns(plan: &PlanNode, ctx: &ColumnScope) -> Result<(), QueryError> {
match plan {
PlanNode::Filter { input, predicate } => {
check_expr_columns(predicate, ctx)?;
check_plan_columns(input, ctx)
}
PlanNode::Project { input, fields } => {
for field in fields {
check_expr_columns(&field.expr, ctx)?;
}
check_plan_columns(input, ctx)
}
PlanNode::GroupBy {
input,
keys,
aggregates,
having,
} => {
for key in keys {
check_expr_columns(&key.expr, ctx)?;
}
for aggregate in aggregates {
check_expr_columns(&aggregate.argument, ctx)?;
}
if let Some(having) = having {
check_expr_columns(having, ctx)?;
}
check_plan_columns(input, ctx)
}
PlanNode::Sort { input, keys } => {
for key in keys {
check_expr_columns(&key.expr, ctx)?;
}
check_plan_columns(input, ctx)
}
PlanNode::Aggregate {
input, argument, ..
} => {
if let Some(argument) = argument {
check_expr_columns(argument, ctx)?;
}
check_plan_columns(input, ctx)
}
PlanNode::NestedLoopJoin {
left, right, on, ..
} => {
if let Some(on) = on {
check_expr_columns(on, ctx)?;
}
check_plan_columns(left, ctx)?;
check_plan_columns(right, ctx)
}
PlanNode::Union { left, right, .. } => {
check_plan_columns(left, ctx)?;
check_plan_columns(right, ctx)
}
PlanNode::NestedProject { input, .. } => check_plan_columns(input, ctx),
PlanNode::Limit { input, .. }
| PlanNode::Offset { input, .. }
| PlanNode::Distinct { input }
| PlanNode::Window { input, .. }
| PlanNode::Update { input, .. }
| PlanNode::Delete { input, .. }
| PlanNode::Explain { input } => check_plan_columns(input, ctx),
_ => Ok(()),
}
}
pub(crate) fn validate_slice_counts(plan: &PlanNode) -> Result<(), QueryError> {
match plan {
PlanNode::Limit { input, count } => {
check_non_negative(count, "limit")?;
validate_slice_counts(input)
}
PlanNode::Offset { input, count } => {
check_non_negative(count, "offset")?;
validate_slice_counts(input)
}
PlanNode::OrderedExprIndexScan { limit, offset, .. } => {
check_non_negative(limit, "limit")?;
if let Some(offset) = offset {
check_non_negative(offset, "offset")?;
}
Ok(())
}
PlanNode::Filter { input, .. }
| PlanNode::Project { input, .. }
| PlanNode::NestedProject { input, .. }
| PlanNode::Sort { input, .. }
| PlanNode::Aggregate { input, .. }
| PlanNode::Distinct { input }
| PlanNode::GroupBy { input, .. }
| PlanNode::Window { input, .. }
| PlanNode::Update { input, .. }
| PlanNode::Delete { input, .. }
| PlanNode::Explain { input } => validate_slice_counts(input),
PlanNode::NestedLoopJoin { left, right, .. } | PlanNode::Union { left, right, .. } => {
validate_slice_counts(left)?;
validate_slice_counts(right)
}
_ => Ok(()),
}
}
fn check_non_negative(count: &Expr, what: &str) -> Result<(), QueryError> {
match count {
Expr::Literal(Literal::Int(value)) if *value < 0 => Err(QueryError::Execution(format!(
"{what} must not be negative, got {value}"
))),
_ => Ok(()),
}
}
pub(crate) fn validate_no_stray_aggregates(plan: &PlanNode) -> Result<(), QueryError> {
match plan {
PlanNode::Project { input, fields } => {
for f in fields {
check_expr_no_aggregate(&f.expr)?;
}
validate_no_stray_aggregates(input)?;
}
PlanNode::Filter { input, predicate } => {
check_expr_no_aggregate(predicate)?;
validate_no_stray_aggregates(input)?;
}
PlanNode::GroupBy {
input,
keys,
aggregates,
having,
} => {
for key in keys {
check_expr_no_aggregate(&key.expr)?;
}
for aggregate in aggregates {
check_expr_no_aggregate(&aggregate.argument)?;
}
if let Some(h) = having {
check_expr_no_aggregate(h)?;
}
validate_no_stray_aggregates(input)?;
}
PlanNode::NestedLoopJoin {
left, right, on, ..
} => {
if let Some(on) = on {
check_expr_no_aggregate(on)?;
}
validate_no_stray_aggregates(left)?;
validate_no_stray_aggregates(right)?;
}
PlanNode::Union { left, right, .. } => {
validate_no_stray_aggregates(left)?;
validate_no_stray_aggregates(right)?;
}
PlanNode::Sort { input, keys } => {
for key in keys {
check_expr_no_aggregate(&key.expr)?;
}
validate_no_stray_aggregates(input)?;
}
PlanNode::Aggregate {
input, argument, ..
} => {
if let Some(argument) = argument {
check_expr_no_aggregate(argument)?;
}
validate_no_stray_aggregates(input)?;
}
PlanNode::Window { input, windows } => {
for window in windows {
for expr in window.args.iter().chain(&window.partition_by) {
check_expr_no_aggregate(expr)?;
}
for key in &window.order_by {
check_expr_no_aggregate(&key.expr)?;
}
}
validate_no_stray_aggregates(input)?;
}
PlanNode::Limit { input, .. }
| PlanNode::Offset { input, .. }
| PlanNode::Distinct { input }
| PlanNode::Update { input, .. }
| PlanNode::Delete { input, .. }
| PlanNode::Explain { input } => {
validate_no_stray_aggregates(input)?;
}
_ => {}
}
Ok(())
}
fn check_expr_no_aggregate(expr: &Expr) -> Result<(), QueryError> {
match expr {
Expr::FunctionCall(..) => Err(QueryError::Execution(
"invalid query: aggregate function in an unsupported position".to_string(),
)),
Expr::BinaryOp(l, _, r) | Expr::Coalesce(l, r) => {
check_expr_no_aggregate(l)?;
check_expr_no_aggregate(r)
}
Expr::UnaryOp(_, inner) | Expr::Cast(inner, _) | Expr::JsonPath { base: inner, .. } => {
check_expr_no_aggregate(inner)
}
Expr::ScalarFunc(_, args) => {
for a in args {
check_expr_no_aggregate(a)?;
}
Ok(())
}
Expr::InList { expr: e, list, .. } => {
check_expr_no_aggregate(e)?;
for item in list {
check_expr_no_aggregate(item)?;
}
Ok(())
}
Expr::InSubquery { expr: e, .. } => check_expr_no_aggregate(e),
Expr::Case { whens, else_expr } => {
for (c, r) in whens {
check_expr_no_aggregate(c)?;
check_expr_no_aggregate(r)?;
}
if let Some(e) = else_expr {
check_expr_no_aggregate(e)?;
}
Ok(())
}
Expr::Window {
args,
partition_by,
order_by,
..
} => {
for expr in args.iter().chain(partition_by) {
check_expr_no_aggregate(expr)?;
}
for key in order_by {
check_expr_no_aggregate(&key.expr)?;
}
Ok(())
}
_ => Ok(()),
}
}