diff --git a/python/pyspark/sql/classic/dataframe.py b/python/pyspark/sql/classic/dataframe.py index eeed6fef44aa4..5f58cceb4393c 100644 --- a/python/pyspark/sql/classic/dataframe.py +++ b/python/pyspark/sql/classic/dataframe.py @@ -2053,6 +2053,7 @@ def replace( def replace( self, to_replace: Dict["LiteralType", "OptionalPrimitiveType"], + *, subset: Optional[List[str]] = ..., ) -> ParentDataFrame: ... @@ -2064,7 +2065,7 @@ def replace( subset: Optional[List[str]] = ..., ) -> ParentDataFrame: ... - def replace( # type: ignore[misc] + def replace( self, to_replace: Union[List["LiteralType"], Dict["LiteralType", "OptionalPrimitiveType"]], value: Optional[ diff --git a/python/pyspark/sql/connect/readwriter.py b/python/pyspark/sql/connect/readwriter.py index 5c2c0c80ccdfb..5d96bd781f616 100644 --- a/python/pyspark/sql/connect/readwriter.py +++ b/python/pyspark/sql/connect/readwriter.py @@ -15,7 +15,7 @@ # limitations under the License. # from typing import Dict -from typing import Optional, Union, List, overload, Tuple, cast, Callable +from typing import Optional, Sequence, Union, List, overload, cast, Callable from typing import TYPE_CHECKING from pyspark.sql.connect.plan import ( @@ -52,7 +52,6 @@ __all__ = ["DataFrameReader", "DataFrameWriter"] PathOrPaths = Union[str, List[str]] -TupleOrListOfString = Union[List[str], Tuple[str, ...]] class OptionUtils: @@ -620,11 +619,11 @@ def options(self, **options: "OptionalPrimitiveType") -> "DataFrameWriter": def partitionBy(self, *cols: str) -> "DataFrameWriter": ... @overload - def partitionBy(self, *cols: List[str]) -> "DataFrameWriter": ... + def partitionBy(self, __cols: Sequence[str]) -> "DataFrameWriter": ... - def partitionBy(self, *cols: Union[str, List[str]]) -> "DataFrameWriter": - if len(cols) == 1 and isinstance(cols[0], (list, tuple)): - cols = cols[0] # type: ignore[assignment] + def partitionBy(self, *cols: Union[str, Sequence[str]]) -> "DataFrameWriter": + if len(cols) == 1 and not isinstance(cols[0], str) and isinstance(cols[0], Sequence): + cols = tuple(cols[0]) self._write.partitioning_cols = cast(List[str], cols) return self @@ -635,10 +634,10 @@ def partitionBy(self, *cols: Union[str, List[str]]) -> "DataFrameWriter": def bucketBy(self, numBuckets: int, col: str, *cols: str) -> "DataFrameWriter": ... @overload - def bucketBy(self, numBuckets: int, col: TupleOrListOfString) -> "DataFrameWriter": ... + def bucketBy(self, numBuckets: int, col: Sequence[str]) -> "DataFrameWriter": ... def bucketBy( - self, numBuckets: int, col: Union[str, TupleOrListOfString], *cols: Optional[str] + self, numBuckets: int, col: Union[str, Sequence[str]], *cols: Optional[str] ) -> "DataFrameWriter": if not isinstance(numBuckets, int): raise PySparkTypeError( @@ -650,7 +649,7 @@ def bucketBy( }, ) - if isinstance(col, (list, tuple)): + if not isinstance(col, str) and isinstance(col, Sequence): if cols: raise PySparkValueError( errorClass="CANNOT_SET_TOGETHER", @@ -659,7 +658,7 @@ def bucketBy( }, ) - col, cols = col[0], col[1:] # type: ignore[assignment] + col, cols = col[0], tuple(col[1:]) for c in cols: if not isinstance(c, str): @@ -691,12 +690,10 @@ def bucketBy( def sortBy(self, col: str, *cols: str) -> "DataFrameWriter": ... @overload - def sortBy(self, col: TupleOrListOfString) -> "DataFrameWriter": ... + def sortBy(self, col: Sequence[str]) -> "DataFrameWriter": ... - def sortBy( - self, col: Union[str, TupleOrListOfString], *cols: Optional[str] - ) -> "DataFrameWriter": - if isinstance(col, (list, tuple)): + def sortBy(self, col: Union[str, Sequence[str]], *cols: Optional[str]) -> "DataFrameWriter": + if not isinstance(col, str) and isinstance(col, Sequence): if cols: raise PySparkValueError( errorClass="CANNOT_SET_TOGETHER", @@ -705,7 +702,7 @@ def sortBy( }, ) - col, cols = col[0], col[1:] # type: ignore[assignment] + col, cols = col[0], tuple(col[1:]) for c in cols: if not isinstance(c, str): @@ -736,11 +733,11 @@ def sortBy( def clusterBy(self, *cols: str) -> "DataFrameWriter": ... @overload - def clusterBy(self, *cols: List[str]) -> "DataFrameWriter": ... + def clusterBy(self, __cols: Sequence[str]) -> "DataFrameWriter": ... - def clusterBy(self, *cols: Union[str, List[str]]) -> "DataFrameWriter": - if len(cols) == 1 and isinstance(cols[0], (list, tuple)): - cols = cols[0] # type: ignore[assignment] + def clusterBy(self, *cols: Union[str, Sequence[str]]) -> "DataFrameWriter": + if len(cols) == 1 and not isinstance(cols[0], str) and isinstance(cols[0], Sequence): + cols = tuple(cols[0]) assert len(cols) > 0, "clusterBy needs one or more clustering columns." self._write.clustering_cols = cast(List[str], cols) return self @@ -752,7 +749,7 @@ def save( path: Optional[str] = None, format: Optional[str] = None, mode: Optional[str] = None, - partitionBy: Optional[Union[str, List[str]]] = None, + partitionBy: Optional[Union[str, Sequence[str]]] = None, **options: "OptionalPrimitiveType", ) -> None: self.mode(mode).options(**options) @@ -785,7 +782,7 @@ def saveAsTable( name: str, format: Optional[str] = None, mode: Optional[str] = None, - partitionBy: Optional[Union[str, List[str]]] = None, + partitionBy: Optional[Union[str, Sequence[str]]] = None, **options: "OptionalPrimitiveType", ) -> None: self.mode(mode).options(**options) @@ -830,7 +827,7 @@ def parquet( self, path: str, mode: Optional[str] = None, - partitionBy: Optional[Union[str, List[str]]] = None, + partitionBy: Optional[Union[str, Sequence[str]]] = None, compression: Optional[str] = None, ) -> None: self.mode(mode) @@ -933,7 +930,7 @@ def orc( self, path: str, mode: Optional[str] = None, - partitionBy: Optional[Union[str, List[str]]] = None, + partitionBy: Optional[Union[str, Sequence[str]]] = None, compression: Optional[str] = None, ) -> None: self.mode(mode) diff --git a/python/pyspark/sql/connect/streaming/readwriter.py b/python/pyspark/sql/connect/streaming/readwriter.py index 130844309ae4c..b307b2d93fb95 100644 --- a/python/pyspark/sql/connect/streaming/readwriter.py +++ b/python/pyspark/sql/connect/streaming/readwriter.py @@ -18,7 +18,7 @@ import re import sys import pickle -from typing import cast, overload, Callable, Dict, List, Optional, TYPE_CHECKING, Union +from typing import cast, overload, Callable, Dict, List, Optional, Sequence, TYPE_CHECKING, Union from pyspark.serializers import CloudPickleSerializer from pyspark.sql.connect.plan import ( @@ -508,11 +508,11 @@ def options(self, **options: "OptionalPrimitiveType") -> "DataStreamWriter": def partitionBy(self, *cols: str) -> "DataStreamWriter": ... @overload - def partitionBy(self, __cols: List[str]) -> "DataStreamWriter": ... + def partitionBy(self, __cols: Sequence[str]) -> "DataStreamWriter": ... - def partitionBy(self, *cols: str) -> "DataStreamWriter": # type: ignore[misc] - if len(cols) == 1 and isinstance(cols[0], (list, tuple)): - cols = cols[0] + def partitionBy(self, *cols: Union[str, Sequence[str]]) -> "DataStreamWriter": + if len(cols) == 1 and not isinstance(cols[0], str) and isinstance(cols[0], Sequence): + cols = tuple(cols[0]) # Clear any existing columns (if any). while len(self._write_proto.partitioning_column_names) > 0: self._write_proto.partitioning_column_names.pop() @@ -525,11 +525,11 @@ def partitionBy(self, *cols: str) -> "DataStreamWriter": # type: ignore[misc] def clusterBy(self, *cols: str) -> "DataStreamWriter": ... @overload - def clusterBy(self, __cols: List[str]) -> "DataStreamWriter": ... + def clusterBy(self, __cols: Sequence[str]) -> "DataStreamWriter": ... - def clusterBy(self, *cols: str) -> "DataStreamWriter": # type: ignore[misc] - if len(cols) == 1 and isinstance(cols[0], (list, tuple)): - cols = cols[0] + def clusterBy(self, *cols: Union[str, Sequence[str]]) -> "DataStreamWriter": + if len(cols) == 1 and not isinstance(cols[0], str) and isinstance(cols[0], Sequence): + cols = tuple(cols[0]) # Clear any existing columns (if any). while len(self._write_proto.clustering_column_names) > 0: self._write_proto.clustering_column_names.pop() @@ -681,7 +681,7 @@ def _start_internal( tableName: Optional[str] = None, format: Optional[str] = None, outputMode: Optional[str] = None, - partitionBy: Optional[Union[str, List[str]]] = None, + partitionBy: Optional[Union[str, Sequence[str]]] = None, queryName: Optional[str] = None, **options: "OptionalPrimitiveType", ) -> StreamingQuery: @@ -727,7 +727,7 @@ def start( path: Optional[str] = None, format: Optional[str] = None, outputMode: Optional[str] = None, - partitionBy: Optional[Union[str, List[str]]] = None, + partitionBy: Optional[Union[str, Sequence[str]]] = None, queryName: Optional[str] = None, **options: "OptionalPrimitiveType", ) -> "StreamingQuery": @@ -748,7 +748,7 @@ def toTable( tableName: str, format: Optional[str] = None, outputMode: Optional[str] = None, - partitionBy: Optional[Union[str, List[str]]] = None, + partitionBy: Optional[Union[str, Sequence[str]]] = None, queryName: Optional[str] = None, **options: "OptionalPrimitiveType", ) -> "StreamingQuery": diff --git a/python/pyspark/sql/dataframe.py b/python/pyspark/sql/dataframe.py index 6c4d32ea1797f..91bb4fca8e9aa 100644 --- a/python/pyspark/sql/dataframe.py +++ b/python/pyspark/sql/dataframe.py @@ -7173,6 +7173,7 @@ def replace( def replace( self, to_replace: Dict["LiteralType", "OptionalPrimitiveType"], + *, subset: Optional[List[str]] = ..., ) -> DataFrame: ... @@ -7184,7 +7185,7 @@ def replace( subset: Optional[List[str]] = ..., ) -> DataFrame: ... - @dispatch_df_method # type: ignore[misc] + @dispatch_df_method def replace( self, to_replace: Union[List["LiteralType"], Dict["LiteralType", "OptionalPrimitiveType"]], diff --git a/python/pyspark/sql/readwriter.py b/python/pyspark/sql/readwriter.py index afe8000b5c456..e0ac0e3aaef1b 100644 --- a/python/pyspark/sql/readwriter.py +++ b/python/pyspark/sql/readwriter.py @@ -15,7 +15,17 @@ # limitations under the License. # import sys -from typing import cast, overload, Dict, Iterable, List, Optional, Tuple, TYPE_CHECKING, Union +from typing import ( + cast, + overload, + Dict, + Iterable, + List, + Optional, + Sequence, + TYPE_CHECKING, + Union, +) from pyspark.util import is_remote_only from pyspark.sql.types import StructType @@ -34,7 +44,6 @@ __all__ = ["DataFrameReader", "DataFrameWriter", "DataFrameWriterV2"] PathOrPaths = Union[str, List[str]] -TupleOrListOfString = Union[List[str], Tuple[str, ...]] class OptionUtils: @@ -1497,9 +1506,9 @@ def options(self, **options: "OptionalPrimitiveType") -> "DataFrameWriter": def partitionBy(self, *cols: str) -> "DataFrameWriter": ... @overload - def partitionBy(self, *cols: List[str]) -> "DataFrameWriter": ... + def partitionBy(self, __cols: Sequence[str]) -> "DataFrameWriter": ... - def partitionBy(self, *cols: Union[str, List[str]]) -> "DataFrameWriter": + def partitionBy(self, *cols: Union[str, Sequence[str]]) -> "DataFrameWriter": """Partitions the output by the given columns on the file system. If specified, the output is laid out on the file system similar @@ -1546,8 +1555,8 @@ def partitionBy(self, *cols: Union[str, List[str]]) -> "DataFrameWriter": """ from pyspark.sql.classic.column import _to_seq - if len(cols) == 1 and isinstance(cols[0], (list, tuple)): - cols = cols[0] # type: ignore[assignment] + if len(cols) == 1 and not isinstance(cols[0], str) and isinstance(cols[0], Sequence): + cols = tuple(cols[0]) self._jwrite = self._jwrite.partitionBy( _to_seq(self._spark._sc, cast(Iterable["ColumnOrName"], cols)) ) @@ -1557,10 +1566,10 @@ def partitionBy(self, *cols: Union[str, List[str]]) -> "DataFrameWriter": def bucketBy(self, numBuckets: int, col: str, *cols: str) -> "DataFrameWriter": ... @overload - def bucketBy(self, numBuckets: int, col: TupleOrListOfString) -> "DataFrameWriter": ... + def bucketBy(self, numBuckets: int, col: Sequence[str]) -> "DataFrameWriter": ... def bucketBy( - self, numBuckets: int, col: Union[str, TupleOrListOfString], *cols: Optional[str] + self, numBuckets: int, col: Union[str, Sequence[str]], *cols: Optional[str] ) -> "DataFrameWriter": """Buckets the output by the given columns. If specified, the output is laid out on the file system similar to Hive's bucketing scheme, @@ -1619,7 +1628,7 @@ def bucketBy( }, ) - if isinstance(col, (list, tuple)): + if not isinstance(col, str) and isinstance(col, Sequence): if cols: raise PySparkValueError( errorClass="CANNOT_SET_TOGETHER", @@ -1628,7 +1637,7 @@ def bucketBy( }, ) - col, cols = col[0], col[1:] # type: ignore[assignment] + col, cols = col[0], tuple(col[1:]) for c in cols: if not isinstance(c, str): @@ -1659,11 +1668,9 @@ def bucketBy( def sortBy(self, col: str, *cols: str) -> "DataFrameWriter": ... @overload - def sortBy(self, col: TupleOrListOfString) -> "DataFrameWriter": ... + def sortBy(self, col: Sequence[str]) -> "DataFrameWriter": ... - def sortBy( - self, col: Union[str, TupleOrListOfString], *cols: Optional[str] - ) -> "DataFrameWriter": + def sortBy(self, col: Union[str, Sequence[str]], *cols: Optional[str]) -> "DataFrameWriter": """Sorts the output in each bucket by the given columns on the file system. .. versionadded:: 2.3.0 @@ -1703,7 +1710,7 @@ def sortBy( """ from pyspark.sql.classic.column import _to_seq - if isinstance(col, (list, tuple)): + if not isinstance(col, str) and isinstance(col, Sequence): if cols: raise PySparkValueError( errorClass="CANNOT_SET_TOGETHER", @@ -1712,7 +1719,7 @@ def sortBy( }, ) - col, cols = col[0], col[1:] # type: ignore[assignment] + col, cols = col[0], tuple(col[1:]) for c in cols: if not isinstance(c, str): @@ -1743,9 +1750,9 @@ def sortBy( def clusterBy(self, *cols: str) -> "DataFrameWriter": ... @overload - def clusterBy(self, *cols: List[str]) -> "DataFrameWriter": ... + def clusterBy(self, __cols: Sequence[str]) -> "DataFrameWriter": ... - def clusterBy(self, *cols: Union[str, List[str]]) -> "DataFrameWriter": + def clusterBy(self, *cols: Union[str, Sequence[str]]) -> "DataFrameWriter": """Clusters the data by the given columns to optimize query performance. .. versionadded:: 4.0.0 @@ -1767,8 +1774,8 @@ def clusterBy(self, *cols: Union[str, List[str]]) -> "DataFrameWriter": """ from pyspark.sql.classic.column import _to_seq - if len(cols) == 1 and isinstance(cols[0], (list, tuple)): - cols = cols[0] # type: ignore[assignment] + if len(cols) == 1 and not isinstance(cols[0], str) and isinstance(cols[0], Sequence): + cols = tuple(cols[0]) assert len(cols) > 0, "clusterBy needs one or more clustering columns." self._jwrite = self._jwrite.clusterBy(cols[0], _to_seq(self._spark._sc, cols[1:])) return self @@ -1778,7 +1785,7 @@ def save( path: Optional[str] = None, format: Optional[str] = None, mode: Optional[str] = None, - partitionBy: Optional[Union[str, List[str]]] = None, + partitionBy: Optional[Union[str, Sequence[str]]] = None, **options: "OptionalPrimitiveType", ) -> None: """Saves the contents of the :class:`DataFrame` to a data source. @@ -1895,7 +1902,7 @@ def saveAsTable( name: str, format: Optional[str] = None, mode: Optional[str] = None, - partitionBy: Optional[Union[str, List[str]]] = None, + partitionBy: Optional[Union[str, Sequence[str]]] = None, **options: "OptionalPrimitiveType", ) -> None: """Saves the content of the :class:`DataFrame` as the specified table. @@ -2039,7 +2046,7 @@ def parquet( self, path: str, mode: Optional[str] = None, - partitionBy: Optional[Union[str, List[str]]] = None, + partitionBy: Optional[Union[str, Sequence[str]]] = None, compression: Optional[str] = None, ) -> None: """Saves the content of the :class:`DataFrame` in Parquet format at the specified path. @@ -2324,7 +2331,7 @@ def orc( self, path: str, mode: Optional[str] = None, - partitionBy: Optional[Union[str, List[str]]] = None, + partitionBy: Optional[Union[str, Sequence[str]]] = None, compression: Optional[str] = None, ) -> None: """Saves the content of the :class:`DataFrame` in ORC format at the specified path. diff --git a/python/pyspark/sql/streaming/readwriter.py b/python/pyspark/sql/streaming/readwriter.py index 6b7faa6222076..9c66531938ee2 100644 --- a/python/pyspark/sql/streaming/readwriter.py +++ b/python/pyspark/sql/streaming/readwriter.py @@ -18,7 +18,7 @@ import re import sys from collections.abc import Iterator -from typing import cast, overload, Any, Callable, List, Optional, TYPE_CHECKING, Union +from typing import cast, overload, Any, Callable, Optional, Sequence, TYPE_CHECKING, Union from pyspark.sql.readwriter import OptionUtils, to_str from pyspark.sql.streaming.query import StreamingQuery @@ -1193,9 +1193,9 @@ def options(self, **options: "OptionalPrimitiveType") -> "DataStreamWriter": def partitionBy(self, *cols: str) -> "DataStreamWriter": ... @overload - def partitionBy(self, __cols: List[str]) -> "DataStreamWriter": ... + def partitionBy(self, __cols: Sequence[str]) -> "DataStreamWriter": ... - def partitionBy(self, *cols: str) -> "DataStreamWriter": # type: ignore[misc] + def partitionBy(self, *cols: Union[str, Sequence[str]]) -> "DataStreamWriter": """Partitions the output by the given columns on the file system. If specified, the output is laid out on the file system similar @@ -1240,8 +1240,8 @@ def partitionBy(self, *cols: str) -> "DataStreamWriter": # type: ignore[misc] """ from pyspark.sql.classic.column import _to_seq - if len(cols) == 1 and isinstance(cols[0], (list, tuple)): - cols = cols[0] + if len(cols) == 1 and not isinstance(cols[0], str) and isinstance(cols[0], Sequence): + cols = tuple(cols[0]) self._jwrite = self._jwrite.partitionBy(_to_seq(self._spark._sc, cols)) return self @@ -1249,9 +1249,9 @@ def partitionBy(self, *cols: str) -> "DataStreamWriter": # type: ignore[misc] def clusterBy(self, *cols: str) -> "DataStreamWriter": ... @overload - def clusterBy(self, __cols: List[str]) -> "DataStreamWriter": ... + def clusterBy(self, __cols: Sequence[str]) -> "DataStreamWriter": ... - def clusterBy(self, *cols: str) -> "DataStreamWriter": # type: ignore[misc] + def clusterBy(self, *cols: Union[str, Sequence[str]]) -> "DataStreamWriter": """Clusters the output by the given columns. If specified, the output is laid out such that records with similar values on the clustering @@ -1297,8 +1297,8 @@ def clusterBy(self, *cols: str) -> "DataStreamWriter": # type: ignore[misc] """ from pyspark.sql.classic.column import _to_seq - if len(cols) == 1 and isinstance(cols[0], (list, tuple)): - cols = cols[0] + if len(cols) == 1 and not isinstance(cols[0], str) and isinstance(cols[0], Sequence): + cols = tuple(cols[0]) self._jwrite = self._jwrite.clusterBy(_to_seq(self._spark._sc, cols)) return self @@ -1754,7 +1754,7 @@ def start( path: Optional[str] = None, format: Optional[str] = None, outputMode: Optional[str] = None, - partitionBy: Optional[Union[str, List[str]]] = None, + partitionBy: Optional[Union[str, Sequence[str]]] = None, queryName: Optional[str] = None, **options: "OptionalPrimitiveType", ) -> "StreamingQuery": @@ -1842,7 +1842,7 @@ def toTable( tableName: str, format: Optional[str] = None, outputMode: Optional[str] = None, - partitionBy: Optional[Union[str, List[str]]] = None, + partitionBy: Optional[Union[str, Sequence[str]]] = None, queryName: Optional[str] = None, **options: "OptionalPrimitiveType", ) -> "StreamingQuery":