Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
6 changes: 5 additions & 1 deletion AGENT.md
Original file line number Diff line number Diff line change
Expand Up @@ -151,7 +151,11 @@ KiteSQL aims to stay easy to build, easy to audit, and easy to understand.

A valid PR should:

- Follow the repository pull request template
- Follow `.github/pull_request_template.md` exactly: preserve its headings, checklist labels,
ordering, and comments. In `Code changes` and `Check List`, only change checkbox states;
do not add, remove, or reword items, or insert explanatory paragraphs.
- Put test results, manual verification steps, compatibility notes, and side-effect explanations
in `Note for reviewer`, not inside the checklist. Keep the PR description concise.
- Describe test coverage and verified behavior instead of listing commands that were run
- Compile cleanly
- Pass all tests via make
Expand Down
15 changes: 15 additions & 0 deletions src/binder/aggregate.rs
Original file line number Diff line number Diff line change
Expand Up @@ -190,6 +190,15 @@ impl<T: Transaction, A: AsRef<[(usize, LogicalType)]>> Binder<'_, '_, T, A> {
if expr.has_agg_call(arena)? {
continue;
}
if !expr.any_referenced_column(arena, |_, _| true)? {
if let Some(position) = unmatched_group_exprs
.iter()
.position(|group_expr| expr.eq_ignore_colref_pos(*group_expr, arena))
{
unmatched_group_exprs.remove(position);
}
continue;
}
let Some(position) = unmatched_group_exprs
.iter()
.position(|group_expr| expr.eq_ignore_colref_pos(*group_expr, arena))
Expand All @@ -202,6 +211,12 @@ impl<T: Transaction, A: AsRef<[(usize, LogicalType)]>> Binder<'_, '_, T, A> {
unmatched_group_exprs.remove(position);
}

unmatched_group_exprs.retain(|expr| {
!matches!(
arena.expression(expr.unpack_alias(arena)),
ScalarExpression::InitValue { .. }
)
});
if !unmatched_group_exprs.is_empty() {
return Err(DatabaseError::AggMiss(
"in the GROUP BY clause the field must be in the select clause".to_string(),
Expand Down
1 change: 1 addition & 0 deletions src/binder/create_view.rs
Original file line number Diff line number Diff line change
Expand Up @@ -91,6 +91,7 @@ impl<T: Transaction, A: AsRef<[(usize, LogicalType)]>> Binder<'_, '_, T, A> {
};
plan = self.bind_project(plan, exprs, arena)?;
let schema = plan.output_schema(arena).clone();
arena.materialize_scalar_queries(&mut plan)?;

Ok(LogicalPlan::new(
Operator::CreateView(CreateViewOperator {
Expand Down
124 changes: 96 additions & 28 deletions src/binder/expr.rs
Original file line number Diff line number Diff line change
Expand Up @@ -18,14 +18,14 @@ use crate::expression;
use crate::expression::agg::AggKind;
use crate::iter_ext::Itertools;

use super::{Binder, BinderContext, QueryBindStep, SubQueryType};
use super::{Binder, BinderContext, BoundScalarQuery, QueryBindStep, SubQueryType};
use crate::expression::function::scala::{ArcScalarFunctionImpl, ScalarFunction};
use crate::expression::function::table::TableFunction;
use crate::expression::function::FunctionSummary;
use crate::expression::{AliasType, ScalarExpression, TypeCast};
use crate::planner::operator::mark_apply::MarkApplyQuantifier;
use crate::planner::operator::scalar_subquery::ScalarSubqueryOperator;
use crate::planner::{ExprRef, LogicalPlan, PlanArena};
use crate::planner::{ExprRef, LogicalPlan, MetaArena, PlanArena};
use crate::storage::Transaction;
use crate::types::value::{DataValue, Utf8Type};
use crate::types::{CharLengthUnits, LogicalType};
Expand Down Expand Up @@ -147,7 +147,7 @@ impl<'a, T: Transaction, A: AsRef<[(usize, LogicalType)]>> Binder<'a, '_, T, A>
{
let mut binder = Binder::new(self.context.fork_empty(), self.args, Some(&self.context));
let sub_query = build(&mut binder, arena)?;
let correlated = binder.context.has_outer_refs();
let correlated = binder.context.has_join_outer_ref();
Ok((sub_query, correlated))
}

Expand Down Expand Up @@ -206,23 +206,42 @@ impl<'a, T: Transaction, A: AsRef<[(usize, LogicalType)]>> Binder<'a, '_, T, A>
&mut PlanArena<'arena>,
) -> Result<LogicalPlan, DatabaseError>,
{
let (sub_query, column, correlated) =
self.bind_subquery_plan_with_output(None, arena, build)?;
let sub_query = ScalarSubqueryOperator::build(sub_query);
let (expr, sub_query) = match self.context.step_now() {
QueryBindStep::Where => (column, sub_query),
QueryBindStep::Project => self.bind_temp_table(column, sub_query, arena)?,
_ => {
return Err(DatabaseError::UnsupportedStmt(
"scalar subqueries can only appear in `WHERE` or SELECT list".to_string(),
))
}
};
self.context.sub_query(SubQueryType::SubQuery {
plan: sub_query,
correlated,
let mut child_context = self.context.fork_empty();
child_context.capture_scalar_outer = true;
let mut binder = Binder::new(child_context, self.args, Some(&self.context));
let mut sub_query = build(&mut binder, arena)?;
let param_bindings = binder.context.scalar_outer_bindings;
if !param_bindings.is_empty() && self.context.capture_scalar_outer {
return Err(DatabaseError::UnsupportedStmt(
"nested correlated scalar queries are not supported yet".into(),
));
}
if !param_bindings.is_empty() && self.context.step_now() != QueryBindStep::Project {
return Err(DatabaseError::UnsupportedStmt(
"correlated scalar subqueries currently require the SELECT list".into(),
));
}
let schema = sub_query.output_schema(arena);
if schema.len() != 1 {
return Err(DatabaseError::MisMatch(
"expects only one expression to be returned",
"the expression returned by the subquery",
));
}
let ty = arena.column(schema[0]).datatype().clone();
let id = arena.alloc_scalar_query_ref(!param_bindings.is_empty());
let value = arena.alloc_expression(if param_bindings.is_empty() {
ScalarExpression::InitValue { id, ty: ty.clone() }
} else {
ScalarExpression::OuterValue { id, ty: ty.clone() }
});
self.context.scalar_queries.push(BoundScalarQuery {
value,
step: self.context.step_now(),
plan: ScalarSubqueryOperator::build(sub_query),
param_bindings,
});
Ok(expr)
Ok(arena.expression(value).clone())
}

pub(crate) fn bind_exists_subquery_plan<'arena, F>(
Expand Down Expand Up @@ -330,6 +349,21 @@ impl<'a, T: Transaction, A: AsRef<[(usize, LogicalType)]>> Binder<'a, '_, T, A>
arena: &mut PlanArena,
) -> Result<ScalarExpression, DatabaseError> {
if table_name.is_none() {
// SELECT expressions and WHERE read input columns, not a same-named output alias.
if bind_table_name.is_none()
&& matches!(
self.context.step_now(),
QueryBindStep::Project | QueryBindStep::Where
)
{
let input = match self.context.using.get(column_name) {
Some(column) => Some(column.visible_expr(arena)?),
None => Self::find_column_in_scope(&self.context, arena, column_name)?,
};
if let Some(input) = input {
return Ok(input);
}
}
if let Some((_, expr)) = self
.context
.expr_aliases
Expand All @@ -349,12 +383,14 @@ impl<'a, T: Transaction, A: AsRef<[(usize, LogicalType)]>> Binder<'a, '_, T, A>
try_default!(&table_name, column_name);
}
if let Some(table) = table_name.or(bind_table_name) {
let mut is_outer = false;
let (source, position_offset) =
match Self::resolve_source_columns_in_scope(&self.context, table) {
Ok(source) => source,
Err(err) => {
if let Some(parent) = self.parent {
self.context.mark_outer_ref();
self.context.mark_join_outer_ref();
is_outer = true;
Self::resolve_source_columns_in_scope(parent, table).map_err(|_| err)?
} else {
return Err(err);
Expand All @@ -365,10 +401,11 @@ impl<'a, T: Transaction, A: AsRef<[(usize, LogicalType)]>> Binder<'a, '_, T, A>
Self::find_column_in_schema(source.schema().iter(), arena, column_name)
.ok_or_else(|| DatabaseError::column_not_found(column_name.to_string()))?;

Ok(ScalarExpression::column_expr(
column,
position_offset + position,
))
let expr = ScalarExpression::column_expr(column, position_offset + position);
if is_outer && self.context.capture_scalar_outer {
return self.capture_scalar_outer(expr, arena);
}
Ok(expr)
} else {
// handle col syntax
let mut find_visible_column =
Expand All @@ -381,8 +418,13 @@ impl<'a, T: Transaction, A: AsRef<[(usize, LogicalType)]>> Binder<'a, '_, T, A>
let mut got_column = find_visible_column(&self.context)?;
if got_column.is_none() {
if let Some(parent) = self.parent {
self.context.mark_outer_ref();
self.context.mark_join_outer_ref();
got_column = find_visible_column(parent)?;
if self.context.capture_scalar_outer {
if let Some(expr) = got_column {
return self.capture_scalar_outer(expr, arena);
}
}
}
}
match got_column {
Expand All @@ -392,6 +434,18 @@ impl<'a, T: Transaction, A: AsRef<[(usize, LogicalType)]>> Binder<'a, '_, T, A>
}
}

fn capture_scalar_outer(
&mut self,
expr: ScalarExpression,
arena: &mut PlanArena,
) -> Result<ScalarExpression, DatabaseError> {
let ty = expr.return_type(arena).into_owned();
let id = arena.alloc_scalar_query_ref(false);
let source = arena.alloc_expression(expr);
self.context.scalar_outer_bindings.push((id, source));
Ok(ScalarExpression::OuterParam { id, ty })
}

pub(crate) fn bind_binary_op_expr(
&mut self,
left_expr: ExprRef,
Expand All @@ -409,10 +463,12 @@ impl<'a, T: Transaction, A: AsRef<[(usize, LogicalType)]>> Binder<'a, '_, T, A>
LogicalType::max_logical_type(&left_ty, &right_ty)?.into_owned()
}
expression::BinaryOperator::Divide => {
if let LogicalType::Decimal(precision, scale) =
LogicalType::max_logical_type(&left_ty, &right_ty)?.into_owned()
let ty = LogicalType::max_logical_type(&left_ty, &right_ty)?.into_owned();
if ty.is_signed_numeric()
|| ty.is_unsigned_numeric()
|| matches!(ty, LogicalType::Decimal(..))
{
LogicalType::Decimal(precision, scale)
ty
} else {
LogicalType::Double
}
Expand Down Expand Up @@ -448,6 +504,18 @@ impl<'a, T: Transaction, A: AsRef<[(usize, LogicalType)]>> Binder<'a, '_, T, A>
op: expression::UnaryOperator,
arena: &mut PlanArena,
) -> Result<ScalarExpression, DatabaseError> {
if op == expression::UnaryOperator::Plus {
let ty = expr.return_type(arena);
if ty.is_numeric()
|| matches!(
ty.as_ref(),
LogicalType::SqlNull | LogicalType::Char(..) | LogicalType::Varchar(..)
)
{
return Ok(arena.expression(expr).clone());
}
return Err(DatabaseError::UnsupportedUnaryOperator(ty.into_owned(), op));
}
let ty = if let expression::UnaryOperator::Not = op {
LogicalType::Boolean
} else {
Expand Down
46 changes: 30 additions & 16 deletions src/binder/mod.rs
Original file line number Diff line number Diff line change
Expand Up @@ -68,7 +68,7 @@ use crate::errors::DatabaseError;
use crate::expression::{ScalarExpression, TypeCast};
use crate::planner::operator::join::JoinType;
use crate::planner::operator::mark_apply::MarkApplyQuantifier;
use crate::planner::{ExprRef, LogicalPlan, PlanArena, PlanRef};
use crate::planner::{ExprRef, LogicalPlan, PlanArena, PlanRef, ScalarQueryRef};
use crate::storage::{TableCache, Transaction, ViewCache};
use crate::types::tuple::Schema;
use crate::types::LogicalType;
Expand Down Expand Up @@ -102,10 +102,6 @@ pub enum QueryBindStep {

#[derive(Debug, Clone, Hash, Eq, PartialEq)]
pub enum SubQueryType {
SubQuery {
plan: LogicalPlan,
correlated: bool,
},
ExistsSubQuery {
plan: LogicalPlan,
correlated: bool,
Expand Down Expand Up @@ -151,8 +147,10 @@ pub(crate) struct CteCheckpoint {

impl BoundSource<'_> {
pub(crate) fn matches_name(&self, table_name: &str) -> bool {
self.table_name.as_ref() == table_name
|| matches!(self.alias.as_ref(), Some(alias) if alias.as_ref() == table_name)
match &self.alias {
Some(alias) => alias.as_ref() == table_name,
None => self.table_name.as_ref() == table_name,
}
}

pub(crate) fn same_binding(
Expand Down Expand Up @@ -237,6 +235,13 @@ impl UsingColumn {
}
}

pub(crate) struct BoundScalarQuery {
value: ExprRef,
step: QueryBindStep,
plan: LogicalPlan,
param_bindings: Vec<(ScalarQueryRef, ExprRef)>,
}

pub struct BinderContext<'a, T: Transaction> {
pub(crate) scala_functions: &'a ScalaFunctions,
pub(crate) table_functions: &'a TableFunctions,
Expand All @@ -256,11 +261,14 @@ pub struct BinderContext<'a, T: Transaction> {
pub(crate) agg_calls: Vec<ExprRef>,
// join
using: HashMap<String, UsingColumn>,

bind_step: QueryBindStep,
has_join_outer_ref: bool,
// subquery
sub_queries: HashMap<QueryBindStep, Vec<SubQueryType>>,
has_outer_refs: bool,
scalar_queries: Vec<BoundScalarQuery>,
capture_scalar_outer: bool,
scalar_outer_bindings: Vec<(ScalarQueryRef, ExprRef)>,

bind_step: QueryBindStep,
pub(crate) allow_default: bool,
}

Expand Down Expand Up @@ -321,7 +329,10 @@ impl<'a, T: Transaction> BinderContext<'a, T> {
using: Default::default(),
bind_step: QueryBindStep::From,
sub_queries: Default::default(),
has_outer_refs: false,
scalar_queries: Vec::new(),
capture_scalar_outer: false,
scalar_outer_bindings: Vec::new(),
has_join_outer_ref: false,
allow_default: false,
}
}
Expand All @@ -347,7 +358,10 @@ impl<'a, T: Transaction> BinderContext<'a, T> {
using: self.using.clone(),
bind_step: self.bind_step,
sub_queries: Default::default(),
has_outer_refs: false,
scalar_queries: Vec::new(),
capture_scalar_outer: false,
scalar_outer_bindings: Vec::new(),
has_join_outer_ref: false,
allow_default: self.allow_default,
}
}
Expand Down Expand Up @@ -428,12 +442,12 @@ impl<'a, T: Transaction> BinderContext<'a, T> {
self.sub_queries.remove(&self.bind_step)
}

pub fn mark_outer_ref(&mut self) {
self.has_outer_refs = true;
pub fn mark_join_outer_ref(&mut self) {
self.has_join_outer_ref = true;
}

pub fn has_outer_refs(&self) -> bool {
self.has_outer_refs
pub fn has_join_outer_ref(&self) -> bool {
self.has_join_outer_ref
}

pub fn table(&self, table_name: TableName) -> Result<Option<&TableCatalog>, DatabaseError> {
Expand Down
Loading
Loading