@@ -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)
107117def 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