Skip to content

Commit 24764bd

Browse files
timsaucerclaude
andcommitted
feat: add null_treatment to lead and lag
IGNORE_NULLS skips null values when counting shift_offset rows, matching SQL LEAD/LAG ... IGNORE NULLS. The argument is appended after order_by, so existing calls are unaffected. Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
1 parent 2747573 commit 24764bd

3 files changed

Lines changed: 55 additions & 4 deletions

File tree

‎crates/core/src/functions.rs‎

Lines changed: 6 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -891,33 +891,35 @@ pub(crate) fn add_builder_fns_to_window(
891891
}
892892

893893
#[pyfunction]
894-
#[pyo3(signature = (arg, shift_offset, default_value=None, partition_by=None, order_by=None))]
894+
#[pyo3(signature = (arg, shift_offset, default_value=None, partition_by=None, order_by=None, null_treatment=None))]
895895
pub fn lead(
896896
arg: PyExpr,
897897
shift_offset: i64,
898898
default_value: Option<PyScalarValue>,
899899
partition_by: Option<Vec<PyExpr>>,
900900
order_by: Option<Vec<PySortExpr>>,
901+
null_treatment: Option<NullTreatment>,
901902
) -> PyDataFusionResult<PyExpr> {
902903
let default_value = default_value.map(|v| v.into());
903904
let window_fn = functions_window::expr_fn::lead(arg.expr, Some(shift_offset), default_value);
904905

905-
add_builder_fns_to_window(window_fn, partition_by, None, order_by, None)
906+
add_builder_fns_to_window(window_fn, partition_by, None, order_by, null_treatment)
906907
}
907908

908909
#[pyfunction]
909-
#[pyo3(signature = (arg, shift_offset, default_value=None, partition_by=None, order_by=None))]
910+
#[pyo3(signature = (arg, shift_offset, default_value=None, partition_by=None, order_by=None, null_treatment=None))]
910911
pub fn lag(
911912
arg: PyExpr,
912913
shift_offset: i64,
913914
default_value: Option<PyScalarValue>,
914915
partition_by: Option<Vec<PyExpr>>,
915916
order_by: Option<Vec<PySortExpr>>,
917+
null_treatment: Option<NullTreatment>,
916918
) -> PyDataFusionResult<PyExpr> {
917919
let default_value = default_value.map(|v| v.into());
918920
let window_fn = functions_window::expr_fn::lag(arg.expr, Some(shift_offset), default_value);
919921

920-
add_builder_fns_to_window(window_fn, partition_by, None, order_by, None)
922+
add_builder_fns_to_window(window_fn, partition_by, None, order_by, null_treatment)
921923
}
922924

923925
#[pyfunction]

‎python/datafusion/functions/__init__.py‎

Lines changed: 28 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -7000,6 +7000,7 @@ def lead(
70007000
default_value: Any | None = None,
70017001
partition_by: list[Expr] | Expr | None = None,
70027002
order_by: list[SortKey] | SortKey | None = None,
7003+
null_treatment: NullTreatment = NullTreatment.RESPECT_NULLS,
70037004
) -> Expr:
70047005
"""Create a lead window function.
70057006
@@ -7030,6 +7031,8 @@ def lead(
70307031
partition_by: Expressions to partition the window frame on.
70317032
order_by: Set ordering within the window frame. Accepts
70327033
column names or expressions.
7034+
null_treatment: Set to ``IGNORE_NULLS`` to skip null values when
7035+
counting ``shift_offset`` rows.
70337036
70347037
Examples:
70357038
>>> ctx = dfn.SessionContext()
@@ -7052,6 +7055,16 @@ def lead(
70527055
... ).alias("lead"))
70537056
>>> result.sort(dfn.col("g"), dfn.col("v")).collect_column("lead").to_pylist()
70547057
[2, 0, 0]
7058+
7059+
>>> df = ctx.from_pydict({"i": [1, 2, 3, 4], "v": [1, None, None, 4]})
7060+
>>> result = df.select(
7061+
... dfn.col("i"),
7062+
... dfn.functions.lead(
7063+
... dfn.col("v"), order_by="i",
7064+
... null_treatment=dfn.common.NullTreatment.IGNORE_NULLS,
7065+
... ).alias("lead"))
7066+
>>> result.sort(dfn.col("i")).collect_column("lead").to_pylist()
7067+
[4, 4, 4, None]
70557068
"""
70567069
if not isinstance(default_value, pa.Scalar) and default_value is not None:
70577070
default_value = pa.scalar(default_value)
@@ -7066,6 +7079,7 @@ def lead(
70667079
default_value,
70677080
partition_by=partition_by_raw,
70687081
order_by=order_by_raw,
7082+
null_treatment=null_treatment.value,
70697083
)
70707084
)
70717085

