Skip to content

Commit e8c2192

Browse files
timsaucerclaude
andcommitted
feat: accept optional arguments upstream supports on trim, array_to_string, substr
- btrim, ltrim, rtrim, and trim take an optional characters argument naming the set to strip. - array_to_string and its aliases take an optional null_string that is written in place of NULL elements. - substr takes an optional length. Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
1 parent 97e6031 commit e8c2192

3 files changed

Lines changed: 166 additions & 25 deletions

File tree

‎crates/core/src/functions.rs‎

Lines changed: 14 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -117,6 +117,20 @@ fn string_to_array(string: PyExpr, delimiter: PyExpr, null_string: Option<PyExpr
117117
.into()
118118
}
119119

120+
#[pyfunction]
121+
#[pyo3(signature = (array, delimiter, null_string=None))]
122+
fn array_to_string(array: PyExpr, delimiter: PyExpr, null_string: Option<PyExpr>) -> PyExpr {
123+
let mut args = vec![array.into(), delimiter.into()];
124+
if let Some(null_string) = null_string {
125+
args.push(null_string.into());
126+
}
127+
Expr::ScalarFunction(datafusion::logical_expr::expr::ScalarFunction::new_udf(
128+
datafusion::functions_nested::string::array_to_string_udf(),
129+
args,
130+
))
131+
.into()
132+
}
133+
120134
#[pyfunction]
121135
#[pyo3(signature = (start, stop, step=None))]
122136
fn gen_series(start: PyExpr, stop: PyExpr, step: Option<PyExpr>) -> PyExpr {
@@ -646,7 +660,6 @@ fn version() -> PyExpr {
646660

647661
// Array Functions
648662
array_fn!(array_append, array element);
649-
array_fn!(array_to_string, array delimiter);
650663
array_fn!(array_dims, array);
651664
array_fn!(array_distinct, array);
652665
array_fn!(array_element, array element);

‎python/datafusion/functions/__init__.py‎

Lines changed: 109 additions & 24 deletions
Original file line numberDiff line numberDiff line change
@@ -486,46 +486,71 @@ def decode(expr: Expr, encoding: Expr | str) -> Expr:
486486
return Expr(f.decode(expr.expr, encoding.expr))
487487

488488

489-
def array_to_string(expr: Expr, delimiter: Expr | str) -> Expr:
489+
def array_to_string(
490+
expr: Expr, delimiter: Expr | str, null_string: Expr | str | None = None
491+
) -> Expr:
490492
"""Converts each element to its text representation.
491493
494+
NULL elements are omitted unless ``null_string`` is given, in which case it
495+
is written in their place.
496+
492497
Examples:
493498
>>> ctx = dfn.SessionContext()
494-
>>> df = ctx.from_pydict({"a": [[1, 2, 3]]})
499+
>>> df = ctx.from_pydict({"a": [[1, None, 3]]})
495500
>>> result = df.select(
496501
... dfn.functions.array_to_string(dfn.col("a"), ",").alias("s"))
497502
>>> result.collect_column("s")[0].as_py()
498-
'1,2,3'
503+
'1,3'
504+
505+
>>> result = df.select(
506+
... dfn.functions.array_to_string(
507+
... dfn.col("a"), ",", null_string="*"
508+
... ).alias("s"))
509+
>>> result.collect_column("s")[0].as_py()
510+
'1,*,3'
499511
"""
500512
delimiter = coerce_to_expr(delimiter)
501-
return Expr(f.array_to_string(expr.expr, delimiter.expr.cast(pa.string())))
513+
null_string = coerce_to_expr_or_none(null_string)
514+
return Expr(
515+
f.array_to_string(
516+
expr.expr,
517+
delimiter.expr.cast(pa.string()),
518+
null_string.expr if null_string is not None else None,
519+
)
520+
)
502521

503522

504-
def array_join(expr: Expr, delimiter: Expr | str) -> Expr:
523+
def array_join(
524+
expr: Expr, delimiter: Expr | str, null_string: Expr | str | None = None
525+
) -> Expr:
505526
"""Converts each element to its text representation.
506527
507528
See Also:
508529
This is an alias for :py:func:`array_to_string`.
509530
"""
510-
return array_to_string(expr, delimiter)
531+
return array_to_string(expr, delimiter, null_string=null_string)
511532

512533

513-
def list_to_string(expr: Expr, delimiter: Expr | str) -> Expr:
534+
def list_to_string(
535+
expr: Expr, delimiter: Expr | str, null_string: Expr | str | None = None
536+
) -> Expr:
514537
"""Converts each element to its text representation.
515538
516539
See Also:
517540
This is an alias for :py:func:`array_to_string`.
518541
"""
519-
return array_to_string(expr, delimiter)
542+
return array_to_string(expr, delimiter, null_string=null_string)
520543

521544

522-
def list_join(expr: Expr, delimiter: Expr | str) -> Expr:
545+
def list_join(
546+
expr: Expr, delimiter: Expr | str, null_string: Expr | str | None = None
547+
) -> Expr:
523548
"""Converts each element to its text representation.
524549
525550
See Also:
526551
This is an alias for :py:func:`array_to_string`.
527552
"""
528-
return array_to_string(expr, delimiter)
553+
return array_to_string(expr, delimiter, null_string=null_string)
529554

530555

531556
def lambda_var(name: str) -> Expr:
@@ -1128,17 +1153,29 @@ def bit_length(arg: Expr) -> Expr:
11281153
return Expr(f.bit_length(arg.expr))
11291154

11301155

1131-
def btrim(arg: Expr) -> Expr:
1132-
"""Removes all characters, spaces by default, from both sides of a string.
1156+
def btrim(arg: Expr, characters: Expr | str | None = None) -> Expr:
1157+
"""Removes ``characters``, spaces by default, from both sides of a string.
11331158
11341159
Examples:
11351160
>>> ctx = dfn.SessionContext()
11361161
>>> df = ctx.from_pydict({"a": [" a "]})
11371162
>>> trim_df = df.select(dfn.functions.btrim(dfn.col("a")).alias("trimmed"))
11381163
>>> trim_df.collect_column("trimmed")[0].as_py()
11391164
'a'
1165+
1166+
Trim a different set of characters:
1167+
1168+
>>> df = ctx.from_pydict({"a": ["xxaxx"]})
1169+
>>> trim_df = df.select(
1170+
... dfn.functions.btrim(dfn.col("a"), characters="x").alias("trimmed")
1171+
... )
1172+
>>> trim_df.collect_column("trimmed")[0].as_py()
1173+
'a'
11401174
"""
1141-
return Expr(f.btrim(arg.expr))
1175+
args = [arg.expr]
1176+
if characters is not None:
1177+
args.append(coerce_to_expr(characters).expr)
1178+
return Expr(f.btrim(*args))
11421179

11431180

11441181
def cbrt(arg: Expr) -> Expr:
@@ -1616,17 +1653,29 @@ def lpad(string: Expr, count: Expr | int, characters: Expr | str | None = None)
16161653
return Expr(f.lpad(string.expr, count.expr, characters.expr))
16171654

16181655

1619-
def ltrim(arg: Expr) -> Expr:
1620-
"""Removes all characters, spaces by default, from the beginning of a string.
1656+
def ltrim(arg: Expr, characters: Expr | str | None = None) -> Expr:
1657+
"""Removes ``characters``, spaces by default, from the beginning of a string.
16211658
16221659
Examples:
16231660
>>> ctx = dfn.SessionContext()
16241661
>>> df = ctx.from_pydict({"a": [" a "]})
16251662
>>> trim_df = df.select(dfn.functions.ltrim(dfn.col("a")).alias("trimmed"))
16261663
>>> trim_df.collect_column("trimmed")[0].as_py()
16271664
'a '
1665+
1666+
Trim a different set of characters:
1667+
1668+
>>> df = ctx.from_pydict({"a": ["xxaxx"]})
1669+
>>> trim_df = df.select(
1670+
... dfn.functions.ltrim(dfn.col("a"), characters="x").alias("trimmed")
1671+
... )
1672+
>>> trim_df.collect_column("trimmed")[0].as_py()
1673+
'axx'
16281674
"""
1629-
return Expr(f.ltrim(arg.expr))
1675+
args = [arg.expr]
1676+
if characters is not None:
1677+
args.append(coerce_to_expr(characters).expr)
1678+
return Expr(f.ltrim(*args))
16301679

16311680

16321681
def md5(arg: Expr) -> Expr:
@@ -2138,17 +2187,29 @@ def rpad(string: Expr, count: Expr | int, characters: Expr | str | None = None)
21382187
return Expr(f.rpad(string.expr, count.expr, characters.expr))
21392188

21402189

2141-
def rtrim(arg: Expr) -> Expr:
2142-
"""Removes all characters, spaces by default, from the end of a string.
2190+
def rtrim(arg: Expr, characters: Expr | str | None = None) -> Expr:
2191+
"""Removes ``characters``, spaces by default, from the end of a string.
21432192
21442193
Examples:
21452194
>>> ctx = dfn.SessionContext()
21462195
>>> df = ctx.from_pydict({"a": [" a "]})
21472196
>>> trim_df = df.select(dfn.functions.rtrim(dfn.col("a")).alias("trimmed"))
21482197
>>> trim_df.collect_column("trimmed")[0].as_py()
21492198
' a'
2199+
2200+
Trim a different set of characters:
2201+
2202+
>>> df = ctx.from_pydict({"a": ["xxaxx"]})
2203+
>>> trim_df = df.select(
2204+
... dfn.functions.rtrim(dfn.col("a"), characters="x").alias("trimmed")
2205+
... )
2206+
>>> trim_df.collect_column("trimmed")[0].as_py()
2207+
'xxa'
21502208
"""
2151-
return Expr(f.rtrim(arg.expr))
2209+
args = [arg.expr]
2210+
if characters is not None:
2211+
args.append(coerce_to_expr(characters).expr)
2212+
return Expr(f.rtrim(*args))
21522213

21532214

21542215
def sha224(arg: Expr) -> Expr:
@@ -2312,8 +2373,10 @@ def strpos(string: Expr, substring: Expr | str) -> Expr:
23122373
return Expr(f.strpos(string.expr, substring.expr))
23132374

23142375

2315-
def substr(string: Expr, position: Expr | int) -> Expr:
2316-
"""Substring from the ``position`` to the end.
2376+
def substr(
2377+
string: Expr, position: Expr | int, length: Expr | int | None = None
2378+
) -> Expr:
2379+
"""Substring from the ``position``, to the end or for ``length`` characters.
23172380
23182381
Examples:
23192382
>>> ctx = dfn.SessionContext()
@@ -2322,7 +2385,17 @@ def substr(string: Expr, position: Expr | int) -> Expr:
23222385
... dfn.functions.substr(dfn.col("a"), 3).alias("s"))
23232386
>>> result.collect_column("s")[0].as_py()
23242387
'llo'
2388+
2389+
>>> result = df.select(
2390+
... dfn.functions.substr(dfn.col("a"), 2, length=3).alias("s"))
2391+
>>> result.collect_column("s")[0].as_py()
2392+
'ell'
2393+
2394+
See Also:
2395+
:py:func:`substring`.
23252396
"""
2397+
if length is not None:
2398+
return substring(string, position, length)
23262399
position = coerce_to_expr(position)
23272400
return Expr(f.substr(string.expr, position.expr))
23282401

@@ -2958,17 +3031,29 @@ def translate(string: Expr, from_val: Expr | str, to_val: Expr | str) -> Expr:
29583031
return Expr(f.translate(string.expr, from_val.expr, to_val.expr))
29593032

29603033

2961-
def trim(arg: Expr) -> Expr:
2962-
"""Removes all characters, spaces by default, from both sides of a string.
3034+
def trim(arg: Expr, characters: Expr | str | None = None) -> Expr:
3035+
"""Removes ``characters``, spaces by default, from both sides of a string.
29633036
29643037
Examples:
29653038
>>> ctx = dfn.SessionContext()
29663039
>>> df = ctx.from_pydict({"a": [" hello "]})
29673040
>>> result = df.select(dfn.functions.trim(dfn.col("a")).alias("t"))
29683041
>>> result.collect_column("t")[0].as_py()
29693042
'hello'
3043+
3044+
Trim a different set of characters:
3045+
3046+
>>> df = ctx.from_pydict({"a": ["xxhelloxx"]})
3047+
>>> result = df.select(
3048+
... dfn.functions.trim(dfn.col("a"), characters="x").alias("t")
3049+
... )
3050+
>>> result.collect_column("t")[0].as_py()
3051+
'hello'
29703052
"""
2971-
return Expr(f.trim(arg.expr))
3053+
args = [arg.expr]
3054+
if characters is not None:
3055+
args.append(coerce_to_expr(characters).expr)
3056+
return Expr(f.trim(*args))
29723057

29733058

29743059
def trunc(num: Expr, precision: Expr | int | None = None) -> Expr:

‎python/tests/test_functions.py‎

Lines changed: 43 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -2547,3 +2547,46 @@ def test_backward_compat_with_lit(self):
25472547
f.split_part(column("a"), literal(","), literal(2)).alias("s")
25482548
).collect()
25492549
assert result[0].column(0)[0].as_py() == "b"
2550+
2551+
2552+
@pytest.mark.parametrize(
2553+
("fn", "expected"),
2554+
[
2555+
(f.btrim, "hi"),
2556+
(f.trim, "hi"),
2557+
(f.ltrim, "hixyx"),
2558+
(f.rtrim, "xyxhi"),
2559+
],
2560+
)
2561+
def test_trim_characters(fn, expected):
2562+
ctx = SessionContext()
2563+
df = ctx.from_pydict({"a": ["xyxhixyx"]})
2564+
assert df.select(fn(column("a"), characters="xy").alias("r")).collect_column(
2565+
"r"
2566+
).to_pylist() == [expected]
2567+
assert df.select(
2568+
fn(column("a"), characters=literal("xy")).alias("r")
2569+
).collect_column("r").to_pylist() == [expected]
2570+
2571+
2572+
@pytest.mark.parametrize(
2573+
"fn", [f.array_to_string, f.array_join, f.list_to_string, f.list_join]
2574+
)
2575+
def test_array_to_string_null_string(fn):
2576+
ctx = SessionContext()
2577+
df = ctx.from_pydict({"a": [[1, None, 3]]})
2578+
without = df.select(fn(column("a"), "-").alias("r")).collect_column("r")
2579+
with_null = df.select(fn(column("a"), "-", null_string="NA").alias("r"))
2580+
assert without.to_pylist() == ["1-3"]
2581+
assert with_null.collect_column("r").to_pylist() == ["1-NA-3"]
2582+
2583+
2584+
def test_substr_length():
2585+
ctx = SessionContext()
2586+
df = ctx.from_pydict({"a": ["hello"]})
2587+
r = df.select(
2588+
f.substr(column("a"), 2).alias("tail"),
2589+
f.substr(column("a"), 2, length=3).alias("mid"),
2590+
f.substr(column("a"), 2, length=literal(3)).alias("mid_expr"),
2591+
).to_pydict()
2592+
assert r == {"tail": ["ello"], "mid": ["ell"], "mid_expr": ["ell"]}

0 commit comments

Comments
 (0)