use crate::{precedence, Error, Expr, Geometry};
use pg_escape::quote_identifier;
use sqlparser::ast::{
Array as SqlArray, BinaryOperator, CastKind,
DataType::{Date, Timestamp},
Expr as SqlExpr,
Expr::{Cast, Nested, Value as ValExpr},
FunctionArgumentList, FunctionArguments, Ident, TimezoneInfo, Value,
};
pub trait ToSqlAst {
fn to_sql_ast(&self) -> Result<SqlExpr, Error>;
fn to_sql(&self) -> Result<String, Error>;
}
fn cast(arg: SqlExpr, data_type: sqlparser::ast::DataType) -> SqlExpr {
Cast {
expr: Box::new(arg),
data_type,
kind: CastKind::Cast,
format: None,
array: false,
}
}
pub(crate) fn func(name: &str, args: Vec<SqlExpr>) -> Result<SqlExpr, Error> {
Ok(SqlExpr::Function(sqlparser::ast::Function {
name: sqlparser::ast::ObjectName(vec![sqlparser::ast::ObjectNamePart::Identifier(
ident_inner(name)?,
)]),
args: FunctionArguments::List(FunctionArgumentList {
duplicate_treatment: None,
args: args
.into_iter()
.map(|arg| {
sqlparser::ast::FunctionArg::Unnamed(sqlparser::ast::FunctionArgExpr::Expr(arg))
})
.collect(),
clauses: vec![],
}),
over: None,
filter: None,
null_treatment: None,
within_group: vec![],
uses_odbc_syntax: false,
parameters: FunctionArguments::None,
}))
}
fn lit_expr(value: &str) -> SqlExpr {
let needs_escaping = value.contains('\'') || value.contains('\\');
ValExpr(if needs_escaping {
Value::EscapedStringLiteral(value.to_string()).into()
} else {
Value::SingleQuotedString(value.to_string()).into()
})
}
fn float_expr(value: &f64) -> SqlExpr {
if value.is_finite() {
return ValExpr(Value::Number(value.to_string(), false).into());
}
let name = if value.is_nan() {
"NaN"
} else if value.is_sign_positive() {
"Infinity"
} else {
"-Infinity"
};
cast(
lit_expr(name),
sqlparser::ast::DataType::Double(sqlparser::ast::ExactNumberInfo::None),
)
}
fn args2ast(args: &[Box<Expr>]) -> Result<Vec<SqlExpr>, Error> {
args.iter()
.map(|arg| arg.to_sql_ast())
.collect::<Result<Vec<_>, _>>()
}
enum Arity {
Exactly(usize),
AtLeast(usize),
Any,
}
fn sql_arity(op: &str) -> Arity {
match op {
"isNull" | "not" => Arity::Exactly(1),
"between" => Arity::Exactly(3),
"in" | "like" | "=" | "a_equals" | "<>" | ">" | ">=" | "<" | "<=" | "^" | "a_contains"
| "a_containedBy" | "a_overlaps" => Arity::Exactly(2),
"+" | "-" | "*" | "/" | "%" => Arity::AtLeast(2),
_ => Arity::Any,
}
}
fn sql_operands(op: &str) -> precedence::Operands {
match op {
"^" => precedence::Operands { first: 0, rest: 0 },
_ => precedence::operands(op),
}
}
fn args2ast_grouped(op: &str, args: &[Box<Expr>]) -> Result<Vec<SqlExpr>, Error> {
let requirement = sql_operands(op);
args.iter()
.enumerate()
.map(|(index, arg)| {
let ast = arg.to_sql_ast()?;
Ok(if requirement.needs_parens(index, arg) {
wrap(ast)
} else {
ast
})
})
.collect::<Result<Vec<_>, _>>()
}
const LIKE_ESCAPE: char = '\\';
fn binop(op: BinaryOperator, args: Vec<SqlExpr>) -> SqlExpr {
let [left, right] = args
.try_into()
.expect("sql_arity checked the operand count");
cmp(left, op, right)
}
fn set_equality(args: Vec<SqlExpr>) -> SqlExpr {
let [left, right] = args
.try_into()
.expect("sql_arity checked the operand count");
wrap(andop(vec![
cmp(left.clone(), BinaryOperator::AtArrow, right.clone()),
cmp(right, BinaryOperator::AtArrow, left),
]))
}
struct Targs {
left_start: SqlExpr,
left_end: SqlExpr,
right_start: SqlExpr,
right_end: SqlExpr,
}
fn lit_or_prop_to_ts(arg: &Expr, unbounded: &str) -> Result<SqlExpr, Error> {
Ok(match arg {
Expr::Property { property } => ident(property)?,
Expr::Literal(v) => cast(
lit_expr(if v == ".." { unbounded } else { v }),
Timestamp(None, TimezoneInfo::WithTimeZone),
),
_ => return Err(Error::OperationError()),
})
}
fn lit_or_prop_to_date(arg: &Expr) -> Result<SqlExpr, Error> {
Ok(match arg {
Expr::Property { property } => ident(property)?,
Expr::Literal(v) => cast(lit_expr(v), Date),
_ => return Err(Error::OperationError()),
})
}
fn interval_bounds(interval: &[Box<Expr>]) -> Result<(&Expr, &Expr), Error> {
match interval {
[start, end] => Ok((start, end)),
_ => Err(Error::InvalidNumberOfArguments {
name: "interval".to_string(),
actual: interval.len(),
expected: 2,
}),
}
}
fn interval_endpoints(interval: &[Box<Expr>]) -> Result<(SqlExpr, SqlExpr), Error> {
let (lo, hi) = interval_bounds(interval)?;
Ok((
lit_or_prop_to_ts(lo, "-infinity")?,
lit_or_prop_to_ts(hi, "infinity")?,
))
}
fn timestamp_literal(ts: jiff::Timestamp) -> SqlExpr {
cast(
lit_expr(&ts.to_string()),
Timestamp(None, TimezoneInfo::WithTimeZone),
)
}
fn t_arg_to_interval(arg: &Expr) -> Result<(SqlExpr, SqlExpr), Error> {
match arg {
Expr::Interval { interval } => interval_endpoints(interval),
Expr::Property { property } => {
let start = ident(property)?;
Ok((start.clone(), start))
}
Expr::Date { date } => {
let day = crate::temporal::DateRange::try_from(Expr::Date { date: date.clone() })?;
Ok((timestamp_literal(day.start), timestamp_literal(day.end)))
}
Expr::Timestamp { timestamp } => {
let start = lit_or_prop_to_ts(timestamp, "infinity")?;
Ok((start.clone(), start))
}
_ => Err(Error::OperationError()),
}
}
fn t_args(args: &[Box<Expr>]) -> Result<Targs, Error> {
let [left, right] = args else {
return Err(Error::InvalidNumberOfArguments {
name: "temporal predicate".to_string(),
actual: args.len(),
expected: 2,
});
};
let (left_start, left_end) = t_arg_to_interval(left)?;
let (right_start, right_end) = t_arg_to_interval(right)?;
Ok(Targs {
left_start,
left_end,
right_start,
right_end,
})
}
fn temporal_sql(op: &str, args: &[Box<Expr>]) -> Result<SqlExpr, Error> {
let t = t_args(args)?;
Ok(match op {
"t_before" => ltop(t.left_end, t.right_start),
"t_after" => ltop(t.right_end, t.left_start),
"t_meets" => eqop(t.left_end, t.right_start),
"t_metBy" => eqop(t.right_end, t.left_start),
"t_overlaps" => wrap(andop(vec![
ltop(t.left_start, t.right_start.clone()),
ltop(t.right_start, t.left_end.clone()),
ltop(t.left_end, t.right_end),
])),
"t_overlappedBy" => wrap(andop(vec![
ltop(t.right_start, t.left_start.clone()),
ltop(t.left_start, t.right_end.clone()),
ltop(t.right_end, t.left_end),
])),
"t_starts" => wrap(andop(vec![
eqop(t.left_start, t.right_start.clone()),
ltop(t.left_end, t.right_end),
])),
"t_startedBy" => wrap(andop(vec![
eqop(t.right_start, t.left_start.clone()),
ltop(t.right_end, t.left_end),
])),
"t_during" => wrap(andop(vec![
gtop(t.left_start, t.right_start),
ltop(t.left_end, t.right_end),
])),
"t_contains" => wrap(andop(vec![
gtop(t.right_start, t.left_start),
ltop(t.right_end, t.left_end),
])),
"t_finishes" => wrap(andop(vec![
eqop(t.left_end, t.right_end),
gtop(t.left_start, t.right_start),
])),
"t_finishedBy" => wrap(andop(vec![
eqop(t.right_end, t.left_end),
gtop(t.right_start, t.left_start),
])),
"t_equals" => wrap(andop(vec![
eqop(t.left_start, t.right_start),
eqop(t.left_end, t.right_end),
])),
"t_disjoint" => wrap(notop(wrap(andop(vec![
lteop(t.left_start, t.right_end),
gteop(t.left_end, t.right_start),
])))),
"t_intersects" => wrap(andop(vec![
lteop(t.left_start, t.right_end),
gteop(t.left_end, t.right_start),
])),
_ => return Err(Error::InvalidOperator(op.to_string())),
})
}
fn chainop(op: BinaryOperator, args: Vec<SqlExpr>) -> Result<SqlExpr, Error> {
let name = op.to_string().to_lowercase();
args.into_iter()
.reduce(|left, right| SqlExpr::BinaryOp {
left: Box::new(left),
op: op.clone(),
right: Box::new(right),
})
.ok_or(Error::InvalidNumberOfArguments {
name,
actual: 0,
expected: 1,
})
}
fn andop(args: Vec<SqlExpr>) -> SqlExpr {
chainop(BinaryOperator::And, args).expect("callers supply at least one operand")
}
fn cmp(left: SqlExpr, op: BinaryOperator, right: SqlExpr) -> SqlExpr {
SqlExpr::BinaryOp {
left: Box::new(left),
op,
right: Box::new(right),
}
}
fn ltop(left: SqlExpr, right: SqlExpr) -> SqlExpr {
cmp(left, BinaryOperator::Lt, right)
}
fn gtop(left: SqlExpr, right: SqlExpr) -> SqlExpr {
cmp(left, BinaryOperator::Gt, right)
}
fn lteop(left: SqlExpr, right: SqlExpr) -> SqlExpr {
cmp(left, BinaryOperator::LtEq, right)
}
fn gteop(left: SqlExpr, right: SqlExpr) -> SqlExpr {
cmp(left, BinaryOperator::GtEq, right)
}
fn eqop(left: SqlExpr, right: SqlExpr) -> SqlExpr {
cmp(left, BinaryOperator::Eq, right)
}
fn notop(arg: SqlExpr) -> SqlExpr {
SqlExpr::UnaryOp {
op: sqlparser::ast::UnaryOperator::Not,
expr: Box::new(arg),
}
}
fn wrap(arg: SqlExpr) -> SqlExpr {
Nested(Box::new(arg))
}
fn ident_inner(property: &str) -> Result<Ident, Error> {
if property.is_empty() {
return Err(Error::EmptySqlIdentifier);
}
let p = quote_identifier(property);
Ok(if p.starts_with('"') && p.ends_with('"') {
Ident::with_quote('"', p[1..p.len() - 1].to_string())
} else {
Ident::new(p)
})
}
fn ident(property: &str) -> Result<SqlExpr, Error> {
Ok(SqlExpr::Identifier(ident_inner(property)?))
}
impl ToSqlAst for Expr {
fn to_sql_ast(&self) -> Result<SqlExpr, Error> {
Ok(match self {
Expr::Bool(v) => ValExpr(Value::Boolean(*v).into()),
Expr::Float(v) => float_expr(v),
Expr::Literal(v) => lit_expr(v),
Expr::Date { ref date } => lit_or_prop_to_date(date.as_ref())?,
Expr::Timestamp { ref timestamp } => lit_or_prop_to_ts(timestamp.as_ref(), "infinity")?,
Expr::Interval { ref interval } => {
let (start, end) = interval_endpoints(interval)?;
SqlExpr::Array(SqlArray {
elem: vec![start, end],
named: true,
})
}
Expr::Null => ValExpr(Value::Null.into()),
Expr::Geometry(v) => match v {
Geometry::GeoJSON(v) => {
let s = lit_expr(&v.to_string());
func("st_geomfromgeojson", vec![s])?
}
Geometry::Wkt(v) => {
let s = lit_expr(&v.to_string());
func("st_geomfromtext", vec![s])?
}
},
Expr::BBox { bbox } => func("st_makeenvelope", args2ast(bbox)?)?,
Expr::Array(ref v) => SqlExpr::Array(SqlArray {
elem: args2ast(v)?,
named: true,
}),
Expr::Property { property } => ident(property)?,
Expr::Operation { op, args } => {
let canonical = crate::expr::canonical_op(op);
let op_str = canonical.as_str();
let a = args2ast_grouped(op_str, args)?;
let required = match sql_arity(op_str) {
Arity::Exactly(expected) if a.len() != expected => Some(expected),
Arity::AtLeast(minimum) if a.len() < minimum => Some(minimum),
_ => None,
};
if let Some(expected) = required {
return Err(Error::InvalidNumberOfArguments {
name: op_str.to_string(),
actual: a.len(),
expected,
});
}
match op_str {
"isNull" => SqlExpr::IsNull(Box::new(a[0].clone())),
"not" => notop(a[0].clone()),
"between" => SqlExpr::Between {
expr: Box::new(a[0].clone()),
negated: false,
low: Box::new(a[1].clone()),
high: Box::new(a[2].clone()),
},
"in" => {
let expr = a[0].clone();
let items = a[1].clone();
SqlExpr::AnyOp {
left: Box::new(expr),
compare_op: BinaryOperator::Eq,
right: Box::new(items),
is_some: true,
}
}
"like" => {
let expr = a[0].clone();
let pattern = a[1].clone();
SqlExpr::Like {
expr: Box::new(expr),
pattern: Box::new(pattern),
escape_char: Some(
Value::SingleQuotedString(LIKE_ESCAPE.to_string()).into(),
),
negated: false,
any: false,
}
}
"accenti" => func("strip_accents", a)?,
"casei" => func("lower", a)?,
"and" => chainop(BinaryOperator::And, a)?,
"or" => chainop(BinaryOperator::Or, a)?,
"=" => binop(BinaryOperator::Eq, a),
"a_equals" => set_equality(a),
"<>" => binop(BinaryOperator::NotEq, a),
">" => binop(BinaryOperator::Gt, a),
">=" => binop(BinaryOperator::GtEq, a),
"<" => binop(BinaryOperator::Lt, a),
"<=" => binop(BinaryOperator::LtEq, a),
"+" => chainop(BinaryOperator::Plus, a)?,
"-" => chainop(BinaryOperator::Minus, a)?,
"*" => chainop(BinaryOperator::Multiply, a)?,
"/" => chainop(BinaryOperator::Divide, a)?,
"%" => chainop(BinaryOperator::Modulo, a)?,
"^" => func("power", a)?,
"s_intersects" => func("st_intersects", a)?,
"s_equals" => func("st_equals", a)?,
"s_within" => func("st_within", a)?,
"s_contains" => func("st_contains", a)?,
"s_crosses" => func("st_crosses", a)?,
"s_overlaps" => func("st_overlaps", a)?,
"s_touches" => func("st_touches", a)?,
"s_disjoint" => func("st_disjoint", a)?,
"a_contains" => binop(BinaryOperator::AtArrow, a),
"a_containedBy" => binop(BinaryOperator::ArrowAt, a),
"a_overlaps" => binop(BinaryOperator::AtAt, a),
name if crate::expr::TEMPORALOPS.contains(&name) => temporal_sql(name, args)?,
_ => func(&canonical, a)?,
}
}
})
}
fn to_sql(&self) -> Result<String, Error> {
Ok(self.to_sql_ast()?.to_string())
}
}
#[cfg(test)]
mod tests {
use super::ToSqlAst;
use crate::Expr;
#[test]
fn test_basic_expression() {
let expr: Expr = "1 + 2 > 4".parse().unwrap();
let sql_ast = expr.to_sql_ast().unwrap();
let sql_str = sql_ast.to_string();
assert_eq!(sql_str, "1 + 2 > 4");
}
#[test]
fn test_t_before_expression() {
let expr: Expr = "t_before(ts_start, DATE('2020-02-01'))".parse().unwrap();
let sql_ast = expr.to_sql_ast().expect("to_sql_ast failed");
let sql_str = sql_ast.to_string();
assert_eq!(
sql_str,
"ts_start < CAST('2020-02-01T00:00:00Z' AS TIMESTAMP WITH TIME ZONE)"
);
}
#[test]
fn test_bbox() {
let expr: Expr = "bbox(1, 2, 3, 4)".parse().unwrap();
assert_eq!(expr.to_sql().unwrap(), "st_makeenvelope(1, 2, 3, 4)");
}
#[test]
fn empty_property_name_is_rejected() {
let expr = Expr::Operation {
op: "=".to_string(),
args: vec![
Box::new(Expr::Property {
property: String::new(),
}),
Box::new(Expr::Float(1.0)),
],
};
assert!(matches!(
expr.to_sql(),
Err(crate::Error::EmptySqlIdentifier)
));
}
#[test]
fn non_finite_numbers_render_as_cast_literals() {
for (value, name) in [
(f64::INFINITY, "Infinity"),
(f64::NEG_INFINITY, "-Infinity"),
(f64::NAN, "NaN"),
] {
let sql = Expr::Float(value).to_sql().expect("renders as SQL");
assert_eq!(sql, format!("CAST('{name}' AS DOUBLE)"));
}
let divided: Expr = "1 / 0".parse().unwrap();
assert_eq!(
divided.reduce(None).unwrap().to_sql().expect("renders"),
"CAST('Infinity' AS DOUBLE)"
);
}
}