use super::Simplify;
use toasty_core::stmt::{self, Expr, ResolvedRef, VisitMut};
impl Simplify<'_> {
pub(super) fn simplify_expr_binary_op(
&mut self,
op: stmt::BinaryOp,
lhs: &mut stmt::Expr,
rhs: &mut stmt::Expr,
) -> Option<stmt::Expr> {
if (op.is_eq() || op.is_ne())
&& let (Some(lhs_len), Some(rhs_len)) = (lhs.record_len(), rhs.record_len())
&& lhs_len != rhs_len
{
return Some(op.is_ne().into());
}
let result = match (&mut *lhs, &mut *rhs) {
(Expr::Reference(lhs), Expr::Reference(rhs))
if lhs == rhs && (op.is_eq() || op.is_ne()) =>
{
if lhs.is_field() {
let field = self.cx.resolve_expr_reference(lhs).as_field_unwrap();
if !field.nullable() {
return Some(op.is_eq().into());
}
}
None
}
(Expr::Record(lhs_rec), Expr::Record(rhs_rec))
if (op.is_eq() || op.is_ne()) && lhs_rec.len() == rhs_rec.len() =>
{
let comparisons: Vec<_> = std::mem::take(&mut lhs_rec.fields)
.into_iter()
.zip(std::mem::take(&mut rhs_rec.fields))
.map(|(l, r)| record_field_comparison(l, op, r))
.collect();
let mut comparison = if op.is_eq() {
Expr::and_from_vec(comparisons)
} else {
Expr::or_from_vec(comparisons)
};
self.visit_expr_mut(&mut comparison);
Some(comparison)
}
(Expr::Record(rec), Expr::Value(stmt::Value::Record(val_rec)))
| (Expr::Value(stmt::Value::Record(val_rec)), Expr::Record(rec))
if (op.is_eq() || op.is_ne()) && rec.len() == val_rec.len() =>
{
let comparisons: Vec<_> = std::mem::take(&mut rec.fields)
.into_iter()
.zip(std::mem::take(&mut val_rec.fields))
.map(|(expr, val)| record_field_comparison(expr, op, Expr::from(val)))
.collect();
let mut comparison = if op.is_eq() {
Expr::and_from_vec(comparisons)
} else {
Expr::or_from_vec(comparisons)
};
self.visit_expr_mut(&mut comparison);
Some(comparison)
}
(Expr::Match(m), _) if !op.is_arithmetic() && m.subject.is_stable() => {
let match_expr = lhs.take();
let other = rhs.take();
Some(self.eliminate_match_in_binary_op(op, match_expr, other, true))
}
(_, Expr::Match(m)) if !op.is_arithmetic() && m.subject.is_stable() => {
let other = lhs.take();
let match_expr = rhs.take();
Some(self.eliminate_match_in_binary_op(op, match_expr, other, false))
}
(Expr::Cast(cast), Expr::Value(value)) | (Expr::Value(value), Expr::Cast(cast))
if (op.is_eq() || op.is_ne()) && cast.from.is_none() && cast.expr.is_column() =>
{
self.strip_decode_cast_comparison(op, cast, value)
}
(Expr::Cast(lhs_cast), Expr::Cast(rhs_cast))
if self.can_strip_decode_cast_column_comparison(op, lhs_cast, rhs_cast) =>
{
self.strip_decode_cast_column_comparison(op, lhs_cast, rhs_cast)
}
(Expr::Project(lhs), Expr::Project(rhs))
if lhs == rhs && lhs.base.is_stable() && (op.is_eq() || op.is_ne()) =>
{
Some(Expr::from(op.is_eq()))
}
_ => None,
};
if result.is_some() {
return result;
}
if self.is_always_null_derived_column(lhs) || self.is_always_null_derived_column(rhs) {
return Some(Expr::null());
}
None
}
fn strip_decode_cast_comparison(
&mut self,
op: stmt::BinaryOp,
cast: &mut stmt::ExprCast,
value: &mut stmt::Value,
) -> Option<Expr> {
let expr_reference = cast.expr.as_expr_reference()?;
let ResolvedRef::Column(column) = self.cx.resolve_expr_reference(expr_reference) else {
return None;
};
let value = column
.ty
.cast(self.cx.schema(), value.take())
.expect("failed to cast value");
Some(Expr::binary_op(cast.expr.take(), op, value))
}
fn can_strip_decode_cast_column_comparison(
&self,
op: stmt::BinaryOp,
lhs: &stmt::ExprCast,
rhs: &stmt::ExprCast,
) -> bool {
if !(op.is_eq() || op.is_ne())
|| lhs.from.is_some()
|| rhs.from.is_some()
|| !lhs.expr.is_column()
|| !rhs.expr.is_column()
|| lhs.ty != rhs.ty
{
return false;
}
let Some(lhs_reference) = lhs.expr.as_expr_reference() else {
return false;
};
let Some(rhs_reference) = rhs.expr.as_expr_reference() else {
return false;
};
let ResolvedRef::Column(lhs_column) = self.cx.resolve_expr_reference(lhs_reference) else {
return false;
};
let ResolvedRef::Column(rhs_column) = self.cx.resolve_expr_reference(rhs_reference) else {
return false;
};
lhs_column.ty == rhs_column.ty && lhs_column.ty.cast_preserves_equality(&lhs.ty)
}
fn strip_decode_cast_column_comparison(
&mut self,
op: stmt::BinaryOp,
lhs: &mut stmt::ExprCast,
rhs: &mut stmt::ExprCast,
) -> Option<Expr> {
Some(Expr::binary_op(lhs.expr.take(), op, rhs.expr.take()))
}
fn is_always_null_derived_column(&self, expr: &Expr) -> bool {
let Expr::Reference(expr_ref) = expr else {
return false;
};
match self.cx.resolve_expr_reference(expr_ref) {
ResolvedRef::Derived(derived_ref) => derived_ref.is_column_always_null(),
_ => false,
}
}
fn eliminate_match_in_binary_op(
&mut self,
op: stmt::BinaryOp,
match_expr: Expr,
other: Expr,
match_on_lhs: bool,
) -> Expr {
self.eliminate_match(match_expr, |arm| {
if match_on_lhs {
Expr::binary_op(arm, op, other.clone())
} else {
Expr::binary_op(other.clone(), op, arm)
}
})
}
pub(super) fn eliminate_match(
&mut self,
match_expr: Expr,
term: impl Fn(Expr) -> Expr,
) -> Expr {
let Expr::Match(match_expr) = match_expr else {
unreachable!()
};
let mut operands = Vec::new();
let patterns: Vec<_> = match_expr.arms.iter().map(|a| a.pattern.clone()).collect();
for arm in match_expr.arms {
let arm_expr = if arm.expr == *match_expr.subject {
Expr::from(arm.pattern.clone())
} else {
arm.expr
};
let guard = Expr::binary_op(
(*match_expr.subject).clone(),
stmt::BinaryOp::Eq,
Expr::from(arm.pattern),
);
let mut term = Expr::and_from_vec(vec![guard, term(arm_expr)]);
self.visit_expr_mut(&mut term);
if is_dead_filter_term(&term) {
continue;
}
operands.push(term);
}
{
let guards: Vec<Expr> = patterns
.into_iter()
.map(|pattern| {
Expr::not(Expr::binary_op(
(*match_expr.subject).clone(),
stmt::BinaryOp::Eq,
Expr::from(pattern),
))
})
.collect();
let mut else_operands = guards;
else_operands.push(term(*match_expr.else_expr));
let mut term = Expr::and_from_vec(else_operands);
self.visit_expr_mut(&mut term);
if !is_dead_filter_term(&term) {
operands.push(term);
}
}
Expr::or_from_vec(operands)
}
}
fn is_dead_filter_term(term: &Expr) -> bool {
if term.is_unsatisfiable() {
return true;
}
if contains_error(term) {
return true;
}
matches!(
term,
Expr::And(and) if and.operands.iter().any(Expr::is_value_null)
)
}
fn contains_error(expr: &Expr) -> bool {
use toasty_core::stmt::Visit;
struct FindError(bool);
impl Visit for FindError {
fn visit_expr_error(&mut self, _: &stmt::ExprError) {
self.0 = true;
}
}
let mut find = FindError(false);
find.visit_expr(expr);
find.0
}
fn record_field_comparison(lhs: Expr, op: stmt::BinaryOp, rhs: Expr) -> Expr {
let other = if lhs.is_value_null() {
rhs
} else if rhs.is_value_null() {
lhs
} else {
return Expr::binary_op(lhs, op, rhs);
};
if op.is_eq() {
Expr::is_null(other)
} else {
Expr::is_not_null(other)
}
}