Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
3 changes: 3 additions & 0 deletions python/pyspark/sql/_typing.pyi
Original file line number Diff line number Diff line change
Expand Up @@ -77,6 +77,9 @@ class SupportsProcess(Protocol):
class SupportsClose(Protocol):
def close(self, error: Exception) -> None: ...

class SupportsOption(Protocol):
def option(self, key: str, value: OptionalPrimitiveType) -> Any: ...

class UserDefinedFunctionLike(Protocol):
func: Callable[..., Any]
evalType: int
Expand Down
4 changes: 4 additions & 0 deletions python/pyspark/sql/connect/_typing.py
Original file line number Diff line number Diff line change
Expand Up @@ -70,6 +70,10 @@
]


class SupportsOption(Protocol):
def option(self, key: str, value: OptionalPrimitiveType) -> Any: ...


class UserDefinedFunctionLike(Protocol):
func: Callable[..., Any]
evalType: int
Expand Down
14 changes: 7 additions & 7 deletions python/pyspark/sql/connect/readwriter.py
Original file line number Diff line number Diff line change
Expand Up @@ -45,7 +45,7 @@

if TYPE_CHECKING:
from pyspark.sql.connect.dataframe import DataFrame
from pyspark.sql.connect._typing import ColumnOrName, OptionalPrimitiveType
from pyspark.sql.connect._typing import ColumnOrName, OptionalPrimitiveType, SupportsOption
from pyspark.sql.connect.session import SparkSession
from pyspark.sql.metrics import ExecutionInfo

Expand All @@ -57,7 +57,7 @@

