From d3503f0c9aff2fa6aa6bbaab6855101bfe0e553b Mon Sep 17 00:00:00 2001 From: ritchie46 Date: Sat, 12 Sep 2026 11:45:54 +0200 Subject: [PATCH 01/15] feat: Close the SQL gaps needed to run all PDS-H queries - `INTERVAL '3' MONTH` / `INTERVAL 3 MONTH` (leading unit) intervals - `YEAR(x)`, `MONTH(x)`, `DAY(x)`, `HOUR(x)`, ... date part functions - `LIKE`/`ILIKE`/`IN`/... predicates inside `JOIN ... ON` - bare ISO date in `TIMESTAMP`/`DATETIME` typed literals - integer literals compared for equality against a String expression (`substring(x, 1, 2) IN (13, 31)`) are compared as strings - literal-only arithmetic (`.06 + 0.01`) is folded exactly so it compares correctly against Decimal columns, matching decimal SQL semantics Closes #29169 Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_01Lj7JbtU82rEjCBvQ98absi --- crates/polars-sql/src/context.rs | 43 ++++- crates/polars-sql/src/functions.rs | 29 +++- crates/polars-sql/src/sql_expr.rs | 192 ++++++++++++++++++--- crates/polars-sql/src/sql_visitors.rs | 16 ++ py-polars/tests/unit/sql/test_joins.py | 52 +++++- py-polars/tests/unit/sql/test_numeric.py | 26 +++ py-polars/tests/unit/sql/test_operators.py | 19 ++ py-polars/tests/unit/sql/test_temporal.py | 98 +++++++++++ 8 files changed, 443 insertions(+), 32 deletions(-) diff --git a/crates/polars-sql/src/context.rs b/crates/polars-sql/src/context.rs index b60430eeba5a..beefd51e55d1 100644 --- a/crates/polars-sql/src/context.rs +++ b/crates/polars-sql/src/context.rs @@ -34,6 +34,7 @@ use crate::sql_visitors::{ QualifyExpression, TableIdentifierCollector, check_for_ambiguous_column_refs, expr_contains_subquery, expr_has_window_functions, expr_references_any_column, expr_refers_to_table, sql_expr_cols_all_in_schema, statement_registers_table, + table_qualified_columns, }; use crate::subquery::{LowerScope, SubqueryBindings, desugar_quantified_subqueries}; use crate::table_functions::PolarsTableFunctions; @@ -3772,12 +3773,48 @@ fn process_join_on( ), }, SQLExpr::Nested(expr) => process_join_on(ctx, expr, tbl_left, tbl_right), - _ => polars_bail!( - SQLInterface: "unsupported join constraint expression: {:?}", sql_expr - ), + // Any other predicate (LIKE, IN, IS NULL, OR, ...) is evaluated on the joined frame. + _ => { + let join_schema = build_join_schema(tbl_left, tbl_right)?; + let suffix = format!(":{}", tbl_right.name); + let predicate = parse_sql_expr(sql_expr, ctx, Some(&join_schema))?; + let predicate = + suffix_right_table_columns(predicate, sql_expr, tbl_left, tbl_right, &suffix)?; + Ok((vec![], vec![], vec![predicate])) + }, } } +/// Rename the columns that `sql_expr` references as `right_table.col` to their merged-schema +/// (suffixed) names when the same column also exists in the left table. +fn suffix_right_table_columns( + expr: Expr, + sql_expr: &SQLExpr, + tbl_left: &TableInfo, + tbl_right: &TableInfo, + suffix: &str, +) -> PolarsResult { + let right_cols = table_qualified_columns(sql_expr, &tbl_right.name); + let left_cols = table_qualified_columns(sql_expr, &tbl_left.name); + let conflicts = |name: &str| { + right_cols.contains(name) + && tbl_left.schema.contains(name) + && tbl_right.schema.contains(name) + }; + if let Some(name) = left_cols.iter().find(|name| conflicts(name)) { + polars_bail!( + SQLInterface: "unsupported join condition: references both '{}.{}' and '{}.{}'", + tbl_left.name, name, tbl_right.name, name + ) + } + Ok(strip_join_aliases(expr).map_expr(|e| match e { + Expr::Column(ref name) if conflicts(name) => { + Expr::Column(PlSmallStr::from_string(format!("{name}{suffix}"))) + }, + other => other, + })) +} + /// Replace aggregates over pre-aggregation columns with references to hoisted /// aggregation outputs, collecting the hoisted aggregates into `agg_out`. /// 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/sql_expr.rs b/crates/polars-sql/src/sql_expr.rs index 790e8f376094..40010f45dc58 100644 --- a/crates/polars-sql/src/sql_expr.rs +++ b/crates/polars-sql/src/sql_expr.rs @@ -663,8 +663,11 @@ impl SQLExprVisitor<'_> { op: &SQLBinaryOperator, right: &SQLExpr, ) -> PolarsResult { + if let Some(folded) = fold_decimal_literal_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 +703,16 @@ 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 + ) { + if let Some(e) = self.int_literal_as_string(&rhs, &lhs) { + rhs = e; + } else if let Some(e) = self.int_literal_as_string(&lhs, &rhs) { + lhs = e; + } + } if matches!(op, SQLBinaryOperator::Plus | SQLBinaryOperator::Minus) && let Some(expr) = self.date_day_offset(&lhs, op, &rhs) @@ -981,7 +994,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 +1016,11 @@ impl SQLExprVisitor<'_> { } } } + if elems.dtype().is_integer() + && dtype_expr_match.is_some_and(|expr| self.is_string_expr(expr)) + { + return elems.cast(&DataType::String); + } Ok(elems) } @@ -1066,6 +1085,19 @@ impl SQLExprVisitor<'_> { }) } + /// `str_expr = 13`: an integer literal tested for equality against a String + /// expression is compared as the string '13'. + fn int_literal_as_string(&self, literal: &Expr, other: &Expr) -> Option { + match literal { + Expr::Literal(LiteralValue::Dyn(DynLiteralValue::Int(n))) + if self.is_string_expr(other) => + { + Some(lit(n.to_string())) + }, + _ => None, + } + } + /// Whether `expr` is known to be `String`; false if the dtype cannot be resolved. fn is_string_expr(&self, expr: &Expr) -> bool { let empty = Schema::default(); @@ -1194,7 +1226,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)) @@ -1492,39 +1524,149 @@ pub fn sql_expr>(s: S) -> PolarsResult { }) } +/// A fixed-point value: `mantissa / 10^scale`. +struct DecimalLiteral { + mantissa: i128, + scale: u32, + has_point: bool, +} + +impl DecimalLiteral { + fn parse(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; + } + Some(Self { + mantissa: format!("{int_part}{frac_part}").parse().ok()?, + scale: frac_part.len() as u32, + has_point: s.contains('.'), + }) + } + + fn rescale(&self, scale: u32) -> Option { + self.mantissa + .checked_mul(10i128.checked_pow(scale - self.scale)?) + } + + fn eval(expr: &SQLExpr) -> Option { + match expr { + SQLExpr::Value(ValueWithSpan { + value: SQLValue::Number(s, _), + .. + }) => Self::parse(s), + SQLExpr::Nested(e) => Self::eval(e), + SQLExpr::UnaryOp { op, expr } => { + let v = Self::eval(expr)?; + match op { + SQLUnaryOperator::Plus => Some(v), + SQLUnaryOperator::Minus => Some(Self { + mantissa: v.mantissa.checked_neg()?, + ..v + }), + _ => None, + } + }, + SQLExpr::BinaryOp { left, op, right } => { + Self::combine(Self::eval(left)?, op, Self::eval(right)?) + }, + _ => None, + } + } + + fn combine(l: Self, op: &SQLBinaryOperator, r: Self) -> Option { + let has_point = l.has_point || r.has_point; + match op { + SQLBinaryOperator::Plus | SQLBinaryOperator::Minus => { + let scale = l.scale.max(r.scale); + let (l, r) = (l.rescale(scale)?, r.rescale(scale)?); + let mantissa = if *op == SQLBinaryOperator::Plus { + l.checked_add(r)? + } else { + l.checked_sub(r)? + }; + Some(Self { + mantissa, + scale, + has_point, + }) + }, + SQLBinaryOperator::Multiply => Some(Self { + mantissa: l.mantissa.checked_mul(r.mantissa)?, + scale: l.scale + r.scale, + has_point, + }), + _ => None, + } + } + + fn to_f64(&self) -> f64 { + let digits = self.mantissa.unsigned_abs().to_string(); + let digits = format!("{:0>width$}", digits, width = self.scale as usize + 1); + let (int_part, frac_part) = digits.split_at(digits.len() - self.scale as usize); + let sign = if self.mantissa < 0 { "-" } else { "" }; + format!("{sign}{int_part}.{frac_part}").parse().unwrap() + } +} + +/// Evaluate arithmetic between numeric literals exactly, so that `.06 + 0.01` yields the +/// float nearest to 0.07 (as it would in decimal SQL) rather than accumulating float error. +/// Integer-only arithmetic is left to the engine. +fn fold_decimal_literal_arithmetic( + left: &SQLExpr, + op: &SQLBinaryOperator, + right: &SQLExpr, +) -> Option { + let value = DecimalLiteral::combine( + DecimalLiteral::eval(left)?, + op, + DecimalLiteral::eval(right)?, + )?; + value.has_point.then(|| lit(value.to_f64())) +} + 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 = 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, + ) => s.clone(), + // "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()) => 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( diff --git a/crates/polars-sql/src/sql_visitors.rs b/crates/polars-sql/src/sql_visitors.rs index 917feac9bf65..2408898c213d 100644 --- a/crates/polars-sql/src/sql_visitors.rs +++ b/crates/polars-sql/src/sql_visitors.rs @@ -52,6 +52,22 @@ pub(crate) fn expr_refers_to_table(expr: &SQLExpr, table_name: &str) -> bool { table_finder.found } +/// Collect the column names that an expression references qualified by the given table +/// (`table_name.col`). +pub(crate) fn table_qualified_columns(expr: &SQLExpr, table_name: &str) -> PlHashSet { + let mut columns = PlHashSet::new(); + let _ = visit_expressions(expr, |e| { + if let SQLExpr::CompoundIdentifier(idents) = e + && idents.len() >= 2 + && idents[0].value.as_str() == table_name + { + columns.insert(idents[1].value.clone()); + } + ControlFlow::<()>::Continue(()) + }); + columns +} + // --------------------------------------------------------------------------- // UnqualifiedColumnsInSchema // --------------------------------------------------------------------------- diff --git a/py-polars/tests/unit/sql/test_joins.py b/py-polars/tests/unit/sql/test_joins.py index 717cb8a683f8..cd6d9e4d5ef1 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,50 @@ def test_join_predicate_operand_spanning_both_sides() -> None: """, compare_with="sqlite", ) + + +@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", + ) + with pytest.raises(SQLInterfaceError, match="references both"): + pl.SQLContext(frames=frames).execute( + f""" + SELECT * FROM customer {join_type} JOIN orders + ON c_key = o_key AND customer.c_name IN (orders.c_name, 'a') + """ + ).collect() diff --git a/py-polars/tests/unit/sql/test_numeric.py b/py-polars/tests/unit/sql/test_numeric.py index 91f4b6540796..c04440c55cba 100644 --- a/py-polars/tests/unit/sql/test_numeric.py +++ b/py-polars/tests/unit/sql/test_numeric.py @@ -8,6 +8,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 +213,31 @@ 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} + + 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..9b1e29fd1318 100644 --- a/py-polars/tests/unit/sql/test_operators.py +++ b/py-polars/tests/unit/sql/test_operators.py @@ -80,6 +80,25 @@ 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"] + + @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", + ) From 4f67cefd64880b42629fc9f2ce36e4336eb4f759 Mon Sep 17 00:00:00 2001 From: ritchie46 Date: Sat, 12 Sep 2026 13:44:12 +0200 Subject: [PATCH 02/15] refactor: Route every non-equi JOIN ON condition through one predicate path Right-table columns are renamed at the SQL level before parsing, so a predicate may reference the same clashing column on both sides. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_01Lj7JbtU82rEjCBvQ98absi --- crates/polars-sql/src/context.rs | 159 +++++++------------------ crates/polars-sql/src/sql_expr.rs | 104 +++++++--------- crates/polars-sql/src/sql_visitors.rs | 16 --- py-polars/tests/unit/sql/test_joins.py | 20 ++-- 4 files changed, 99 insertions(+), 200 deletions(-) diff --git a/crates/polars-sql/src/context.rs b/crates/polars-sql/src/context.rs index beefd51e55d1..9f8b9e856cb2 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,7 +20,7 @@ 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}; @@ -34,7 +34,6 @@ use crate::sql_visitors::{ QualifyExpression, TableIdentifierCollector, check_for_ambiguous_column_refs, expr_contains_subquery, expr_has_window_functions, expr_references_any_column, expr_refers_to_table, sql_expr_cols_all_in_schema, statement_registers_table, - table_qualified_columns, }; use crate::subquery::{LowerScope, SubqueryBindings, desugar_quantified_subqueries}; use crate::table_functions::PolarsTableFunctions; @@ -3695,9 +3694,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, @@ -3728,91 +3727,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), - // Any other predicate (LIKE, IN, IS NULL, OR, ...) is evaluated on the joined frame. - _ => { - let join_schema = build_join_schema(tbl_left, tbl_right)?; - let suffix = format!(":{}", tbl_right.name); - let predicate = parse_sql_expr(sql_expr, ctx, Some(&join_schema))?; - let predicate = - suffix_right_table_columns(predicate, sql_expr, tbl_left, tbl_right, &suffix)?; - Ok((vec![], vec![], vec![predicate])) - }, + _ => process_join_predicate(ctx, sql_expr, tbl_left, tbl_right), } } -/// Rename the columns that `sql_expr` references as `right_table.col` to their merged-schema -/// (suffixed) names when the same column also exists in the left table. -fn suffix_right_table_columns( - expr: Expr, +/// Parse a non-equi join condition 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, - suffix: &str, -) -> PolarsResult { - let right_cols = table_qualified_columns(sql_expr, &tbl_right.name); - let left_cols = table_qualified_columns(sql_expr, &tbl_left.name); - let conflicts = |name: &str| { - right_cols.contains(name) - && tbl_left.schema.contains(name) - && tbl_right.schema.contains(name) - }; - if let Some(name) = left_cols.iter().find(|name| conflicts(name)) { - polars_bail!( - SQLInterface: "unsupported join condition: references both '{}.{}' and '{}.{}'", - tbl_left.name, name, tbl_right.name, name - ) +) -> 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 = Schema::clone(&tbl_left.schema); + for (name, dtype) in tbl_right.schema.iter() { + let name = if conflicts(name) { + PlSmallStr::from_string(format!("{name}{suffix}")) + } else { + name.clone() + }; + joined_schema.insert(name, dtype.clone()); } - Ok(strip_join_aliases(expr).map_expr(|e| match e { - Expr::Column(ref name) if conflicts(name) => { - Expr::Column(PlSmallStr::from_string(format!("{name}{suffix}"))) - }, - other => other, - })) + + // `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 @@ -3903,41 +3863,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 { diff --git a/crates/polars-sql/src/sql_expr.rs b/crates/polars-sql/src/sql_expr.rs index 40010f45dc58..37cd67375f9f 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; @@ -707,10 +708,19 @@ impl SQLExprVisitor<'_> { op, SQLBinaryOperator::Eq | SQLBinaryOperator::NotEq | SQLBinaryOperator::Spaceship ) { - if let Some(e) = self.int_literal_as_string(&rhs, &lhs) { - rhs = e; - } else if let Some(e) = self.int_literal_as_string(&lhs, &rhs) { - lhs = e; + // `str_expr = 13` compares against the string '13' + match (&lhs, &rhs) { + (Expr::Literal(LiteralValue::Dyn(DynLiteralValue::Int(n))), other) + if self.is_string_expr(other) => + { + lhs = lit(n.to_string()) + }, + (other, Expr::Literal(LiteralValue::Dyn(DynLiteralValue::Int(n)))) + if self.is_string_expr(other) => + { + rhs = lit(n.to_string()) + }, + _ => {}, } } @@ -1085,19 +1095,6 @@ impl SQLExprVisitor<'_> { }) } - /// `str_expr = 13`: an integer literal tested for equality against a String - /// expression is compared as the string '13'. - fn int_literal_as_string(&self, literal: &Expr, other: &Expr) -> Option { - match literal { - Expr::Literal(LiteralValue::Dyn(DynLiteralValue::Int(n))) - if self.is_string_expr(other) => - { - Some(lit(n.to_string())) - }, - _ => None, - } - } - /// Whether `expr` is known to be `String`; false if the dtype cannot be resolved. fn is_string_expr(&self, expr: &Expr) -> bool { let empty = Schema::default(); @@ -1528,7 +1525,6 @@ pub fn sql_expr>(s: S) -> PolarsResult { struct DecimalLiteral { mantissa: i128, scale: u32, - has_point: bool, } impl DecimalLiteral { @@ -1544,7 +1540,6 @@ impl DecimalLiteral { Some(Self { mantissa: format!("{int_part}{frac_part}").parse().ok()?, scale: frac_part.len() as u32, - has_point: s.contains('.'), }) } @@ -1571,62 +1566,51 @@ impl DecimalLiteral { _ => None, } }, - SQLExpr::BinaryOp { left, op, right } => { - Self::combine(Self::eval(left)?, op, Self::eval(right)?) - }, + SQLExpr::BinaryOp { left, op, right } => Self::combine(left, op, right), _ => None, } } - fn combine(l: Self, op: &SQLBinaryOperator, r: Self) -> Option { - let has_point = l.has_point || r.has_point; - match op { - SQLBinaryOperator::Plus | SQLBinaryOperator::Minus => { - let scale = l.scale.max(r.scale); - let (l, r) = (l.rescale(scale)?, r.rescale(scale)?); - let mantissa = if *op == SQLBinaryOperator::Plus { - l.checked_add(r)? - } else { - l.checked_sub(r)? - }; - Some(Self { - mantissa, - scale, - has_point, - }) - }, - SQLBinaryOperator::Multiply => Some(Self { + fn combine(left: &SQLExpr, op: &SQLBinaryOperator, right: &SQLExpr) -> Option { + if !matches!( + op, + SQLBinaryOperator::Plus | SQLBinaryOperator::Minus | SQLBinaryOperator::Multiply + ) { + return None; + } + let (l, r) = (Self::eval(left)?, Self::eval(right)?); + if *op == SQLBinaryOperator::Multiply { + return Some(Self { mantissa: l.mantissa.checked_mul(r.mantissa)?, scale: l.scale + r.scale, - has_point, - }), - _ => None, + }); } + let scale = l.scale.max(r.scale); + let (l, r) = (l.rescale(scale)?, r.rescale(scale)?); + let mantissa = if *op == SQLBinaryOperator::Plus { + l.checked_add(r)? + } else { + l.checked_sub(r)? + }; + Some(Self { mantissa, scale }) } fn to_f64(&self) -> f64 { - let digits = self.mantissa.unsigned_abs().to_string(); - let digits = format!("{:0>width$}", digits, width = self.scale as usize + 1); - let (int_part, frac_part) = digits.split_at(digits.len() - self.scale as usize); - let sign = if self.mantissa < 0 { "-" } else { "" }; - format!("{sign}{int_part}.{frac_part}").parse().unwrap() + format!("{}e-{}", self.mantissa, self.scale) + .parse() + .unwrap() } } -/// Evaluate arithmetic between numeric literals exactly, so that `.06 + 0.01` yields the -/// float nearest to 0.07 (as it would in decimal SQL) rather than accumulating float error. -/// Integer-only arithmetic is left to the engine. +/// Evaluate `+`/`-`/`*` between numeric literals exactly; integer-only arithmetic is +/// left to the engine. fn fold_decimal_literal_arithmetic( left: &SQLExpr, op: &SQLBinaryOperator, right: &SQLExpr, ) -> Option { - let value = DecimalLiteral::combine( - DecimalLiteral::eval(left)?, - op, - DecimalLiteral::eval(right)?, - )?; - value.has_point.then(|| lit(value.to_f64())) + let value = DecimalLiteral::combine(left, op, right)?; + (value.scale > 0).then(|| lit(value.to_f64())) } pub(crate) fn interval_to_duration(interval: &Interval, fixed: bool) -> PolarsResult { @@ -1636,7 +1620,7 @@ pub(crate) fn interval_to_duration(interval: &Interval, fixed: bool) -> PolarsRe { polars_bail!(SQLSyntax: "unsupported interval syntax ('{}')", interval) } - let s = match (&*interval.value, &interval.leading_field) { + 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) }, @@ -1646,7 +1630,7 @@ pub(crate) fn interval_to_duration(interval: &Interval, fixed: bool) -> PolarsRe .. }), None, - ) => s.clone(), + ) => Cow::Borrowed(s), // "INTERVAL '3' MONTH" and "INTERVAL 3 MONTH": the value is a bare count of the unit ( SQLExpr::Value(ValueWithSpan { @@ -1654,7 +1638,7 @@ pub(crate) fn interval_to_duration(interval: &Interval, fixed: bool) -> PolarsRe .. }), Some(unit), - ) if n.bytes().all(|b| b.is_ascii_digit()) => format!("{n} {unit}"), + ) if n.bytes().all(|b| b.is_ascii_digit()) => Cow::Owned(format!("{n} {unit}")), _ => polars_bail!(SQLSyntax: "invalid interval {:?}", interval), }; if s.contains('-') { diff --git a/crates/polars-sql/src/sql_visitors.rs b/crates/polars-sql/src/sql_visitors.rs index 2408898c213d..917feac9bf65 100644 --- a/crates/polars-sql/src/sql_visitors.rs +++ b/crates/polars-sql/src/sql_visitors.rs @@ -52,22 +52,6 @@ pub(crate) fn expr_refers_to_table(expr: &SQLExpr, table_name: &str) -> bool { table_finder.found } -/// Collect the column names that an expression references qualified by the given table -/// (`table_name.col`). -pub(crate) fn table_qualified_columns(expr: &SQLExpr, table_name: &str) -> PlHashSet { - let mut columns = PlHashSet::new(); - let _ = visit_expressions(expr, |e| { - if let SQLExpr::CompoundIdentifier(idents) = e - && idents.len() >= 2 - && idents[0].value.as_str() == table_name - { - columns.insert(idents[1].value.clone()); - } - ControlFlow::<()>::Continue(()) - }); - columns -} - // --------------------------------------------------------------------------- // UnqualifiedColumnsInSchema // --------------------------------------------------------------------------- diff --git a/py-polars/tests/unit/sql/test_joins.py b/py-polars/tests/unit/sql/test_joins.py index cd6d9e4d5ef1..fba271d7df52 100644 --- a/py-polars/tests/unit/sql/test_joins.py +++ b/py-polars/tests/unit/sql/test_joins.py @@ -1990,10 +1990,16 @@ def test_join_on_pattern_predicates(join_type: str) -> None: """, compare_with="duckdb", ) - with pytest.raises(SQLInterfaceError, match="references both"): - pl.SQLContext(frames=frames).execute( - f""" - SELECT * FROM customer {join_type} JOIN orders - ON c_key = o_key AND customer.c_name IN (orders.c_name, 'a') - """ - ).collect() + # 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", + ) From 95ec1d687472def4f3481d4d1b4223203512fdbb Mon Sep 17 00:00:00 2001 From: ritchie46 Date: Sat, 12 Sep 2026 18:51:09 +0200 Subject: [PATCH 03/15] fix: Apply string/integer literal coercion on every SQL equality path Constant WHERE conditions are now evaluated through the expression visitor instead of comparing raw SQL literals, and the OR-chain IN fallback shares the same coercion as binary comparisons. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_01Lj7JbtU82rEjCBvQ98absi --- crates/polars-sql/src/context.rs | 48 +++++++--------------- crates/polars-sql/src/sql_expr.rs | 43 ++++++++++--------- crates/polars-sql/src/sql_visitors.rs | 28 ++++++++++--- py-polars/tests/unit/sql/test_operators.py | 32 +++++++++++++++ 4 files changed, 92 insertions(+), 59 deletions(-) diff --git a/crates/polars-sql/src/context.rs b/crates/polars-sql/src/context.rs index 9f8b9e856cb2..409f438a63fc 100644 --- a/crates/polars-sql/src/context.rs +++ b/crates/polars-sql/src/context.rs @@ -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()); + // shortcut filter evaluation for a constant condition (eg: "WHERE 1 = 1") + if !expr_references_any_column(expr) && !expr_contains_subquery(expr) { + let satisfied = evaluate_constant_predicate(self, expr)?; + return Ok(if satisfied == (filter_mode == FilterMode::KeepTrue) { + lf + } else { + lf.clear() + }); } // Lower eligible `[NOT] EXISTS` / `[NOT] IN (subquery)` conjuncts @@ -2233,7 +2219,7 @@ impl SQLContext { // 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 satisfied = evaluate_constant_predicate(self, expr)?; let builder = tbl_left .frame .clone() @@ -3863,21 +3849,15 @@ fn suffix_conflicting_columns( }) } -/// 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 { +/// Evaluate a column-free (constant) condition to a definite true/false verdict; +/// SQL treats an unknown (NULL) condition the same as false for matching. +fn evaluate_constant_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")]) + .select([predicate.cast(DataType::Boolean).alias("predicate")]) .collect()?; - Ok(df - .column("_constant_join_predicate")? - .bool()? - .get(0) - .unwrap_or(false)) + Ok(df.column("predicate")?.bool()?.get(0).unwrap_or(false)) } fn process_join_constraint( diff --git a/crates/polars-sql/src/sql_expr.rs b/crates/polars-sql/src/sql_expr.rs index 37cd67375f9f..afef52cbd8ac 100644 --- a/crates/polars-sql/src/sql_expr.rs +++ b/crates/polars-sql/src/sql_expr.rs @@ -708,20 +708,7 @@ impl SQLExprVisitor<'_> { op, SQLBinaryOperator::Eq | SQLBinaryOperator::NotEq | SQLBinaryOperator::Spaceship ) { - // `str_expr = 13` compares against the string '13' - match (&lhs, &rhs) { - (Expr::Literal(LiteralValue::Dyn(DynLiteralValue::Int(n))), other) - if self.is_string_expr(other) => - { - lhs = lit(n.to_string()) - }, - (other, Expr::Literal(LiteralValue::Dyn(DynLiteralValue::Int(n)))) - if self.is_string_expr(other) => - { - rhs = lit(n.to_string()) - }, - _ => {}, - } + (lhs, rhs) = self.convert_int_literal_for_string(lhs, rhs); } if matches!(op, SQLBinaryOperator::Plus | SQLBinaryOperator::Minus) @@ -990,13 +977,14 @@ impl SQLExprVisitor<'_> { negated: bool, ) -> PolarsResult { 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); - for e in elements { + let mut membership: Option = None; + for e in list { 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); + let eq = lhs.eq(rhs); + membership = Some(membership.map_or(eq.clone(), |m| m.or(eq))); } + let membership = membership.unwrap(); Ok(if negated { membership.not() } else { @@ -1095,6 +1083,23 @@ impl SQLExprVisitor<'_> { }) } + /// `str_expr = 13` compares against the string '13'. + fn convert_int_literal_for_string(&self, lhs: Expr, rhs: Expr) -> (Expr, Expr) { + match (&lhs, &rhs) { + (Expr::Literal(LiteralValue::Dyn(DynLiteralValue::Int(n))), other) + if self.is_string_expr(other) => + { + (lit(n.to_string()), rhs) + }, + (other, Expr::Literal(LiteralValue::Dyn(DynLiteralValue::Int(n)))) + if self.is_string_expr(other) => + { + (lhs, lit(n.to_string())) + }, + _ => (lhs, rhs), + } + } + /// Whether `expr` is known to be `String`; false if the dtype cannot be resolved. fn is_string_expr(&self, expr: &Expr) -> bool { let empty = Schema::default(); 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/py-polars/tests/unit/sql/test_operators.py b/py-polars/tests/unit/sql/test_operators.py index 9b1e29fd1318..0894e7146fb4 100644 --- a/py-polars/tests/unit/sql/test_operators.py +++ b/py-polars/tests/unit/sql/test_operators.py @@ -98,6 +98,38 @@ def test_string_compared_with_integer_literals() -> None: 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( "in_clause", From 0f92f9ba32bed0b54fd1b4aaf383c1d1108df3b6 Mon Sep 17 00:00:00 2001 From: ritchie46 Date: Sat, 12 Sep 2026 19:39:38 +0200 Subject: [PATCH 04/15] fix: Only shortcut WHERE/ON conditions that are independent of the frame Window, aggregate and selector expressions name no column but still read the input; the constant shortcut now checks the parsed expression. Equi-join keys get the same string/integer literal coercion as `=`. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_01Lj7JbtU82rEjCBvQ98absi --- crates/polars-sql/src/context.rs | 54 ++++++++++++++++------ crates/polars-sql/src/sql_expr.rs | 14 ++++++ py-polars/tests/unit/sql/test_operators.py | 27 +++++++++++ 3 files changed, 80 insertions(+), 15 deletions(-) diff --git a/crates/polars-sql/src/context.rs b/crates/polars-sql/src/context.rs index 409f438a63fc..f0dedb4c07fe 100644 --- a/crates/polars-sql/src/context.rs +++ b/crates/polars-sql/src/context.rs @@ -27,8 +27,8 @@ 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, + order_by_sort_options, parse_sql_array, parse_sql_equality_operands, parse_sql_expr, + resolve_compound_identifier, to_sql_interface_err, }; use crate::sql_visitors::{ QualifyExpression, TableIdentifierCollector, check_for_ambiguous_column_refs, @@ -2147,8 +2147,7 @@ impl SQLContext { }; // shortcut filter evaluation for a constant condition (eg: "WHERE 1 = 1") - if !expr_references_any_column(expr) && !expr_contains_subquery(expr) { - let satisfied = evaluate_constant_predicate(self, expr)?; + if let Some(satisfied) = evaluate_constant_predicate(self, expr, &schema)? { return Ok(if satisfied == (filter_mode == FilterMode::KeepTrue) { lf } else { @@ -2216,10 +2215,8 @@ impl SQLContext { 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_predicate(self, expr)?; + let join_schema = build_join_schema(tbl_left, tbl_right)?; + if let Some(satisfied) = evaluate_constant_predicate(self, expr, &join_schema)? { let builder = tbl_left .frame .clone() @@ -3594,8 +3591,9 @@ 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_on, right_on) = + parse_sql_equality_operands(expr_left, expr_right, ctx, Some(join_schema))?; + let (left_on, right_on) = (strip_join_aliases(left_on), strip_join_aliases(right_on)); // 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 @@ -3849,15 +3847,41 @@ fn suffix_conflicting_columns( }) } -/// Evaluate a column-free (constant) condition to a definite true/false verdict; -/// SQL treats an unknown (NULL) condition the same as false for matching. -fn evaluate_constant_predicate(ctx: &mut SQLContext, expr: &SQLExpr) -> PolarsResult { - let predicate = parse_sql_expr(expr, ctx, None)?; +/// Evaluate a condition that does not depend on the input frame (eg: `1 = 1`) to a +/// definite true/false verdict; SQL treats an unknown (NULL) condition the same as +/// false for matching. Returns `None` for any other condition. +fn evaluate_constant_predicate( + ctx: &mut SQLContext, + expr: &SQLExpr, + schema: &Schema, +) -> PolarsResult> { + if expr_references_any_column(expr) || expr_contains_subquery(expr) { + return Ok(None); + } + let predicate = parse_sql_expr(expr, ctx, Some(schema))?; + // only literals and operations over them; anything reading the frame (a selector, + // `len()`, an aggregation or window) has no value without it + let is_constant = predicate.into_iter().all(|e| { + matches!( + e, + Expr::Literal(_) + | Expr::BinaryExpr { .. } + | Expr::Cast { .. } + | Expr::Ternary { .. } + | Expr::Function { .. } + | Expr::Alias(..) + ) + }); + if !is_constant { + return Ok(None); + } let df = DataFrame::empty() .lazy() .select([predicate.cast(DataType::Boolean).alias("predicate")]) .collect()?; - Ok(df.column("predicate")?.bool()?.get(0).unwrap_or(false)) + Ok(Some( + df.column("predicate")?.bool()?.get(0).unwrap_or(false), + )) } fn process_join_constraint( diff --git a/crates/polars-sql/src/sql_expr.rs b/crates/polars-sql/src/sql_expr.rs index afef52cbd8ac..af1e6a37da0c 100644 --- a/crates/polars-sql/src/sql_expr.rs +++ b/crates/polars-sql/src/sql_expr.rs @@ -1667,6 +1667,20 @@ pub(crate) fn parse_sql_expr( visitor.visit_expr(expr) } +/// Parse the two operands of an equality (eg: a join key pair), applying the same +/// literal coercion as `=` in an expression. +pub(crate) fn parse_sql_equality_operands( + left: &SQLExpr, + right: &SQLExpr, + ctx: &mut SQLContext, + active_schema: Option<&Schema>, +) -> PolarsResult<(Expr, Expr)> { + let mut visitor = SQLExprVisitor { ctx, active_schema }; + let lhs = visitor.visit_expr(left)?; + let rhs = visitor.visit_expr(right)?; + Ok(visitor.convert_int_literal_for_string(lhs, rhs)) +} + pub(crate) fn parse_sql_array(expr: &SQLExpr, ctx: &mut SQLContext) -> PolarsResult { match expr { SQLExpr::Array(arr) => { diff --git a/py-polars/tests/unit/sql/test_operators.py b/py-polars/tests/unit/sql/test_operators.py index 0894e7146fb4..607aad12430c 100644 --- a/py-polars/tests/unit/sql/test_operators.py +++ b/py-polars/tests/unit/sql/test_operators.py @@ -131,6 +131,33 @@ def test_constant_where_condition(condition: str, keeps_rows: bool) -> None: assert res.height == (3 if keeps_rows else 0) +@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)] + + @pytest.mark.parametrize( "in_clause", [ From 3c95e32fbb93eb3c466a311a2377d5b63f8b5ee1 Mon Sep 17 00:00:00 2001 From: ritchie46 Date: Sat, 12 Sep 2026 20:01:26 +0200 Subject: [PATCH 05/15] fix: Resolve join key dtypes against the operand's own table schema Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_01Lj7JbtU82rEjCBvQ98absi --- crates/polars-sql/src/context.rs | 24 +++++++-- crates/polars-sql/src/sql_expr.rs | 57 +++++++++++----------- py-polars/tests/unit/sql/test_operators.py | 11 +++++ 3 files changed, 59 insertions(+), 33 deletions(-) diff --git a/crates/polars-sql/src/context.rs b/crates/polars-sql/src/context.rs index f0dedb4c07fe..e3fa7686f374 100644 --- a/crates/polars-sql/src/context.rs +++ b/crates/polars-sql/src/context.rs @@ -27,7 +27,7 @@ use sqlparser::parser::{Parser, ParserOptions}; use crate::function_registry::{DefaultFunctionRegistry, FunctionRegistry}; use crate::sql_expr::{ - order_by_sort_options, parse_sql_array, parse_sql_equality_operands, parse_sql_expr, + 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::{ @@ -3591,9 +3591,25 @@ 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, right_on) = - parse_sql_equality_operands(expr_left, expr_right, ctx, Some(join_schema))?; - let (left_on, right_on) = (strip_join_aliases(left_on), strip_join_aliases(right_on)); + 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))?); + + // an operand's dtype comes from the table it names; the merged schema keeps the left + // dtype for a column name that exists in both tables + let operand_schema = |expr: &SQLExpr| -> &Schema { + match ( + expr_refers_to_table(expr, &tbl_left.name), + expr_refers_to_table(expr, &tbl_right.name), + ) { + (true, false) => &tbl_left.schema, + (false, true) => &tbl_right.schema, + _ => join_schema, + } + }; + let (left_on, right_on) = convert_int_literal_for_string( + (left_on, operand_schema(expr_left)), + (right_on, operand_schema(expr_right)), + ); // 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 diff --git a/crates/polars-sql/src/sql_expr.rs b/crates/polars-sql/src/sql_expr.rs index af1e6a37da0c..c8c99416cd96 100644 --- a/crates/polars-sql/src/sql_expr.rs +++ b/crates/polars-sql/src/sql_expr.rs @@ -1083,28 +1083,15 @@ impl SQLExprVisitor<'_> { }) } - /// `str_expr = 13` compares against the string '13'. fn convert_int_literal_for_string(&self, lhs: Expr, rhs: Expr) -> (Expr, Expr) { - match (&lhs, &rhs) { - (Expr::Literal(LiteralValue::Dyn(DynLiteralValue::Int(n))), other) - if self.is_string_expr(other) => - { - (lit(n.to_string()), rhs) - }, - (other, Expr::Literal(LiteralValue::Dyn(DynLiteralValue::Int(n)))) - if self.is_string_expr(other) => - { - (lhs, lit(n.to_string())) - }, - _ => (lhs, rhs), - } + let empty = Schema::default(); + let schema = self.active_schema.unwrap_or(&empty); + convert_int_literal_for_string((lhs, schema), (rhs, schema)) } - /// Whether `expr` is known to be `String`; false if the dtype cannot be resolved. fn is_string_expr(&self, expr: &Expr) -> bool { let empty = Schema::default(); - let schema = self.active_schema.unwrap_or(&empty); - matches!(expr.to_field(schema), Ok(fld) if fld.dtype == DataType::String) + is_string_expr(expr, self.active_schema.unwrap_or(&empty)) } /// Visit a SQL literal. @@ -1667,18 +1654,30 @@ pub(crate) fn parse_sql_expr( visitor.visit_expr(expr) } -/// Parse the two operands of an equality (eg: a join key pair), applying the same -/// literal coercion as `=` in an expression. -pub(crate) fn parse_sql_equality_operands( - left: &SQLExpr, - right: &SQLExpr, - ctx: &mut SQLContext, - active_schema: Option<&Schema>, -) -> PolarsResult<(Expr, Expr)> { - let mut visitor = SQLExprVisitor { ctx, active_schema }; - let lhs = visitor.visit_expr(left)?; - let rhs = visitor.visit_expr(right)?; - Ok(visitor.convert_int_literal_for_string(lhs, rhs)) +/// 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 { diff --git a/py-polars/tests/unit/sql/test_operators.py b/py-polars/tests/unit/sql/test_operators.py index 607aad12430c..3ff6749927de 100644 --- a/py-polars/tests/unit/sql/test_operators.py +++ b/py-polars/tests/unit/sql/test_operators.py @@ -157,6 +157,17 @@ def test_join_key_string_compared_with_integer_literal() -> None: ) 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)] + @pytest.mark.parametrize( "in_clause", From a485dbbcb738057496afe351e10f99d48ddaa2a3 Mon Sep 17 00:00:00 2001 From: ritchie46 Date: Sat, 12 Sep 2026 20:09:30 +0200 Subject: [PATCH 06/15] fix: Parse each join key operand against its own table schema Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_01Lj7JbtU82rEjCBvQ98absi --- crates/polars-sql/src/context.rs | 14 ++++++-------- py-polars/tests/unit/sql/test_operators.py | 11 +++++++++++ 2 files changed, 17 insertions(+), 8 deletions(-) diff --git a/crates/polars-sql/src/context.rs b/crates/polars-sql/src/context.rs index e3fa7686f374..98e73d3c52e3 100644 --- a/crates/polars-sql/src/context.rs +++ b/crates/polars-sql/src/context.rs @@ -3591,10 +3591,7 @@ 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))?); - - // an operand's dtype comes from the table it names; the merged schema keeps the left + // 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 = |expr: &SQLExpr| -> &Schema { match ( @@ -3606,10 +3603,11 @@ fn determine_left_right_join_on( _ => join_schema, } }; - let (left_on, right_on) = convert_int_literal_for_string( - (left_on, operand_schema(expr_left)), - (right_on, operand_schema(expr_right)), - ); + let (left_schema, right_schema) = (operand_schema(expr_left), operand_schema(expr_right)); + 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 diff --git a/py-polars/tests/unit/sql/test_operators.py b/py-polars/tests/unit/sql/test_operators.py index 3ff6749927de..e5290826cf82 100644 --- a/py-polars/tests/unit/sql/test_operators.py +++ b/py-polars/tests/unit/sql/test_operators.py @@ -168,6 +168,17 @@ def test_join_key_string_compared_with_integer_literal() -> None: ) 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", From 735b751436aeb5ea34cd524a0cdb3318e7868717 Mon Sep 17 00:00:00 2001 From: ritchie46 Date: Sat, 12 Sep 2026 20:23:47 +0200 Subject: [PATCH 07/15] refactor: Tidy the SQL join and constant-predicate helpers --- crates/polars-sql/src/context.rs | 56 +++++++++++++------------------ crates/polars-sql/src/sql_expr.rs | 20 +++++------ 2 files changed, 31 insertions(+), 45 deletions(-) diff --git a/crates/polars-sql/src/context.rs b/crates/polars-sql/src/context.rs index 98e73d3c52e3..93f548877cc2 100644 --- a/crates/polars-sql/src/context.rs +++ b/crates/polars-sql/src/context.rs @@ -2147,7 +2147,7 @@ impl SQLContext { }; // shortcut filter evaluation for a constant condition (eg: "WHERE 1 = 1") - if let Some(satisfied) = evaluate_constant_predicate(self, expr, &schema)? { + if let Some(satisfied) = evaluate_constant_predicate(self, expr)? { return Ok(if satisfied == (filter_mode == FilterMode::KeepTrue) { lf } else { @@ -2215,8 +2215,7 @@ impl SQLContext { join_type: JoinType, ) -> PolarsResult { if let JoinConstraint::On(expr) = constraint { - let join_schema = build_join_schema(tbl_left, tbl_right)?; - if let Some(satisfied) = evaluate_constant_predicate(self, expr, &join_schema)? { + if let Some(satisfied) = evaluate_constant_predicate(self, expr)? { let builder = tbl_left .frame .clone() @@ -3591,19 +3590,24 @@ 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_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 = |expr: &SQLExpr| -> &Schema { - match ( - expr_refers_to_table(expr, &tbl_left.name), - expr_refers_to_table(expr, &tbl_right.name), - ) { + 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(expr_left), operand_schema(expr_right)); + 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) = @@ -3619,14 +3623,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 @@ -3732,8 +3728,8 @@ fn process_join_on( } } -/// Parse a non-equi join condition into a `join_where` predicate over the joined frame, -/// where right-table columns that also exist in the left table carry a suffix. +/// 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, @@ -3743,14 +3739,12 @@ fn process_join_predicate( let suffix = format!(":{}", tbl_right.name); let conflicts = |name: &str| tbl_left.schema.contains(name) && tbl_right.schema.contains(name); - let mut joined_schema = Schema::clone(&tbl_left.schema); - for (name, dtype) in tbl_right.schema.iter() { - let name = if conflicts(name) { - PlSmallStr::from_string(format!("{name}{suffix}")) - } else { - name.clone() - }; - joined_schema.insert(name, dtype.clone()); + 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 @@ -3864,15 +3858,11 @@ fn suffix_conflicting_columns( /// Evaluate a condition that does not depend on the input frame (eg: `1 = 1`) to a /// definite true/false verdict; SQL treats an unknown (NULL) condition the same as /// false for matching. Returns `None` for any other condition. -fn evaluate_constant_predicate( - ctx: &mut SQLContext, - expr: &SQLExpr, - schema: &Schema, -) -> PolarsResult> { +fn evaluate_constant_predicate(ctx: &mut SQLContext, expr: &SQLExpr) -> PolarsResult> { if expr_references_any_column(expr) || expr_contains_subquery(expr) { return Ok(None); } - let predicate = parse_sql_expr(expr, ctx, Some(schema))?; + let predicate = parse_sql_expr(expr, ctx, None)?; // only literals and operations over them; anything reading the frame (a selector, // `len()`, an aggregation or window) has no value without it let is_constant = predicate.into_iter().all(|e| { diff --git a/crates/polars-sql/src/sql_expr.rs b/crates/polars-sql/src/sql_expr.rs index c8c99416cd96..2566587a1a33 100644 --- a/crates/polars-sql/src/sql_expr.rs +++ b/crates/polars-sql/src/sql_expr.rs @@ -977,14 +977,15 @@ impl SQLExprVisitor<'_> { negated: bool, ) -> PolarsResult { polars_ensure!(!list.is_empty(), SQLSyntax: "IN list must not be empty"); - let mut membership: Option = None; - for e in list { + let mut elements = list.iter(); + let first = self.visit_expr(elements.next().unwrap())?; + 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)?; let (lhs, rhs) = self.convert_int_literal_for_string(expr.clone(), e); - let eq = lhs.eq(rhs); - membership = Some(membership.map_or(eq.clone(), |m| m.or(eq))); + membership = membership.or(lhs.eq(rhs)); } - let membership = membership.unwrap(); Ok(if negated { membership.not() } else { @@ -1015,7 +1016,7 @@ impl SQLExprVisitor<'_> { } } if elems.dtype().is_integer() - && dtype_expr_match.is_some_and(|expr| self.is_string_expr(expr)) + && dtype_expr_match.is_some_and(|expr| self.expr_dtype(expr) == Some(DataType::String)) { return elems.cast(&DataType::String); } @@ -1071,7 +1072,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); @@ -1089,11 +1090,6 @@ impl SQLExprVisitor<'_> { convert_int_literal_for_string((lhs, schema), (rhs, schema)) } - fn is_string_expr(&self, expr: &Expr) -> bool { - let empty = Schema::default(); - is_string_expr(expr, self.active_schema.unwrap_or(&empty)) - } - /// Visit a SQL literal. /// /// e.g. 1, 'foo', 1.0, NULL From c72c69d5d6ad592a759dff0d9270ed27280f2989 Mon Sep 17 00:00:00 2001 From: ritchie46 Date: Sun, 13 Sep 2026 09:06:25 +0200 Subject: [PATCH 08/15] refactor: Split SQL literal folding into syntax policy and exact arithmetic `polars_compute::decimal::exact` owns checked fixed-point scalar arithmetic on `(mantissa, scale)` pairs; `polars_sql::literal_folding` owns which SQL forms qualify, reading the literal spelling, and the single conversion to Float64. --- Cargo.lock | 1 + crates/polars-compute/src/decimal.rs | 72 ++++++++++++++++++ crates/polars-sql/Cargo.toml | 1 + crates/polars-sql/src/lib.rs | 1 + crates/polars-sql/src/literal_folding.rs | 73 ++++++++++++++++++ crates/polars-sql/src/sql_expr.rs | 95 +----------------------- py-polars/tests/unit/sql/test_numeric.py | 63 ++++++++++++++++ 7 files changed, 213 insertions(+), 93 deletions(-) create mode 100644 crates/polars-sql/src/literal_folding.rs 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-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/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..8e653c6d17ad --- /dev/null +++ b/crates/polars-sql/src/literal_folding.rs @@ -0,0 +1,73 @@ +//! Exact folding of arithmetic between numeric SQL literals. +//! +//! `.06 + 0.01` is computed on the literal spellings as fixed-point values, so the +//! result is the float nearest to `0.07` rather than the float sum of two rounded +//! floats. Intermediate arithmetic is exact within `i128`; the result is still 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 2566587a1a33..bf3a2e8cff5c 100644 --- a/crates/polars-sql/src/sql_expr.rs +++ b/crates/polars-sql/src/sql_expr.rs @@ -33,6 +33,7 @@ use sqlparser::tokenizer::Token; use crate::SQLContext; use crate::functions::SQLFunctionVisitor; +use crate::literal_folding::try_fold_decimal_arithmetic; 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, @@ -664,7 +665,7 @@ impl SQLExprVisitor<'_> { op: &SQLBinaryOperator, right: &SQLExpr, ) -> PolarsResult { - if let Some(folded) = fold_decimal_literal_arithmetic(left, op, right) { + if let Some(folded) = try_fold_decimal_arithmetic(left, op, right) { return Ok(folded); } // need special handling for interval offsets and comparisons @@ -1509,98 +1510,6 @@ pub fn sql_expr>(s: S) -> PolarsResult { }) } -/// A fixed-point value: `mantissa / 10^scale`. -struct DecimalLiteral { - mantissa: i128, - scale: u32, -} - -impl DecimalLiteral { - fn parse(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; - } - Some(Self { - mantissa: format!("{int_part}{frac_part}").parse().ok()?, - scale: frac_part.len() as u32, - }) - } - - fn rescale(&self, scale: u32) -> Option { - self.mantissa - .checked_mul(10i128.checked_pow(scale - self.scale)?) - } - - fn eval(expr: &SQLExpr) -> Option { - match expr { - SQLExpr::Value(ValueWithSpan { - value: SQLValue::Number(s, _), - .. - }) => Self::parse(s), - SQLExpr::Nested(e) => Self::eval(e), - SQLExpr::UnaryOp { op, expr } => { - let v = Self::eval(expr)?; - match op { - SQLUnaryOperator::Plus => Some(v), - SQLUnaryOperator::Minus => Some(Self { - mantissa: v.mantissa.checked_neg()?, - ..v - }), - _ => None, - } - }, - SQLExpr::BinaryOp { left, op, right } => Self::combine(left, op, right), - _ => None, - } - } - - fn combine(left: &SQLExpr, op: &SQLBinaryOperator, right: &SQLExpr) -> Option { - if !matches!( - op, - SQLBinaryOperator::Plus | SQLBinaryOperator::Minus | SQLBinaryOperator::Multiply - ) { - return None; - } - let (l, r) = (Self::eval(left)?, Self::eval(right)?); - if *op == SQLBinaryOperator::Multiply { - return Some(Self { - mantissa: l.mantissa.checked_mul(r.mantissa)?, - scale: l.scale + r.scale, - }); - } - let scale = l.scale.max(r.scale); - let (l, r) = (l.rescale(scale)?, r.rescale(scale)?); - let mantissa = if *op == SQLBinaryOperator::Plus { - l.checked_add(r)? - } else { - l.checked_sub(r)? - }; - Some(Self { mantissa, scale }) - } - - fn to_f64(&self) -> f64 { - format!("{}e-{}", self.mantissa, self.scale) - .parse() - .unwrap() - } -} - -/// Evaluate `+`/`-`/`*` between numeric literals exactly; integer-only arithmetic is -/// left to the engine. -fn fold_decimal_literal_arithmetic( - left: &SQLExpr, - op: &SQLBinaryOperator, - right: &SQLExpr, -) -> Option { - let value = DecimalLiteral::combine(left, op, right)?; - (value.scale > 0).then(|| lit(value.to_f64())) -} - pub(crate) fn interval_to_duration(interval: &Interval, fixed: bool) -> PolarsResult { if interval.last_field.is_some() || interval.leading_precision.is_some() diff --git a/py-polars/tests/unit/sql/test_numeric.py b/py-polars/tests/unit/sql/test_numeric.py index c04440c55cba..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 @@ -238,6 +239,68 @@ def test_decimal_literal_arithmetic_is_exact() -> None: 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: From 1d41858494315b5a896ce69264a2d55a4e62c0ca Mon Sep 17 00:00:00 2001 From: ritchie46 Date: Sun, 13 Sep 2026 09:21:50 +0200 Subject: [PATCH 09/15] refactor: Plan constant SQL predicates instead of evaluating them at translation A WHERE condition is always translated to a filter; the planner folds a constant one. A column-free JOIN ON condition becomes a boolean join key against `true` whenever the planner proves it input-independent, so true, false and NULL conditions all defer to normal execution for every join type. An inner join on the same non-null constant key on both sides lowers to a cross join. --- crates/polars-plan/src/dsl/meta.rs | 7 ++ .../src/plans/aexpr/properties/general.rs | 19 +++++ .../src/plans/conversion/dsl_to_ir/join.rs | 20 +++++ crates/polars-sql/src/context.rs | 80 +++++-------------- py-polars/tests/unit/sql/test_joins.py | 63 +++++++++++++++ py-polars/tests/unit/sql/test_operators.py | 61 ++++++++++++++ 6 files changed, 190 insertions(+), 60 deletions(-) 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..9c64df67e43b 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 @@ -156,6 +156,26 @@ 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. + 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 + && 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/src/context.rs b/crates/polars-sql/src/context.rs index 93f548877cc2..e3513b03d8b7 100644 --- a/crates/polars-sql/src/context.rs +++ b/crates/polars-sql/src/context.rs @@ -2146,15 +2146,6 @@ impl SQLContext { Some(s) => s, }; - // shortcut filter evaluation for a constant condition (eg: "WHERE 1 = 1") - if let Some(satisfied) = evaluate_constant_predicate(self, expr)? { - return Ok(if satisfied == (filter_mode == FilterMode::KeepTrue) { - lf - } else { - lf.clear() - }); - } - // Lower eligible `[NOT] EXISTS` / `[NOT] IN (subquery)` conjuncts // to semi / anti joins; whatever remains goes through the ordinary // filter path below. @@ -2214,29 +2205,31 @@ impl SQLContext { constraint: &JoinConstraint, join_type: JoinType, ) -> PolarsResult { - if let JoinConstraint::On(expr) = constraint { - if let Some(satisfied) = evaluate_constant_predicate(self, expr)? { - let builder = tbl_left + // 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. Null keys do not + // match, which is SQL's treatment of an unknown condition. + if let JoinConstraint::On(expr) = constraint + && !expr_references_any_column(expr) + && !expr_contains_subquery(expr) + { + let predicate = parse_sql_expr(expr, self, None)?; + if predicate + .clone() + .meta() + .is_input_independent_scalar() + .unwrap_or(false) + { + return tbl_left .frame .clone() .join_builder() .with(tbl_right.frame.clone()) + .left_on([predicate.cast(DataType::Boolean)]) + .right_on([lit(true)]) + .how(join_type) .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()? - }); + .coalesce(JoinCoalesce::KeepColumns) + .finish(); } } let (left_on, right_on, predicates) = @@ -3855,39 +3848,6 @@ fn suffix_conflicting_columns( }) } -/// Evaluate a condition that does not depend on the input frame (eg: `1 = 1`) to a -/// definite true/false verdict; SQL treats an unknown (NULL) condition the same as -/// false for matching. Returns `None` for any other condition. -fn evaluate_constant_predicate(ctx: &mut SQLContext, expr: &SQLExpr) -> PolarsResult> { - if expr_references_any_column(expr) || expr_contains_subquery(expr) { - return Ok(None); - } - let predicate = parse_sql_expr(expr, ctx, None)?; - // only literals and operations over them; anything reading the frame (a selector, - // `len()`, an aggregation or window) has no value without it - let is_constant = predicate.into_iter().all(|e| { - matches!( - e, - Expr::Literal(_) - | Expr::BinaryExpr { .. } - | Expr::Cast { .. } - | Expr::Ternary { .. } - | Expr::Function { .. } - | Expr::Alias(..) - ) - }); - if !is_constant { - return Ok(None); - } - let df = DataFrame::empty() - .lazy() - .select([predicate.cast(DataType::Boolean).alias("predicate")]) - .collect()?; - Ok(Some( - df.column("predicate")?.bool()?.get(0).unwrap_or(false), - )) -} - fn process_join_constraint( constraint: &JoinConstraint, tbl_left: &TableInfo, diff --git a/py-polars/tests/unit/sql/test_joins.py b/py-polars/tests/unit/sql/test_joins.py index fba271d7df52..9f461b82196e 100644 --- a/py-polars/tests/unit/sql/test_joins.py +++ b/py-polars/tests/unit/sql/test_joins.py @@ -1952,6 +1952,69 @@ def test_join_predicate_operand_spanning_both_sides() -> None: ) +@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", + ], +) +@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") + + +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 + + @pytest.mark.parametrize("join_type", ["INNER", "LEFT"]) def test_join_on_pattern_predicates(join_type: str) -> None: frames = { diff --git a/py-polars/tests/unit/sql/test_operators.py b/py-polars/tests/unit/sql/test_operators.py index e5290826cf82..6bfcea4d3dfb 100644 --- a/py-polars/tests/unit/sql/test_operators.py +++ b/py-polars/tests/unit/sql/test_operators.py @@ -131,6 +131,67 @@ def test_constant_where_condition(condition: str, keeps_rows: bool) -> None: 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), + ("UPPER('x') = 'X'", True), + ("CASE WHEN 1 < 2 THEN FALSE ELSE TRUE END", 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"), [ From 2f6361cbea933e62120d9f6eee0219a39640d5c6 Mon Sep 17 00:00:00 2001 From: ritchie46 Date: Sun, 13 Sep 2026 09:38:49 +0200 Subject: [PATCH 10/15] fix: Keep constant membership, validation and boolean coercion on the lazy predicate path A literal on the left of `IN` is compared as an OR-chain, which the planner folds; a constant WHERE condition is cast to boolean lazily; the inner join to cross join rewrite only applies to unvalidated joins. --- .../src/plans/conversion/dsl_to_ir/join.rs | 4 +- crates/polars-sql/src/context.rs | 56 ++++++++++++------- crates/polars-sql/src/sql_expr.rs | 10 ++-- py-polars/tests/unit/sql/test_joins.py | 17 ++++++ py-polars/tests/unit/sql/test_operators.py | 6 ++ 5 files changed, 68 insertions(+), 25 deletions(-) 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 9c64df67e43b..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; @@ -157,7 +158,7 @@ pub fn resolve_join( 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. + // 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()), @@ -166,6 +167,7 @@ pub fn resolve_join( _ => false, }; if options.args.how == JoinType::Inner + && options.args.validation == JoinValidation::ManyToMany && left_on .iter() .zip(&right_on) diff --git a/crates/polars-sql/src/context.rs b/crates/polars-sql/src/context.rs index e3513b03d8b7..99285ea4f012 100644 --- a/crates/polars-sql/src/context.rs +++ b/crates/polars-sql/src/context.rs @@ -2146,6 +2146,16 @@ impl SQLContext { Some(s) => s, }; + // 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 // to semi / anti joins; whatever remains goes through the ordinary // filter path below. @@ -2198,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, @@ -2209,28 +2234,19 @@ impl SQLContext { // every row, or none: join on the condition itself as a boolean key. Null keys do not // match, which is SQL's treatment of an unknown condition. if let JoinConstraint::On(expr) = constraint - && !expr_references_any_column(expr) - && !expr_contains_subquery(expr) + && let Some(predicate) = self.input_independent_predicate(expr)? { - let predicate = parse_sql_expr(expr, self, None)?; - if predicate + return tbl_left + .frame .clone() - .meta() - .is_input_independent_scalar() - .unwrap_or(false) - { - return tbl_left - .frame - .clone() - .join_builder() - .with(tbl_right.frame.clone()) - .left_on([predicate.cast(DataType::Boolean)]) - .right_on([lit(true)]) - .how(join_type) - .suffix(format!(":{}", tbl_right.name)) - .coalesce(JoinCoalesce::KeepColumns) - .finish(); - } + .join_builder() + .with(tbl_right.frame.clone()) + .left_on([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)?; diff --git a/crates/polars-sql/src/sql_expr.rs b/crates/polars-sql/src/sql_expr.rs index bf3a2e8cff5c..f8f6974733be 100644 --- a/crates/polars-sql/src/sql_expr.rs +++ b/crates/polars-sql/src/sql_expr.rs @@ -295,10 +295,12 @@ impl SQLExprVisitor<'_> { } => { let expr = self.visit_expr(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, an aggregate on the left, or a literal on the + // left (a constant, which the planner folds as an OR-chain but not as a set + // membership) falls back to an OR-chain of equality comparisons. + let use_or_chain = matches!(expr, Expr::Literal(_)) + || 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() diff --git a/py-polars/tests/unit/sql/test_joins.py b/py-polars/tests/unit/sql/test_joins.py index 9f461b82196e..5c914040d63f 100644 --- a/py-polars/tests/unit/sql/test_joins.py +++ b/py-polars/tests/unit/sql/test_joins.py @@ -1977,6 +1977,9 @@ def test_join_predicate_operand_spanning_both_sides() -> None: "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')", ], ) @pytest.mark.parametrize("empty_side", [None, "a", "b"]) @@ -2015,6 +2018,20 @@ def test_join_on_constant_true_plans_cross_join() -> None: 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 = { diff --git a/py-polars/tests/unit/sql/test_operators.py b/py-polars/tests/unit/sql/test_operators.py index 6bfcea4d3dfb..2de6224be581 100644 --- a/py-polars/tests/unit/sql/test_operators.py +++ b/py-polars/tests/unit/sql/test_operators.py @@ -139,8 +139,14 @@ def test_constant_where_condition(condition: str, keeps_rows: bool) -> None: ("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), ], ) @pytest.mark.parametrize("empty", [False, True]) From 4906f71655132fc6696dc4e529563bae07211e8d Mon Sep 17 00:00:00 2001 From: ritchie46 Date: Sun, 13 Sep 2026 09:48:28 +0200 Subject: [PATCH 11/15] fix: Compare any constant left operand of IN as an OR-chain --- crates/polars-sql/src/sql_expr.rs | 18 +++++++++++++----- py-polars/tests/unit/sql/test_joins.py | 3 +++ 2 files changed, 16 insertions(+), 5 deletions(-) diff --git a/crates/polars-sql/src/sql_expr.rs b/crates/polars-sql/src/sql_expr.rs index f8f6974733be..5fbfb39d2818 100644 --- a/crates/polars-sql/src/sql_expr.rs +++ b/crates/polars-sql/src/sql_expr.rs @@ -34,6 +34,7 @@ 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, @@ -293,13 +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, an aggregate on the left, or a literal on the - // left (a constant, which the planner folds as an OR-chain but not as a set + // use. A non-literal element, an aggregate on the left, or a constant on + // the left (which the planner folds as an OR-chain but not as a set // membership) falls back to an OR-chain of equality comparisons. - let use_or_chain = matches!(expr, Expr::Literal(_)) - || has_expr(&expr, |e| matches!(e, Expr::Agg(_) | Expr::Len)); + 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 { diff --git a/py-polars/tests/unit/sql/test_joins.py b/py-polars/tests/unit/sql/test_joins.py index 5c914040d63f..630875ddefaf 100644 --- a/py-polars/tests/unit/sql/test_joins.py +++ b/py-polars/tests/unit/sql/test_joins.py @@ -1980,6 +1980,9 @@ def test_join_predicate_operand_spanning_both_sides() -> None: "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)", ], ) @pytest.mark.parametrize("empty_side", [None, "a", "b"]) From f35484349819b63fb2b7b36ee9f9494b95c106f0 Mon Sep 17 00:00:00 2001 From: ritchie46 Date: Sun, 13 Sep 2026 09:57:07 +0200 Subject: [PATCH 12/15] fix: Build SQL array literals as scalar list values An `ARRAY[...]` literal is one list value, not a one-row Series, so the planner classifies expressions over it as input-independent. --- crates/polars-sql/src/sql_expr.rs | 14 +++++------ py-polars/tests/unit/sql/test_array.py | 2 +- py-polars/tests/unit/sql/test_joins.py | 29 ++++++++++++++++++++++ py-polars/tests/unit/sql/test_operators.py | 2 ++ 4 files changed, 39 insertions(+), 8 deletions(-) diff --git a/crates/polars-sql/src/sql_expr.rs b/crates/polars-sql/src/sql_expr.rs index 5fbfb39d2818..223bf4d11f31 100644 --- a/crates/polars-sql/src/sql_expr.rs +++ b/crates/polars-sql/src/sql_expr.rs @@ -1044,14 +1044,14 @@ 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 dtype = DataType::List(Box::new(elems.dtype().clone())); + lit(Scalar::new(dtype, AnyValue::List(elems))) } else { - elems - }; - Ok(lit(res)) + lit(elems) + }) } /// Visit a SQL `CAST` or `TRY_CAST` expression. diff --git a/py-polars/tests/unit/sql/test_array.py b/py-polars/tests/unit/sql/test_array.py index 4c4fc5354ed4..48c4b49d278e 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 diff --git a/py-polars/tests/unit/sql/test_joins.py b/py-polars/tests/unit/sql/test_joins.py index 630875ddefaf..b556d9f74959 100644 --- a/py-polars/tests/unit/sql/test_joins.py +++ b/py-polars/tests/unit/sql/test_joins.py @@ -1983,6 +1983,8 @@ def test_join_predicate_operand_spanning_both_sides() -> None: "(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"]) @@ -2007,6 +2009,33 @@ def test_join_on_constant_condition( 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]}), diff --git a/py-polars/tests/unit/sql/test_operators.py b/py-polars/tests/unit/sql/test_operators.py index 2de6224be581..6e1b08394d5b 100644 --- a/py-polars/tests/unit/sql/test_operators.py +++ b/py-polars/tests/unit/sql/test_operators.py @@ -147,6 +147,8 @@ def test_constant_where_condition(condition: str, keeps_rows: bool) -> None: ("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]) From 17332f72d1a514ee101bbf843e45a848c148789a Mon Sep 17 00:00:00 2001 From: ritchie46 Date: Sun, 13 Sep 2026 10:08:36 +0200 Subject: [PATCH 13/15] fix: Keep the empty default column name of SQL array literals --- crates/polars-sql/src/context.rs | 2 +- crates/polars-sql/src/sql_expr.rs | 3 ++- py-polars/tests/unit/sql/test_array.py | 7 +++++++ 3 files changed, 10 insertions(+), 2 deletions(-) diff --git a/crates/polars-sql/src/context.rs b/crates/polars-sql/src/context.rs index 99285ea4f012..a8ecee4f5307 100644 --- a/crates/polars-sql/src/context.rs +++ b/crates/polars-sql/src/context.rs @@ -2241,7 +2241,7 @@ impl SQLContext { .clone() .join_builder() .with(tbl_right.frame.clone()) - .left_on([predicate.cast(DataType::Boolean)]) + .left_on([strip_join_aliases(predicate).cast(DataType::Boolean)]) .right_on([lit(true)]) .how(join_type) .suffix(format!(":{}", tbl_right.name)) diff --git a/crates/polars-sql/src/sql_expr.rs b/crates/polars-sql/src/sql_expr.rs index 223bf4d11f31..b1d59cdca895 100644 --- a/crates/polars-sql/src/sql_expr.rs +++ b/crates/polars-sql/src/sql_expr.rs @@ -1047,8 +1047,9 @@ impl SQLExprVisitor<'_> { // 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))) + lit(Scalar::new(dtype, AnyValue::List(elems))).alias(name) } else { lit(elems) }) diff --git a/py-polars/tests/unit/sql/test_array.py b/py-polars/tests/unit/sql/test_array.py index 48c4b49d278e..acc8b98c02d2 100644 --- a/py-polars/tests/unit/sql/test_array.py +++ b/py-polars/tests/unit/sql/test_array.py @@ -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]]} From 2a37aa6a2c4400e41ca6686e1e5aa81815f89b34 Mon Sep 17 00:00:00 2001 From: ritchie46 Date: Sun, 13 Sep 2026 10:13:19 +0200 Subject: [PATCH 14/15] docs: Tighten comments on the SQL constant handling --- crates/polars-sql/src/context.rs | 4 ++-- crates/polars-sql/src/literal_folding.rs | 7 +++---- crates/polars-sql/src/sql_expr.rs | 5 ++--- 3 files changed, 7 insertions(+), 9 deletions(-) diff --git a/crates/polars-sql/src/context.rs b/crates/polars-sql/src/context.rs index a8ecee4f5307..19af92462e9c 100644 --- a/crates/polars-sql/src/context.rs +++ b/crates/polars-sql/src/context.rs @@ -2231,8 +2231,8 @@ impl SQLContext { join_type: JoinType, ) -> PolarsResult { // 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. Null keys do not - // match, which is SQL's treatment of an unknown condition. + // 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)? { diff --git a/crates/polars-sql/src/literal_folding.rs b/crates/polars-sql/src/literal_folding.rs index 8e653c6d17ad..ccc11e2aaa94 100644 --- a/crates/polars-sql/src/literal_folding.rs +++ b/crates/polars-sql/src/literal_folding.rs @@ -1,9 +1,8 @@ //! Exact folding of arithmetic between numeric SQL literals. //! -//! `.06 + 0.01` is computed on the literal spellings as fixed-point values, so the -//! result is the float nearest to `0.07` rather than the float sum of two rounded -//! floats. Intermediate arithmetic is exact within `i128`; the result is still the -//! ordinary `Float64` literal. Integer-only arithmetic is left to the engine. +//! `.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}; diff --git a/crates/polars-sql/src/sql_expr.rs b/crates/polars-sql/src/sql_expr.rs index b1d59cdca895..3ce136ae0547 100644 --- a/crates/polars-sql/src/sql_expr.rs +++ b/crates/polars-sql/src/sql_expr.rs @@ -297,9 +297,8 @@ impl SQLExprVisitor<'_> { 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, an aggregate on the left, or a constant on - // the left (which the planner folds as an OR-chain but not as a set - // membership) falls back to an OR-chain of equality comparisons. + // 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() From 4baf69d1eaac9cccbde3872c442b20d27d54eca6 Mon Sep 17 00:00:00 2001 From: ritchie46 Date: Sun, 13 Sep 2026 10:30:55 +0200 Subject: [PATCH 15/15] test: Select the SQL array literal from a one-row frame --- crates/polars-sql/tests/functions_string.rs | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) 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();