From 3e6cdb6e06440642795a3e22fec7388cc75f0051 Mon Sep 17 00:00:00 2001 From: Spenser Sun Date: Sun, 2 Aug 2026 08:06:02 +0000 Subject: [PATCH] [SPARK-58500][PYTHON] Accept tuple in DataFrame.describe and selectExpr Widen the ``describe`` and ``selectExpr`` type annotations from ``Union[str, List[str]]`` to ``Union[str, Sequence[str]]`` (overloads + impl) across the base, classic, and connect layers, and change the runtime unwrap check from ``isinstance(x[0], list)`` to a Sequence check so a single tuple of column names / expressions is unpacked like a list. Normalizes the overload shape (describe had none; connect selectExpr had none) and drops the ``# type: ignore[assignment]`` on the unwrap. --- python/pyspark/sql/classic/dataframe.py | 20 ++++++++++------ python/pyspark/sql/connect/dataframe.py | 24 ++++++++++++++----- python/pyspark/sql/dataframe.py | 12 +++++++--- .../sql/tests/connect/test_connect_basic.py | 9 +++++++ .../sql/tests/connect/test_connect_stat.py | 4 ++++ 5 files changed, 53 insertions(+), 16 deletions(-) diff --git a/python/pyspark/sql/classic/dataframe.py b/python/pyspark/sql/classic/dataframe.py index eeed6fef44aa4..f02926dc3408e 100644 --- a/python/pyspark/sql/classic/dataframe.py +++ b/python/pyspark/sql/classic/dataframe.py @@ -1033,9 +1033,15 @@ def _jcols_ordinal( _cols.append(c) # type: ignore[arg-type] return self._jseq(_cols, _to_java_column) - def describe(self, *cols: Union[str, List[str]]) -> ParentDataFrame: - if len(cols) == 1 and isinstance(cols[0], list): - cols = cols[0] # type: ignore[assignment] + @overload + def describe(self, *cols: str) -> ParentDataFrame: ... + + @overload + def describe(self, __cols: Sequence[str]) -> ParentDataFrame: ... + + def describe(self, *cols: Union[str, Sequence[str]]) -> ParentDataFrame: + if len(cols) == 1 and not isinstance(cols[0], str) and isinstance(cols[0], Sequence): + cols = tuple(cols[0]) jdf = self._jdf.describe(self._jseq(cols)) return DataFrame(jdf, self.sparkSession) @@ -1116,11 +1122,11 @@ def select(self, *cols: Union[Sequence["ColumnOrName"], "ColumnOrName"]) -> Pare def selectExpr(self, *expr: str) -> ParentDataFrame: ... @overload - def selectExpr(self, *expr: List[str]) -> ParentDataFrame: ... + def selectExpr(self, __expr: Sequence[str]) -> ParentDataFrame: ... - def selectExpr(self, *expr: Union[str, List[str]]) -> ParentDataFrame: - if len(expr) == 1 and isinstance(expr[0], list): - expr = expr[0] # type: ignore[assignment] + def selectExpr(self, *expr: Union[str, Sequence[str]]) -> ParentDataFrame: + if len(expr) == 1 and not isinstance(expr[0], str) and isinstance(expr[0], Sequence): + expr = tuple(expr[0]) jdf = self._jdf.selectExpr(self._jseq(expr)) return DataFrame(jdf, self.sparkSession) diff --git a/python/pyspark/sql/connect/dataframe.py b/python/pyspark/sql/connect/dataframe.py index 093489757115b..3e7c8ff2e6af2 100644 --- a/python/pyspark/sql/connect/dataframe.py +++ b/python/pyspark/sql/connect/dataframe.py @@ -300,10 +300,16 @@ def select(self, *cols: Union[Sequence["ColumnOrName"], "ColumnOrName"]) -> Pare session=self._session, ) - def selectExpr(self, *expr: Union[str, List[str]]) -> ParentDataFrame: + @overload + def selectExpr(self, *expr: str) -> ParentDataFrame: ... + + @overload + def selectExpr(self, __expr: Sequence[str]) -> ParentDataFrame: ... + + def selectExpr(self, *expr: Union[str, Sequence[str]]) -> ParentDataFrame: sql_expr = [] - if len(expr) == 1 and isinstance(expr[0], list): - expr = expr[0] # type: ignore[assignment] + if len(expr) == 1 and not isinstance(expr[0], str) and isinstance(expr[0], Sequence): + expr = tuple(expr[0]) for element in expr: if isinstance(element, str): sql_expr.append(F.expr(element)) @@ -1601,9 +1607,15 @@ def summary(self, *statistics: str) -> ParentDataFrame: session=self._session, ) - def describe(self, *cols: Union[str, List[str]]) -> ParentDataFrame: - if len(cols) == 1 and isinstance(cols[0], list): - cols = cols[0] # type: ignore[assignment] + @overload + def describe(self, *cols: str) -> ParentDataFrame: ... + + @overload + def describe(self, __cols: Sequence[str]) -> ParentDataFrame: ... + + def describe(self, *cols: Union[str, Sequence[str]]) -> ParentDataFrame: + if len(cols) == 1 and not isinstance(cols[0], str) and isinstance(cols[0], Sequence): + cols = tuple(cols[0]) _cols = [] for column in cols: diff --git a/python/pyspark/sql/dataframe.py b/python/pyspark/sql/dataframe.py index 6c4d32ea1797f..348dbe0c02c9d 100644 --- a/python/pyspark/sql/dataframe.py +++ b/python/pyspark/sql/dataframe.py @@ -3360,8 +3360,14 @@ def _get_col( orderBy = sort + @overload + def describe(self, *cols: str) -> "DataFrame": ... + + @overload + def describe(self, __cols: Sequence[str]) -> "DataFrame": ... + @dispatch_df_method - def describe(self, *cols: Union[str, List[str]]) -> "DataFrame": + def describe(self, *cols: Union[str, Sequence[str]]) -> "DataFrame": """Computes basic statistics for numeric and string columns. .. versionadded:: 1.3.1 @@ -3777,10 +3783,10 @@ def select(self, *cols: Union[Sequence["ColumnOrName"], "ColumnOrName"]) -> "Dat def selectExpr(self, *expr: str) -> "DataFrame": ... @overload - def selectExpr(self, *expr: List[str]) -> "DataFrame": ... + def selectExpr(self, __expr: Sequence[str]) -> "DataFrame": ... @dispatch_df_method - def selectExpr(self, *expr: Union[str, List[str]]) -> "DataFrame": + def selectExpr(self, *expr: Union[str, Sequence[str]]) -> "DataFrame": """Projects a set of SQL expressions and returns a new :class:`DataFrame`. This is a variant of :func:`select` that accepts SQL expressions. diff --git a/python/pyspark/sql/tests/connect/test_connect_basic.py b/python/pyspark/sql/tests/connect/test_connect_basic.py index b8d258267d267..9f5571d0eac1e 100755 --- a/python/pyspark/sql/tests/connect/test_connect_basic.py +++ b/python/pyspark/sql/tests/connect/test_connect_basic.py @@ -842,6 +842,15 @@ def test_select_expr(self): .toPandas(), ) + self.assert_eq( + self.connect.read.table(self.tbl_name) + .selectExpr(("id * 2", "cast(name as long) as name")) + .toPandas(), + self.spark.read.table(self.tbl_name) + .selectExpr(("id * 2", "cast(name as long) as name")) + .toPandas(), + ) + def test_select_star(self): data = [Row(a=1, b=Row(c=2, d=Row(e=3)))] diff --git a/python/pyspark/sql/tests/connect/test_connect_stat.py b/python/pyspark/sql/tests/connect/test_connect_stat.py index 3d05cb7b4bb5c..f971e1363fa0d 100644 --- a/python/pyspark/sql/tests/connect/test_connect_stat.py +++ b/python/pyspark/sql/tests/connect/test_connect_stat.py @@ -169,6 +169,10 @@ def test_describe(self): self.connect.read.table(self.tbl_name).describe(["id", "name"]).toPandas(), self.spark.read.table(self.tbl_name).describe(["id", "name"]).toPandas(), ) + self.assert_eq( + self.connect.read.table(self.tbl_name).describe(("id", "name")).toPandas(), + self.spark.read.table(self.tbl_name).describe(("id", "name")).toPandas(), + ) def test_stat_cov(self): # SPARK-41067: Test the stat.cov method