use core::fmt;
use std::fmt::{Display, Formatter};
use derive_more::{From, TryUnwrap, Unwrap};
use extract::ExtractFunction;
use literal::PatternVariableReference;
use crate::sql::expression::apply::{
BinaryOperatorApply, FunctionCallApply, Lambda, UnaryOperatorApply,
};
use crate::sql::expression::identifier::Identifier;
use crate::sql::expression::literal::{
ArrayLiteral, BinaryLiteral, BooleanLiteral, ColumnReference, DecimalLiteral, IntegerLiteral,
IntervalLiteral, NullLiteral, RowLiteral, ScientificLiteral, StringLiteral, TimestampLiteral,
TupleLiteral, UnicodeStringLiteral,
};
use crate::sql::expression::regexp::{RegexpCountFunction, RegexpExtractFunction};
use crate::sql::query::window::WindowExpression;
use crate::sql::types::SQLType;
use crate::write_utils::{maybe_newline, maybe_pad, newline_or_space, Indent};
use super::query::window::{NamedWindowReference, WindowReference, WindowSpecification};
pub mod apply;
pub mod extract;
pub mod identifier;
pub mod literal;
pub mod operator;
pub mod precedence;
pub mod regexp;
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum Leaf {
BinaryLiteral(BinaryLiteral),
BooleanLiteral(BooleanLiteral),
ColumnReference(ColumnReference),
PatternVariableReference(PatternVariableReference),
DecimalLiteral(DecimalLiteral),
IntegerLiteral(IntegerLiteral),
NullLiteral(NullLiteral),
SQLIntervalLiteral(IntervalLiteral),
ScientificLiteral(ScientificLiteral),
StringLiteral(StringLiteral),
TimestampLiteral(TimestampLiteral),
UnicodeStringLiteral(UnicodeStringLiteral),
}
impl Leaf {
pub fn fmt_indented(&self, f: &mut Formatter<'_>, indentation: Indent) -> fmt::Result {
match self {
Leaf::BinaryLiteral(l) => l.fmt_indented(f, indentation),
Leaf::BooleanLiteral(l) => l.fmt_indented(f, indentation),
Leaf::ColumnReference(l) => l.fmt_indented(f, indentation),
Leaf::PatternVariableReference(l) => l.fmt_indented(f, indentation),
Leaf::DecimalLiteral(l) => l.fmt_indented(f, indentation),
Leaf::IntegerLiteral(l) => l.fmt_indented(f, indentation),
Leaf::NullLiteral(l) => l.fmt_indented(f, indentation),
Leaf::SQLIntervalLiteral(l) => l.fmt_indented(f, indentation),
Leaf::ScientificLiteral(l) => l.fmt_indented(f, indentation),
Leaf::StringLiteral(l) => l.fmt_indented(f, indentation),
Leaf::TimestampLiteral(l) => l.fmt_indented(f, indentation),
Leaf::UnicodeStringLiteral(l) => l.fmt_indented(f, indentation),
}
}
}
impl Display for Leaf {
fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
self.fmt_indented(f, Indent::default())
}
}
impl From<BinaryLiteral> for SQLExpression {
fn from(l: BinaryLiteral) -> Self {
SQLExpression::Leaf(Leaf::BinaryLiteral(l))
}
}
impl From<BooleanLiteral> for SQLExpression {
fn from(l: BooleanLiteral) -> Self {
SQLExpression::Leaf(Leaf::BooleanLiteral(l))
}
}
impl From<ColumnReference> for SQLExpression {
fn from(l: ColumnReference) -> Self {
SQLExpression::Leaf(Leaf::ColumnReference(l))
}
}
impl From<PatternVariableReference> for SQLExpression {
fn from(l: PatternVariableReference) -> Self {
SQLExpression::Leaf(Leaf::PatternVariableReference(l))
}
}
impl From<DecimalLiteral> for SQLExpression {
fn from(l: DecimalLiteral) -> Self {
SQLExpression::Leaf(Leaf::DecimalLiteral(l))
}
}
impl From<IntegerLiteral> for SQLExpression {
fn from(l: IntegerLiteral) -> Self {
SQLExpression::Leaf(Leaf::IntegerLiteral(l))
}
}
impl From<NullLiteral> for SQLExpression {
fn from(l: NullLiteral) -> Self {
SQLExpression::Leaf(Leaf::NullLiteral(l))
}
}
impl From<IntervalLiteral> for SQLExpression {
fn from(l: IntervalLiteral) -> Self {
SQLExpression::Leaf(Leaf::SQLIntervalLiteral(l))
}
}
impl From<ScientificLiteral> for SQLExpression {
fn from(l: ScientificLiteral) -> Self {
SQLExpression::Leaf(Leaf::ScientificLiteral(l))
}
}
impl From<StringLiteral> for SQLExpression {
fn from(l: StringLiteral) -> Self {
SQLExpression::Leaf(Leaf::StringLiteral(l))
}
}
impl From<TimestampLiteral> for SQLExpression {
fn from(l: TimestampLiteral) -> Self {
SQLExpression::Leaf(Leaf::TimestampLiteral(l))
}
}
impl From<UnicodeStringLiteral> for SQLExpression {
fn from(l: UnicodeStringLiteral) -> Self {
SQLExpression::Leaf(Leaf::UnicodeStringLiteral(l))
}
}
#[derive(Debug, Clone, From, TryUnwrap, Unwrap, PartialEq, Eq)]
pub enum SQLExpression {
ArrayLiteral(ArrayLiteral),
Case(Case),
Cast(Cast),
Dot(Dot),
FunctionCallApply(FunctionCallApply),
Lambda(Lambda),
Leaf(Leaf),
OrderByExpression(OrderByExpression),
RegexpCountFunction(RegexpCountFunction),
RegexpExtractFunction(RegexpExtractFunction),
ExtractFunction(ExtractFunction),
RowLiteral(RowLiteral),
BinaryOperatorApply(BinaryOperatorApply),
SQLIndexLookup(IndexLookup),
UnaryOperatorApply(UnaryOperatorApply),
TryCast(TryCast),
WindowExpression(WindowExpression),
TupleLiteral(TupleLiteral),
}
impl SQLExpression {
pub fn get_column_references(&self) -> Vec<ColumnReference> {
match self {
SQLExpression::Leaf(Leaf::ColumnReference(c)) => vec![c.clone()],
SQLExpression::Leaf(_) => vec![],
SQLExpression::BinaryOperatorApply(BinaryOperatorApply { left, right, .. }) => {
let mut ret = left.get_column_references();
ret.extend(right.get_column_references());
ret
}
SQLExpression::UnaryOperatorApply(u) => u.operand.get_column_references(),
SQLExpression::Cast(Cast { expression, .. }) => expression.get_column_references(),
SQLExpression::RowLiteral(rl) => rl
.values
.iter()
.flat_map(|e| e.get_column_references())
.collect(),
SQLExpression::RegexpExtractFunction(r) => r.value.get_column_references(),
SQLExpression::RegexpCountFunction(r) => r.value.get_column_references(),
SQLExpression::Case(c) => {
let mut ret = c
.else_expression
.as_ref()
.map(|e| e.get_column_references())
.unwrap_or(vec![]);
for (condition, result) in c.when_expressions.iter() {
ret.extend(condition.get_column_references());
ret.extend(result.get_column_references());
}
ret
}
SQLExpression::FunctionCallApply(function_call_apply) => function_call_apply
.arguments
.iter()
.flat_map(|arg| arg.get_column_references())
.collect(),
SQLExpression::Dot(dot) => dot.expression.as_ref().get_column_references(),
SQLExpression::ArrayLiteral(a) => a
.elements
.iter()
.flat_map(|e| e.get_column_references())
.collect(),
SQLExpression::TryCast(TryCast { expression, .. }) => {
expression.get_column_references()
}
SQLExpression::SQLIndexLookup(s) => {
let mut ret = s.expression.get_column_references();
ret.extend(s.index.get_column_references());
ret
}
SQLExpression::Lambda(Lambda { body, .. }) => body.get_column_references(),
SQLExpression::WindowExpression(w) => w.expression.get_column_references(),
SQLExpression::OrderByExpression(o) => o.expression.get_column_references(),
SQLExpression::ExtractFunction(ef) => ef.value.get_column_references(),
SQLExpression::TupleLiteral(t) => t
.elements
.iter()
.flat_map(|e| e.get_column_references())
.collect(),
}
}
pub fn dot(self, rhs: Identifier) -> SQLExpression {
SQLExpression::Dot(Dot::new(self, rhs))
}
pub fn index(self, index: SQLExpression) -> SQLExpression {
SQLExpression::FunctionCallApply(FunctionCallApply::with_two("element_at", self, index))
}
pub fn cast(self, type_: SQLType) -> SQLExpression {
SQLExpression::Cast(Cast::new(self, type_))
}
pub fn fmt_indented(&self, f: &mut Formatter<'_>, indentation: Indent) -> fmt::Result {
match self {
SQLExpression::ArrayLiteral(a) => a.fmt_indented(f, indentation),
SQLExpression::Case(c) => c.fmt_indented(f, indentation),
SQLExpression::Cast(c) => c.fmt_indented(f, indentation),
SQLExpression::Dot(d) => d.fmt_indented(f, indentation),
SQLExpression::FunctionCallApply(fc) => fc.fmt_indented(f, indentation),
SQLExpression::Lambda(l) => l.fmt_indented(f, indentation),
SQLExpression::Leaf(l) => l.fmt_indented(f, indentation),
SQLExpression::OrderByExpression(o) => o.fmt_indented(f, indentation),
SQLExpression::RegexpCountFunction(r) => r.fmt_indented(f, indentation),
SQLExpression::RegexpExtractFunction(r) => r.fmt_indented(f, indentation),
SQLExpression::RowLiteral(r) => r.fmt_indented(f, indentation),
SQLExpression::BinaryOperatorApply(b) => b.fmt_indented(f, indentation),
SQLExpression::SQLIndexLookup(i) => i.fmt_indented(f, indentation),
SQLExpression::UnaryOperatorApply(u) => u.fmt_indented(f, indentation),
SQLExpression::TryCast(t) => t.fmt_indented(f, indentation),
SQLExpression::WindowExpression(w) => w.fmt_indented(f, indentation),
SQLExpression::ExtractFunction(ef) => ef.fmt_indented(f, indentation),
SQLExpression::TupleLiteral(t) => t.fmt_indented(f, indentation),
}
}
pub fn subelements(&self) -> usize {
match self {
SQLExpression::Leaf(_) => 0,
SQLExpression::ArrayLiteral(array_literal) => array_literal.subelements(),
SQLExpression::RowLiteral(row_literal) => row_literal.subelements(),
SQLExpression::WindowExpression(window_expression) => window_expression.subelements(),
SQLExpression::Case(case) => case.subelements(),
SQLExpression::FunctionCallApply(function_call_apply) => {
function_call_apply.subelements()
}
SQLExpression::Cast(cast) => cast.expression.subelements(),
SQLExpression::TryCast(try_cast) => try_cast.expression.subelements(),
SQLExpression::Dot(dot) => dot.expression.subelements(),
SQLExpression::Lambda(lambda) => lambda.body.subelements(),
SQLExpression::OrderByExpression(order_by_expression) => {
order_by_expression.expression.subelements()
}
SQLExpression::RegexpCountFunction(regexp_count_function) => {
regexp_count_function.value.subelements()
}
SQLExpression::RegexpExtractFunction(regexp_extract_function) => {
regexp_extract_function.value.subelements()
}
SQLExpression::BinaryOperatorApply(binary_operator_apply) => {
binary_operator_apply.left.subelements() + binary_operator_apply.right.subelements()
}
SQLExpression::SQLIndexLookup(index_lookup) => index_lookup.expression.subelements(),
SQLExpression::UnaryOperatorApply(unary_operator_apply) => {
unary_operator_apply.operand.subelements()
}
SQLExpression::ExtractFunction(ef) => ef.value.subelements(),
SQLExpression::TupleLiteral(tuple_literal) => tuple_literal.subelements(),
}
}
pub fn map_leaves<F>(self, f: F) -> Self
where
F: Clone + FnOnce(Leaf) -> SQLExpression,
{
match self {
SQLExpression::ArrayLiteral(array_literal) => ArrayLiteral::new(
array_literal
.elements
.into_iter()
.map(|e| e.map_leaves(f.clone()))
.collect(),
)
.into(),
SQLExpression::Case(case) => Case {
when_expressions: case
.when_expressions
.into_iter()
.map(|(condition, value)| {
(
Box::new(condition.map_leaves(f.clone())),
Box::new(value.map_leaves(f.clone())),
)
})
.collect(),
else_expression: case
.else_expression
.map(|else_expression| Box::new(else_expression.map_leaves(f))),
}
.into(),
SQLExpression::Cast(cast) => Cast {
expression: Box::new(cast.expression.map_leaves(f)),
to: cast.to,
}
.into(),
SQLExpression::Dot(dot) => Dot {
expression: Box::new(dot.expression.map_leaves(f)),
identifier: dot.identifier,
}
.into(),
SQLExpression::Lambda(lambda) => Lambda {
arguments: lambda.arguments,
body: Box::new(lambda.body.map_leaves(f)),
}
.into(),
SQLExpression::OrderByExpression(order_by_expression) => OrderByExpression {
expression: Box::new(order_by_expression.expression.map_leaves(f)),
direction: order_by_expression.direction,
}
.into(),
SQLExpression::RegexpCountFunction(regexp_count_function) => RegexpCountFunction {
value: Box::new(regexp_count_function.value.map_leaves(f)),
regex: regexp_count_function.regex,
}
.into(),
SQLExpression::RegexpExtractFunction(regexp_extract_function) => {
RegexpExtractFunction {
value: Box::new(regexp_extract_function.value.map_leaves(f)),
regex: regexp_extract_function.regex,
index: regexp_extract_function.index,
}
.into()
}
SQLExpression::RowLiteral(row_literal) => RowLiteral {
values: row_literal
.values
.into_iter()
.map(|value| value.map_leaves(f.clone()))
.collect(),
}
.into(),
SQLExpression::BinaryOperatorApply(binary_operator_apply) => BinaryOperatorApply {
operator: binary_operator_apply.operator,
left: Box::new(binary_operator_apply.left.map_leaves(f.clone())),
right: Box::new(binary_operator_apply.right.map_leaves(f)),
}
.into(),
SQLExpression::SQLIndexLookup(index_lookup) => IndexLookup {
expression: Box::new(index_lookup.expression.map_leaves(f)),
index: index_lookup.index,
}
.into(),
SQLExpression::UnaryOperatorApply(unary_operator_apply) => UnaryOperatorApply {
operator: unary_operator_apply.operator,
operand: Box::new(unary_operator_apply.operand.map_leaves(f)),
}
.into(),
SQLExpression::TryCast(try_cast) => TryCast {
expression: Box::new(try_cast.expression.map_leaves(f)),
to: try_cast.to,
}
.into(),
SQLExpression::WindowExpression(window_expression) => WindowExpression {
expression: Box::new(window_expression.expression.map_leaves(f.clone())),
window_reference: match window_expression.window_reference {
WindowReference::NamedWindowReference(named_window_reference) => {
NamedWindowReference {
name: named_window_reference.name,
}
.into()
}
WindowReference::WindowSpecification(window_specification) => {
WindowSpecification {
partition_by: window_specification
.partition_by
.into_iter()
.map(|e| e.map_leaves(f.clone()))
.collect(),
order_by: window_specification
.order_by
.into_iter()
.map(|ws| {
OrderByExpression {
expression: ws.expression.map_leaves(f.clone()).into(),
direction: ws.direction,
}
.into()
})
.collect(),
frame: window_specification.frame,
}
.into()
}
},
}
.into(),
SQLExpression::FunctionCallApply(function_call_apply) => FunctionCallApply {
function_name: function_call_apply.function_name.clone(),
arguments: function_call_apply
.arguments
.into_iter()
.map(|argument| argument.map_leaves(f.clone()))
.collect(),
named_arguments: function_call_apply
.named_arguments
.into_iter()
.map(|(name, exp)| (name, exp.map_leaves(f.clone())))
.collect(),
order_by: function_call_apply.order_by,
ignore_nulls: function_call_apply.ignore_nulls,
distinct: function_call_apply.distinct,
}
.into(),
SQLExpression::Leaf(leaf) => f(leaf),
SQLExpression::ExtractFunction(ef) => ExtractFunction {
field: ef.field.clone(),
value: Box::new(ef.value.map_leaves(f.clone())),
}
.into(),
SQLExpression::TupleLiteral(tuple_literal) => TupleLiteral::new(
tuple_literal
.elements
.into_iter()
.map(|e| e.map_leaves(f.clone()))
.collect(),
)
.into(),
}
}
}
impl Display for SQLExpression {
fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
self.fmt_indented(f, Indent::default())
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct Case {
pub when_expressions: Vec<(Box<SQLExpression>, Box<SQLExpression>)>,
pub else_expression: Option<Box<SQLExpression>>,
}
impl Case {
pub fn new(
when_expressions: &[(SQLExpression, SQLExpression)],
else_expression: Option<SQLExpression>,
) -> Self {
Self {
when_expressions: when_expressions
.iter()
.map(|(c, r)| (Box::new(c.clone()), Box::new(r.clone())))
.collect(),
else_expression: else_expression.map(|e| Box::new(e)),
}
}
pub fn fmt_indented(&self, f: &mut Formatter<'_>, indentation: Indent) -> fmt::Result {
write!(f, "CASE")?;
newline_or_space(f, indentation)?;
let when_expr_count = self.when_expressions.len();
for (i, (when, then)) in self.when_expressions.iter().enumerate() {
maybe_pad(f, indentation.nested())?;
write!(f, "WHEN ")?;
when.fmt_indented(f, indentation.nested())?;
write!(f, " THEN ")?;
then.fmt_indented(f, indentation.nested())?;
if i != when_expr_count - 1 {
newline_or_space(f, indentation)?;
}
}
if let Some(else_expr) = &self.else_expression {
newline_or_space(f, indentation)?;
maybe_pad(f, indentation.nested())?;
write!(f, "ELSE ")?;
else_expr.fmt_indented(f, indentation.nested())?;
}
newline_or_space(f, indentation)?;
maybe_pad(f, indentation)?;
write!(f, "END")
}
fn subelements(&self) -> usize {
self.when_expressions
.iter()
.map(|(when, then)| when.subelements() + then.subelements() + 1)
.sum::<usize>()
+ self
.else_expression
.as_ref()
.map_or(0, |else_expr| else_expr.subelements() + 1)
}
}
impl Display for Case {
fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
self.fmt_indented(f, Indent::default())
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct Cast {
pub expression: Box<SQLExpression>,
pub to: SQLType,
}
impl Cast {
pub fn new(expression: SQLExpression, to: SQLType) -> Self {
Self {
expression: Box::new(expression),
to,
}
}
pub fn fmt_indented(&self, f: &mut Formatter<'_>, indentation: Indent) -> fmt::Result {
let indentation = match indentation {
Indent::Pretty { .. }
if self.to.subfields() > 3 || self.expression.subelements() > 3 =>
{
indentation
}
_ => Indent::Compact,
};
write!(f, "CAST(")?;
maybe_newline(f, indentation)?;
maybe_pad(f, indentation.nested())?;
self.expression.fmt_indented(f, indentation.nested())?;
newline_or_space(f, indentation)?;
maybe_pad(f, indentation.nested())?;
write!(f, "AS")?;
newline_or_space(f, indentation)?;
maybe_pad(f, indentation.nested())?;
self.to.fmt_indented(f, indentation.nested())?;
maybe_newline(f, indentation)?;
maybe_pad(f, indentation)?;
write!(f, ")")?;
Ok(())
}
}
impl Display for Cast {
fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
self.fmt_indented(f, Indent::default())
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct TryCast {
pub expression: Box<SQLExpression>,
pub to: SQLType,
}
impl TryCast {
pub fn new(expression: SQLExpression, to: SQLType) -> Self {
Self {
expression: Box::new(expression),
to,
}
}
pub fn fmt_indented(&self, f: &mut Formatter<'_>, indentation: Indent) -> fmt::Result {
let indentation = match indentation {
Indent::Pretty { .. }
if self.to.subfields() > 3 || self.expression.subelements() > 3 =>
{
indentation
}
_ => Indent::Compact,
};
write!(f, "TRY_CAST(")?;
maybe_newline(f, indentation)?;
maybe_pad(f, indentation.nested())?;
self.expression.fmt_indented(f, indentation.nested())?;
newline_or_space(f, indentation)?;
maybe_pad(f, indentation.nested())?;
write!(f, "AS")?;
newline_or_space(f, indentation)?;
maybe_pad(f, indentation.nested())?;
self.to.fmt_indented(f, indentation.nested())?;
maybe_newline(f, indentation)?;
maybe_pad(f, indentation)?;
write!(f, ")")?;
Ok(())
}
}
impl Display for TryCast {
fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
self.fmt_indented(f, Indent::default())
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct Dot {
pub expression: Box<SQLExpression>,
pub identifier: Identifier,
}
impl Dot {
pub fn new(expression: SQLExpression, identifier: Identifier) -> Self {
Self {
expression: Box::new(expression),
identifier,
}
}
pub fn to_sql_expression(self) -> SQLExpression {
SQLExpression::Dot(self)
}
pub fn fmt_indented(&self, f: &mut Formatter<'_>, _indentation: Indent) -> fmt::Result {
write!(f, "{}.{}", self.expression, self.identifier.to_string())
}
}
impl Display for Dot {
fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
self.fmt_indented(f, Indent::default())
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct IndexLookup {
pub expression: Box<SQLExpression>,
pub index: Box<SQLExpression>,
}
impl IndexLookup {
pub fn new(expression: SQLExpression, index: SQLExpression) -> Self {
Self {
expression: Box::new(expression),
index: Box::new(index),
}
}
pub fn to_sql_expression(self) -> SQLExpression {
SQLExpression::SQLIndexLookup(self)
}
pub fn fmt_indented(&self, f: &mut Formatter<'_>, _indentation: Indent) -> fmt::Result {
write!(f, "{}[{}]", self.expression, self.index)
}
}
impl Display for IndexLookup {
fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
self.fmt_indented(f, Indent::default())
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct OrderByExpression {
pub expression: Box<SQLExpression>,
pub direction: Direction,
}
impl OrderByExpression {
pub fn new(expression: SQLExpression, direction: Direction) -> Self {
Self {
expression: Box::new(expression),
direction,
}
}
pub fn fmt_indented(&self, f: &mut Formatter<'_>, _indentation: Indent) -> fmt::Result {
write!(
f,
"{} {}",
self.expression,
match self.direction {
Direction::ASC => "ASC",
Direction::DESC => "DESC",
}
)
}
}
impl Display for OrderByExpression {
fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
self.fmt_indented(f, Indent::default())
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum Direction {
ASC,
DESC,
}