Sitelet https://github.com/apache/datafusion-python/pull/1786/files
Skip to content
Merged
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
4 changes: 4 additions & 0 deletions .github/workflows/build.yml
Original file line number Diff line number Diff line change
Expand Up @@ -81,6 +81,10 @@ jobs:
uv run --no-project ruff check --output-format=github python/
uv run --no-project ruff format --check python/

- name: Run ty
run: |
uv run --no-project ty check --output-format github
Comment thread
timsaucer marked this conversation as resolved.

- name: Run codespell
run: |
uv run --no-project codespell --toml pyproject.toml
Expand Down
7 changes: 7 additions & 0 deletions .pre-commit-config.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -30,6 +30,13 @@ repos:
- id: ruff-format
- repo: local
hooks:
- id: ty
name: ty
description: Type check python/datafusion with ty.
entry: uv run --no-project ty check
pass_filenames: false
types: [file, python]
language: system
- id: rust-fmt
name: Rust fmt
description: Run cargo fmt on files included in the commit. rustfmt should be installed before-hand.
Expand Down
12 changes: 12 additions & 0 deletions pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -181,6 +181,16 @@ extend-allowed-calls = ["datafusion.lit", "lit"]

# CI and pre-commit invoke codespell with different paths, so we have a little
# redundancy here, and we intentionally drop python in the path.
[tool.ty.src]
# Only the published package is type checked for now. Tests and examples
# can be added once the package itself is clean.
include = ["python/datafusion"]

[tool.ty.analysis]
# `_internal` is the compiled PyO3 extension and ships no `.pyi` stubs.
# pandas and polars are optional and only imported under TYPE_CHECKING.
allowed-unresolved-imports = ["datafusion._internal", "pandas", "polars"]

[tool.codespell]
skip = [
"*/tests/test_functions.py",
Expand Down Expand Up @@ -219,6 +229,8 @@ dev = [
"pyyaml>=6.0.3",
"ruff>=0.15.1",
"toml>=0.10.2",
# Pinned exactly: ty is pre-1.0 and new releases can add diagnostics.
"ty==0.0.84",
]
# Release tooling only. Kept out of `dev` because pygithub pulls in
# cryptography, which ships no free-threaded wheel and fails to build
Expand Down
6 changes: 1 addition & 5 deletions python/datafusion/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -57,13 +57,9 @@

from __future__ import annotations

import importlib.metadata as importlib_metadata
from typing import Any

try:
import importlib.metadata as importlib_metadata
except ImportError:
import importlib_metadata # type: ignore[import]

# Public submodules
from . import functions, ipc, object_store, substrait, unparser

Expand Down
9 changes: 5 additions & 4 deletions python/datafusion/catalog.py
Original file line number Diff line number Diff line change
Expand Up @@ -19,6 +19,7 @@

from __future__ import annotations

import sys
from abc import ABC, abstractmethod
from typing import TYPE_CHECKING, Any, Protocol

Expand All @@ -31,10 +32,10 @@
from datafusion.context import TableProviderExportable
from datafusion.expr import CreateExternalTable

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


__all__ = [
Expand Down
14 changes: 7 additions & 7 deletions python/datafusion/context.py
Original file line number Diff line number Diff line change
Expand Up @@ -44,14 +44,15 @@

from __future__ import annotations

import sys
import uuid
import warnings
from typing import TYPE_CHECKING, Any, Protocol

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 urllib.parse import urlparse
Expand Down Expand Up @@ -85,7 +86,6 @@

if TYPE_CHECKING:
import pathlib
import sys
from collections.abc import Iterable, Sequence

import pandas as pd
Expand Down Expand Up @@ -1292,7 +1292,7 @@ def register_csv(
delimiter=delimiter,
schema_infer_max_records=schema_infer_max_records,
file_extension=file_extension,
file_compression_type=file_compression_type,
file_compression_type=file_compression_type or "",
Comment thread
timsaucer marked this conversation as resolved.
)
)

Expand Down Expand Up @@ -2196,7 +2196,7 @@ def read_csv(
schema_infer_max_records=schema_infer_max_records,
file_extension=file_extension,
table_partition_cols=table_partition_cols,
file_compression_type=file_compression_type,
file_compression_type=file_compression_type or "",
Comment thread
timsaucer marked this conversation as resolved.
)
)

Expand Down
24 changes: 13 additions & 11 deletions python/datafusion/dataframe.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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)

Copy link
Copy Markdown
Member Author

Choose a reason for hiding this comment

The 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."
Expand Down Expand Up @@ -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`.
Expand Down Expand Up @@ -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
)
)

Expand Down
39 changes: 30 additions & 9 deletions python/datafusion/expr.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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

Copy link
Copy Markdown
Member Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

These overloads make it so type checkers understand that if a None if fed in you get a None out, otherwise you get a valid list of Expr



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):
Expand All @@ -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:
Expand Down Expand Up @@ -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.
Expand All @@ -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.
Expand Down
24 changes: 12 additions & 12 deletions python/datafusion/functions/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -871,8 +871,8 @@ def concat(*args: Expr) -> Expr:
>>> result.collect_column("c")[0].as_py()
'hello world'
"""
args = [arg.expr for arg in args]
return Expr(f.concat(args))
raw_args = [arg.expr for arg in args]
return Expr(f.concat(raw_args))


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


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


def cos(arg: Expr) -> Expr:
Expand Down Expand Up @@ -3105,8 +3105,8 @@ def make_array(*args: Expr) -> Expr:
>>> result.collect_column("arr")[0].as_py()
[1, 2, 3]
"""
args = [arg.expr for arg in args]
return Expr(f.make_array(args))
raw_args = [arg.expr for arg in args]
return Expr(f.make_array(raw_args))


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


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


def array_cat(*args: Expr) -> Expr:
Expand Down
6 changes: 3 additions & 3 deletions python/datafusion/functions/spark.py
Original file line number Diff line number Diff line change
Expand Up @@ -56,14 +56,14 @@ def _filter_raw(filter: Expr | None) -> Any:
return filter.expr if filter is not None else None


def _coerce_i32(value: Expr | int | None) -> Expr | None:
"""Coerce a native ``int`` to an int32 literal, passing ``Expr``/``None`` through.
def _coerce_i32(value: Expr | int) -> Expr:
"""Coerce a native ``int`` to an int32 literal, passing ``Expr`` through.

Several Spark datetime and interval builders require 32-bit integer
inputs, so a bare ``int`` must become an int32 literal rather than the
int64 default that :meth:`Expr.literal` would produce.
"""
Comment thread
timsaucer marked this conversation as resolved.
Comment thread
timsaucer marked this conversation as resolved.
if value is None or isinstance(value, Expr):
if isinstance(value, Expr):
return value
return Expr.literal(pa.scalar(value, type=pa.int32()))

Expand Down
2 changes: 1 addition & 1 deletion python/datafusion/plan.py
Original file line number Diff line number Diff line change
Expand Up @@ -143,7 +143,7 @@ def to_proto(self) -> bytes:
)
return self.to_bytes()

def __eq__(self, other: LogicalPlan) -> bool:
def __eq__(self, other: object) -> bool:
"""Test equality."""
if not isinstance(other, LogicalPlan):
return False
Expand Down
Loading
Loading