Repository navigation
ci: add ty type checker and fix existing type errors #1786
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
Changes from all commits
File filter
Filter by extension
Conversations
Jump to
Diff view
Diff view
There are no files selected for viewing
| Original file line number | Diff line number | Diff line change |
|---|---|---|
|
|
@@ -44,19 +44,21 @@ | |
|
|
||
| from __future__ import annotations | ||
|
|
||
| import sys | ||
| import warnings | ||
| from collections.abc import AsyncIterator, Iterable, Iterator, Sequence | ||
| from typing import ( | ||
| TYPE_CHECKING, | ||
| Any, | ||
| Literal, | ||
| cast, | ||
| overload, | ||
| ) | ||
|
|
||
| try: | ||
| from warnings import deprecated # Python 3.13+ | ||
| except ImportError: | ||
| from typing_extensions import deprecated # Python 3.12 | ||
| if sys.version_info >= (3, 13): | ||
| from warnings import deprecated | ||
| else: | ||
| from typing_extensions import deprecated | ||
|
|
||
| from datafusion._internal import DataFrame as DataFrameInternal | ||
| from datafusion._internal import DataFrameWriteOptions as DataFrameWriteOptionsInternal | ||
|
|
@@ -1158,8 +1160,8 @@ def join( | |
| if left_on is not None or right_on is not None: | ||
| error_msg = "`left_on` or `right_on` should not provided with `on`" | ||
| raise ValueError(error_msg) | ||
| left_on = on | ||
| right_on = on | ||
| # The legacy ``(left, right)`` tuple form was consumed above. | ||
| left_on = right_on = cast("str | Sequence[str]", on) | ||
|
Member
Author
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. This is a no-op at runtime but will make it so type checkers recognize that we've already verified the type. |
||
| elif left_on is not None or right_on is not None: | ||
| if left_on is None or right_on is None: | ||
| error_msg = "`left_on` and `right_on` should both be provided." | ||
|
|
@@ -1344,10 +1346,11 @@ def repartition_by_hash(self, *exprs: Expr | str, num: int) -> DataFrame: | |
| Returns: | ||
| Repartitioned DataFrame. | ||
| """ | ||
| exprs = [self.parse_sql_expr(e) if isinstance(e, str) else e for e in exprs] | ||
| exprs = expr_list_to_raw_expr_list(exprs) | ||
| raw_exprs = expr_list_to_raw_expr_list( | ||
| [self.parse_sql_expr(e) if isinstance(e, str) else e for e in exprs] | ||
| ) | ||
|
|
||
| return DataFrame(self.df.repartition_by_hash(*exprs, num=num)) | ||
| return DataFrame(self.df.repartition_by_hash(*raw_exprs, num=num)) | ||
|
|
||
| def union(self, other: DataFrame, distinct: bool = False) -> DataFrame: | ||
| """Calculate the union of two :py:class:`DataFrame`. | ||
|
|
@@ -1833,10 +1836,9 @@ def unnest_columns( | |
| >>> df.unnest_columns("a", recursions=[("a", "a", 1)]).to_pydict() | ||
| {'a': [1, 2, 3], 'b': ['x', 'x', 'y']} | ||
| """ | ||
| columns = list(columns) | ||
| return DataFrame( | ||
| self.df.unnest_columns( | ||
| columns, preserve_nulls=preserve_nulls, recursions=recursions | ||
| list(columns), preserve_nulls=preserve_nulls, recursions=recursions | ||
| ) | ||
| ) | ||
|
|
||
|
|
||
| Original file line number | Diff line number | Diff line change |
|---|---|---|
|
|
@@ -46,13 +46,14 @@ | |
|
|
||
| from __future__ import annotations | ||
|
|
||
| import sys | ||
| from collections.abc import Callable, Iterable, Sequence | ||
| from typing import TYPE_CHECKING, Any, ClassVar | ||
| from typing import TYPE_CHECKING, Any, ClassVar, overload | ||
|
|
||
| try: | ||
| from warnings import deprecated # Python 3.13+ | ||
| except ImportError: | ||
| from typing_extensions import deprecated # Python 3.12 | ||
| if sys.version_info >= (3, 13): | ||
| from warnings import deprecated | ||
| else: | ||
| from typing_extensions import deprecated | ||
|
|
||
| import pyarrow as pa | ||
|
|
||
|
|
@@ -114,7 +115,7 @@ def _create_external_table_location(self: Any) -> str: | |
| return locations[0] if locations else "" | ||
|
|
||
|
|
||
| CreateExternalTable.location = _create_external_table_location | ||
| CreateExternalTable.location = _create_external_table_location # ty: ignore[deprecated] | ||
|
|
||
| CreateFunction = expr_internal.CreateFunction | ||
| CreateFunctionBody = expr_internal.CreateFunctionBody | ||
|
|
@@ -410,8 +411,18 @@ def _to_raw_expr(value: Expr | str) -> expr_internal.Expr: | |
| raise TypeError(error) | ||
|
|
||
|
|
||
| @overload | ||
| def expr_list_to_raw_expr_list(expr_list: None) -> None: ... | ||
|
|
||
|
|
||
| @overload | ||
| def expr_list_to_raw_expr_list( | ||
| expr_list: Sequence[Expr | str] | Expr | str, | ||
| ) -> list[expr_internal.Expr]: ... | ||
|
Comment on lines
+414
to
+421
Member
Author
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. These overloads make it so type checkers understand that if a |
||
|
|
||
|
|
||
| def expr_list_to_raw_expr_list( | ||
| expr_list: list[Expr] | Expr | None, | ||
| expr_list: Sequence[Expr | str] | Expr | str | None, | ||
| ) -> list[expr_internal.Expr] | None: | ||
| """Convert a sequence of expressions or column names to raw expressions.""" | ||
| if isinstance(expr_list, Expr | str): | ||
|
|
@@ -428,6 +439,16 @@ def sort_or_default(e: Expr | SortExpr) -> expr_internal.SortExpr: | |
| return SortExpr(e, ascending=True, nulls_first=False).raw_sort | ||
|
|
||
|
|
||
| @overload | ||
| def sort_list_to_raw_sort_list(sort_list: None) -> None: ... | ||
|
|
||
|
|
||
| @overload | ||
| def sort_list_to_raw_sort_list( | ||
| sort_list: Sequence[SortKey] | SortKey, | ||
| ) -> list[expr_internal.SortExpr]: ... | ||
|
|
||
|
|
||
| def sort_list_to_raw_sort_list( | ||
| sort_list: Sequence[SortKey] | SortKey | None, | ||
| ) -> list[expr_internal.SortExpr] | None: | ||
|
|
@@ -771,7 +792,7 @@ def __getitem__(self, key: str | int) -> Expr: | |
| return Expr(functions_internal.array_slice(self.expr, start, stop, step)) | ||
| return Expr(self.expr.__getitem__(key)) | ||
|
|
||
| def __eq__(self, rhs: object) -> Expr: | ||
| def __eq__(self, rhs: object) -> Expr: # ty: ignore[invalid-method-override] | ||
| """Equal to. | ||
|
|
||
| Accepts either an expression or any valid PyArrow scalar literal value. | ||
|
|
@@ -782,7 +803,7 @@ def __eq__(self, rhs: object) -> Expr: | |
| rhs = Expr.literal(rhs) | ||
| return Expr(self.expr.__eq__(rhs.expr)) | ||
|
|
||
| def __ne__(self, rhs: object) -> Expr: | ||
| def __ne__(self, rhs: object) -> Expr: # ty: ignore[invalid-method-override] | ||
| """Not equal to. | ||
|
|
||
| Accepts either an expression or any valid PyArrow scalar literal value. | ||
|
|
||
Uh oh!
There was an error while loading. Please reload this page.