Skip to content

Commit 97e6031

Browse files
timsaucerclaude
andcommitted
feat(spark): add pyspark aliases and optional-length substr
- Add getbit, dateadd, datediff, datepart, sha, ceiling, printf, char_length, and character_length as aliases of their Spark primaries. - Add substr, whose len argument is optional as in pyspark. It calls the UDF directly because upstream expr_fn::substring always takes a length. - Rename the spark.last_day parameter from col to date to match pyspark, with an upgrade-guide note for keyword callers. Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
1 parent 8027fcb commit 97e6031

4 files changed

Lines changed: 175 additions & 2 deletions

File tree

‎crates/core/src/spark_functions.rs‎

Lines changed: 13 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -254,6 +254,18 @@ spark_udf_vec!(format_string, udf::string::format_string);
254254
spark_expr_fn!(quote, arg1);
255255
spark_expr_fn!(space, arg1);
256256
spark_expr_fn!(substring, str pos length);
257+
/// `substr(str, pos, len=None)`. Upstream `expr_fn::substring` always takes a
258+
/// length, so call the UDF directly to allow the two-argument form.
259+
#[pyfunction]
260+
#[pyo3(signature = (str, pos, len=None))]
261+
fn substr(str: PyExpr, pos: PyExpr, len: Option<PyExpr>) -> PyExpr {
262+
let args: Vec<Expr> = [Some(str), Some(pos), len]
263+
.into_iter()
264+
.flatten()
265+
.map(Into::into)
266+
.collect();
267+
Expr::ScalarFunction(ScalarFunction::new_udf(udf::string::substring(), args)).into()
268+
}
257269
spark_expr_fn!(unbase64, str);
258270
spark_expr_fn!(soundex, str);
259271
spark_expr_fn!(is_valid_utf8, str);
@@ -372,6 +384,7 @@ pub(crate) fn init_module(m: &Bound<'_, PyModule>) -> PyResult<()> {
372384
m.add_wrapped(wrap_pyfunction!(format_string))?;
373385
m.add_wrapped(wrap_pyfunction!(space))?;
374386
m.add_wrapped(wrap_pyfunction!(substring))?;
387+
m.add_wrapped(wrap_pyfunction!(substr))?;
375388
m.add_wrapped(wrap_pyfunction!(unbase64))?;
376389
m.add_wrapped(wrap_pyfunction!(soundex))?;
377390
m.add_wrapped(wrap_pyfunction!(is_valid_utf8))?;

‎docs/source/user-guide/upgrade-guides.md‎

Lines changed: 11 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -218,6 +218,17 @@ f.bit_and(column("a"), filter=my_filter) # after
218218
Passing `filter` to `mean` previously raised a `TypeError`, whether passed
219219
positionally or by keyword; it now works.
220220

221+
### `spark.last_day` renamed its parameter
222+
223+
The parameter of {py:func}`datafusion.functions.spark.last_day` is now named
224+
`date`, matching `pyspark.sql.functions.last_day`. Positional calls are
225+
unaffected; update any call passing it by keyword.
226+
227+
```python
228+
spark.last_day(col=d) # before
229+
spark.last_day(date=d) # after
230+
```
231+
221232
### Changes to the `datafusion-python-util` crate
222233

223234
Extension libraries written in Rust usually depend on the

‎python/datafusion/functions/spark.py‎

Lines changed: 115 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -356,6 +356,15 @@ def bit_get(col: Expr, pos: Expr | str) -> Expr:
356356
return Expr(_f.bit_get(col.expr, _to_raw_expr(pos)))
357357

358358

359+
def getbit(col: Expr, pos: Expr | str) -> Expr:
360+
"""Spark ``getbit``: returns the bit (0 or 1) at ``pos``.
361+
362+
See Also:
363+
This is an alias for :py:func:`bit_get`.
364+
"""
365+
return bit_get(col, pos)
366+
367+
359368
def bit_count(col: Expr) -> Expr:
360369
"""Spark ``bit_count``: number of bits set in the integer's binary form.
361370
@@ -539,6 +548,15 @@ def date_add(start: Expr, days: Expr | int) -> Expr:
539548
return Expr(_f.date_add(start.expr, _coerce_i32(days).expr))
540549

541550

551+
def dateadd(start: Expr, days: Expr | int) -> Expr:
552+
"""Spark ``dateadd``: date + N days.
553+
554+
See Also:
555+
This is an alias for :py:func:`date_add`.
556+
"""
557+
return date_add(start, days)
558+
559+
542560
def date_sub(start: Expr, days: Expr | int) -> Expr:
543561
"""Spark ``date_sub``: date - N days.
544562
@@ -627,7 +645,7 @@ def second(col: Expr) -> Expr:
627645
return Expr(_f.second(col.expr))
628646

629647

630-
def last_day(col: Expr) -> Expr:
648+
def last_day(date: Expr) -> Expr:
631649
"""Spark ``last_day``: last day of the month containing the date.
632650
633651
Examples:
@@ -640,7 +658,7 @@ def last_day(col: Expr) -> Expr:
640658
>>> r.collect_column("v")[0].as_py()
641659
datetime.date(2020, 1, 31)
642660
"""
643-
return Expr(_f.last_day(col.expr))
661+
return Expr(_f.last_day(date.expr))
644662

645663

646664
def make_dt_interval(
@@ -754,6 +772,15 @@ def date_diff(end: Expr, start: Expr) -> Expr:
754772
return Expr(_f.date_diff(end.expr, start.expr))
755773

756774

775+
def datediff(end: Expr, start: Expr) -> Expr:
776+
"""Spark ``datediff``: number of days from ``start`` to ``end``.
777+
778+
See Also:
779+
This is an alias for :py:func:`date_diff`.
780+
"""
781+
return date_diff(end, start)
782+
783+
757784
def date_trunc(format: Expr | str, timestamp: Expr) -> Expr:
758785
"""Spark ``date_trunc``: truncate timestamp to unit ``fmt``.
759786
@@ -832,6 +859,15 @@ def date_part(field: Expr | str, source: Expr) -> Expr:
832859
return Expr(_f.date_part(coerce_to_expr(field).expr, source.expr))
833860

834861

862+
def datepart(field: Expr | str, source: Expr) -> Expr:
863+
"""Spark ``datepart``: extract ``field`` from a date/time/timestamp.
864+
865+
See Also:
866+
This is an alias for :py:func:`date_part`.
867+
"""
868+
return date_part(field, source)
869+
870+
835871
def from_utc_timestamp(timestamp: Expr, tz: Expr | str) -> Expr:
836872
"""Spark ``from_utc_timestamp``: interpret ``ts`` as UTC, convert to ``tz``.
837873
@@ -991,6 +1027,15 @@ def sha1(col: Expr) -> Expr:
9911027
return Expr(_f.sha1(col.expr))
9921028

9931029

1030+
def sha(col: Expr) -> Expr:
1031+
"""Spark ``sha``: SHA-1 hash as a hex string.
1032+
1033+
See Also:
1034+
This is an alias for :py:func:`sha1`.
1035+
"""
1036+
return sha1(col)
1037+
1038+
9941039
def sha2(col: Expr, numBits: Expr | int) -> Expr: # noqa: N803
9951040
"""Spark ``sha2``: SHA-2 family hash (224, 256, 384, 512). Bit length 0 = 256.
9961041
@@ -1175,6 +1220,15 @@ def ceil(col: Expr) -> Expr:
11751220
return Expr(_f.ceil(col.expr))
11761221

11771222

1223+
def ceiling(col: Expr) -> Expr:
1224+
"""Spark ``ceiling``: smallest integer ≥ arg.
1225+
1226+
See Also:
1227+
This is an alias for :py:func:`ceil`.
1228+
"""
1229+
return ceil(col)
1230+
1231+
11781232
def expm1(col: Expr) -> Expr:
11791233
"""Spark ``expm1``: exp(arg) - 1.
11801234
@@ -1562,6 +1616,24 @@ def length(col: Expr) -> Expr:
15621616
return Expr(_f.length(col.expr))
15631617

15641618

1619+
def character_length(col: Expr) -> Expr:
1620+
"""Spark ``character_length``: character length of a string, or bytes of binary.
1621+
1622+
See Also:
1623+
This is an alias for :py:func:`length`.
1624+
"""
1625+
return length(col)
1626+
1627+
1628+
def char_length(col: Expr) -> Expr:
1629+
"""Spark ``char_length``: character length of a string, or bytes of binary.
1630+
1631+
See Also:
1632+
This is an alias for :py:func:`length`.
1633+
"""
1634+
return length(col)
1635+
1636+
15651637
def like(
15661638
str: Expr,
15671639
pattern: Expr | str,
@@ -1639,6 +1711,15 @@ def format_string(format: str | Expr, *cols: Expr) -> Expr:
16391711
return Expr(_f.format_string(fmt_expr.expr, *[c.expr for c in cols]))
16401712

16411713

1714+
def printf(format: str | Expr, *cols: Expr) -> Expr:
1715+
"""Spark ``printf``: printf-style format string.
1716+
1717+
See Also:
1718+
This is an alias for :py:func:`format_string`.
1719+
"""
1720+
return format_string(format, *cols)
1721+
1722+
16421723
def space(col: Expr | int) -> Expr:
16431724
"""Spark ``space``: string of n spaces.
16441725
@@ -1673,6 +1754,28 @@ def substring(str: Expr, pos: Expr | int, len: Expr | int) -> Expr:
16731754
)
16741755

16751756

1757+
def substr(str: Expr, pos: Expr | int, len: Expr | int | None = None) -> Expr:
1758+
"""Spark ``substr``: 1-indexed substring, to the end when ``len`` is omitted.
1759+
1760+
Same as :py:func:`substring` except that ``len`` is optional. ``pos`` and
1761+
``len`` accept native ``int`` values or :class:`Expr`.
1762+
1763+
Examples:
1764+
>>> ctx = dfn.SessionContext()
1765+
>>> df = ctx.from_pydict({"x": [1]})
1766+
>>> r = df.select(dfn.functions.spark.substr(dfn.lit("hello"), 2).alias("v"))
1767+
>>> r.collect_column("v")[0].as_py()
1768+
'ello'
1769+
1770+
>>> r = df.select(
1771+
... dfn.functions.spark.substr(dfn.lit("hello"), 2, len=3).alias("v"))
1772+
>>> r.collect_column("v")[0].as_py()
1773+
'ell'
1774+
"""
1775+
len_raw = coerce_to_expr(len).expr if len is not None else None
1776+
return Expr(_f.substr(str.expr, coerce_to_expr(pos).expr, len_raw))
1777+
1778+
16761779
def unbase64(col: Expr) -> Expr:
16771780
"""Spark ``unbase64``: decode a base64 string to binary.
16781781
@@ -1867,7 +1970,10 @@ def url_encode(str: Expr) -> Expr:
18671970
"bitmap_count",
18681971
"bitwise_not",
18691972
"ceil",
1973+
"ceiling",
18701974
"char",
1975+
"char_length",
1976+
"character_length",
18711977
"collect_list",
18721978
"collect_set",
18731979
"concat",
@@ -1880,12 +1986,16 @@ def url_encode(str: Expr) -> Expr:
18801986
"date_part",
18811987
"date_sub",
18821988
"date_trunc",
1989+
"dateadd",
1990+
"datediff",
1991+
"datepart",
18831992
"elt",
18841993
"expm1",
18851994
"factorial",
18861995
"floor",
18871996
"format_string",
18881997
"from_utc_timestamp",
1998+
"getbit",
18891999
"hex",
18902000
"hour",
18912001
"hypot",
@@ -1914,11 +2024,13 @@ def url_encode(str: Expr) -> Expr:
19142024
"pmod",
19152025
"pow",
19162026
"power",
2027+
"printf",
19172028
"quote",
19182029
"rint",
19192030
"round",
19202031
"sec",
19212032
"second",
2033+
"sha",
19222034
"sha1",
19232035
"sha2",
19242036
"shiftleft",
@@ -1932,6 +2044,7 @@ def url_encode(str: Expr) -> Expr:
19322044
"space",
19332045
"spark_cast",
19342046
"str_to_map",
2047+
"substr",
19352048
"substring",
19362049
"time_trunc",
19372050
"to_utc_timestamp",

‎python/tests/test_spark_functions.py‎

Lines changed: 36 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -500,3 +500,39 @@ def test_sql_concat_semantics_override():
500500
ctx2.sql("SELECT concat('a', NULL, 'b') AS c").collect_column("c")[0].as_py()
501501
)
502502
assert spark_out is None
503+
504+
505+
@pytest.mark.parametrize(
506+
("alias_fn", "primary_fn", "args"),
507+
[
508+
(spark.getbit, spark.bit_get, lambda: (lit(5), lit(0))),
509+
(spark.dateadd, spark.date_add, lambda: (_ts().cast(pa.date32()), 3)),
510+
(
511+
spark.datediff,
512+
spark.date_diff,
513+
lambda: (_ts().cast(pa.date32()), lit("2020-01-01").cast(pa.date32())),
514+
),
515+
(spark.datepart, spark.date_part, lambda: ("YEAR", _ts())),
516+
(spark.sha, spark.sha1, lambda: (lit("abc"),)),
517+
(spark.ceiling, spark.ceil, lambda: (lit(1.2),)),
518+
(spark.printf, spark.format_string, lambda: ("%d-%s", lit(42), lit("hi"))),
519+
(spark.char_length, spark.length, lambda: (lit("hello"),)),
520+
(spark.character_length, spark.length, lambda: (lit("hello"),)),
521+
(spark.power, spark.pow, lambda: (lit(2), lit(3))),
522+
(spark.substr, spark.substring, lambda: (lit("hello"), 2, 3)),
523+
],
524+
)
525+
def test_aliases_match_primary(df, alias_fn, primary_fn, args):
526+
assert _val(df, alias_fn(*args())) == _val(df, primary_fn(*args()))
527+
528+
529+
def test_substr_without_len(df):
530+
assert _val(df, spark.substr(lit("hello"), 2)) == "ello"
531+
assert _val(df, spark.substr(lit("hello"), -3)) == "llo"
532+
533+
534+
def test_last_day_date_keyword(df):
535+
import datetime as dt
536+
537+
d = lit(pa.scalar(dt.date(2024, 2, 10), type=pa.date32()))
538+
assert _val(df, spark.last_day(date=d)) == dt.date(2024, 2, 29)

0 commit comments

Comments
 (0)