class OptionUtils:
def _set_opts(
self,
self: "SupportsOption",
schema: Optional[Union[StructType, str]] = None,
**options: "OptionalPrimitiveType",
) -> None:
Expand All @@ -68,7 +68,7 @@ def _set_opts(
self.schema(schema) # type: ignore[attr-defined]
for k, v in options.items():
if v is not None:
self.option(k, v) # type: ignore[attr-defined]
self.option(k, v)


class DataFrameReader(OptionUtils):
Expand Down Expand Up @@ -130,15 +130,15 @@ def load(
self.schema(schema)
self.options(**options)

paths = path
if isinstance(path, str):
paths = [path]
paths: Optional[List[str]] = None
if path is not None:
paths = [path] if isinstance(path, str) else path

plan = DataSource(
format=self._format,
schema=self._schema,
options=self._options,
paths=paths, # type: ignore[arg-type]
paths=paths,
)
return self._df(plan)

Expand Down
18 changes: 11 additions & 7 deletions python/pyspark/sql/connect/streaming/readwriter.py
Original file line number Diff line number Diff line change
Expand Up @@ -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 (
Expand Down Expand Up @@ -510,13 +510,15 @@ def partitionBy(self, *cols: str) -> "DataStreamWriter": ...
@overload
def partitionBy(self, __cols: List[str]) -> "DataStreamWriter": ...

def partitionBy(self, *cols: str) -> "DataStreamWriter": # type: ignore[misc]
def partitionBy(self, *cols: Union[str, List[str]]) -> "DataStreamWriter":
if len(cols) == 1 and isinstance(cols[0], (list, tuple)):
cols = cols[0]
columns: Sequence[str] = cols[0]
else:
columns = cast("Sequence[str]", cols)
# Clear any existing columns (if any).
while len(self._write_proto.partitioning_column_names) > 0:
self._write_proto.partitioning_column_names.pop()
self._write_proto.partitioning_column_names.extend(cast(List[str], cols))
self._write_proto.partitioning_column_names.extend(columns)
return self

partitionBy.__doc__ = PySparkDataStreamWriter.partitionBy.__doc__
Expand All @@ -527,13 +529,15 @@ def clusterBy(self, *cols: str) -> "DataStreamWriter": ...
@overload
def clusterBy(self, __cols: List[str]) -> "DataStreamWriter": ...

def clusterBy(self, *cols: str) -> "DataStreamWriter": # type: ignore[misc]
def clusterBy(self, *cols: Union[str, List[str]]) -> "DataStreamWriter":
if len(cols) == 1 and isinstance(cols[0], (list, tuple)):
cols = cols[0]
columns: Sequence[str] = cols[0]
else:
columns = cast("Sequence[str]", cols)
# Clear any existing columns (if any).
while len(self._write_proto.clustering_column_names) > 0:
self._write_proto.clustering_column_names.pop()
self._write_proto.clustering_column_names.extend(cast(List[str], cols))
self._write_proto.clustering_column_names.extend(columns)
return self

clusterBy.__doc__ = PySparkDataStreamWriter.clusterBy.__doc__
Expand Down
6 changes: 3 additions & 3 deletions python/pyspark/sql/readwriter.py
Original file line number Diff line number Diff line change
Expand Up @@ -26,7 +26,7 @@
if TYPE_CHECKING:
from py4j.java_gateway import JavaObject
from pyspark.core.rdd import RDD
from pyspark.sql._typing import OptionalPrimitiveType, ColumnOrName
from pyspark.sql._typing import OptionalPrimitiveType, ColumnOrName, SupportsOption
from pyspark.sql.session import SparkSession
from pyspark.sql.dataframe import DataFrame
from pyspark.sql.streaming import StreamingQuery
Expand All @@ -39,7 +39,7 @@

class OptionUtils:
def _set_opts(
self,
self: "SupportsOption",
schema: Optional[Union[StructType, str]] = None,
**options: "OptionalPrimitiveType",
) -> None:
Expand All @@ -50,7 +50,7 @@ def _set_opts(
self.schema(schema) # type: ignore[attr-defined]
for k, v in options.items():
if v is not None:
self.option(k, v) # type: ignore[attr-defined]
self.option(k, v)


class DataFrameReader(OptionUtils):
Expand Down
18 changes: 11 additions & 7 deletions python/pyspark/sql/streaming/readwriter.py
Original file line number Diff line number Diff line change
Expand Up @@ -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, List, Optional, Sequence, TYPE_CHECKING, Union

from pyspark.sql.readwriter import OptionUtils, to_str
from pyspark.sql.streaming.query import StreamingQuery
Expand Down Expand Up @@ -1195,7 +1195,7 @@ def partitionBy(self, *cols: str) -> "DataStreamWriter": ...
@overload
def partitionBy(self, __cols: List[str]) -> "DataStreamWriter": ...

def partitionBy(self, *cols: str) -> "DataStreamWriter": # type: ignore[misc]
def partitionBy(self, *cols: Union[str, List[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
Expand Down Expand Up @@ -1241,8 +1241,10 @@ 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]
self._jwrite = self._jwrite.partitionBy(_to_seq(self._spark._sc, cols))
columns: Sequence[str] = cols[0]
else:
columns = cast("Sequence[str]", cols)
self._jwrite = self._jwrite.partitionBy(_to_seq(self._spark._sc, columns))
return self

@overload
Expand All @@ -1251,7 +1253,7 @@ def clusterBy(self, *cols: str) -> "DataStreamWriter": ...
@overload
def clusterBy(self, __cols: List[str]) -> "DataStreamWriter": ...

def clusterBy(self, *cols: str) -> "DataStreamWriter": # type: ignore[misc]
def clusterBy(self, *cols: Union[str, List[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
Expand Down Expand Up @@ -1298,8 +1300,10 @@ 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]
self._jwrite = self._jwrite.clusterBy(_to_seq(self._spark._sc, cols))
columns: Sequence[str] = cols[0]
else:
columns = cast("Sequence[str]", cols)
self._jwrite = self._jwrite.clusterBy(_to_seq(self._spark._sc, columns))
return self

def queryName(self, queryName: str) -> "DataStreamWriter":
Expand Down