diff --git a/Cargo.lock b/Cargo.lock index 73a4ad91aa16..1a1f218078c5 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -3767,6 +3767,7 @@ version = "0.55.1" dependencies = [ "bitflags", "hex", + "polars-compute", "polars-core", "polars-error", "polars-lazy", diff --git a/crates/polars-compute/src/decimal.rs b/crates/polars-compute/src/decimal.rs index adf4ede49bd5..0f4f85fc4309 100644 --- a/crates/polars-compute/src/decimal.rs +++ b/crates/polars-compute/src/decimal.rs @@ -465,6 +465,55 @@ pub fn dec128_to_i128(x: i128, s: usize) -> i128 { if s == 0 { x } else { div_128_pow10(x, s) } } +/// Returns `x * 10^e`, or None on overflow; `e` is not bounded by `DEC128_MAX_PREC`. +#[inline] +pub fn i128_mul_pow10(x: i128, mut e: usize) -> Option { + let mut r = x; + while e > DEC128_MAX_PREC { + r = r.checked_mul(POW10_I128[DEC128_MAX_PREC])?; + e -= DEC128_MAX_PREC; + } + r.checked_mul(POW10_I128[e]) +} + +/// Exact scalar arithmetic on `(mantissa, scale)` fixed-point values, bounded only by +/// `i128`: the result scale is the larger input scale for `add`/`sub` and the sum of the +/// input scales for `mul`, so no rounding takes place. Returns None on overflow. +pub mod exact { + use super::i128_mul_pow10; + + fn align(l: (i128, usize), r: (i128, usize)) -> Option<(i128, i128, usize)> { + let s = l.1.max(r.1); + Some(( + i128_mul_pow10(l.0, s - l.1)?, + i128_mul_pow10(r.0, s - r.1)?, + s, + )) + } + + #[inline] + pub fn add(l: (i128, usize), r: (i128, usize)) -> Option<(i128, usize)> { + let (l, r, s) = align(l, r)?; + Some((l.checked_add(r)?, s)) + } + + #[inline] + pub fn sub(l: (i128, usize), r: (i128, usize)) -> Option<(i128, usize)> { + let (l, r, s) = align(l, r)?; + Some((l.checked_sub(r)?, s)) + } + + #[inline] + pub fn mul(l: (i128, usize), r: (i128, usize)) -> Option<(i128, usize)> { + Some((l.0.checked_mul(r.0)?, l.1 + r.1)) + } + + #[inline] + pub fn neg(x: (i128, usize)) -> Option<(i128, usize)> { + Some((x.0.checked_neg()?, x.1)) + } +} + /// Converts an i128 to a Decimal128 with the given precision and scale, /// returning None if the value doesn't fit. #[inline] @@ -886,6 +935,29 @@ mod test { use super::*; + #[test] + fn test_exact_scalar_arithmetic() { + // 0.06 + 0.01 = 0.07, exact at scale 2 + assert_eq!(exact::add((6, 2), (1, 2)), Some((7, 2))); + // 1.5 - 3 aligns to the larger scale + assert_eq!(exact::sub((15, 1), (3, 0)), Some((-15, 1))); + // 1.10 * 1.10 = 1.2100 at the summed scale + assert_eq!(exact::mul((110, 2), (110, 2)), Some((12100, 4))); + assert_eq!(exact::neg((5, 1)), Some((-5, 1))); + + // scales beyond DEC128_MAX_PREC stay exact + assert_eq!(exact::mul((1, 38), (1, 38)), Some((1, 76))); + assert_eq!(exact::add((0, 0), (1, 76)), Some((1, 76))); + assert_eq!(i128_mul_pow10(0, 200), Some(0)); + + // overflow is reported rather than wrapped + assert_eq!(exact::add((i128::MAX, 0), (1, 0)), None); + assert_eq!(exact::add((1, 0), (1, 39)), None); + assert_eq!(exact::mul((i128::MAX, 0), (2, 0)), None); + assert_eq!(exact::neg((i128::MIN, 0)), None); + assert_eq!(i128_mul_pow10(1, 39), None); + } + fn bigdecimal_to_dec128(x: &BigDecimal, p: usize, s: usize) -> Option { let n = x .with_scale_round(s as i64, RoundingMode::HalfEven) diff --git a/crates/polars-plan/src/dsl/meta.rs b/crates/polars-plan/src/dsl/meta.rs index 9ffef892ca91..377688e2a1c3 100644 --- a/crates/polars-plan/src/dsl/meta.rs +++ b/crates/polars-plan/src/dsl/meta.rs @@ -159,4 +159,11 @@ impl MetaNameSpace { let ae = expr_arena.get(e_ir.node()); Ok(is_row_separable(&mut stack, ae, &expr_arena)) } + + /// Indicate if this expression yields one scalar value that does not depend on any + /// input frame (see [`is_input_independent_scalar_rec`]). + pub fn is_input_independent_scalar(self) -> PolarsResult { + let (e_ir, expr_arena) = self.into_expr_ir()?; + Ok(is_input_independent_scalar_rec(e_ir.node(), &expr_arena)) + } } diff --git a/crates/polars-plan/src/plans/aexpr/properties/general.rs b/crates/polars-plan/src/plans/aexpr/properties/general.rs index 2e4bbef9ced9..d02a80dcaf96 100644 --- a/crates/polars-plan/src/plans/aexpr/properties/general.rs +++ b/crates/polars-plan/src/plans/aexpr/properties/general.rs @@ -193,6 +193,25 @@ pub fn is_elementwise_rec(node: Node, expr_arena: &Arena) -> bool { property_rec(node, expr_arena, is_elementwise) } +/// Whether `node` yields one scalar value that does not depend on any input frame: a scalar +/// literal, or elementwise operations (binary, cast, ternary, elementwise function) over such +/// values. Columns, the frame length, aggregations, windows, nested evaluations and user +/// functions are excluded. +pub fn is_input_independent_scalar_rec(node: Node, expr_arena: &Arena) -> bool { + property_rec(node, expr_arena, |stack, ae, _| { + let independent = match ae { + AExpr::Literal(lv) => lv.is_scalar(), + AExpr::BinaryExpr { .. } | AExpr::Cast { .. } | AExpr::Ternary { .. } => true, + AExpr::Function { options, .. } => options.is_elementwise(), + _ => false, + }; + if independent { + ae.inputs_rev(stack); + } + independent + }) +} + /// Checks if the top-level expression node is row-separable. If this is the case, then `stack` will /// be extended further with any nested expression nodes. pub fn is_row_separable(stack: &mut UnitVec, ae: &AExpr, expr_arena: &Arena) -> bool { diff --git a/crates/polars-plan/src/plans/conversion/dsl_to_ir/join.rs b/crates/polars-plan/src/plans/conversion/dsl_to_ir/join.rs index 1d9902aa0076..c3985360b28b 100644 --- a/crates/polars-plan/src/plans/conversion/dsl_to_ir/join.rs +++ b/crates/polars-plan/src/plans/conversion/dsl_to_ir/join.rs @@ -3,6 +3,7 @@ use either::Either; use polars_core::chunked_array::cast::CastOptions; use polars_core::error::feature_gated; use polars_core::utils::{get_numeric_upcast_supertype_lossless, try_get_supertype}; +use polars_ops::prelude::JoinValidation; use polars_utils::format_pl_smallstr; use polars_utils::itertools::Itertools; @@ -156,6 +157,27 @@ pub fn resolve_join( let schema_left = ctxt.lp_arena.get(input_left).schema(ctxt.lp_arena); let schema_right = ctxt.lp_arena.get(input_right).schema(ctxt.lp_arena); + // Inner-joining on the same non-null constant on both sides pairs every row with every + // row: a cross join (unless the key multiplicity is to be validated). + let same_constant_key = |l: &ExprIR, r: &ExprIR| match ( + ctxt.expr_arena.get(l.node()), + ctxt.expr_arena.get(r.node()), + ) { + (AExpr::Literal(l), AExpr::Literal(r)) => l.is_scalar() && !l.is_null() && l == r, + _ => false, + }; + if options.args.how == JoinType::Inner + && options.args.validation == JoinValidation::ManyToMany + && left_on + .iter() + .zip(&right_on) + .all(|(l, r)| same_constant_key(l, r)) + { + options.args.how = JoinType::Cross; + left_on.clear(); + right_on.clear(); + } + // # Resolve scalars // // Scalars need to be expanded. We translate them to temporary columns added with diff --git a/crates/polars-sql/Cargo.toml b/crates/polars-sql/Cargo.toml index d6f318652fa8..e3f8aab10682 100644 --- a/crates/polars-sql/Cargo.toml +++ b/crates/polars-sql/Cargo.toml @@ -9,6 +9,7 @@ repository = { workspace = true } description = "SQL transpiler for Polars. Converts SQL to Polars logical plans" [dependencies] +polars-compute = { workspace = true, features = ["dtype-decimal"] } polars-core = { workspace = true, features = ["rows"] } polars-error = { workspace = true } polars-lazy = { workspace = true, features = ["abs", "binary_encoding", "concat_str", "cov", "cross_join", "cum_agg", "dtype-array", "dtype-date", "dtype-decimal", "dtype-struct", "is_in", "list_eval", "log", "meta", "offset_by", "range", "regex", "round_series", "sign", "string_normalize", "string_pad", "string_reverse", "strings", "timezones", "trigonometry"] } diff --git a/crates/polars-sql/src/context.rs b/crates/polars-sql/src/context.rs index b60430eeba5a..19af92462e9c 100644 --- a/crates/polars-sql/src/context.rs +++ b/crates/polars-sql/src/context.rs @@ -1,5 +1,5 @@ use std::borrow::Cow; -use std::ops::Deref; +use std::ops::{ControlFlow, Deref}; use std::sync::{Arc, RwLock}; use polars_core::frame::row::Row; @@ -20,15 +20,15 @@ use sqlparser::ast::{ OrderByKind, Query, RenameSelectItem, Select, SelectFlavor, SelectItem, SelectItemQualifiedWildcardKind, SetExpr, SetOperator, SetQuantifier, Statement, TableAlias, TableFactor, TableWithJoins, Truncate, UnaryOperator as SQLUnaryOperator, Value as SQLValue, - ValueWithSpan, Values, Visit, WildcardAdditionalOptions, WindowSpec, + ValueWithSpan, Values, Visit, WildcardAdditionalOptions, WindowSpec, visit_expressions_mut, }; use sqlparser::dialect::GenericDialect; use sqlparser::parser::{Parser, ParserOptions}; use crate::function_registry::{DefaultFunctionRegistry, FunctionRegistry}; use crate::sql_expr::{ - order_by_sort_options, parse_sql_array, parse_sql_expr, resolve_compound_identifier, - to_sql_interface_err, + convert_int_literal_for_string, order_by_sort_options, parse_sql_array, parse_sql_expr, + resolve_compound_identifier, to_sql_interface_err, }; use crate::sql_visitors::{ QualifyExpression, TableIdentifierCollector, check_for_ambiguous_column_refs, @@ -2146,28 +2146,14 @@ impl SQLContext { Some(s) => s, }; - // shortcut filter evaluation if given expression is just TRUE or FALSE - let (all_true, all_false) = match expr { - SQLExpr::Value(ValueWithSpan { - value: SQLValue::Boolean(b), - .. - }) => (*b, !*b), - SQLExpr::BinaryOp { left, op, right } => match (&**left, &**right, op) { - (SQLExpr::Value(a), SQLExpr::Value(b), SQLBinaryOperator::Eq) => { - (a.value == b.value, a.value != b.value) - }, - (SQLExpr::Value(a), SQLExpr::Value(b), SQLBinaryOperator::NotEq) => { - (a.value != b.value, a.value == b.value) - }, - _ => (false, false), - }, - _ => (false, false), - }; - let removing = filter_mode == FilterMode::RemoveTrue; - if (all_true && !removing) || (all_false && removing) { - return Ok(lf); - } else if (all_false && !removing) || (all_true && removing) { - return Ok(lf.clear()); + // A condition that reads no input (eg: "WHERE 1 = 1") is accepted as any type + // that casts to boolean; the planner folds it. + if let Some(predicate) = self.input_independent_predicate(expr)? { + let predicate = predicate.cast(DataType::Boolean); + return Ok(match filter_mode { + FilterMode::KeepTrue => lf.filter(predicate), + FilterMode::RemoveTrue => lf.remove(predicate), + }); } // Lower eligible `[NOT] EXISTS` / `[NOT] IN (subquery)` conjuncts @@ -2222,6 +2208,21 @@ impl SQLContext { Ok(lf) } + /// Parse a condition that yields one value independent of any input frame (a literal or + /// elementwise operations over literals); `None` for any other condition. + fn input_independent_predicate(&mut self, expr: &SQLExpr) -> PolarsResult> { + if expr_references_any_column(expr) || expr_contains_subquery(expr) { + return Ok(None); + } + let predicate = parse_sql_expr(expr, self, None)?; + let independent = predicate + .clone() + .meta() + .is_input_independent_scalar() + .unwrap_or(false); + Ok(independent.then_some(predicate)) + } + pub(super) fn process_join( &mut self, tbl_left: &TableInfo, @@ -2229,33 +2230,23 @@ impl SQLContext { constraint: &JoinConstraint, join_type: JoinType, ) -> PolarsResult { - if let JoinConstraint::On(expr) = constraint { - // A subquery references no column of its own, so it needs excluding here - // as well as it would otherwise read as a constant predicate. - if !expr_references_any_column(expr) && !expr_contains_subquery(expr) { - let satisfied = evaluate_constant_join_predicate(self, expr)?; - let builder = tbl_left - .frame - .clone() - .join_builder() - .with(tbl_right.frame.clone()) - .suffix(format!(":{}", tbl_right.name)) - .coalesce(JoinCoalesce::KeepColumns); - - // Only INNER: an always-true outer join still has to emit null-extended - // left rows when the right side is empty, which a cross join would not. - return Ok(if satisfied && join_type == JoinType::Inner { - builder.how(JoinType::Cross).finish()? - } else { - // Match every row against every row, or none against none. - let right_key = if satisfied { lit(1i32) } else { lit(2i32) }; - builder - .left_on([lit(1i32)]) - .right_on([right_key]) - .how(join_type) - .finish()? - }); - } + // A condition that reads no input (eg: `ON TRUE`, `ON 1 = 1`) pairs every row with + // every row, or none: join on the condition itself as a boolean key (a null key + // matches nothing, like an unknown condition). + if let JoinConstraint::On(expr) = constraint + && let Some(predicate) = self.input_independent_predicate(expr)? + { + return tbl_left + .frame + .clone() + .join_builder() + .with(tbl_right.frame.clone()) + .left_on([strip_join_aliases(predicate).cast(DataType::Boolean)]) + .right_on([lit(true)]) + .how(join_type) + .suffix(format!(":{}", tbl_right.name)) + .coalesce(JoinCoalesce::KeepColumns) + .finish(); } let (left_on, right_on, predicates) = process_join_constraint(constraint, tbl_left, tbl_right, self)?; @@ -3608,8 +3599,28 @@ fn determine_left_right_join_on( ) -> PolarsResult<(Vec, Vec)> { // parse, removing any aliases that may have been added by `resolve_column` // (called inside `parse_sql_expr`) as we need the actual/underlying col - let left_on = strip_join_aliases(parse_sql_expr(expr_left, ctx, Some(join_schema))?); - let right_on = strip_join_aliases(parse_sql_expr(expr_right, ctx, Some(join_schema))?); + let left_refs = ( + expr_refers_to_table(expr_left, &tbl_left.name), + expr_refers_to_table(expr_left, &tbl_right.name), + ); + let right_refs = ( + expr_refers_to_table(expr_right, &tbl_left.name), + expr_refers_to_table(expr_right, &tbl_right.name), + ); + // an operand's dtypes come from the table it names; the merged schema keeps the left + // dtype for a column name that exists in both tables + let operand_schema = |refs: (bool, bool)| -> &Schema { + match refs { + (true, false) => &tbl_left.schema, + (false, true) => &tbl_right.schema, + _ => join_schema, + } + }; + let (left_schema, right_schema) = (operand_schema(left_refs), operand_schema(right_refs)); + let left_on = strip_join_aliases(parse_sql_expr(expr_left, ctx, Some(left_schema))?); + let right_on = strip_join_aliases(parse_sql_expr(expr_right, ctx, Some(right_schema))?); + let (left_on, right_on) = + convert_int_literal_for_string((left_on, left_schema), (right_on, right_schema)); // a constant operand is a literal, or any other expression referencing no column (such as // `UPPER('it')`); it can be evaluated against either input, so it has no table affinity @@ -3621,14 +3632,6 @@ fn determine_left_right_join_on( // ------------------------------------------------------------------ // simple/typical case: can fully resolve SQL-level table references // ------------------------------------------------------------------ - let left_refs = ( - expr_refers_to_table(expr_left, &tbl_left.name), - expr_refers_to_table(expr_left, &tbl_right.name), - ); - let right_refs = ( - expr_refers_to_table(expr_right, &tbl_left.name), - expr_refers_to_table(expr_right, &tbl_right.name), - ); // if the SQL-level references unambiguously indicate table ownership, we're done match (left_refs, right_refs) { // standard: left expr → left table, right expr → right table @@ -3694,9 +3697,9 @@ fn determine_left_right_join_on( /// Returns `(left_on, right_on, join_where_predicates)`. /// /// - Equi-conditions (`=`) are returned as paired `left_on`/`right_on` entries. -/// - Non-equi conditions (`<`, `<=`, `>`, `>=`, `!=`) are returned as `join_where` predicates -/// that reference columns using their merged-schema names (right columns that conflict with the -/// left schema are suffixed). +/// - Any other condition is returned as a `join_where` predicate that references columns +/// using their merged-schema names (right columns that conflict with the left schema are +/// suffixed). fn process_join_on( ctx: &mut SQLContext, sql_expr: &SQLExpr, @@ -3727,57 +3730,52 @@ fn process_join_on( )?; Ok((l, r, vec![])) }, - SQLBinaryOperator::Lt - | SQLBinaryOperator::LtEq - | SQLBinaryOperator::Gt - | SQLBinaryOperator::GtEq - | SQLBinaryOperator::NotEq => { - let join_schema = build_join_schema(tbl_left, tbl_right)?; - let suffix = format!(":{}", tbl_right.name); - - // Parse both operands and suffix each independently based on whether - // it references the right table (preserving SQL operand order). - let lhs = suffix_if_right_table( - parse_sql_expr(left, ctx, Some(&join_schema))?, - left, - tbl_left, - tbl_right, - &suffix, - ); - let rhs = suffix_if_right_table( - parse_sql_expr(right, ctx, Some(&join_schema))?, - right, - tbl_left, - tbl_right, - &suffix, - ); - - let polars_op = match op { - SQLBinaryOperator::Lt => Operator::Lt, - SQLBinaryOperator::LtEq => Operator::LtEq, - SQLBinaryOperator::Gt => Operator::Gt, - SQLBinaryOperator::GtEq => Operator::GtEq, - SQLBinaryOperator::NotEq => Operator::NotEq, - _ => unreachable!(), - }; - let predicate = Expr::BinaryExpr { - left: Arc::new(lhs), - op: polars_op, - right: Arc::new(rhs), - }; - Ok((vec![], vec![], vec![predicate])) - }, - _ => polars_bail!( - SQLInterface: "unsupported join constraint operator '{:?}'", op - ), + _ => process_join_predicate(ctx, sql_expr, tbl_left, tbl_right), }, SQLExpr::Nested(expr) => process_join_on(ctx, expr, tbl_left, tbl_right), - _ => polars_bail!( - SQLInterface: "unsupported join constraint expression: {:?}", sql_expr - ), + _ => process_join_predicate(ctx, sql_expr, tbl_left, tbl_right), } } +/// Parse a join condition other than a plain equality into a `join_where` predicate over +/// the joined frame, where right-table columns that also exist in the left table carry a suffix. +fn process_join_predicate( + ctx: &mut SQLContext, + sql_expr: &SQLExpr, + tbl_left: &TableInfo, + tbl_right: &TableInfo, +) -> PolarsResult<(Vec, Vec, Vec)> { + let suffix = format!(":{}", tbl_right.name); + let conflicts = |name: &str| tbl_left.schema.contains(name) && tbl_right.schema.contains(name); + + let mut joined_schema = build_join_schema(tbl_left, tbl_right)?; + for (name, dtype) in tbl_right.schema.iter().filter(|(name, _)| conflicts(name)) { + joined_schema.insert( + PlSmallStr::from_string(format!("{name}{suffix}")), + dtype.clone(), + ); + } + + // `right_table.col` -> `col:right_table` when the name is also a left column + let mut sql_expr = sql_expr.clone(); + let _ = visit_expressions_mut(&mut sql_expr, |e| { + if let SQLExpr::CompoundIdentifier(idents) = e + && idents.len() >= 2 + && idents[0].value == tbl_right.name + && conflicts(&idents[1].value) + { + let suffixed = Ident::new(format!("{}{suffix}", idents[1].value)); + idents.splice(0..2, [suffixed]); + if idents.len() == 1 { + *e = SQLExpr::Identifier(idents.pop().unwrap()); + } + } + ControlFlow::<()>::Continue(()) + }); + let predicate = strip_join_aliases(parse_sql_expr(&sql_expr, ctx, Some(&joined_schema))?); + Ok((vec![], vec![], vec![predicate])) +} + /// Replace aggregates over pre-aggregation columns with references to hoisted /// aggregation outputs, collecting the hoisted aggregates into `agg_out`. /// @@ -3866,58 +3864,6 @@ fn suffix_conflicting_columns( }) } -/// Suffix conflicting column names in `expr` if the SQL-level expression references the right -/// table. Uses table qualifiers first, falling back to schema membership when unqualified. -fn suffix_if_right_table( - expr: Expr, - sql_expr: &SQLExpr, - tbl_left: &TableInfo, - tbl_right: &TableInfo, - suffix: &str, -) -> Expr { - // Strip any alias added by resolve_column - let expr = match expr { - Expr::Alias(inner, _) => Arc::unwrap_or_clone(inner), - e => e, - }; - - let refs_left = expr_refers_to_table(sql_expr, &tbl_left.name); - let refs_right = expr_refers_to_table(sql_expr, &tbl_right.name); - - let is_right = if refs_right && !refs_left { - true - } else if refs_left { - false - } else { - // Unqualified: check schema membership - !expr_cols_all_in_schema(&expr, &tbl_left.schema) - && expr_cols_all_in_schema(&expr, &tbl_right.schema) - }; - - if is_right { - suffix_conflicting_columns(expr, tbl_left, tbl_right, suffix) - } else { - expr - } -} - -/// Evaluate a column-free (constant) join ON-expression to a definite true/false -/// verdict; SQL treats an unknown (NULL) condition the same as false for matching. -fn evaluate_constant_join_predicate(ctx: &mut SQLContext, expr: &SQLExpr) -> PolarsResult { - let predicate = parse_sql_expr(expr, ctx, None)?; - let df = DataFrame::empty() - .lazy() - .select([predicate - .cast(DataType::Boolean) - .alias("_constant_join_predicate")]) - .collect()?; - Ok(df - .column("_constant_join_predicate")? - .bool()? - .get(0) - .unwrap_or(false)) -} - fn process_join_constraint( constraint: &JoinConstraint, tbl_left: &TableInfo, diff --git a/crates/polars-sql/src/functions.rs b/crates/polars-sql/src/functions.rs index 92c9b1fba4c2..4364a247f8e2 100644 --- a/crates/polars-sql/src/functions.rs +++ b/crates/polars-sql/src/functions.rs @@ -313,6 +313,12 @@ pub(crate) enum PolarsSQLFunctions { /// SELECT DATE_PART('year', col1) FROM df; /// SELECT DATE_PART('day', col1) FROM df; DatePart, + /// SQL date part accessor functions ('YEAR', 'MONTH', 'DAY', 'HOUR', etc). + /// Shorthand for DATE_PART with a fixed part. + /// ```sql + /// SELECT YEAR(col1), MONTH(col1), DAYOFWEEK(col1) FROM df; + /// ``` + DatePartOf(DateTimeField), /// SQL 'strftime' function. /// Converts a datetime to a string using a format string. /// ```sql @@ -861,6 +867,10 @@ impl PolarsSQLFunctions { "covar_samp", "date", "date_part", + "day", + "dayofmonth", + "dayofweek", + "dayofyear", "degrees", "dense_rank", "ends_with", @@ -869,6 +879,7 @@ impl PolarsSQLFunctions { "first_value", "floor", "greatest", + "hour", "if", "ifnull", "initcap", @@ -889,9 +900,10 @@ impl PolarsSQLFunctions { "ltrim", "max", "median", - "quantile_disc", "min", + "minute", "mod", + "month", "nullif", "octet_length", "pi", @@ -899,6 +911,7 @@ impl PolarsSQLFunctions { "power", "quantile_cont", "quantile_disc", + "quarter", "radians", "rank", "regexp_like", @@ -909,6 +922,7 @@ impl PolarsSQLFunctions { "row_number", "rpad", "rtrim", + "second", "sign", "sin", "sind", @@ -931,6 +945,8 @@ impl PolarsSQLFunctions { "var", "var_samp", "variance", + "week", + "year", ] } } @@ -1008,6 +1024,16 @@ impl PolarsSQLFunctions { // ---- "date" => Self::Date, "date_part" => Self::DatePart, + "year" => Self::DatePartOf(DateTimeField::Year), + "quarter" => Self::DatePartOf(DateTimeField::Quarter), + "month" => Self::DatePartOf(DateTimeField::Month), + "week" => Self::DatePartOf(DateTimeField::IsoWeek), + "day" | "dayofmonth" => Self::DatePartOf(DateTimeField::Day), + "dayofweek" => Self::DatePartOf(DateTimeField::DayOfWeek), + "dayofyear" => Self::DatePartOf(DateTimeField::DayOfYear), + "hour" => Self::DatePartOf(DateTimeField::Hour), + "minute" => Self::DatePartOf(DateTimeField::Minute), + "second" => Self::DatePartOf(DateTimeField::Second), "strftime" => Self::Strftime, "timestamp" | "datetime" => Self::Timestamp, @@ -1285,6 +1311,7 @@ impl SQLFunctionVisitor<'_> { }, } }), + DatePartOf(field) => self.try_visit_unary(|e| parse_extract_date_part(e, &field)), Strftime => { let args = extract_args(function)?; match args.len() { diff --git a/crates/polars-sql/src/lib.rs b/crates/polars-sql/src/lib.rs index 372084e4d2cf..12e5630a3b45 100644 --- a/crates/polars-sql/src/lib.rs +++ b/crates/polars-sql/src/lib.rs @@ -5,6 +5,7 @@ mod context; pub mod function_registry; mod functions; pub mod keywords; +mod literal_folding; mod resolver; mod sql_expr; mod sql_visitors; diff --git a/crates/polars-sql/src/literal_folding.rs b/crates/polars-sql/src/literal_folding.rs new file mode 100644 index 000000000000..ccc11e2aaa94 --- /dev/null +++ b/crates/polars-sql/src/literal_folding.rs @@ -0,0 +1,72 @@ +//! Exact folding of arithmetic between numeric SQL literals. +//! +//! `.06 + 0.01` is computed on the literal spellings as fixed-point values (exact within +//! `i128`) and converted once to the ordinary `Float64` literal. Integer-only arithmetic +//! is left to the engine. + +use polars_compute::decimal::exact; +use polars_plan::prelude::{Expr, lit}; +use sqlparser::ast::{ + BinaryOperator as SQLBinaryOperator, Expr as SQLExpr, UnaryOperator as SQLUnaryOperator, + Value as SQLValue, ValueWithSpan, +}; + +/// A literal as `(mantissa, scale)`: `mantissa / 10^scale`. +type Fixed = (i128, usize); + +/// Read a numeric literal spelling (digits with an optional `.`) without rounding. +fn parse_spelling(s: &str) -> Option { + let (int_part, frac_part) = s.split_once('.').unwrap_or((s, "")); + if !int_part + .bytes() + .chain(frac_part.bytes()) + .all(|b| b.is_ascii_digit()) + { + return None; + } + let mantissa = format!("{int_part}{frac_part}").parse().ok()?; + Some((mantissa, frac_part.len())) +} + +fn eval(expr: &SQLExpr) -> Option { + match expr { + SQLExpr::Value(ValueWithSpan { + value: SQLValue::Number(s, _), + .. + }) => parse_spelling(s), + SQLExpr::Nested(e) => eval(e), + SQLExpr::UnaryOp { op, expr } => match op { + SQLUnaryOperator::Plus => eval(expr), + SQLUnaryOperator::Minus => exact::neg(eval(expr)?), + _ => None, + }, + SQLExpr::BinaryOp { left, op, right } => combine(left, op, right), + _ => None, + } +} + +fn combine(left: &SQLExpr, op: &SQLBinaryOperator, right: &SQLExpr) -> Option { + let f = match op { + SQLBinaryOperator::Plus => exact::add, + SQLBinaryOperator::Minus => exact::sub, + SQLBinaryOperator::Multiply => exact::mul, + _ => return None, + }; + f(eval(left)?, eval(right)?) +} + +/// One correctly rounded conversion of the exact result. +fn to_f64((mantissa, scale): Fixed) -> f64 { + format!("{mantissa}e-{scale}").parse().unwrap() +} + +/// Fold `left right` when both sides are numeric literal arithmetic with a +/// decimal point; `None` leaves the expression to ordinary translation. +pub(crate) fn try_fold_decimal_arithmetic( + left: &SQLExpr, + op: &SQLBinaryOperator, + right: &SQLExpr, +) -> Option { + let value = combine(left, op, right)?; + (value.1 > 0).then(|| lit(to_f64(value))) +} diff --git a/crates/polars-sql/src/sql_expr.rs b/crates/polars-sql/src/sql_expr.rs index 790e8f376094..3ce136ae0547 100644 --- a/crates/polars-sql/src/sql_expr.rs +++ b/crates/polars-sql/src/sql_expr.rs @@ -6,6 +6,7 @@ //! - all Polars SQL keywords [`all_keywords`] //! - all Polars SQL functions [`all_functions`] +use std::borrow::Cow; use std::fmt::Display; use std::ops::Div; @@ -32,6 +33,8 @@ use sqlparser::tokenizer::Token; use crate::SQLContext; use crate::functions::SQLFunctionVisitor; +use crate::literal_folding::try_fold_decimal_arithmetic; +use crate::sql_visitors::expr_references_any_column; use crate::subquery::is_correlated_subquery; use crate::types::{ bitstring_to_bytes_literal, is_iso_date, is_iso_datetime, is_iso_time, map_sql_dtype_to_polars, @@ -291,12 +294,20 @@ impl SQLExprVisitor<'_> { list, negated, } => { - let expr = self.visit_expr(expr)?; + let sql_expr = expr; + let expr = self.visit_expr(sql_expr)?; // Prefer the all-literal `is_in` fast path, which predicate pushdown can - // use. A non-literal element, or an aggregate on the left, falls back to an - // OR-chain of equality comparisons. - let expr_is_aggregate = has_expr(&expr, |e| matches!(e, Expr::Agg(_) | Expr::Len)); - let elements = if expr_is_aggregate { + // use. A non-literal element, or an aggregate or constant on the left, falls + // back to an OR-chain of equality comparisons (which the planner can fold). + let is_constant = !expr_references_any_column(sql_expr) + && expr + .clone() + .meta() + .is_input_independent_scalar() + .unwrap_or(false); + let use_or_chain = + is_constant || has_expr(&expr, |e| matches!(e, Expr::Agg(_) | Expr::Len)); + let elements = if use_or_chain { None } else { self.array_expr_to_series(list).ok() @@ -663,8 +674,11 @@ impl SQLExprVisitor<'_> { op: &SQLBinaryOperator, right: &SQLExpr, ) -> PolarsResult { + if let Some(folded) = try_fold_decimal_arithmetic(left, op, right) { + return Ok(folded); + } // need special handling for interval offsets and comparisons - let (lhs, mut rhs) = match (left, op, right) { + let (mut lhs, mut rhs) = match (left, op, right) { (_, SQLBinaryOperator::Minus, SQLExpr::Interval(v)) => { let duration = interval_to_duration(v, false)?; return Ok(self @@ -700,6 +714,12 @@ impl SQLExprVisitor<'_> { _ => (self.visit_expr(left)?, self.visit_expr(right)?), }; rhs = self.convert_temporal_strings(&lhs, &rhs); + if matches!( + op, + SQLBinaryOperator::Eq | SQLBinaryOperator::NotEq | SQLBinaryOperator::Spaceship + ) { + (lhs, rhs) = self.convert_int_literal_for_string(lhs, rhs); + } if matches!(op, SQLBinaryOperator::Plus | SQLBinaryOperator::Minus) && let Some(expr) = self.date_day_offset(&lhs, op, &rhs) @@ -969,10 +989,12 @@ impl SQLExprVisitor<'_> { polars_ensure!(!list.is_empty(), SQLSyntax: "IN list must not be empty"); let mut elements = list.iter(); let first = self.visit_expr(elements.next().unwrap())?; - let mut membership = expr.clone().eq(first); + let (lhs, rhs) = self.convert_int_literal_for_string(expr.clone(), first); + let mut membership = lhs.eq(rhs); for e in elements { let e = self.visit_expr(e)?; - membership = membership.or(expr.clone().eq(e)); + let (lhs, rhs) = self.convert_int_literal_for_string(expr.clone(), e); + membership = membership.or(lhs.eq(rhs)); } Ok(if negated { membership.not() @@ -981,7 +1003,8 @@ impl SQLExprVisitor<'_> { }) } - /// Handle implicit temporal strings, eg: "dt IN ('2024-04-30','2024-05-01')". + /// Handle implicit temporal strings, eg: "dt IN ('2024-04-30','2024-05-01')", and + /// integer literals tested against a String expression, eg: "str IN (13, 31)". /// (not yet as versatile as the temporal string conversions in visit_binary_op) fn cast_array_elements_for( &self, @@ -1002,6 +1025,11 @@ impl SQLExprVisitor<'_> { } } } + if elems.dtype().is_integer() + && dtype_expr_match.is_some_and(|expr| self.expr_dtype(expr) == Some(DataType::String)) + { + return elems.cast(&DataType::String); + } Ok(elems) } @@ -1015,14 +1043,15 @@ impl SQLExprVisitor<'_> { let elems = self.array_expr_to_series(elements)?; let elems = self.cast_array_elements_for(elems, dtype_expr_match)?; - // if we are parsing the list as an element in a series, implode. - // otherwise, return the series as-is. - let res = if result_as_element { - elems.implode()?.into_series() + // if we are parsing the list as an element in a series, the result is one + // (scalar) list value; otherwise, return the series as-is. + Ok(if result_as_element { + let name = elems.name().clone(); + let dtype = DataType::List(Box::new(elems.dtype().clone())); + lit(Scalar::new(dtype, AnyValue::List(elems))).alias(name) } else { - elems - }; - Ok(lit(res)) + lit(elems) + }) } /// Visit a SQL `CAST` or `TRY_CAST` expression. @@ -1054,7 +1083,7 @@ impl SQLExprVisitor<'_> { if matches!( polars_type, DataType::Date | DataType::Time | DataType::Datetime(_, _) - ) && self.is_string_expr(&expr) + ) && self.expr_dtype(&expr) == Some(DataType::String) && let Some(parsed) = parse_string_as_temporal(expr.clone(), &polars_type, strict) { return Ok(parsed); @@ -1066,11 +1095,10 @@ impl SQLExprVisitor<'_> { }) } - /// Whether `expr` is known to be `String`; false if the dtype cannot be resolved. - fn is_string_expr(&self, expr: &Expr) -> bool { + fn convert_int_literal_for_string(&self, lhs: Expr, rhs: Expr) -> (Expr, Expr) { let empty = Schema::default(); let schema = self.active_schema.unwrap_or(&empty); - matches!(expr.to_field(schema), Ok(fld) if fld.dtype == DataType::String) + convert_int_literal_for_string((lhs, schema), (rhs, schema)) } /// Visit a SQL literal. @@ -1194,7 +1222,7 @@ impl SQLExprVisitor<'_> { _ => "DATETIME", }; polars_ensure!( - is_iso_datetime(value), + is_iso_datetime(value) || is_iso_date(value), SQLSyntax: "invalid {} literal '{}'", fn_name, value, ); Ok(DataType::Datetime(timeunit_from_precision(prec)?, None)) @@ -1494,37 +1522,42 @@ pub fn sql_expr>(s: S) -> PolarsResult { pub(crate) fn interval_to_duration(interval: &Interval, fixed: bool) -> PolarsResult { if interval.last_field.is_some() - || interval.leading_field.is_some() || interval.leading_precision.is_some() || interval.fractional_seconds_precision.is_some() { polars_bail!(SQLSyntax: "unsupported interval syntax ('{}')", interval) } - let s = match &*interval.value { - SQLExpr::UnaryOp { .. } => { + let s: Cow = match (&*interval.value, &interval.leading_field) { + (SQLExpr::UnaryOp { .. }, _) => { polars_bail!(SQLSyntax: "unary ops are not valid on interval strings; found {}", interval.value) }, - SQLExpr::Value(ValueWithSpan { - value: SQLValue::SingleQuotedString(s), - .. - }) => Some(s), - _ => None, + ( + SQLExpr::Value(ValueWithSpan { + value: SQLValue::SingleQuotedString(s), + .. + }), + None, + ) => Cow::Borrowed(s), + // "INTERVAL '3' MONTH" and "INTERVAL 3 MONTH": the value is a bare count of the unit + ( + SQLExpr::Value(ValueWithSpan { + value: SQLValue::SingleQuotedString(n) | SQLValue::Number(n, _), + .. + }), + Some(unit), + ) if n.bytes().all(|b| b.is_ascii_digit()) => Cow::Owned(format!("{n} {unit}")), + _ => polars_bail!(SQLSyntax: "invalid interval {:?}", interval), }; - match s { - Some(s) if s.contains('-') => { - polars_bail!(SQLInterface: "minus signs are not yet supported in interval strings; found '{}'", s) - }, - Some(s) => { - // years, quarters, and months do not have a fixed duration; these - // interval parts can only be used with respect to a reference point - let duration = Duration::parse_interval(s); - if fixed && duration.months() != 0 { - polars_bail!(SQLSyntax: "fixed-duration interval cannot contain years, quarters, or months; found {}", s) - }; - Ok(duration) - }, - None => polars_bail!(SQLSyntax: "invalid interval {:?}", interval), + if s.contains('-') { + polars_bail!(SQLInterface: "minus signs are not yet supported in interval strings; found '{}'", s) } + // years, quarters, and months do not have a fixed duration; these + // interval parts can only be used with respect to a reference point + let duration = Duration::try_parse_interval(&s)?; + if fixed && duration.months() != 0 { + polars_bail!(SQLSyntax: "fixed-duration interval cannot contain years, quarters, or months; found {}", s) + }; + Ok(duration) } pub(crate) fn parse_sql_expr( @@ -1536,6 +1569,32 @@ pub(crate) fn parse_sql_expr( visitor.visit_expr(expr) } +/// Whether `expr` is known to be `String`; false if the dtype cannot be resolved. +fn is_string_expr(expr: &Expr, schema: &Schema) -> bool { + matches!(expr.to_field(schema), Ok(fld) if fld.dtype == DataType::String) +} + +/// `str_expr = 13` compares against the string '13'; each operand's dtype is +/// resolved against its own schema. +pub(crate) fn convert_int_literal_for_string( + (lhs, lhs_schema): (Expr, &Schema), + (rhs, rhs_schema): (Expr, &Schema), +) -> (Expr, Expr) { + match (&lhs, &rhs) { + (Expr::Literal(LiteralValue::Dyn(DynLiteralValue::Int(n))), other) + if is_string_expr(other, rhs_schema) => + { + (lit(n.to_string()), rhs) + }, + (other, Expr::Literal(LiteralValue::Dyn(DynLiteralValue::Int(n)))) + if is_string_expr(other, lhs_schema) => + { + (lhs, lit(n.to_string())) + }, + _ => (lhs, rhs), + } +} + pub(crate) fn parse_sql_array(expr: &SQLExpr, ctx: &mut SQLContext) -> PolarsResult { match expr { SQLExpr::Array(arr) => { diff --git a/crates/polars-sql/src/sql_visitors.rs b/crates/polars-sql/src/sql_visitors.rs index 917feac9bf65..3263feb416cc 100644 --- a/crates/polars-sql/src/sql_visitors.rs +++ b/crates/polars-sql/src/sql_visitors.rs @@ -7,8 +7,8 @@ use std::ops::ControlFlow; use polars_core::prelude::*; use sqlparser::ast::{ - Expr as SQLExpr, ObjectName, Query, SetExpr, Statement, TableFactor, Visit, - Visitor as SQLVisitor, visit_expressions, + Expr as SQLExpr, FunctionArg, FunctionArgExpr, FunctionArguments, ObjectName, Query, SetExpr, + Statement, TableFactor, Visit, Visitor as SQLVisitor, visit_expressions, }; use sqlparser::keywords::ALL_KEYWORDS; @@ -324,10 +324,26 @@ impl SQLVisitor for ColumnRefFinder { type Break = (); fn pre_visit_expr(&mut self, expr: &SQLExpr) -> ControlFlow<()> { - if matches!( - expr, - SQLExpr::Identifier(_) | SQLExpr::CompoundIdentifier(_) - ) { + let is_column_ref = match expr { + SQLExpr::Identifier(_) + | SQLExpr::CompoundIdentifier(_) + | SQLExpr::Wildcard(_) + | SQLExpr::QualifiedWildcard(..) => true, + // wildcard arguments, eg: COUNT(*) / COLUMNS(*) + SQLExpr::Function(func) => match &func.args { + FunctionArguments::List(args) => args.args.iter().any(|arg| { + matches!( + arg, + FunctionArg::Unnamed( + FunctionArgExpr::Wildcard | FunctionArgExpr::QualifiedWildcard(_) + ) + ) + }), + _ => false, + }, + _ => false, + }; + if is_column_ref { ControlFlow::Break(()) } else { ControlFlow::Continue(()) diff --git a/crates/polars-sql/tests/functions_string.rs b/crates/polars-sql/tests/functions_string.rs index e79494dca954..aec2743ba8d2 100644 --- a/crates/polars-sql/tests/functions_string.rs +++ b/crates/polars-sql/tests/functions_string.rs @@ -118,7 +118,7 @@ fn test_array_to_string() { #[test] fn test_array_literal() { let mut context = SQLContext::new(); - context.register("df", DataFrame::empty().lazy()); + context.register("df", df! {"x" => &[0]}.unwrap().lazy()); let sql = "SELECT [100,200,300] AS arr FROM df"; let df_sql = context.execute(sql).unwrap().collect().unwrap(); diff --git a/py-polars/tests/unit/sql/test_array.py b/py-polars/tests/unit/sql/test_array.py index 4c4fc5354ed4..acc8b98c02d2 100644 --- a/py-polars/tests/unit/sql/test_array.py +++ b/py-polars/tests/unit/sql/test_array.py @@ -47,7 +47,7 @@ def test_array_agg(sort_order: str | None, limit: int | None, expected: Any) -> def test_array_literals() -> None: - with pl.SQLContext(df=None, eager=True) as ctx: + with pl.SQLContext(df=pl.DataFrame({"x": [0]}), eager=True) as ctx: res = ctx.execute( """ SELECT @@ -358,3 +358,10 @@ def test_array_typed_literals_mixed_error() -> None: match="expected consistent dtypes", ): pl.sql("SELECT ARRAY[DATE '2024-01-01', TIME '12:00:00']").collect() + + +def test_array_literal_default_name() -> None: + res = pl.sql("SELECT ARRAY[1, 2]", eager=True) + assert res.columns == [""] + res = pl.sql('SELECT t."" AS arr FROM (SELECT ARRAY[1, 2]) t', eager=True) + assert res.to_dict(as_series=False) == {"arr": [[1, 2]]} diff --git a/py-polars/tests/unit/sql/test_joins.py b/py-polars/tests/unit/sql/test_joins.py index 717cb8a683f8..b556d9f74959 100644 --- a/py-polars/tests/unit/sql/test_joins.py +++ b/py-polars/tests/unit/sql/test_joins.py @@ -9,6 +9,7 @@ import polars as pl from polars.exceptions import ( ColumnNotFoundError, + ComputeError, InvalidOperationError, SQLInterfaceError, SQLSyntaxError, @@ -1834,9 +1835,7 @@ def test_join_on_invalid_expr() -> None: "df1": pl.DataFrame({"a": [1, 2, 3]}), "df2": pl.DataFrame({"a": [2, 3, 9]}), } - with pytest.raises( - SQLInterfaceError, match="unsupported join constraint expression" - ): + with pytest.raises(ComputeError, match="predicates must resolve to boolean"): pl.SQLContext(frames, eager=True).execute( "SELECT * FROM df1 JOIN df2 ON (df1.a)" ) @@ -1951,3 +1950,168 @@ def test_join_predicate_operand_spanning_both_sides() -> None: """, compare_with="sqlite", ) + + +@pytest.mark.parametrize( + "join_type", + [ + "INNER JOIN", + "LEFT JOIN", + "RIGHT JOIN", + "FULL OUTER JOIN", + "SEMI JOIN", + "ANTI JOIN", + ], +) +@pytest.mark.parametrize( + "condition", + [ + "TRUE", + "FALSE", + "NULL", + "1 = 1", + "1 = 0", + "NULL = NULL", + "1 < 2", + "'13' = 13", + "UPPER('x') = 'X'", + "(1 = 1) AND (2 > 1)", + "CASE WHEN 1 = 1 THEN TRUE ELSE FALSE END", + "1 IN (1, 2)", + "3 NOT IN (1, 2)", + "'b' IN ('a', 'c')", + "(1 + 0) IN (1, 2)", + "UPPER('a') IN ('A', 'B')", + "CAST(1 AS INT) NOT IN (1, 2)", + "ARRAY_LENGTH(ARRAY[1, 2]) IN (2, 3)", + "ARRAY_CONTAINS(ARRAY[1, 2], 1)", + ], +) +@pytest.mark.parametrize("empty_side", [None, "a", "b"]) +def test_join_on_constant_condition( + join_type: str, condition: str, empty_side: str | None +) -> None: + frames = { + "a": pl.DataFrame({"k": [1, 2], "x": ["p", "q"]}), + "b": pl.DataFrame({"k": [2, 3], "y": ["r", "s"]}), + } + if empty_side: + frames[empty_side] = frames[empty_side].clear() + + if "SEMI" in join_type or "ANTI" in join_type: + query = f"SELECT a.k, a.x FROM a {join_type} b ON {condition} ORDER BY 1, 2" + else: + query = f""" + SELECT a.k, a.x, b.k AS bk, b.y + FROM a {join_type} b ON {condition} + ORDER BY 1, 2, 3, 4 + """ + assert_sql_matches(frames, query=query, compare_with="duckdb") + + +@pytest.mark.parametrize( + "join_type", + [ + "INNER JOIN", + "LEFT JOIN", + "RIGHT JOIN", + "FULL OUTER JOIN", + "SEMI JOIN", + "ANTI JOIN", + ], +) +def test_join_on_constant_any_condition(join_type: str) -> None: + # DuckDB does not support ANY(array) outside inner joins; compare with TRUE/FALSE + frames = { + "a": pl.DataFrame({"k": [1, 2]}), + "b": pl.DataFrame({"v": ["r", "s"]}), + } + ctx = pl.SQLContext(frames=frames) + for condition, verdict in [ + ("1 = ANY(ARRAY[1, 2])", "TRUE"), + ("3 = ANY(ARRAY[1, 2])", "FALSE"), + ]: + res = ctx.execute(f"SELECT * FROM a {join_type} b ON {condition}").collect() + expected = ctx.execute(f"SELECT * FROM a {join_type} b ON {verdict}").collect() + assert_frame_equal(res, expected, check_row_order=False) + + +def test_join_on_constant_true_plans_cross_join() -> None: + frames = { + "a": pl.LazyFrame({"k": [1, 2]}), + "b": pl.LazyFrame({"v": ["r", "s"]}), + } + ctx = pl.SQLContext(frames=frames) + for condition in ["TRUE", "1 = 1", "1 < 2"]: + plan = ctx.execute(f"SELECT * FROM a JOIN b ON {condition}").explain() + assert plan.startswith("CROSS JOIN") + # an always-true outer join is not a cross join: it must keep unmatched rows + plan = ctx.execute("SELECT * FROM a LEFT JOIN b ON TRUE").explain() + assert "CROSS JOIN" not in plan + + +def test_constant_key_join_keeps_validation() -> None: + a = pl.LazyFrame({"k": [1, 2]}) + b = pl.LazyFrame({"v": ["r", "s"]}) + assert a.join(b, left_on=pl.lit(1), right_on=pl.lit(1)).collect().height == 4 + for validate in ["1:1", "1:m", "m:1"]: + with pytest.raises(ComputeError, match="join keys did not fulfill"): + a.join( + b, + left_on=pl.lit(1), + right_on=pl.lit(1), + validate=validate, # type: ignore[arg-type] + ).collect() + + +@pytest.mark.parametrize("join_type", ["INNER", "LEFT"]) +def test_join_on_pattern_predicates(join_type: str) -> None: + frames = { + "customer": pl.DataFrame({"c_key": [1, 2, 3], "c_name": ["a", "b", "c"]}), + "orders": pl.DataFrame( + { + "o_key": [1, 1, 2, 3], + "o_comment": ["special requests", "no", "special packages", "no"], + "c_name": ["x", "x", "y", "z"], + } + ), + } + assert_sql_matches( + frames, + query=f""" + SELECT c_key, COUNT(o_key) AS n_orders + FROM customer + {join_type} JOIN orders + ON c_key = o_key AND o_comment NOT LIKE '%special%requests%' + GROUP BY c_key + ORDER BY c_key + """, + compare_with="duckdb", + ) + # predicates on a right-table column that also exists on the left + assert_sql_matches( + frames, + query=f""" + SELECT customer.c_name, orders.c_name AS o_name, o_comment + FROM customer + {join_type} JOIN orders + ON customer.c_key = orders.o_key + AND orders.c_name IN ('x', 'z') + AND o_comment ILIKE 'NO%' + ORDER BY 1, 2, 3 + """, + compare_with="duckdb", + ) + # the same clashing column name on both sides of one predicate + assert_sql_matches( + frames, + query=f""" + SELECT c_key, orders.c_name AS o_name + FROM customer + {join_type} JOIN orders + ON c_key = o_key AND customer.c_name < orders.c_name + AND orders.c_name NOT IN (customer.c_name, 'y') + ORDER BY 1, 2 + """, + compare_with="duckdb", + ) diff --git a/py-polars/tests/unit/sql/test_numeric.py b/py-polars/tests/unit/sql/test_numeric.py index 91f4b6540796..f83436f58bf0 100644 --- a/py-polars/tests/unit/sql/test_numeric.py +++ b/py-polars/tests/unit/sql/test_numeric.py @@ -1,5 +1,6 @@ from __future__ import annotations +import re from decimal import Decimal as D from typing import TYPE_CHECKING @@ -8,6 +9,7 @@ import polars as pl from polars.exceptions import SQLInterfaceError, SQLSyntaxError from polars.testing import assert_frame_equal, assert_series_equal +from tests.unit.sql import assert_sql_matches if TYPE_CHECKING: from polars._typing import PolarsDataType @@ -212,6 +214,93 @@ def test_stddev_variance() -> None: ) +def test_decimal_literal_arithmetic_is_exact() -> None: + df = pl.DataFrame( + { + "disc": [D("0.04"), D("0.05"), D("0.06"), D("0.07"), D("0.08")], + "qty": [1, 2, 3, 4, 5], + }, + schema={"disc": pl.Decimal(15, 2), "qty": pl.Int64}, + ) + assert_sql_matches( + df, + query=""" + SELECT disc, qty + FROM self + WHERE disc BETWEEN .06 - 0.01 AND .06 + 0.01 + AND disc <= (0.03 + .01) * 2 - -0.01 + ORDER BY disc + """, + expected={"disc": [D("0.05"), D("0.06"), D("0.07")], "qty": [2, 3, 4]}, + compare_with="duckdb", + ) + res = df.sql("SELECT 0.1 + 0.2 AS x, 1 + 2 AS y, 2 * 1.5 AS z FROM self LIMIT 1") + assert res.row(0) == (0.3, 3, 3.0) + assert res.schema == {"x": pl.Float64, "y": pl.Int32, "z": pl.Float64} + + +@pytest.mark.parametrize( + "expr", + [ + "0.1 + 0.2", + "-.5 + .5", + "-(0.5) * 2", + "+1.5 + 1", + "-(-1.5)", + "1.5 - 3", + "3 * 0.1", + "1.10 * 1.10", + "(1.5 + 0.5) * (2 - 0.5)", + "1.5 + 2 * 0.25", + "0.00000000000000000000000000000000000001 * 0.00000000000000000000000000000000000001", + "12345678901234567890.123456789 + 0.000000001", + "99999999999999999999999999999999999999.9 + 0.1", + ], +) +def test_literal_arithmetic_folds_exactly(expr: str) -> None: + # literal-only `+`/`-`/`*` is computed exactly, then converted to Float64 once; + # the reference evaluates the same expression with Python's exact Decimal + decimal_expr = re.sub(r"\d*\.\d+|\d+", lambda m: f"D('{m.group()}')", expr) + expected = float(eval(decimal_expr)) + res = pl.sql(f"SELECT {expr} AS x", eager=True) + assert res.schema == {"x": pl.Float64} + assert res.item() == expected + + +@pytest.mark.parametrize( + ("expr", "expected", "dtype"), + [ + # integer-only arithmetic stays on the ordinary path + ("1 + 2", 3, pl.Int32), + # eligible children fold even when the parent cannot + ("0.1 + 0.2 + a", 1.3, pl.Float64), + ("0.1 + 0.2 = 0.3", True, pl.Boolean), + # division and casts are left to the engine + ("0.5 / 0.25", 2.0, pl.Float64), + ("1.5 + 0.5 / 2", 1.75, pl.Float64), + ("1.5 + CAST(1 AS FLOAT)", 2.5, pl.Float64), + ("0.0 - 0.0", 0.0, pl.Float64), + # beyond the exact domain: falls back to float arithmetic + ( + "170141183460469231731687303715884105727.0 + 1.0", + 1.7014118346046923e38, + pl.Float64, + ), + ], +) +def test_literal_arithmetic_fallback( + expr: str, expected: object, dtype: pl.DataType +) -> None: + res = pl.DataFrame({"a": [1]}).sql(f"SELECT {expr} AS x FROM self") + assert res.schema == {"x": dtype} + assert res.item() == expected + + +def test_literal_scientific_notation_unsupported() -> None: + with pytest.raises(SQLInterfaceError, match="cannot parse literal"): + pl.sql("SELECT 1e2 + 0.5 AS x", eager=True) + + def test_int_div_true_division() -> None: df = pl.DataFrame({"num": [1], "denum": [3]}) with pl.SQLContext(df=df, eager=True) as ctx: diff --git a/py-polars/tests/unit/sql/test_operators.py b/py-polars/tests/unit/sql/test_operators.py index e745d441c7a0..6e1b08394d5b 100644 --- a/py-polars/tests/unit/sql/test_operators.py +++ b/py-polars/tests/unit/sql/test_operators.py @@ -80,6 +80,175 @@ def test_equal_not_equal() -> None: } +def test_string_compared_with_integer_literals() -> None: + # integer literals tested for (in)equality against a string are compared as strings + df = pl.DataFrame({"phone": ["13-123", "31-456", "22-789", "22"]}) + res = df.sql( + """ + SELECT phone + FROM self + WHERE SUBSTRING(phone, 1, 2) IN (13, 31) + OR phone = 22 + OR 13 <> SUBSTRING(phone, 1, 2) + ORDER BY phone + """ + ) + assert res.to_series().to_list() == ["13-123", "22", "22-789", "31-456"] + + res = df.sql("SELECT phone FROM self WHERE phone NOT IN (22, 13)") + assert res.to_series().to_list() == ["13-123", "31-456", "22-789"] + + # aggregate / non-literal IN lists take the OR-chain path + res = df.sql( + """ + SELECT + MAX(phone) IN (31, 13) AS agg_in, + MIN(phone) IN (31, MAX(phone)) AS agg_in_expr + FROM self + """ + ) + assert res.row(0) == (False, False) + res = df.sql("SELECT phone IN (22, LEFT(phone, 2)) AS x FROM self") + assert res.to_series().to_list() == [False, False, False, True] + + +@pytest.mark.parametrize( + ("condition", "keeps_rows"), + [ + ("1 = 1", True), + ("1 = 1.0", True), + ("'13' = 13", True), + ("('13' = 13)", True), + ("'13' <> 13", False), + ("1 < 2 AND 'a' = 'b'", False), + ("NULL = NULL", False), + ("NULL IS NULL", True), + ], +) +def test_constant_where_condition(condition: str, keeps_rows: bool) -> None: + df = pl.DataFrame({"a": [1, 2, 3]}) + res = df.sql(f"SELECT a FROM self WHERE {condition}") + assert res.height == (3 if keeps_rows else 0) + + +@pytest.mark.parametrize( + ("condition", "verdict"), + [ + ("TRUE", True), + ("FALSE", False), + ("1 = 1", True), + ("1 = 0", False), + ("NULL = NULL", None), + ("NULL", None), + ("UPPER('x') = 'X'", True), + ("CASE WHEN 1 < 2 THEN FALSE ELSE TRUE END", False), + # non-boolean constants are cast to boolean + ("1", True), + ("0", False), + ("1 + 1", True), + ("2 IN (1, 2)", True), + ("ARRAY_LENGTH(ARRAY[1, 2])", True), + ("ARRAY_CONTAINS(ARRAY[1, 2], 3)", False), + ], +) +@pytest.mark.parametrize("empty", [False, True]) +def test_constant_condition_select_and_delete( + condition: str, verdict: bool | None, empty: bool +) -> None: + df = pl.DataFrame({"a": [1, 2, 3], "b": ["x", "y", "z"]}) + if empty: + df = df.clear() + ctx = pl.SQLContext(frames={"tbl": df.lazy()}) + + selected = ctx.execute(f"SELECT * FROM tbl WHERE {condition}").collect() + deleted = ctx.execute(f"DELETE FROM tbl WHERE {condition}").collect() + assert selected.schema == df.schema + assert deleted.schema == df.schema + # an unknown condition keeps nothing in SELECT and removes nothing in DELETE + assert selected.height == (df.height if verdict else 0) + assert deleted.height == (0 if verdict else df.height) + + +def test_constant_where_condition_is_planned_not_executed() -> None: + # translating the query must not run the frame or any user function in it + def boom(df: pl.DataFrame) -> pl.DataFrame: + msg = "executed" + raise RuntimeError(msg) + + lf = pl.LazyFrame({"a": [1, 2, 3]}).map_batches(boom, schema={"a": pl.Int64}) + ctx = pl.SQLContext(frames={"tbl": lf, "other": lf}) + for query in [ + "SELECT a FROM tbl WHERE 1 = 1 AND UPPER('x') = 'X'", + "SELECT * FROM tbl JOIN other ON 1 = 1", + "SELECT * FROM tbl LEFT JOIN other ON FALSE", + ]: + planned = ctx.execute(query) + with pytest.raises(RuntimeError, match="executed"): + planned.collect() + + # a false condition never reads the input at all + res = ctx.execute("SELECT a FROM tbl WHERE 1 = 0").collect() + assert res.schema == {"a": pl.Int64} + assert res.height == 0 + + # constant conditions are folded by the planner: a true filter disappears + lf = pl.LazyFrame({"a": [1, 2, 3]}) + ctx = pl.SQLContext(frames={"tbl": lf}) + plan = ctx.execute("SELECT a FROM tbl WHERE 1 = 1").explain() + assert "FILTER" not in plan + plan = ctx.execute("SELECT a FROM tbl WHERE 1 = 0").explain() + assert "FILTER" not in plan + + +@pytest.mark.parametrize( + ("condition", "expected"), + [ + ("ROW_NUMBER() OVER () <= 2", [1, 2]), + ("COUNT(1) > 1", [1, 2, 3]), + ("COLUMNS('^a$') > 1", [2, 3]), + ], +) +def test_where_condition_without_column_names( + condition: str, expected: list[int] +) -> None: + df = pl.DataFrame({"a": [1, 2, 3]}) + res = df.sql(f"SELECT a FROM self WHERE {condition}") + assert res.to_series().to_list() == expected + + +def test_join_key_string_compared_with_integer_literal() -> None: + frames = { + "a": pl.DataFrame({"s": ["13", "31", "22"]}), + "b": pl.DataFrame({"k": [1, 2]}), + } + res = pl.SQLContext(frames=frames).execute( + "SELECT s, k FROM a JOIN b ON a.s = 13 AND b.k = 2", eager=True + ) + assert res.rows() == [("13", 2)] + + # a clashing column name with a different dtype per table + frames = { + "a": pl.DataFrame({"k": [1, 2], "x": ["13", "31"]}), + "b": pl.DataFrame({"k": [1, 2], "x": [13, 31]}), + } + res = pl.SQLContext(frames=frames).execute( + "SELECT a.k, a.x, b.x AS bx FROM a JOIN b ON a.k = b.k AND b.x = 13 AND a.x = 13", + eager=True, + ) + assert res.rows() == [(1, "13", 13)] + + # a nested comparison inside a join key operand resolves against its own table too + frames = { + "a": pl.DataFrame({"k": [1, 2], "x": ["13", "31"], "flag": [True, True]}), + "b": pl.DataFrame({"k": [1, 2], "x": [13, 31]}), + } + res = pl.SQLContext(frames=frames).execute( + "SELECT a.k FROM a JOIN b ON a.k = b.k AND a.flag = (b.x = 13)", + eager=True, + ) + assert res.rows() == [(1,)] + + @pytest.mark.parametrize( "in_clause", [ diff --git a/py-polars/tests/unit/sql/test_temporal.py b/py-polars/tests/unit/sql/test_temporal.py index 7fd1c4ab1488..5bc803454d60 100644 --- a/py-polars/tests/unit/sql/test_temporal.py +++ b/py-polars/tests/unit/sql/test_temporal.py @@ -573,3 +573,101 @@ def test_date_arithmetic_leaves_other_dtypes_alone() -> None: res = ctx.execute("SELECT a + 5 AS x, dt - dt AS y FROM tbl") assert res.schema["x"] == pl.Int64 assert res.schema["y"] == pl.Duration("us") + + +@pytest.mark.parametrize("fn", ["TIMESTAMP", "DATETIME"]) +def test_typed_timestamp_literal_bare_date(fn: str) -> None: + df = pl.DataFrame( + { + "ts": [ + datetime(1993, 12, 31, 23, 59), + datetime(1994, 1, 1), + datetime(1994, 1, 2), + ] + } + ) + assert_sql_matches( + df, + query=f"SELECT ts FROM self WHERE ts >= {fn} '1994-01-01' ORDER BY ts", + expected={"ts": [datetime(1994, 1, 1), datetime(1994, 1, 2)]}, + compare_with="duckdb", + ) + + +def test_interval_leading_field() -> None: + df = pl.DataFrame( + { + "dt": [date(1994, 1, 1), date(1994, 3, 31), date(1994, 4, 1)], + "dtm": [ + datetime(1994, 1, 1), + datetime(1994, 1, 1, 2), + datetime(1994, 1, 1, 3), + ], + } + ) + assert_sql_matches( + df, + query=""" + SELECT + dt + INTERVAL '3' MONTH AS m3, + dt + INTERVAL 1 YEAR AS y1, + dt - INTERVAL '2' DAY AS d2, + dtm + INTERVAL '90' MINUTE AS min90 + FROM self + WHERE dt < DATE '1994-01-01' + INTERVAL '3' MONTH + AND dtm < TIMESTAMP '1994-01-01' + INTERVAL 2 HOUR + ORDER BY dt + """, + expected={ + "m3": [date(1994, 4, 1)], + "y1": [date(1995, 1, 1)], + "d2": [date(1993, 12, 30)], + "min90": [datetime(1994, 1, 1, 1, 30)], + }, + compare_with="duckdb", + ) + with pytest.raises(SQLSyntaxError, match="invalid interval"): + df.sql("SELECT dt + INTERVAL '1 2' MONTH FROM self") + + +def test_date_part_functions() -> None: + df = pl.DataFrame( + { + "dtm": [ + datetime(1994, 3, 6, 13, 45, 30), + datetime(2000, 12, 31, 0, 1, 2), + ] + } + ) + assert_sql_matches( + df, + query=""" + SELECT + YEAR(dtm) AS y, + QUARTER(dtm) AS q, + MONTH(dtm) AS mo, + WEEK(dtm) AS w, + DAY(dtm) AS d, + DAYOFWEEK(dtm) AS dow, + DAYOFYEAR(dtm) AS doy, + HOUR(dtm) AS h, + MINUTE(dtm) AS mi, + SECOND(dtm) AS s + FROM self + WHERE YEAR(dtm) IN (1994, 2000) + ORDER BY dtm + """, + expected={ + "y": [1994, 2000], + "q": [1, 4], + "mo": [3, 12], + "w": [9, 52], + "d": [6, 31], + "dow": [0, 0], + "doy": [65, 366], + "h": [13, 0], + "mi": [45, 1], + "s": [30, 2], + }, + compare_with="duckdb", + )