use inillucent_value::Collation;
use super::{comparison_rules, result_collation, Binder, BoundExpr, SubqueryKind};
use crate::ast::{self, BinaryOp, Expr, ExprId, InRhs, SelectBody};
use crate::bind::refused;
use crate::diagnostic::{ParseError, ParseErrorKind};
use crate::lexer::Span;
impl Binder<'_> {
pub(super) fn bind_row_in(
&mut self,
parts: &[ExprId],
rhs: &InRhs,
negated: bool,
span: Span,
) -> Result<BoundExpr, ParseError> {
let mut bound_lefts = Vec::with_capacity(parts.len());
for part in parts {
bound_lefts.push(self.bind_expr(*part)?);
}
self.bind_row_in_bound(bound_lefts, rhs, negated, span)
}
pub(super) fn bind_row_in_bound(
&mut self,
bound_lefts: Vec<BoundExpr>,
rhs: &InRhs,
negated: bool,
span: Span,
) -> Result<BoundExpr, ParseError> {
let Some(rows) = self.row_value_list(rhs) else {
return match rhs {
InRhs::Select(select) => {
self.bind_row_in_query(bound_lefts, *select, negated, span)
}
_ => Err(ParseError::new(
ParseErrorKind::Unsupported("a row value IN a query rather than a value list"),
span,
)),
};
};
let mut bound_rows = Vec::with_capacity(rows.len());
for row in &rows {
if row.len() != bound_lefts.len() {
let plural = if row.len() == 1 { "" } else { "s" };
let _ = span;
return Err(ParseError::new(
ParseErrorKind::Refused(format!(
"IN(...) element has {} term{plural} - expected {}",
row.len(),
bound_lefts.len()
)),
Span::default(),
));
}
let mut bound_rights = Vec::with_capacity(row.len());
for value in row {
bound_rights.push(self.bind_expr(*value)?);
}
bound_rows.push(bound_rights);
}
let rules = match bound_rows.len() {
0 | 1 => None,
_ => Some(row_in_rules(
&bound_lefts,
&bound_rows,
matches!(rhs, InRhs::Select(_)),
)),
};
let mut chain: Option<BoundExpr> = None;
for bound_rights in &bound_rows {
let one = match &rules {
Some(rules) => equality_chain_under(&bound_lefts, bound_rights, rules),
None => equality_chain(&bound_lefts, bound_rights),
};
chain = Some(match chain {
None => one,
Some(held) => BoundExpr::Or(Box::new(held), Box::new(one)),
});
} let bound = chain.unwrap_or(BoundExpr::Integer(0));
Ok(if negated {
BoundExpr::Not(Box::new(bound))
} else {
bound
})
}
fn row_value_list(&self, rhs: &InRhs) -> Option<Vec<Vec<ExprId>>> {
match rhs {
InRhs::Select(select) => {
let held = self.ast.select(*select)?;
if !held.compounds.is_empty() || !held.with.ctes.is_empty() {
return None;
}
let core = self.ast.core(held.first)?;
match &core.body {
SelectBody::Values(rows) => Some(rows.clone()),
_ => None,
}
}
InRhs::List(items) => {
let mut rows = Vec::with_capacity(items.len());
for item in items {
rows.push(self.row_value_parts(*item)?);
}
Some(rows)
}
InRhs::Table { .. } => None,
}
}
pub(super) fn row_value_parts(&self, id: ExprId) -> Option<Vec<ExprId>> {
match self.ast.expr(id)? {
Expr::RowValue(parts) => Some(parts.clone()),
_ => None,
}
}
pub(super) fn bind_row_against_query(
&mut self,
op: BinaryOp,
lefts: &[ExprId],
select: ast::SelectId,
span: Span,
) -> Result<BoundExpr, ParseError> {
let block = self.bind_value_subquery(select, span)?;
if block.columns.len() != lefts.len() {
return Err(misused(span));
}
let mut bound_lefts = Vec::with_capacity(lefts.len());
for part in lefts {
bound_lefts.push(self.bind_expr(*part)?);
}
let bound_rights = self.query_columns(block);
compare_bound_rows(op, &bound_lefts, &bound_rights, span)
}
pub(crate) fn bind_query_columns(
&mut self,
select: ast::SelectId,
span: Span,
) -> Result<Vec<BoundExpr>, ParseError> {
let block = self.bind_value_subquery(select, span)?;
Ok(self.query_columns(block))
}
fn query_columns(&mut self, block: crate::bind::BoundSelect) -> Vec<BoundExpr> {
let mut bound_rights = Vec::with_capacity(block.columns.len());
for at in 0..block.columns.len() {
let mut one = block.clone();
one.columns = block
.columns
.get(at..at.saturating_add(1))
.map_or_else(Vec::new, <[crate::bind::BoundResultColumn]>::to_vec);
let collation = one
.columns
.first()
.map(|column| result_collation(&column.expr))
.unwrap_or(Collation::Binary);
bound_rights.push(BoundExpr::Subquery {
id: self.next_subquery_id(),
kind: SubqueryKind::Scalar,
negated: false,
operand: None,
block: Box::new(one),
affinity: None,
collation,
});
}
bound_rights
}
pub(super) fn bind_row_comparison(
&mut self,
op: BinaryOp,
lefts: &[ExprId],
rights: &[ExprId],
span: Span,
) -> Result<BoundExpr, ParseError> {
if lefts.len() != rights.len() || lefts.is_empty() {
return Err(misused(span));
}
let mut bound_lefts = Vec::with_capacity(lefts.len());
let mut bound_rights = Vec::with_capacity(rights.len());
for (left, right) in lefts.iter().zip(rights.iter()) {
bound_lefts.push(self.bind_expr(*left)?);
bound_rights.push(self.bind_expr(*right)?);
}
match op {
BinaryOp::Equal => Ok(equality_chain(&bound_lefts, &bound_rights)),
BinaryOp::NotEqual => Ok(BoundExpr::Not(Box::new(equality_chain(
&bound_lefts,
&bound_rights,
)))),
BinaryOp::Less | BinaryOp::LessEqual | BinaryOp::Greater | BinaryOp::GreaterEqual => {
Ok(lexicographic_chain(op, &bound_lefts, &bound_rights, 0))
}
_ => Err(misused(span)),
}
}
}
impl Binder<'_> {
pub(super) fn row_query(&self, id: ExprId) -> Option<ast::SelectId> {
let Some(Expr::Subquery(select)) = self.ast.expr(id) else {
return None;
};
let width = self.query_width(*select)?;
(width > 1).then_some(*select)
}
fn query_width(&self, select: ast::SelectId) -> Option<usize> {
let held = self.ast.select(select)?;
match &self.ast.core(held.first)?.body {
SelectBody::Values(rows) => rows.first().map(Vec::len),
SelectBody::Select { columns, .. } => {
let star = columns
.iter()
.any(|column| matches!(self.ast.expr(column.expr), Some(Expr::Star { .. })));
(!star).then_some(columns.len())
}
}
}
pub(super) fn bind_row_query_versus(
&mut self,
op: BinaryOp,
select: ast::SelectId,
right: ExprId,
span: Span,
) -> Result<BoundExpr, ParseError> {
let lefts = self.bind_query_columns(select, span)?;
let rights = self.bind_row_operand(right, span)?;
if lefts.len() != rights.len() {
return Err(misused(span));
}
compare_bound_rows(op, &lefts, &rights, span)
}
pub(super) fn bind_row_query_is(
&mut self,
negated: bool,
select: ast::SelectId,
right: ExprId,
span: Span,
) -> Result<BoundExpr, ParseError> {
let lefts = self.bind_query_columns(select, span)?;
let rights = self.bind_row_operand(right, span)?;
if lefts.len() != rights.len() {
return Err(misused(span));
}
Ok(is_chain(negated, &lefts, &rights))
}
fn bind_row_operand(&mut self, id: ExprId, span: Span) -> Result<Vec<BoundExpr>, ParseError> {
if let Some(parts) = self.row_value_parts(id) {
let mut bound = Vec::with_capacity(parts.len());
for part in &parts {
bound.push(self.bind_expr(*part)?);
}
return Ok(bound);
}
match self.row_query(id) {
Some(select) => self.bind_query_columns(select, span),
None => Err(misused(span)),
}
}
fn bind_row_versus(
&mut self,
op: BinaryOp,
lefts: &[ExprId],
right: ExprId,
span: Span,
) -> Result<BoundExpr, ParseError> {
if let Some(rights) = self.row_value_parts(right) {
return self.bind_row_comparison(op, lefts, &rights, span);
}
if let Some(Expr::Subquery(select)) = self.ast.expr(right) {
let select = *select;
return self.bind_row_against_query(op, lefts, select, span);
}
Err(misused(span))
}
pub(super) fn bind_row_is(
&mut self,
negated: bool,
lefts: &[ExprId],
rights: &[ExprId],
span: Span,
) -> Result<BoundExpr, ParseError> {
if lefts.len() != rights.len() || lefts.is_empty() {
return Err(misused(span));
}
let mut bound_lefts = Vec::with_capacity(lefts.len());
let mut bound_rights = Vec::with_capacity(rights.len());
for (left, right) in lefts.iter().zip(rights.iter()) {
bound_lefts.push(self.bind_expr(*left)?);
bound_rights.push(self.bind_expr(*right)?);
}
Ok(is_chain(negated, &bound_lefts, &bound_rights))
}
pub(super) fn bind_row_between(
&mut self,
negated: bool,
parts: &[ExprId],
low: ExprId,
high: ExprId,
span: Span,
) -> Result<BoundExpr, ParseError> {
let above = self.bind_row_versus(BinaryOp::GreaterEqual, parts, low, span)?;
let below = self.bind_row_versus(BinaryOp::LessEqual, parts, high, span)?;
let both = BoundExpr::And(Box::new(above), Box::new(below));
Ok(match negated {
true => BoundExpr::Not(Box::new(both)),
false => both,
})
}
pub(super) fn bind_row_case(
&mut self,
parts: &[ExprId],
branches: &[(ExprId, ExprId)],
otherwise: Option<ExprId>,
span: Span,
) -> Result<BoundExpr, ParseError> {
let mut bound_branches = Vec::with_capacity(branches.len());
for (when, then) in branches {
let test = self.bind_row_versus(BinaryOp::Equal, parts, *when, span)?;
bound_branches.push((test, self.bind_expr(*then)?));
}
let otherwise = match otherwise {
Some(expr) => Some(Box::new(self.bind_expr(expr)?)),
None => None,
};
Ok(BoundExpr::Case {
operand: None,
branches: bound_branches,
otherwise,
comparisons: Vec::new(),
})
}
}
fn is_chain(negated: bool, lefts: &[BoundExpr], rights: &[BoundExpr]) -> BoundExpr {
let mut chain: Option<BoundExpr> = None;
for (left, right) in lefts.iter().zip(rights.iter()) {
let (affinity, collation) = comparison_rules(left, right);
let one = BoundExpr::Is {
negated: false,
left: Box::new(left.clone()),
right: Box::new(right.clone()),
affinity,
collation,
};
chain = Some(match chain {
None => one,
Some(held) => BoundExpr::And(Box::new(held), Box::new(one)),
});
}
let chain = chain.unwrap_or(BoundExpr::Null);
match negated {
true => BoundExpr::Not(Box::new(chain)),
false => chain,
}
}
pub(super) fn query_arity_refusal(found: usize, expected: usize, span: Span) -> ParseError {
refused(
format!("sub-select returns {found} columns - expected {expected}"),
span,
)
}
pub(super) fn misused(span: Span) -> ParseError {
let _ = span;
ParseError::new(
ParseErrorKind::Refused("row value misused".to_string()),
Span::default(),
)
}
pub(super) fn equality_chain(lefts: &[BoundExpr], rights: &[BoundExpr]) -> BoundExpr {
let rules: Vec<_> = lefts
.iter()
.zip(rights.iter())
.map(|(left, right)| comparison_rules(left, right))
.collect();
equality_chain_under(lefts, rights, &rules)
}
pub(super) fn row_in_rules(
lefts: &[BoundExpr],
rows: &[Vec<BoundExpr>],
values_clause: bool,
) -> Vec<(Option<inillucent_value::Affinity>, Collation)> {
let (Some(first), Some(last)) = (rows.first(), rows.last()) else {
return Vec::new();
};
let mut rules = Vec::with_capacity(lefts.len());
for (at, left) in lefts.iter().enumerate() {
let (Some(first), Some(last)) = (first.get(at), last.get(at)) else {
rules.push((None, Collation::Binary));
continue;
};
let (mut affinity, mut collation) = comparison_rules(left, last);
if values_clause {
let implicit = match last.explicit_collation() {
Some(_) => None,
None => last.collation(),
};
collation = left
.explicit_collation()
.or_else(|| first.explicit_collation())
.or_else(|| left.collation())
.or(implicit)
.unwrap_or(Collation::Binary);
if matches!(last, BoundExpr::Cast { .. }) {
affinity = left.affinity();
}
}
rules.push((affinity, collation));
}
rules
}
pub(super) fn equality_chain_under(
lefts: &[BoundExpr],
rights: &[BoundExpr],
rules: &[(Option<inillucent_value::Affinity>, Collation)],
) -> BoundExpr {
let mut chain: Option<BoundExpr> = None;
for ((left, right), (affinity, collation)) in lefts.iter().zip(rights.iter()).zip(rules.iter())
{
let (affinity, collation) = (*affinity, *collation);
let one = BoundExpr::Compare {
op: BinaryOp::Equal,
left: Box::new(left.clone()),
right: Box::new(right.clone()),
affinity,
collation,
};
chain = Some(match chain {
None => one,
Some(held) => BoundExpr::And(Box::new(held), Box::new(one)),
});
}
chain.unwrap_or(BoundExpr::Null)
}
fn lexicographic_chain(
op: BinaryOp,
lefts: &[BoundExpr],
rights: &[BoundExpr],
at: usize,
) -> BoundExpr {
let (Some(left), Some(right)) = (lefts.get(at), rights.get(at)) else {
return BoundExpr::Null;
};
let (affinity, collation) = comparison_rules(left, right);
let last = at.saturating_add(1) >= lefts.len();
let strict = match op {
BinaryOp::LessEqual if !last => BinaryOp::Less,
BinaryOp::GreaterEqual if !last => BinaryOp::Greater,
other => other,
};
let decided = BoundExpr::Compare {
op: strict,
left: Box::new(left.clone()),
right: Box::new(right.clone()),
affinity,
collation,
};
if last {
return decided;
}
let same = BoundExpr::Compare {
op: BinaryOp::Equal,
left: Box::new(left.clone()),
right: Box::new(right.clone()),
affinity,
collation,
};
BoundExpr::Or(
Box::new(decided),
Box::new(BoundExpr::And(
Box::new(same),
Box::new(lexicographic_chain(op, lefts, rights, at.saturating_add(1))),
)),
)
}
fn compare_bound_rows(
op: BinaryOp,
lefts: &[BoundExpr],
rights: &[BoundExpr],
span: Span,
) -> Result<BoundExpr, ParseError> {
match op {
BinaryOp::Equal => Ok(equality_chain(lefts, rights)),
BinaryOp::NotEqual => Ok(BoundExpr::Not(Box::new(equality_chain(lefts, rights)))),
BinaryOp::Less | BinaryOp::LessEqual | BinaryOp::Greater | BinaryOp::GreaterEqual => {
Ok(lexicographic_chain(op, lefts, rights, 0))
}
_ => Err(misused(span)),
}
}