@@ -7076,6 +7090,7 @@ def lag(
70767090
default_value: Any | None = None,
70777091
partition_by: list[Expr] | Expr | None = None,
70787092
order_by: list[SortKey] | SortKey | None = None,
7093+
null_treatment: NullTreatment = NullTreatment.RESPECT_NULLS,
70797094
) -> Expr:
70807095
"""Create a lag window function.
70817096
@@ -7103,6 +7118,8 @@ def lag(
71037118
partition_by: Expressions to partition the window frame on.
71047119
order_by: Set ordering within the window frame. Accepts
71057120
column names or expressions.
7121+
null_treatment: Set to ``IGNORE_NULLS`` to skip null values when
7122+
counting ``shift_offset`` rows.
71067123
71077124
Examples:
71087125
>>> ctx = dfn.SessionContext()
@@ -7125,6 +7142,16 @@ def lag(
71257142
... ).alias("lag"))
71267143
>>> result.sort(dfn.col("g"), dfn.col("v")).collect_column("lag").to_pylist()
71277144
[0, 1, 0]
7145+
7146+
>>> df = ctx.from_pydict({"i": [1, 2, 3, 4], "v": [1, None, None, 4]})
7147+
>>> result = df.select(
7148+
... dfn.col("i"),
7149+
... dfn.functions.lag(
7150+
... dfn.col("v"), order_by="i",
7151+
... null_treatment=dfn.common.NullTreatment.IGNORE_NULLS,
7152+
... ).alias("lag"))
7153+
>>> result.sort(dfn.col("i")).collect_column("lag").to_pylist()
7154+
[None, 1, 1, 1]
71287155
"""
71297156
if not isinstance(default_value, pa.Scalar):
71307157
default_value = pa.scalar(default_value)
@@ -7139,6 +7166,7 @@ def lag(
71397166
default_value,
71407167
partition_by=partition_by_raw,
71417168
order_by=order_by_raw,
7169+
null_treatment=null_treatment.value,
71427170
)
71437171
)
71447172

‎python/tests/test_dataframe.py‎

Lines changed: 21 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -46,6 +46,7 @@
4646
from datafusion import (
4747
functions as f,
4848
)
49+
from datafusion.common import NullTreatment
4950
from datafusion.dataframe import DataFrameWriteOptions
5051
from datafusion.dataframe_formatter import (
5152
DataFrameHtmlFormatter,
@@ -1077,6 +1078,26 @@ def test_distinct():
10771078
),
10781079
[-1, -1, None, 7, -1, -1, None],
10791080
),
1081+
(
1082+
"lead_ignore_nulls",
1083+
f.lead(
1084+
column("b"),
1085+
order_by=column("a"),
1086+
partition_by=column("c"),
1087+
null_treatment=NullTreatment.IGNORE_NULLS,
1088+
),
1089+
[7, 7, 8, None, 9, 9, None],
1090+
),
1091+
(
1092+
"lag_ignore_nulls",
1093+
f.lag(
1094+
column("b"),
1095+
order_by=column("a"),
1096+
partition_by=column("c"),
1097+
null_treatment=NullTreatment.IGNORE_NULLS,
1098+
),
1099+
[None, 7, 7, 7, None, 9, 9],
1100+
),
10801101
(
10811102
"first_value",
10821103
f.first_value(column("a")).over(

0 commit comments

Comments
 (0)