Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
8 changes: 6 additions & 2 deletions crates/polars-plan/src/plans/conversion/dsl_to_ir/join.rs
Original file line number Diff line number Diff line change
Expand Up @@ -659,8 +659,12 @@ fn build_upcast_node_list(
right.to_dtype(&ToFieldContext::new(expr_arena, schema_merged))?;
if dtype_left != dtype_right {
// Ensure that we have a lossless cast between the two types.
let dt = if dtype_left.is_primitive_numeric()
|| dtype_right.is_primitive_numeric()
// Decimal has no lossless numeric upcast.
let either_decimal =
dtype_left.is_decimal() || dtype_right.is_decimal();
let dt = if !either_decimal
&& (dtype_left.is_primitive_numeric()
|| dtype_right.is_primitive_numeric())
{
get_numeric_upcast_supertype_lossless(&dtype_left, &dtype_right)
.ok_or(PolarsError::SchemaMismatch(
Expand Down
41 changes: 41 additions & 0 deletions crates/polars-sql/src/sql_expr.rs
Original file line number Diff line number Diff line change
Expand Up @@ -11,6 +11,7 @@ use std::ops::Div;

use polars_core::prelude::*;
use polars_lazy::prelude::*;
use polars_plan::dsl::functions::{DurationArgs, duration};
use polars_plan::plans::DynLiteralValue;
use polars_plan::prelude::{has_expr, typed_lit};
use polars_time::Duration;
Expand Down Expand Up @@ -608,6 +609,40 @@ impl SQLExprVisitor<'_> {
Ok(expr)
}

/// Best-effort dtype for an expression; `None` if it cannot be resolved.
fn expr_dtype(&self, expr: &Expr) -> Option<DataType> {
let empty = Schema::default();
let schema = self.active_schema.unwrap_or(&empty);
expr.to_field(schema).ok().map(|fld| fld.dtype)
}

/// `date + n` / `date - n` shift the date by a whole number of days.
fn date_day_offset(&self, lhs: &Expr, op: &SQLBinaryOperator, rhs: &Expr) -> Option<Expr> {
let subtract = matches!(op, SQLBinaryOperator::Minus);
let left_dtype = self.expr_dtype(lhs);
let right_dtype = self.expr_dtype(rhs);
let is_date = |dtype: &Option<DataType>| matches!(dtype, Some(DataType::Date));
let is_int =
|dtype: &Option<DataType>| dtype.as_ref().is_some_and(|dtype| dtype.is_integer());

let (date, days) = if is_date(&left_dtype) && is_int(&right_dtype) {
(lhs, rhs)
} else if !subtract && is_date(&right_dtype) && is_int(&left_dtype) {
(rhs, lhs)
} else {
return None;
};
let offset = duration(DurationArgs {
days: days.clone(),
..Default::default()
});
Some(if subtract {
date.clone() - offset
} else {
date.clone() + offset
})
}

/// Visit a SQL binary operator.
///
/// e.g. "column + 1", "column1 <= column2"
Expand Down Expand Up @@ -655,6 +690,12 @@ impl SQLExprVisitor<'_> {
};
rhs = self.convert_temporal_strings(&lhs, &rhs);

if matches!(op, SQLBinaryOperator::Plus | SQLBinaryOperator::Minus)
&& let Some(expr) = self.date_day_offset(&lhs, op, &rhs)
{
return Ok(expr);
}

Ok(match op {
// ----
// Bitwise operators
Expand Down
21 changes: 21 additions & 0 deletions py-polars/tests/unit/operations/test_inequality_join.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
from __future__ import annotations

from datetime import datetime
from decimal import Decimal
from typing import TYPE_CHECKING, Any

import hypothesis.strategies as st
Expand Down Expand Up @@ -972,3 +973,23 @@ def test_cross_join_validity_bitmap_offset_26925(
)

assert_frame_equal(actual, expected, check_exact=True)


def test_join_where_decimal_vs_float() -> None:
left = pl.LazyFrame(
{"k": [1, 1, 2], "amount": ["1.50", "4.00", "9.00"]}
).with_columns(pl.col("amount").cast(pl.Decimal(38, 2)))
right = pl.LazyFrame({"k": [1, 2], "limit": [3.0, 10.0]})

actual = left.join_where(
right,
pl.col("k") == pl.col("k_right"),
pl.col("amount") <= pl.col("limit"),
).collect()

assert actual.sort("amount").to_dict(as_series=False) == {
"k": [1, 2],
"amount": [Decimal("1.50"), Decimal("9.00")],
"k_right": [1, 2],
"limit": [3.0, 10.0],
}
43 changes: 43 additions & 0 deletions py-polars/tests/unit/sql/test_temporal.py
Original file line number Diff line number Diff line change
Expand Up @@ -530,3 +530,46 @@ def test_typed_timestamp_literal_precision(precision: int, time_unit: str) -> No
f"SELECT TIMESTAMP({precision}) '2020-01-01 08:00:00.123' AS x FROM tbl"
)
assert res.schema["x"] == pl.Datetime(time_unit) # type: ignore[arg-type]


def test_date_plus_integer_days() -> None:
df = pl.DataFrame(
{
"dt": [date(2020, 1, 1), date(2020, 2, 28), date(2021, 2, 28)],
"n": [5, 2, 2],
}
)
with pl.SQLContext(frames={"tbl": df}, eager=True) as ctx:
res = ctx.execute(
"""
SELECT
dt + 5 AS plus_lit,
5 + dt AS lit_plus,
dt - 5 AS minus_lit,
dt + n AS plus_col,
dt - n AS minus_col
FROM tbl
"""
)
assert res.to_dict(as_series=False) == {
"plus_lit": [date(2020, 1, 6), date(2020, 3, 4), date(2021, 3, 5)],
"lit_plus": [date(2020, 1, 6), date(2020, 3, 4), date(2021, 3, 5)],
"minus_lit": [date(2019, 12, 27), date(2020, 2, 23), date(2021, 2, 23)],
"plus_col": [date(2020, 1, 6), date(2020, 3, 1), date(2021, 3, 2)],
"minus_col": [date(2019, 12, 27), date(2020, 2, 26), date(2021, 2, 26)],
}


def test_date_integer_arithmetic_in_filter() -> None:
df = pl.DataFrame({"dt": [date(2020, 1, 1), date(2020, 1, 8), date(2020, 1, 15)]})
with pl.SQLContext(frames={"tbl": df}, eager=True) as ctx:
res = ctx.execute("SELECT dt FROM tbl WHERE dt > DATE '2020-01-01' + 7")
assert res.to_series().to_list() == [date(2020, 1, 15)]


def test_date_arithmetic_leaves_other_dtypes_alone() -> None:
df = pl.DataFrame({"a": [1, 2], "dt": [date(2020, 1, 1), date(2020, 3, 5)]})
with pl.SQLContext(frames={"tbl": df}, eager=True) as ctx:
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")
Loading