use crate::result::QueryError;
use powdb_storage::catalog::Catalog;
use powdb_storage::types::*;
use crate::executor::compiled::*;
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 mut found: Option<TypeId> = None;
for (n, t) in scope {
if n == name {
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();
collect_rebound_names(plan, &mut rebound);
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 ctx = ColumnScope {
known,
rebound,
scope,
};
check_plan_columns(plan, &ctx)
}
struct ColumnScope {
known: std::collections::HashSet<String>,
rebound: std::collections::HashSet<String>,
scope: Vec<(String, TypeId)>,
}
fn collect_rebound_names(plan: &PlanNode, out: &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());
}
for aggregate in aggregates {
out.insert(aggregate.output_name.clone());
}
}
PlanNode::Window { windows, .. } => {
for window in windows {
out.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),
PlanNode::NestedLoopJoin { left, right, .. } | PlanNode::Union { left, right, .. } => {
collect_rebound_names(left, out);
collect_rebound_names(right, out);
}
_ => {}
}
}
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)
))
}
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(),
});
}
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));
}
}
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(_, args) => {
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(()),
}
}