Skip to content

Commit 32b1bbf

Browse files
Update test_aggregation.py
1 parent 9840d67 commit 32b1bbf

1 file changed

Lines changed: 61 additions & 0 deletions

File tree

‎python/tests/test_aggregation.py‎

Lines changed: 61 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -102,6 +102,16 @@ def df_aggregate_100():
102102
(f.var(column("a")), lambda a, b, c, d: np.array(np.var(a, ddof=1))),
103103
(f.var_pop(column("b")), lambda a, b, c, d: np.array(np.var(b, ddof=0))),
104104
(f.var_samp(column("c")), lambda a, b, c, d: np.array(np.var(c, ddof=1))),
105+
# Column names in place of column(...)
106+
(f.avg("a"), lambda a, b, c, d: np.array(np.average(a))),
107+
(f.corr("a", "b"), lambda a, b, c, d: np.array(np.corrcoef(a, b)[0][1])),
108+
(f.count("a"), lambda a, b, c, d: pa.array([len(a)])),
109+
(f.covar("a", "b"), lambda a, b, c, d: np.array(np.cov(a, b, ddof=1)[0][1])),
110+
(f.max("a"), lambda a, b, c, d: np.array(np.max(a))),
111+
(f.mean("b"), lambda a, b, c, d: np.array(np.mean(b))),
112+
(f.sum("b"), lambda a, b, c, d: np.array(np.sum(b.to_pylist()))),
113+
(f.stddev_samp("c"), lambda a, b, c, d: np.array(np.std(c, ddof=1))),
114+
(f.var("a"), lambda a, b, c, d: np.array(np.var(a, ddof=1))),
105115
],
106116
)
107117
def test_aggregation_stats(df, agg_expr, calc_expected):
@@ -496,3 +506,54 @@ def test_string_agg(name, expr, result) -> None:
496506
}
497507
df.show()
498508
assert df.collect()[0].to_pydict() == expected
509+
510+
511+
@pytest.mark.parametrize(
512+
("by_name", "by_expr"),
513+
[
514+
(f.count(["e"]), f.count([column("e")])),
515+
(f.regr_slope("c", "a"), f.regr_slope(column("c"), column("a"))),
516+
(
517+
f.approx_percentile_cont("b", 0.5),
518+
f.approx_percentile_cont(column("b"), 0.5),
519+
),
520+
(
521+
f.approx_percentile_cont_with_weight("b", "a", 0.5),
522+
f.approx_percentile_cont_with_weight(column("b"), column("a"), 0.5),
523+
),
524+
(f.percentile_cont("c", 0.5), f.percentile_cont(column("c"), 0.5)),
525+
(
526+
f.first_value("a", order_by="c"),
527+
f.first_value(column("a"), order_by="c"),
528+
),
529+
(f.array_agg("a", order_by="a"), f.array_agg(column("a"), order_by="a")),
530+
(f.bool_and("d"), f.bool_and(column("d"))),
531+
],
532+
)
533+
def test_aggregate_accepts_column_name(df, by_name, by_expr) -> None:
534+
result = df.aggregate([], [by_name.alias("v")]).collect()[0]
535+
expected = df.aggregate([], [by_expr.alias("v")]).collect()[0]
536+
assert result == expected
537+
538+
539+
def test_string_agg_accepts_column_name() -> None:
540+
ctx = SessionContext()
541+
df = ctx.from_pydict({"a": ["one", "two", "three"], "b": [2, 0, 1]})
542+
543+
result = df.aggregate([], [f.string_agg("a", ",", order_by="b").alias("v")])
544+
545+
assert result.collect()[0].to_pydict() == {"v": ["two,three,one"]}
546+
547+
548+
@pytest.mark.parametrize(
549+
"build",
550+
[
551+
lambda: f.sum(1),
552+
lambda: f.corr("a", 1),
553+
lambda: f.count(["a", 1]),
554+
lambda: f.approx_percentile_cont_with_weight("a", 1, 0.5),
555+
],
556+
)
557+
def test_aggregate_rejects_non_column_input(build) -> None:
558+
with pytest.raises(TypeError, match="Expected Expr or column name"):
559+
build()

0 commit comments

Comments
 (0)