Skip to content

Commit 5febeb3

Browse files
timsaucerclaude
andauthored
ci: add ty type checker and fix existing type errors (#1786)
Add ty (pinned to 0.0.84, since it is pre-1.0) to the dev dependency group, configure it in pyproject.toml to check python/datafusion, and run it in the lint-python CI job and as a local pre-commit hook. datafusion._internal ships no stubs, and pandas/polars are optional TYPE_CHECKING-only imports, so they are allowed to stay unresolved. Fix the diagnostics ty reported: - Make the internal udtf decorator helper require `name`, matching the public overloads. `@udtf()` previously passed None into Rust and failed with "'None' is not an instance of 'str'". - Give AggregateUDF.__init__ defaults matching its FFI overload, and add the same FFI overload to ScalarUDF and WindowUDF. - Add None/non-None overloads to expr_list_to_raw_expr_list and sort_list_to_raw_sort_list so callers that unpack the result type check. - Stop rebinding typed *args and parameters to values of other types. - Import warnings.deprecated behind a sys.version_info check and drop the unreachable importlib_metadata fallback. - Fix smaller annotation mismatches (LogicalPlan.__eq__, spark._coerce_i32, CSV file_compression_type). Co-authored-by: Claude Opus 5.5 (1M context) <noreply@anthropic.com>
1 parent 6c5d9ff commit 5febeb3

13 files changed

Lines changed: 184 additions & 67 deletions

File tree

‎.github/workflows/build.yml‎

Lines changed: 4 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -81,6 +81,10 @@ jobs:
8181
uv run --no-project ruff check --output-format=github python/
8282
uv run --no-project ruff format --check python/
8383
84+
- name: Run ty
85+
run: |
86+
uv run --no-project ty check --output-format github
87+
8488
- name: Run codespell
8589
run: |
8690
uv run --no-project codespell --toml pyproject.toml

‎.pre-commit-config.yaml‎

Lines changed: 7 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -30,6 +30,13 @@ repos:
3030
- id: ruff-format
3131
- repo: local
3232
hooks:
33+
- id: ty
34+
name: ty
35+
description: Type check python/datafusion with ty.
36+
entry: uv run --no-project ty check
37+
pass_filenames: false
38+
types: [file, python]
39+
language: system
3340
- id: rust-fmt
3441
name: Rust fmt
3542
description: Run cargo fmt on files included in the commit. rustfmt should be installed before-hand.

‎pyproject.toml‎

Lines changed: 12 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -181,6 +181,16 @@ extend-allowed-calls = ["datafusion.lit", "lit"]
181181

182182
# CI and pre-commit invoke codespell with different paths, so we have a little
183183
# redundancy here, and we intentionally drop python in the path.
184+
[tool.ty.src]
185+
# Only the published package is type checked for now. Tests and examples
186+
# can be added once the package itself is clean.
187+
include = ["python/datafusion"]
188+
189+
[tool.ty.analysis]
190+
# `_internal` is the compiled PyO3 extension and ships no `.pyi` stubs.
191+
# pandas and polars are optional and only imported under TYPE_CHECKING.
192+
allowed-unresolved-imports = ["datafusion._internal", "pandas", "polars"]
193+
184194
[tool.codespell]
185195
skip = [
186196
"*/tests/test_functions.py",
@@ -219,6 +229,8 @@ dev = [
219229
"pyyaml>=6.0.3",
220230
"ruff>=0.15.1",
221231
"toml>=0.10.2",
232+
# Pinned exactly: ty is pre-1.0 and new releases can add diagnostics.
233+
"ty==0.0.84",
222234
]
223235
# Release tooling only. Kept out of `dev` because pygithub pulls in
224236
# cryptography, which ships no free-threaded wheel and fails to build

‎python/datafusion/__init__.py‎

Lines changed: 1 addition & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -57,13 +57,9 @@
5757

5858
from __future__ import annotations
5959

60+
import importlib.metadata as importlib_metadata
6061
from typing import Any
6162

62-
try:
63-
import importlib.metadata as importlib_metadata
64-
except ImportError:
65-
import importlib_metadata # type: ignore[import]
66-
6763
# Public submodules
6864
from . import functions, ipc, object_store, substrait, unparser
6965

‎python/datafusion/catalog.py‎

Lines changed: 5 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -19,6 +19,7 @@
1919

2020
from __future__ import annotations
2121

22+
import sys
2223
from abc import ABC, abstractmethod
2324
from typing import TYPE_CHECKING, Any, Protocol
2425

@@ -31,10 +32,10 @@
3132
from datafusion.context import TableProviderExportable
3233
from datafusion.expr import CreateExternalTable
3334

34-
try:
35-
from warnings import deprecated # Python 3.13+
36-
except ImportError:
37-
from typing_extensions import deprecated # Python 3.12
35+
if sys.version_info >= (3, 13):
36+
from warnings import deprecated
37+
else:
38+
from typing_extensions import deprecated
3839

3940

4041
__all__ = [

‎python/datafusion/context.py‎

Lines changed: 7 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -44,14 +44,15 @@
4444

4545
from __future__ import annotations
4646

47+
import sys
4748
import uuid
4849
import warnings
4950
from typing import TYPE_CHECKING, Any, Protocol
5051

51-
try:
52-
from warnings import deprecated # Python 3.13+
53-
except ImportError:
54-
from typing_extensions import deprecated # Python 3.12
52+
if sys.version_info >= (3, 13):
53+
from warnings import deprecated
54+
else:
55+
from typing_extensions import deprecated
5556

5657

5758
from urllib.parse import urlparse
@@ -85,7 +86,6 @@
8586

8687
if TYPE_CHECKING:
8788
import pathlib
88-
import sys
8989
from collections.abc import Iterable, Sequence
9090

9191
import pandas as pd
@@ -1292,7 +1292,7 @@ def register_csv(
12921292
delimiter=delimiter,
12931293
schema_infer_max_records=schema_infer_max_records,
12941294
file_extension=file_extension,
1295-
file_compression_type=file_compression_type,
1295+
file_compression_type=file_compression_type or "",
12961296
)
12971297
)
12981298

@@ -2196,7 +2196,7 @@ def read_csv(
21962196
schema_infer_max_records=schema_infer_max_records,
21972197
file_extension=file_extension,
21982198
table_partition_cols=table_partition_cols,
2199-
file_compression_type=file_compression_type,
2199+
file_compression_type=file_compression_type or "",
22002200
)
22012201
)
22022202

‎python/datafusion/dataframe.py‎

Lines changed: 13 additions & 11 deletions
Original file line numberDiff line numberDiff line change
@@ -44,19 +44,21 @@
4444

4545
from __future__ import annotations
4646

47+
import sys
4748
import warnings
4849
from collections.abc import AsyncIterator, Iterable, Iterator, Sequence
4950
from typing import (
5051
TYPE_CHECKING,
5152
Any,
5253
Literal,
54+
cast,
5355
overload,
5456
)
5557

56-
try:
57-
from warnings import deprecated # Python 3.13+
58-
except ImportError:
59-
from typing_extensions import deprecated # Python 3.12
58+
if sys.version_info >= (3, 13):
59+
from warnings import deprecated
60+
else:
61+
from typing_extensions import deprecated
6062

6163
from datafusion._internal import DataFrame as DataFrameInternal
6264
from datafusion._internal import DataFrameWriteOptions as DataFrameWriteOptionsInternal
@@ -1158,8 +1160,8 @@ def join(
11581160
if left_on is not None or right_on is not None:
11591161
error_msg = "`left_on` or `right_on` should not provided with `on`"
11601162
raise ValueError(error_msg)
1161-
left_on = on
1162-
right_on = on
1163+
# The legacy ``(left, right)`` tuple form was consumed above.
1164+
left_on = right_on = cast("str | Sequence[str]", on)
11631165
elif left_on is not None or right_on is not None:
11641166
if left_on is None or right_on is None:
11651167
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:
13441346
Returns:
13451347
Repartitioned DataFrame.
13461348
"""
1347-
exprs = [self.parse_sql_expr(e) if isinstance(e, str) else e for e in exprs]
1348-
exprs = expr_list_to_raw_expr_list(exprs)
1349+
raw_exprs = expr_list_to_raw_expr_list(
1350+
[self.parse_sql_expr(e) if isinstance(e, str) else e for e in exprs]
1351+
)
13491352

1350-
return DataFrame(self.df.repartition_by_hash(*exprs, num=num))
1353+
return DataFrame(self.df.repartition_by_hash(*raw_exprs, num=num))
13511354

13521355
def union(self, other: DataFrame, distinct: bool = False) -> DataFrame:
13531356
"""Calculate the union of two :py:class:`DataFrame`.
@@ -1833,10 +1836,9 @@ def unnest_columns(
18331836
>>> df.unnest_columns("a", recursions=[("a", "a", 1)]).to_pydict()
18341837
{'a': [1, 2, 3], 'b': ['x', 'x', 'y']}
18351838
"""
1836-
columns = list(columns)
18371839
return DataFrame(
18381840
self.df.unnest_columns(
1839-
columns, preserve_nulls=preserve_nulls, recursions=recursions
1841+
list(columns), preserve_nulls=preserve_nulls, recursions=recursions
18401842
)
18411843
)
18421844

‎python/datafusion/expr.py‎

Lines changed: 30 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -46,13 +46,14 @@
4646

4747
from __future__ import annotations
4848

49+
import sys
4950
from collections.abc import Callable, Iterable, Sequence
50-
from typing import TYPE_CHECKING, Any, ClassVar
51+
from typing import TYPE_CHECKING, Any, ClassVar, overload
5152

52-
try:
53-
from warnings import deprecated # Python 3.13+
54-
except ImportError:
55-
from typing_extensions import deprecated # Python 3.12
53+
if sys.version_info >= (3, 13):
54+
from warnings import deprecated
55+
else:
56+
from typing_extensions import deprecated
5657

5758
import pyarrow as pa
5859

@@ -114,7 +115,7 @@ def _create_external_table_location(self: Any) -> str:
114115
return locations[0] if locations else ""
115116

116117

117-
CreateExternalTable.location = _create_external_table_location
118+
CreateExternalTable.location = _create_external_table_location # ty: ignore[deprecated]
118119

119120
CreateFunction = expr_internal.CreateFunction
120121
CreateFunctionBody = expr_internal.CreateFunctionBody
@@ -410,8 +411,18 @@ def _to_raw_expr(value: Expr | str) -> expr_internal.Expr:
410411
raise TypeError(error)
411412

412413

414+
@overload
415+
def expr_list_to_raw_expr_list(expr_list: None) -> None: ...
416+
417+
418+
@overload
419+
def expr_list_to_raw_expr_list(
420+
expr_list: Sequence[Expr | str] | Expr | str,
421+
) -> list[expr_internal.Expr]: ...
422+
423+
413424
def expr_list_to_raw_expr_list(
414-
expr_list: list[Expr] | Expr | None,
425+
expr_list: Sequence[Expr | str] | Expr | str | None,
415426
) -> list[expr_internal.Expr] | None:
416427
"""Convert a sequence of expressions or column names to raw expressions."""
417428
if isinstance(expr_list, Expr | str):
@@ -428,6 +439,16 @@ def sort_or_default(e: Expr | SortExpr) -> expr_internal.SortExpr:
428439
return SortExpr(e, ascending=True, nulls_first=False).raw_sort
429440

430441

442+
@overload
443+
def sort_list_to_raw_sort_list(sort_list: None) -> None: ...
444+
445+
446+
@overload
447+
def sort_list_to_raw_sort_list(
448+
sort_list: Sequence[SortKey] | SortKey,
449+
) -> list[expr_internal.SortExpr]: ...
450+
451+
431452
def sort_list_to_raw_sort_list(
432453
sort_list: Sequence[SortKey] | SortKey | None,
433454
) -> list[expr_internal.SortExpr] | None:
@@ -771,7 +792,7 @@ def __getitem__(self, key: str | int) -> Expr:
771792
return Expr(functions_internal.array_slice(self.expr, start, stop, step))
772793
return Expr(self.expr.__getitem__(key))
773794

774-
def __eq__(self, rhs: object) -> Expr:
795+
def __eq__(self, rhs: object) -> Expr: # ty: ignore[invalid-method-override]
775796
"""Equal to.
776797
777798
Accepts either an expression or any valid PyArrow scalar literal value.
@@ -782,7 +803,7 @@ def __eq__(self, rhs: object) -> Expr:
782803
rhs = Expr.literal(rhs)
783804
return Expr(self.expr.__eq__(rhs.expr))
784805

785-
def __ne__(self, rhs: object) -> Expr:
806+
def __ne__(self, rhs: object) -> Expr: # ty: ignore[invalid-method-override]
786807
"""Not equal to.
787808
788809
Accepts either an expression or any valid PyArrow scalar literal value.

‎python/datafusion/functions/__init__.py‎

Lines changed: 12 additions & 12 deletions
Original file line numberDiff line numberDiff line change
@@ -871,8 +871,8 @@ def concat(*args: Expr) -> Expr:
871871
>>> result.collect_column("c")[0].as_py()
872872
'hello world'
873873
"""
874-
args = [arg.expr for arg in args]
875-
return Expr(f.concat(args))
874+
raw_args = [arg.expr for arg in args]
875+
return Expr(f.concat(raw_args))
876876

877877

878878
def concat_ws(separator: str, *args: Expr) -> Expr:
@@ -888,8 +888,8 @@ def concat_ws(separator: str, *args: Expr) -> Expr:
888888
>>> result.collect_column("c")[0].as_py()
889889
'hello-world'
890890
"""
891-
args = [arg.expr for arg in args]
892-
return Expr(f.concat_ws(separator, args))
891+
raw_args = [arg.expr for arg in args]
892+
return Expr(f.concat_ws(separator, raw_args))
893893

894894

895895
def order_by(expr: Expr, ascending: bool = True, nulls_first: bool = False) -> SortExpr:
@@ -1272,8 +1272,8 @@ def coalesce(*args: Expr) -> Expr:
12721272
>>> result.collect_column("c")[0].as_py()
12731273
2
12741274
"""
1275-
args = [arg.expr for arg in args]
1276-
return Expr(f.coalesce(*args))
1275+
raw_args = [arg.expr for arg in args]
1276+
return Expr(f.coalesce(*raw_args))
12771277

12781278

12791279
def cos(arg: Expr) -> Expr:
@@ -3105,8 +3105,8 @@ def make_array(*args: Expr) -> Expr:
31053105
>>> result.collect_column("arr")[0].as_py()
31063106
[1, 2, 3]
31073107
"""
3108-
args = [arg.expr for arg in args]
3109-
return Expr(f.make_array(args))
3108+
raw_args = [arg.expr for arg in args]
3109+
return Expr(f.make_array(raw_args))
31103110

31113111

31123112
def make_list(*args: Expr) -> Expr:
@@ -3240,8 +3240,8 @@ def struct(*args: Expr) -> Expr:
32403240
>>> result.collect_column("s")[0].as_py() == {"c0": 1, "c1": 2}
32413241
True
32423242
"""
3243-
args = [arg.expr for arg in args]
3244-
return Expr(f.struct(*args))
3243+
raw_args = [arg.expr for arg in args]
3244+
return Expr(f.struct(*raw_args))
32453245

32463246

32473247
def named_struct(name_pairs: list[tuple[str, Expr]]) -> Expr:
@@ -3762,8 +3762,8 @@ def array_concat(*args: Expr) -> Expr:
37623762
>>> result.collect_column("result")[0].as_py()
37633763
[1, 2, 3, 4]
37643764
"""
3765-
args = [arg.expr for arg in args]
3766-
return Expr(f.array_concat(args))
3765+
raw_args = [arg.expr for arg in args]
3766+
return Expr(f.array_concat(raw_args))
37673767

37683768

37693769
def array_cat(*args: Expr) -> Expr:

‎python/datafusion/functions/spark.py‎

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -56,14 +56,14 @@ def _filter_raw(filter: Expr | None) -> Any:
5656
return filter.expr if filter is not None else None
5757

5858

59-
def _coerce_i32(value: Expr | int | None) -> Expr | None:
60-
"""Coerce a native ``int`` to an int32 literal, passing ``Expr``/``None`` through.
59+
def _coerce_i32(value: Expr | int) -> Expr:
60+
"""Coerce a native ``int`` to an int32 literal, passing ``Expr`` through.
6161
6262
Several Spark datetime and interval builders require 32-bit integer
6363
inputs, so a bare ``int`` must become an int32 literal rather than the
6464
int64 default that :meth:`Expr.literal` would produce.
6565
"""
66-
if value is None or isinstance(value, Expr):
66+
if isinstance(value, Expr):
6767
return value
6868
return Expr.literal(pa.scalar(value, type=pa.int32()))
6969

0 commit comments

Comments
 (0)