From 868baf5b985a9972f1cd9d688ac01a9cd6e566a0 Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Fri, 2 Oct 2026 02:05:21 +0900 Subject: [PATCH] Mark every checkable override with @override Mark the methods that override a base class method with pyathena.util.override, and enable mypy's explicit-override error code so that mypy reports an override without the decorator. This covers the synchronous and asyncio cursors and result sets, converters, readers and the formatter, the SQLAlchemy dialects, compilers, types and test-suite requirements, and the filesystem helpers whose bases are typed. mypy skips both checks for an unannotated function: it does not report a marked method whose base method is missing, and it does not report an unannotated property override without the decorator. 79 of the marked methods are unannotated, mostly SQLAlchemy overrides. The fsspec overrides stay unmarked: fsspec has no type information, so mypy rejects @override on them. mypy cannot accept @override on a property that overrides a property with a setter (python/mypy#15900); those seven getters keep the decorator and ignore the error. Co-Authored-By: Claude Opus 5.5 --- pyathena/__init__.py | 4 +++ pyathena/aio/arrow/cursor.py | 1 + pyathena/aio/cursor.py | 3 +- pyathena/aio/pandas/cursor.py | 1 + pyathena/aio/polars/cursor.py | 1 + pyathena/aio/s3fs/cursor.py | 1 + pyathena/aio/spark/cursor.py | 1 + pyathena/aio/sqlalchemy/arrow.py | 4 ++- pyathena/aio/sqlalchemy/base.py | 8 ++++- pyathena/aio/sqlalchemy/pandas.py | 4 ++- pyathena/aio/sqlalchemy/polars.py | 4 ++- pyathena/aio/sqlalchemy/rest.py | 2 ++ pyathena/aio/sqlalchemy/s3fs.py | 3 ++ pyathena/arrow/async_cursor.py | 7 +++- pyathena/arrow/converter.py | 3 ++ pyathena/arrow/cursor.py | 3 ++ pyathena/arrow/result_set.py | 5 ++- pyathena/async_cursor.py | 4 +++ pyathena/converter.py | 3 +- pyathena/cursor.py | 5 ++- pyathena/filesystem/s3_executor.py | 6 ++++ pyathena/filesystem/s3_object.py | 14 ++++++++ pyathena/formatter.py | 2 ++ pyathena/pandas/async_cursor.py | 7 +++- pyathena/pandas/converter.py | 3 ++ pyathena/pandas/cursor.py | 3 ++ pyathena/pandas/reader.py | 4 +++ pyathena/pandas/result_set.py | 6 +++- pyathena/polars/async_cursor.py | 7 +++- pyathena/polars/converter.py | 4 +++ pyathena/polars/cursor.py | 3 ++ pyathena/polars/result_set.py | 6 +++- pyathena/result_set.py | 10 +++++- pyathena/s3fs/async_cursor.py | 7 +++- pyathena/s3fs/converter.py | 2 ++ pyathena/s3fs/cursor.py | 3 ++ pyathena/s3fs/reader.py | 6 ++++ pyathena/s3fs/result_set.py | 5 ++- pyathena/spark/async_cursor.py | 3 ++ pyathena/spark/common.py | 7 +++- pyathena/spark/cursor.py | 3 ++ pyathena/sqlalchemy/array.py | 15 +++++++++ pyathena/sqlalchemy/arrow.py | 4 ++- pyathena/sqlalchemy/base.py | 17 ++++++++++ pyathena/sqlalchemy/compiler.py | 50 ++++++++++++++++++++++++++++- pyathena/sqlalchemy/map.py | 3 ++ pyathena/sqlalchemy/pandas.py | 4 ++- pyathena/sqlalchemy/polars.py | 4 ++- pyathena/sqlalchemy/requirements.py | 41 +++++++++++++++++++++++ pyathena/sqlalchemy/rest.py | 2 ++ pyathena/sqlalchemy/s3fs.py | 3 ++ pyathena/sqlalchemy/struct.py | 4 +++ pyathena/sqlalchemy/temporal.py | 7 ++++ pyathena/sqlalchemy/types.py | 2 ++ pyathena/sqlalchemy/util.py | 5 ++- pyproject.toml | 2 ++ 56 files changed, 319 insertions(+), 22 deletions(-) diff --git a/pyathena/__init__.py b/pyathena/__init__.py index acba8de5a..1e73d195d 100644 --- a/pyathena/__init__.py +++ b/pyathena/__init__.py @@ -5,6 +5,7 @@ from pyathena.error import * # noqa: F403 from pyathena.options import ExecuteOptions as ExecuteOptions +from pyathena.util import override if TYPE_CHECKING: from pyathena.aio.connection import AioConnection @@ -34,16 +35,19 @@ class DBAPITypeObject(frozenset[str]): https://www.python.org/dev/peps/pep-0249/#type-objects-and-constructors """ + @override def __eq__(self, other: object): if isinstance(other, frozenset): return frozenset.__eq__(self, other) return other in self + @override def __ne__(self, other: object): if isinstance(other, frozenset): return frozenset.__ne__(self, other) return other not in self + @override def __hash__(self): return frozenset.__hash__(self) diff --git a/pyathena/aio/arrow/cursor.py b/pyathena/aio/arrow/cursor.py index 83c18876f..2932b24fd 100644 --- a/pyathena/aio/arrow/cursor.py +++ b/pyathena/aio/arrow/cursor.py @@ -73,6 +73,7 @@ def __init__( self._result_set: AthenaArrowResultSet | None = None @staticmethod + @override def get_default_converter( unload: bool = False, ) -> DefaultArrowTypeConverter | DefaultArrowUnloadTypeConverter | Any: diff --git a/pyathena/aio/cursor.py b/pyathena/aio/cursor.py index 083d35970..38ef885fd 100644 --- a/pyathena/aio/cursor.py +++ b/pyathena/aio/cursor.py @@ -66,7 +66,8 @@ def __init__( self._result_set: AthenaAioResultSet | None = None self._result_set_class = AthenaAioResultSet - @property + @property # type: ignore[explicit-override] # python/mypy#15900 + @override def arraysize(self) -> int: return self._arraysize diff --git a/pyathena/aio/pandas/cursor.py b/pyathena/aio/pandas/cursor.py index ca2edd768..863fe5168 100644 --- a/pyathena/aio/pandas/cursor.py +++ b/pyathena/aio/pandas/cursor.py @@ -86,6 +86,7 @@ def __init__( self._result_set: AthenaPandasResultSet | None = None @staticmethod + @override def get_default_converter( unload: bool = False, ) -> DefaultPandasTypeConverter | Any: diff --git a/pyathena/aio/polars/cursor.py b/pyathena/aio/polars/cursor.py index 8d5da6524..ffe330187 100644 --- a/pyathena/aio/polars/cursor.py +++ b/pyathena/aio/polars/cursor.py @@ -79,6 +79,7 @@ def __init__( self._result_set: AthenaPolarsResultSet | None = None @staticmethod + @override def get_default_converter( unload: bool = False, ) -> DefaultPolarsTypeConverter | DefaultPolarsUnloadTypeConverter | Any: diff --git a/pyathena/aio/s3fs/cursor.py b/pyathena/aio/s3fs/cursor.py index 79c52dfe6..5850ac460 100644 --- a/pyathena/aio/s3fs/cursor.py +++ b/pyathena/aio/s3fs/cursor.py @@ -72,6 +72,7 @@ def __init__( self._result_set: AthenaS3FSResultSet | None = None @staticmethod + @override def get_default_converter( unload: bool = False, ) -> DefaultS3FSTypeConverter: diff --git a/pyathena/aio/spark/cursor.py b/pyathena/aio/spark/cursor.py index e2c09880e..6c8ece8aa 100644 --- a/pyathena/aio/spark/cursor.py +++ b/pyathena/aio/spark/cursor.py @@ -50,6 +50,7 @@ class AioSparkCursor(SparkBaseCursor, WithCalculationExecution): """ @property + @override def calculation_execution(self) -> AthenaCalculationExecution | None: return self._calculation_execution diff --git a/pyathena/aio/sqlalchemy/arrow.py b/pyathena/aio/sqlalchemy/arrow.py index 2e9f9cebe..f65671137 100644 --- a/pyathena/aio/sqlalchemy/arrow.py +++ b/pyathena/aio/sqlalchemy/arrow.py @@ -8,7 +8,7 @@ from typing import TYPE_CHECKING from pyathena.aio.sqlalchemy.base import AthenaAioDialect -from pyathena.util import strtobool +from pyathena.util import override, strtobool if TYPE_CHECKING: from types import ModuleType @@ -43,6 +43,7 @@ class AthenaAioArrowDialect(AthenaAioDialect): driver = "aioarrow" supports_statement_cache = True + @override def create_connect_args(self, url): from pyathena.aio.arrow.cursor import AioArrowCursor @@ -57,5 +58,6 @@ def create_connect_args(self, url): return [[], opts] @classmethod + @override def import_dbapi(cls) -> "ModuleType": return super().import_dbapi() diff --git a/pyathena/aio/sqlalchemy/base.py b/pyathena/aio/sqlalchemy/base.py index d46d89089..6c65dc3b6 100644 --- a/pyathena/aio/sqlalchemy/base.py +++ b/pyathena/aio/sqlalchemy/base.py @@ -31,7 +31,7 @@ ProgrammingError, ) from pyathena.sqlalchemy.base import AthenaDialect -from pyathena.util import RetryConfig +from pyathena.util import RetryConfig, override if TYPE_CHECKING: from types import ModuleType @@ -140,6 +140,7 @@ def __init__(self, dbapi: AsyncAdapt_pyathena_dbapi, connection: AioConnection) self._connection = connection # type: ignore[assignment] @property + @override def driver_connection(self) -> AioConnection: return self._connection # type: ignore[return-value] @@ -225,21 +226,26 @@ class AthenaAioDialect(AthenaDialect): supports_statement_cache = True @classmethod + @override def get_pool_class(cls, url: URL) -> type: return pool.AsyncAdaptedQueuePool @classmethod + @override def import_dbapi(cls) -> ModuleType: return AsyncAdapt_pyathena_dbapi() # type: ignore[return-value] @classmethod + @override def dbapi(cls) -> ModuleType: # type: ignore[override] return AsyncAdapt_pyathena_dbapi() # type: ignore[return-value] + @override def create_connect_args(self, url: URL) -> tuple[tuple[str], MutableMapping[str, Any]]: opts = self._create_connect_args(url) self._connect_options = opts return cast(tuple[str], ()), opts + @override def get_driver_connection(self, connection: Any) -> Any: return connection diff --git a/pyathena/aio/sqlalchemy/pandas.py b/pyathena/aio/sqlalchemy/pandas.py index bfd2fd28b..6f10eaf6d 100644 --- a/pyathena/aio/sqlalchemy/pandas.py +++ b/pyathena/aio/sqlalchemy/pandas.py @@ -8,7 +8,7 @@ from typing import TYPE_CHECKING from pyathena.aio.sqlalchemy.base import AthenaAioDialect -from pyathena.util import strtobool +from pyathena.util import override, strtobool if TYPE_CHECKING: from types import ModuleType @@ -45,6 +45,7 @@ class AthenaAioPandasDialect(AthenaAioDialect): driver = "aiopandas" supports_statement_cache = True + @override def create_connect_args(self, url): from pyathena.aio.pandas.cursor import AioPandasCursor @@ -63,5 +64,6 @@ def create_connect_args(self, url): return [[], opts] @classmethod + @override def import_dbapi(cls) -> "ModuleType": return super().import_dbapi() diff --git a/pyathena/aio/sqlalchemy/polars.py b/pyathena/aio/sqlalchemy/polars.py index f7eff642f..e4f0218f7 100644 --- a/pyathena/aio/sqlalchemy/polars.py +++ b/pyathena/aio/sqlalchemy/polars.py @@ -8,7 +8,7 @@ from typing import TYPE_CHECKING from pyathena.aio.sqlalchemy.base import AthenaAioDialect -from pyathena.util import strtobool +from pyathena.util import override, strtobool if TYPE_CHECKING: from types import ModuleType @@ -43,6 +43,7 @@ class AthenaAioPolarsDialect(AthenaAioDialect): driver = "aiopolars" supports_statement_cache = True + @override def create_connect_args(self, url): from pyathena.aio.polars.cursor import AioPolarsCursor @@ -57,5 +58,6 @@ def create_connect_args(self, url): return [[], opts] @classmethod + @override def import_dbapi(cls) -> "ModuleType": return super().import_dbapi() diff --git a/pyathena/aio/sqlalchemy/rest.py b/pyathena/aio/sqlalchemy/rest.py index a20f9131e..74d3adb8f 100644 --- a/pyathena/aio/sqlalchemy/rest.py +++ b/pyathena/aio/sqlalchemy/rest.py @@ -8,6 +8,7 @@ from typing import TYPE_CHECKING from pyathena.aio.sqlalchemy.base import AthenaAioDialect +from pyathena.util import override if TYPE_CHECKING: from types import ModuleType @@ -39,5 +40,6 @@ class AthenaAioRestDialect(AthenaAioDialect): supports_statement_cache = True @classmethod + @override def import_dbapi(cls) -> "ModuleType": return super().import_dbapi() diff --git a/pyathena/aio/sqlalchemy/s3fs.py b/pyathena/aio/sqlalchemy/s3fs.py index d7a042156..39d8e034b 100644 --- a/pyathena/aio/sqlalchemy/s3fs.py +++ b/pyathena/aio/sqlalchemy/s3fs.py @@ -8,6 +8,7 @@ from typing import TYPE_CHECKING from pyathena.aio.sqlalchemy.base import AthenaAioDialect +from pyathena.util import override if TYPE_CHECKING: from types import ModuleType @@ -37,6 +38,7 @@ class AthenaAioS3FSDialect(AthenaAioDialect): driver = "aios3fs" supports_statement_cache = True + @override def create_connect_args(self, url): from pyathena.aio.s3fs.cursor import AioS3FSCursor @@ -46,5 +48,6 @@ def create_connect_args(self, url): return [[], opts] @classmethod + @override def import_dbapi(cls) -> "ModuleType": return super().import_dbapi() diff --git a/pyathena/arrow/async_cursor.py b/pyathena/arrow/async_cursor.py index 3deeac12c..ab2a18eef 100644 --- a/pyathena/arrow/async_cursor.py +++ b/pyathena/arrow/async_cursor.py @@ -15,6 +15,7 @@ from pyathena.common import CursorIterator from pyathena.model import AthenaQueryExecution from pyathena.options import ExecuteOptions +from pyathena.util import override _logger = logging.getLogger(__name__) @@ -130,6 +131,7 @@ def __init__( self._request_timeout = request_timeout @staticmethod + @override def get_default_converter( unload: bool = False, ) -> DefaultArrowTypeConverter | DefaultArrowUnloadTypeConverter | Any: @@ -137,7 +139,8 @@ def get_default_converter( return DefaultArrowUnloadTypeConverter() return DefaultArrowTypeConverter() - @property + @property # type: ignore[explicit-override] # python/mypy#15900 + @override def arraysize(self) -> int: return self._arraysize @@ -147,6 +150,7 @@ def arraysize(self, value: int) -> None: raise ProgrammingError("arraysize must be a positive integer value.") self._arraysize = value + @override def _collect_result_set( self, query_id: str, @@ -171,6 +175,7 @@ def _collect_result_set( **kwargs, ) + @override def execute( self, operation: str, diff --git a/pyathena/arrow/converter.py b/pyathena/arrow/converter.py index cd9c55892..cdb475eb2 100644 --- a/pyathena/arrow/converter.py +++ b/pyathena/arrow/converter.py @@ -14,6 +14,7 @@ _to_json, _to_time, ) +from pyathena.util import override _logger = logging.getLogger(__name__) @@ -90,6 +91,7 @@ def _dtypes(self) -> dict[str, type[Any]]: } return self.__dtypes + @override def convert(self, type_: str, value: str | None, type_hint: str | None = None) -> Any | None: converter = self.get(type_) return converter(value) @@ -114,6 +116,7 @@ def __init__(self) -> None: default=_to_default, ) + @override def convert(self, type_: str, value: str | None, type_hint: str | None = None) -> Any | None: converter = self.get(type_) return converter(value) diff --git a/pyathena/arrow/cursor.py b/pyathena/arrow/cursor.py index 5ad7cd603..757d98da8 100644 --- a/pyathena/arrow/cursor.py +++ b/pyathena/arrow/cursor.py @@ -14,6 +14,7 @@ from pyathena.model import AthenaQueryExecution from pyathena.options import ExecuteOptions from pyathena.result_set import WithFetch +from pyathena.util import override if TYPE_CHECKING: import polars as pl @@ -116,6 +117,7 @@ def __init__( self._request_timeout = request_timeout @staticmethod + @override def get_default_converter( unload: bool = False, ) -> DefaultArrowTypeConverter | DefaultArrowUnloadTypeConverter | Any: @@ -123,6 +125,7 @@ def get_default_converter( return DefaultArrowUnloadTypeConverter() return DefaultArrowTypeConverter() + @override def execute( self, operation: str, diff --git a/pyathena/arrow/result_set.py b/pyathena/arrow/result_set.py index 7fa0a7d9e..2b091fa6d 100644 --- a/pyathena/arrow/result_set.py +++ b/pyathena/arrow/result_set.py @@ -14,7 +14,7 @@ from pyathena.error import ProgrammingError from pyathena.model import AthenaQueryExecution from pyathena.result_set import AthenaResultSet -from pyathena.util import RetryConfig, parse_output_location +from pyathena.util import RetryConfig, override, parse_output_location if TYPE_CHECKING: import polars as pl @@ -213,6 +213,7 @@ def converters(self) -> dict[str, Callable[[str | None], Any | None]]: description = self.description if self.description else [] return {d[0]: self._converter.get(d[1]) for d in description} + @override def _fetch(self) -> None: try: rows = next(self._batches) @@ -227,6 +228,7 @@ def _fetch(self) -> None: ] self._rows.extend(processed_rows) + @override def fetchone( self, ) -> tuple[Any | None, ...] | dict[Any, Any | None] | None: @@ -383,6 +385,7 @@ def as_polars(self) -> pl.DataFrame: "polars is required for as_polars(). Install it with: pip install polars" ) from e + @override def close(self) -> None: import pyarrow as pa diff --git a/pyathena/async_cursor.py b/pyathena/async_cursor.py index a829ac145..620df9d68 100644 --- a/pyathena/async_cursor.py +++ b/pyathena/async_cursor.py @@ -11,6 +11,7 @@ from pyathena.model import AthenaQueryExecution from pyathena.options import ExecuteOptions from pyathena.result_set import AthenaDictResultSet, AthenaResultSet +from pyathena.util import override _logger = logging.getLogger(__name__) @@ -98,6 +99,7 @@ def arraysize(self, value: int) -> None: ) self._arraysize = value + @override def close(self, wait: bool = False) -> None: self._executor.shutdown(wait=wait) @@ -160,6 +162,7 @@ def _collect_result_set( result_set_type_hints=result_set_type_hints, ) + @override def execute( self, operation: str, @@ -231,6 +234,7 @@ def execute( self._collect_result_set, query_id, options.result_set_type_hints ) + @override def executemany( self, operation: str, diff --git a/pyathena/converter.py b/pyathena/converter.py index 03cabb232..7336bfaea 100644 --- a/pyathena/converter.py +++ b/pyathena/converter.py @@ -19,7 +19,7 @@ TypeSignatureParser, _split_array_items, ) -from pyathena.util import strtobool +from pyathena.util import override, strtobool _logger = logging.getLogger(__name__) @@ -657,6 +657,7 @@ def _normalize_hive_syntax(type_str: str) -> str: lambda m: DefaultTypeConverter._HIVE_REPLACEMENTS[m.group()], type_str ) + @override def convert(self, type_: str, value: str | None, type_hint: str | None = None) -> Any | None: """Convert a string value to the appropriate Python type. diff --git a/pyathena/cursor.py b/pyathena/cursor.py index eb302b306..2e789520a 100644 --- a/pyathena/cursor.py +++ b/pyathena/cursor.py @@ -9,6 +9,7 @@ from pyathena.model import AthenaQueryExecution from pyathena.options import ExecuteOptions from pyathena.result_set import AthenaDictResultSet, AthenaResultSet, WithFetch +from pyathena.util import override _logger = logging.getLogger(__name__) @@ -70,7 +71,8 @@ def __init__( ) self._result_set_class = AthenaResultSet - @property + @property # type: ignore[explicit-override] # python/mypy#15900 + @override def arraysize(self) -> int: return self._arraysize @@ -82,6 +84,7 @@ def arraysize(self, value: int) -> None: ) self._arraysize = value + @override def execute( self, operation: str, diff --git a/pyathena/filesystem/s3_executor.py b/pyathena/filesystem/s3_executor.py index b8c02013f..603005555 100644 --- a/pyathena/filesystem/s3_executor.py +++ b/pyathena/filesystem/s3_executor.py @@ -14,6 +14,8 @@ from concurrent.futures.thread import ThreadPoolExecutor from typing import Any, TypeVar +from pyathena.util import override + T = TypeVar("T") @@ -53,9 +55,11 @@ class S3ThreadPoolExecutor(S3Executor): def __init__(self, max_workers: int) -> None: self._executor = ThreadPoolExecutor(max_workers=max_workers) + @override def submit(self, fn: Callable[..., T], *args: Any, **kwargs: Any) -> Future[T]: return self._executor.submit(fn, *args, **kwargs) + @override def shutdown(self, wait: bool = True) -> None: self._executor.shutdown(wait=wait) @@ -81,6 +85,7 @@ class S3AioExecutor(S3Executor): def __init__(self, loop: asyncio.AbstractEventLoop | None = None) -> None: self._loop = loop + @override def submit(self, fn: Callable[..., T], *args: Any, **kwargs: Any) -> Future[T]: if self._loop is not None and self._loop.is_running(): return asyncio.run_coroutine_threadsafe( @@ -91,6 +96,7 @@ def submit(self, fn: Callable[..., T], *args: Any, **kwargs: Any) -> Future[T]: "Use S3ThreadPoolExecutor for synchronous usage." ) + @override def shutdown(self, wait: bool = True) -> None: # No resources to release — work is dispatched to the event loop. pass diff --git a/pyathena/filesystem/s3_object.py b/pyathena/filesystem/s3_object.py index 192b73730..392fe0358 100644 --- a/pyathena/filesystem/s3_object.py +++ b/pyathena/filesystem/s3_object.py @@ -6,6 +6,8 @@ from datetime import datetime from typing import Any +from pyathena.util import override + _logger = logging.getLogger(__name__) _API_FIELD_TO_S3_OBJECT_PROPERTY = { @@ -135,30 +137,38 @@ def __init__( else: self.name = f"{self.get('bucket')}/{self.get('key')}" + @override def get(self, key: str, default: Any = None) -> Any: return super().get(key, default) + @override def __getitem__(self, item: str) -> Any: return self.__dict__.get(item) def __getattr__(self, item: str): return self.get(item) + @override def __setitem__(self, key: str, value: Any) -> None: self.__dict__[key] = value + @override def __setattr__(self, attr: str, value: Any) -> None: self[attr] = value + @override def __delitem__(self, key: str) -> None: del self.__dict__[key] + @override def __iter__(self) -> Iterator[str]: return iter(self.__dict__.keys()) + @override def __len__(self) -> int: return len(self.__dict__) + @override def __str__(self): return str(self.__dict__) @@ -227,15 +237,19 @@ def __init__(self, response: dict[str, Any]) -> None: self._version_id: str | None = response.get("VersionId") self._user_metadata: dict[str, str] = response.get("Metadata", {}) + @override def __getitem__(self, key: str) -> str: return self._user_metadata[key] + @override def __iter__(self) -> Iterator[str]: return iter(self._user_metadata) + @override def __len__(self) -> int: return len(self._user_metadata) + @override def __repr__(self) -> str: return f"{self.__class__.__name__}({self._user_metadata!r})" diff --git a/pyathena/formatter.py b/pyathena/formatter.py index 98a051def..b1c686e7a 100644 --- a/pyathena/formatter.py +++ b/pyathena/formatter.py @@ -14,6 +14,7 @@ from pyathena.error import ProgrammingError from pyathena.model import AthenaCompression, AthenaFileFormat +from pyathena.util import override _logger = logging.getLogger(__name__) @@ -396,6 +397,7 @@ class DefaultParameterFormatter(Formatter): def __init__(self) -> None: super().__init__(mappings=deepcopy(_DEFAULT_FORMATTERS), default=None) + @override def format(self, operation: str, parameters: dict[str, Any] | None = None) -> str: if not operation or not operation.strip(): raise ProgrammingError("Query is none or empty.") diff --git a/pyathena/pandas/async_cursor.py b/pyathena/pandas/async_cursor.py index 4cd443b95..a98927a83 100644 --- a/pyathena/pandas/async_cursor.py +++ b/pyathena/pandas/async_cursor.py @@ -16,6 +16,7 @@ DefaultPandasUnloadTypeConverter, ) from pyathena.pandas.result_set import AthenaPandasResultSet +from pyathena.util import override _logger = logging.getLogger(__name__) @@ -99,6 +100,7 @@ def __init__( self._chunksize = chunksize @staticmethod + @override def get_default_converter( unload: bool = False, ) -> DefaultPandasTypeConverter | Any: @@ -106,7 +108,8 @@ def get_default_converter( return DefaultPandasUnloadTypeConverter() return DefaultPandasTypeConverter() - @property + @property # type: ignore[explicit-override] # python/mypy#15900 + @override def arraysize(self) -> int: return self._arraysize @@ -116,6 +119,7 @@ def arraysize(self, value: int) -> None: raise ProgrammingError("arraysize must be a positive integer value.") self._arraysize = value + @override def _collect_result_set( self, query_id: str, @@ -146,6 +150,7 @@ def _collect_result_set( **kwargs, ) + @override def execute( self, operation: str, diff --git a/pyathena/pandas/converter.py b/pyathena/pandas/converter.py index 87417a87d..54f97ca9b 100644 --- a/pyathena/pandas/converter.py +++ b/pyathena/pandas/converter.py @@ -13,6 +13,7 @@ _to_default, _to_json, ) +from pyathena.util import override _logger = logging.getLogger(__name__) @@ -80,6 +81,7 @@ def _dtypes(self) -> dict[str, type[Any]]: } return self.__dtypes + @override def convert(self, type_: str, value: str | None, type_hint: str | None = None) -> Any | None: converter = self.get(type_) return converter(value) @@ -104,6 +106,7 @@ def __init__(self) -> None: default=_to_default, ) + @override def convert(self, type_: str, value: str | None, type_hint: str | None = None) -> Any | None: converter = self.get(type_) return converter(value) diff --git a/pyathena/pandas/cursor.py b/pyathena/pandas/cursor.py index 883172c05..f574a2d30 100644 --- a/pyathena/pandas/cursor.py +++ b/pyathena/pandas/cursor.py @@ -19,6 +19,7 @@ ) from pyathena.pandas.result_set import AthenaPandasResultSet, PandasDataFrameIterator from pyathena.result_set import WithFetch +from pyathena.util import override if TYPE_CHECKING: from pandas import DataFrame @@ -129,6 +130,7 @@ def __init__( self._auto_optimize_chunksize = auto_optimize_chunksize @staticmethod + @override def get_default_converter( unload: bool = False, ) -> DefaultPandasTypeConverter | Any: @@ -136,6 +138,7 @@ def get_default_converter( return DefaultPandasUnloadTypeConverter() return DefaultPandasTypeConverter() + @override def execute( self, operation: str, diff --git a/pyathena/pandas/reader.py b/pyathena/pandas/reader.py index e7a1fa546..bfdcc4777 100644 --- a/pyathena/pandas/reader.py +++ b/pyathena/pandas/reader.py @@ -12,6 +12,7 @@ from typing import Any from pyathena.s3fs.reader import AthenaCSVReader +from pyathena.util import override _BINARY_NULL = "__PYATHENA_BINARY_NULL__" _CSV_FIELD = re.compile(r'(?:^|,)(?P"[^"]*(?:""[^"]*)*"|[^,]*)') @@ -32,9 +33,11 @@ def __init__(self, stream: Any, binary_columns: set[int]) -> None: self._header = True self._buffer = b"" + @override def readable(self) -> bool: return True + @override def readinto(self, buffer: Any) -> int: if self.closed: raise ValueError("I/O operation on closed file.") @@ -64,6 +67,7 @@ def readinto(self, buffer: Any) -> int: self._buffer = self._buffer[size:] return size + @override def close(self) -> None: try: self._reader.close() diff --git a/pyathena/pandas/result_set.py b/pyathena/pandas/result_set.py index d6299b189..1b36ad027 100644 --- a/pyathena/pandas/result_set.py +++ b/pyathena/pandas/result_set.py @@ -22,7 +22,7 @@ from pyathena.model import AthenaQueryExecution from pyathena.pandas.reader import _BINARY_NULL, BinaryCSVReader from pyathena.result_set import AthenaResultSet -from pyathena.util import RetryConfig, parse_output_location +from pyathena.util import RetryConfig, override, parse_output_location if TYPE_CHECKING: from pandas import DataFrame @@ -88,6 +88,7 @@ def __init__( self._trunc_date = trunc_date self._csv_stream = csv_stream + @override def __next__(self) -> DataFrame: """Get the next DataFrame chunk. @@ -104,6 +105,7 @@ def __next__(self) -> DataFrame: self.close() raise + @override def __iter__(self) -> PandasDataFrameIterator: """Return self as iterator.""" return self @@ -478,6 +480,7 @@ def _trunc_date(self, df: DataFrame) -> DataFrame: df.isetitem(df.columns.get_loc(time_col), truncated[time_col]) return df + @override def fetchone( self, ) -> tuple[Any | None, ...] | dict[Any, Any | None] | None: @@ -837,6 +840,7 @@ def iter_chunks(self) -> PandasDataFrameIterator: """ return self._df_iter + @override def close(self) -> None: import pandas as pd diff --git a/pyathena/polars/async_cursor.py b/pyathena/polars/async_cursor.py index cab288362..719475547 100644 --- a/pyathena/polars/async_cursor.py +++ b/pyathena/polars/async_cursor.py @@ -15,6 +15,7 @@ DefaultPolarsUnloadTypeConverter, ) from pyathena.polars.result_set import AthenaPolarsResultSet +from pyathena.util import override _logger = logging.getLogger(__name__) @@ -125,6 +126,7 @@ def __init__( self._chunksize = chunksize @staticmethod + @override def get_default_converter( unload: bool = False, ) -> DefaultPolarsTypeConverter | DefaultPolarsUnloadTypeConverter | Any: @@ -140,7 +142,8 @@ def get_default_converter( return DefaultPolarsUnloadTypeConverter() return DefaultPolarsTypeConverter() - @property + @property # type: ignore[explicit-override] # python/mypy#15900 + @override def arraysize(self) -> int: """Get the number of rows to fetch per batch.""" return self._arraysize @@ -159,6 +162,7 @@ def arraysize(self, value: int) -> None: raise ProgrammingError("arraysize must be a positive integer value.") self._arraysize = value + @override def _collect_result_set( self, query_id: str, @@ -185,6 +189,7 @@ def _collect_result_set( **kwargs, ) + @override def execute( self, operation: str, diff --git a/pyathena/polars/converter.py b/pyathena/polars/converter.py index c6c83340d..3b3835288 100644 --- a/pyathena/polars/converter.py +++ b/pyathena/polars/converter.py @@ -20,6 +20,7 @@ _to_json, _to_time, ) +from pyathena.util import override _logger = logging.getLogger(__name__) @@ -93,6 +94,7 @@ def _dtypes(self) -> dict[str, Any]: } return self.__dtypes + @override def get_dtype(self, type_: str, precision: int = 0, scale: int = 0) -> Any: """Get the Polars data type for a given Athena type. @@ -110,6 +112,7 @@ def get_dtype(self, type_: str, precision: int = 0, scale: int = 0) -> Any: return pl.Decimal(precision=precision, scale=scale) return self._types.get(type_) + @override def convert(self, type_: str, value: str | None, type_hint: str | None = None) -> Any | None: converter = self.get(type_) return converter(value) @@ -134,6 +137,7 @@ def __init__(self) -> None: default=_to_default, ) + @override def convert(self, type_: str, value: str | None, type_hint: str | None = None) -> Any | None: converter = self.get(type_) return converter(value) diff --git a/pyathena/polars/cursor.py b/pyathena/polars/cursor.py index a92ed580a..233cdd7c7 100644 --- a/pyathena/polars/cursor.py +++ b/pyathena/polars/cursor.py @@ -19,6 +19,7 @@ ) from pyathena.polars.result_set import AthenaPolarsResultSet from pyathena.result_set import WithFetch +from pyathena.util import override if TYPE_CHECKING: import polars as pl @@ -128,6 +129,7 @@ def __init__( self._chunksize = chunksize @staticmethod + @override def get_default_converter( unload: bool = False, ) -> DefaultPolarsTypeConverter | DefaultPolarsUnloadTypeConverter | Any: @@ -143,6 +145,7 @@ def get_default_converter( return DefaultPolarsUnloadTypeConverter() return DefaultPolarsTypeConverter() + @override def execute( self, operation: str, diff --git a/pyathena/polars/result_set.py b/pyathena/polars/result_set.py index b7703c25c..4fc9e1014 100644 --- a/pyathena/polars/result_set.py +++ b/pyathena/polars/result_set.py @@ -23,7 +23,7 @@ from pyathena.model import AthenaQueryExecution from pyathena.polars.util import to_column_info from pyathena.result_set import AthenaResultSet -from pyathena.util import RetryConfig +from pyathena.util import RetryConfig, override if TYPE_CHECKING: import polars as pl @@ -86,6 +86,7 @@ def __init__( self._converters = converters self._column_names = column_names + @override def __next__(self) -> pl.DataFrame: """Get the next DataFrame chunk. @@ -101,6 +102,7 @@ def __next__(self) -> pl.DataFrame: self.close() raise + @override def __iter__(self) -> PolarsDataFrameIterator: """Return self as iterator.""" return self @@ -355,6 +357,7 @@ def _create_dataframe_iterator(self) -> PolarsDataFrameIterator: return PolarsDataFrameIterator(reader, self.converters, self._get_column_names()) + @override def fetchone( self, ) -> tuple[Any | None, ...] | dict[Any, Any | None] | None: @@ -684,6 +687,7 @@ def iter_chunks(self) -> PolarsDataFrameIterator: """ return self._df_iter + @override def close(self) -> None: """Close the result set and release resources.""" import polars as pl diff --git a/pyathena/result_set.py b/pyathena/result_set.py index 8bfe56922..dade34f42 100644 --- a/pyathena/result_set.py +++ b/pyathena/result_set.py @@ -13,7 +13,7 @@ from pyathena.converter import Converter, DefaultTypeConverter from pyathena.error import DataError, OperationalError, ProgrammingError from pyathena.model import AthenaQueryExecution -from pyathena.util import RetryConfig, parse_output_location, retry_api_call +from pyathena.util import RetryConfig, override, parse_output_location, retry_api_call if TYPE_CHECKING: from pyathena.connection import Connection @@ -429,6 +429,7 @@ def _pre_fetch(self) -> None: offset = 1 if rows and self._is_first_row_column_labels(rows) else 0 self._process_rows(rows, offset) + @override def fetchone( self, ) -> tuple[Any | None, ...] | dict[Any, Any | None] | None: @@ -441,6 +442,7 @@ def fetchone( self._rownumber += 1 return self._rows.popleft() + @override def fetchmany( self, size: int | None = None ) -> list[tuple[Any | None, ...] | dict[Any, Any | None]]: @@ -464,6 +466,7 @@ def fetchmany( break return rows + @override def fetchall( self, ) -> list[tuple[Any | None, ...] | dict[Any, Any | None]]: @@ -753,6 +756,7 @@ class AthenaDictResultSet(AthenaResultSet): # You can override this to use OrderedDict or other dict-like types. dict_type: type[Any] = dict + @override def _get_rows( self, offset: int, @@ -1137,6 +1141,7 @@ class WithFetch(WithResultSet, BaseCursor, CursorIterator): format-specific helpers. """ + @override def fetchone( self, ) -> tuple[Any | None, ...] | dict[Any, Any | None] | None: @@ -1153,6 +1158,7 @@ def fetchone( result_set = cast(AthenaResultSet, self.result_set) return result_set.fetchone() + @override def fetchmany( self, size: int | None = None ) -> list[tuple[Any | None, ...] | dict[Any, Any | None]]: @@ -1172,6 +1178,7 @@ def fetchmany( result_set = cast(AthenaResultSet, self.result_set) return result_set.fetchmany(size) + @override def fetchall( self, ) -> list[tuple[Any | None, ...] | dict[Any, Any | None]]: @@ -1188,6 +1195,7 @@ def fetchall( result_set = cast(AthenaResultSet, self.result_set) return result_set.fetchall() + @override def executemany( self, operation: str, diff --git a/pyathena/s3fs/async_cursor.py b/pyathena/s3fs/async_cursor.py index 8fd600656..67ed69638 100644 --- a/pyathena/s3fs/async_cursor.py +++ b/pyathena/s3fs/async_cursor.py @@ -12,6 +12,7 @@ from pyathena.options import ExecuteOptions from pyathena.s3fs.converter import DefaultS3FSTypeConverter from pyathena.s3fs.result_set import AthenaS3FSResultSet, CSVReaderType +from pyathena.util import override _logger = logging.getLogger(__name__) @@ -108,6 +109,7 @@ def __init__( self._csv_reader = csv_reader @staticmethod + @override def get_default_converter( unload: bool = False, ) -> DefaultS3FSTypeConverter: @@ -121,7 +123,8 @@ def get_default_converter( """ return DefaultS3FSTypeConverter() - @property + @property # type: ignore[explicit-override] # python/mypy#15900 + @override def arraysize(self) -> int: """Get the number of rows to fetch at a time.""" return self._arraysize @@ -140,6 +143,7 @@ def arraysize(self, value: int) -> None: raise ProgrammingError("arraysize must be a positive integer value.") self._arraysize = value + @override def _collect_result_set( self, query_id: str, @@ -172,6 +176,7 @@ def _collect_result_set( **kwargs, ) + @override def execute( self, operation: str, diff --git a/pyathena/s3fs/converter.py b/pyathena/s3fs/converter.py index a33c5e739..168029609 100644 --- a/pyathena/s3fs/converter.py +++ b/pyathena/s3fs/converter.py @@ -16,6 +16,7 @@ Converter, _to_default, ) +from pyathena.util import override if TYPE_CHECKING: from pyathena.converter import DefaultTypeConverter @@ -55,6 +56,7 @@ def __init__(self) -> None: ) self._default_type_converter: DefaultTypeConverter | None = None + @override def convert(self, type_: str, value: str | None, type_hint: str | None = None) -> Any | None: """Convert a string value to the appropriate Python type. diff --git a/pyathena/s3fs/cursor.py b/pyathena/s3fs/cursor.py index 6052b66e1..1cb04d74a 100644 --- a/pyathena/s3fs/cursor.py +++ b/pyathena/s3fs/cursor.py @@ -11,6 +11,7 @@ from pyathena.result_set import WithFetch from pyathena.s3fs.converter import DefaultS3FSTypeConverter from pyathena.s3fs.result_set import AthenaS3FSResultSet, CSVReaderType +from pyathena.util import override _logger = logging.getLogger(__name__) @@ -106,6 +107,7 @@ def __init__( self._csv_reader = csv_reader @staticmethod + @override def get_default_converter( unload: bool = False, ) -> DefaultS3FSTypeConverter: @@ -119,6 +121,7 @@ def get_default_converter( """ return DefaultS3FSTypeConverter() + @override def execute( self, operation: str, diff --git a/pyathena/s3fs/reader.py b/pyathena/s3fs/reader.py index daee58d7d..607cb9445 100644 --- a/pyathena/s3fs/reader.py +++ b/pyathena/s3fs/reader.py @@ -11,6 +11,8 @@ from collections.abc import Iterator from typing import Any +from pyathena.util import override + class DefaultCSVReader(Iterator[list[str]]): """CSV reader using Python's standard csv module. @@ -43,10 +45,12 @@ def __init__(self, file_obj: Any, delimiter: str = ",") -> None: self._file: Any | None = file_obj self._reader = csv.reader(file_obj, delimiter=delimiter) + @override def __iter__(self) -> DefaultCSVReader: """Iterate over rows in the CSV file.""" return self + @override def __next__(self) -> list[str]: """Read and parse the next line. @@ -114,10 +118,12 @@ def __init__(self, file_obj: Any, delimiter: str = ",") -> None: self._file: Any | None = file_obj self._delimiter = delimiter + @override def __iter__(self) -> AthenaCSVReader: """Iterate over rows in the CSV file.""" return self + @override def __next__(self) -> list[str | None]: """Read and parse the next line. diff --git a/pyathena/s3fs/result_set.py b/pyathena/s3fs/result_set.py index 16bc2d985..e0ee25bf3 100644 --- a/pyathena/s3fs/result_set.py +++ b/pyathena/s3fs/result_set.py @@ -19,7 +19,7 @@ from pyathena.model import AthenaQueryExecution from pyathena.result_set import AthenaResultSet from pyathena.s3fs.reader import AthenaCSVReader, DefaultCSVReader -from pyathena.util import RetryConfig, parse_output_location +from pyathena.util import RetryConfig, override, parse_output_location if TYPE_CHECKING: from pyathena.connection import Connection @@ -151,6 +151,7 @@ def _init_csv_reader(self) -> None: _logger.exception(f"Failed to open {path}.") raise OperationalError(*e.args) from e + @override def _fetch(self) -> None: """Fetch next batch of rows from CSV.""" if not self._csv_reader: @@ -203,6 +204,7 @@ def _fetch(self) -> None: self._rows.append(converted_row) rows_fetched += 1 + @override def fetchone( self, ) -> tuple[Any | None, ...] | dict[Any, Any | None] | None: @@ -220,6 +222,7 @@ def fetchone( self._rownumber += 1 return self._rows.popleft() + @override def close(self) -> None: """Close the result set and release resources.""" super().close() diff --git a/pyathena/spark/async_cursor.py b/pyathena/spark/async_cursor.py index fffd6c6a1..8d12fd56c 100644 --- a/pyathena/spark/async_cursor.py +++ b/pyathena/spark/async_cursor.py @@ -12,6 +12,7 @@ from pyathena.model import AthenaCalculationExecution from pyathena.spark.common import SparkBaseCursor +from pyathena.util import override if TYPE_CHECKING: from pyathena.model import AthenaQueryExecution @@ -117,6 +118,7 @@ def __init__( **kwargs, ) + @override def close(self, wait: bool = False) -> None: """Close the cursor, then shut down the executor. @@ -161,6 +163,7 @@ def poll(self, query_id: str) -> "Future[AthenaCalculationExecution]": "Future[AthenaCalculationExecution]", self._executor.submit(self._poll, query_id) ) + @override def execute( self, operation: str, diff --git a/pyathena/spark/common.py b/pyathena/spark/common.py index 016373139..d4871ef9e 100644 --- a/pyathena/spark/common.py +++ b/pyathena/spark/common.py @@ -27,7 +27,7 @@ AthenaQueryExecution, AthenaSessionStatus, ) -from pyathena.util import parse_output_location, retry_api_call +from pyathena.util import override, parse_output_location, retry_api_call _logger = logging.getLogger(__name__) @@ -309,6 +309,7 @@ def _terminate_session_by_id(self, session_id: str) -> None: _logger.exception(f"Failed to terminate session: {session_id}.") raise OperationalError(*e.args) from e + @override def _poll_until_terminal( self, query_id: str ) -> AthenaQueryExecution | AthenaCalculationExecution: @@ -334,6 +335,7 @@ def _poll_until_terminal( return self._get_calculation_execution(query_id) time.sleep(self._poll_interval) + @override def _poll(self, query_id: str) -> AthenaQueryExecution | AthenaCalculationExecution: """Wait for a calculation execution to reach a terminal state. @@ -493,6 +495,7 @@ def _wait_for_calculation_start(future: Future[str]) -> str: wait((future,), timeout=_INTERRUPT_CHECK_INTERVAL) return future.result() + @override def _cancel(self, query_id: str) -> None: """Stop a calculation execution with ``StopCalculationExecution``. @@ -514,6 +517,7 @@ def _cancel(self, query_id: str) -> None: _logger.exception("Failed to cancel calculation.") raise OperationalError(*e.args) from e + @override def close(self) -> None: """Close the cursor, terminating its Spark session if configured to. @@ -529,6 +533,7 @@ def close(self) -> None: # Terminated; later calls do nothing. self._terminate_session_on_close = False + @override def executemany( self, operation: str, diff --git a/pyathena/spark/cursor.py b/pyathena/spark/cursor.py index 7ca174421..151c24858 100644 --- a/pyathena/spark/cursor.py +++ b/pyathena/spark/cursor.py @@ -13,6 +13,7 @@ from pyathena import OperationalError, ProgrammingError from pyathena.model import AthenaCalculationExecution, AthenaCalculationExecutionStatus from pyathena.spark.common import SparkBaseCursor, WithCalculationExecution +from pyathena.util import override _logger = logging.getLogger(__name__) @@ -64,6 +65,7 @@ class SparkCursor(SparkBaseCursor, WithCalculationExecution): """ @property + @override def calculation_execution(self) -> AthenaCalculationExecution | None: return self._calculation_execution @@ -96,6 +98,7 @@ def get_std_error(self) -> str | None: return None return self._read_s3_file_as_text(self._calculation_execution.std_error_s3_uri) + @override def execute( self, operation: str, diff --git a/pyathena/sqlalchemy/array.py b/pyathena/sqlalchemy/array.py index 7f4046ff1..3a91a1b51 100644 --- a/pyathena/sqlalchemy/array.py +++ b/pyathena/sqlalchemy/array.py @@ -19,6 +19,7 @@ from pyathena.sqlalchemy.map import AthenaMap from pyathena.sqlalchemy.struct import AthenaStruct from pyathena.sqlalchemy.temporal import AthenaDate, AthenaTimestamp +from pyathena.util import override # SQLAlchemy 2.0.0's ARRAY comparator is not generic at runtime. if TYPE_CHECKING: @@ -56,6 +57,7 @@ class AthenaArray(sqltypes.ARRAY[Any]): class Comparator(_ArrayComparatorBase): """Build array indexing expressions with inclusive SQL slice bounds.""" + @override def _setup_getitem(self, index): if isinstance(index, slice): if index.step is not None and (type(index.step) is not int or index.step != 1): @@ -92,19 +94,23 @@ def __init__( else: super().__init__(item_type or sqltypes.String(), as_tuple, dimensions, zero_indexes) + @override def bind_expression(self, bindvalue): """Cast a bound ARRAY value to its declared Athena element type.""" # The cast also gives empty arrays and NULL-only arrays their element type. return cast(bindvalue, self)._annotate({"_pyathena_array_bind": True}) + @override def bind_processor(self, dialect): """Return a processor that marks native ARRAY, MAP, and ROW parameters.""" return _ArrayValueProcessor(self, dialect).bind + @override def literal_processor(self, dialect): """Return a processor that renders typed Athena array literals.""" return _ArrayValueProcessor(self, dialect).literal + @override def column_expression(self, colexpr): """Project the outer ARRAY result as JSON while retaining its Python type.""" return ( @@ -113,6 +119,7 @@ def column_expression(self, colexpr): else _ArrayJSONProjection(colexpr, self) ) + @override def result_processor(self, dialect, coltype): """Return a processor that restores the declared Python element types.""" return _ArrayValueProcessor(self, dialect).result @@ -124,11 +131,13 @@ class _ArraySliceStepType(types.TypeDecorator[int]): impl = types.Integer cache_ok = True + @override def process_bind_param(self, value, dialect): if type(value) is not int or value != 1: raise ValueError("Athena ARRAY slices support only step=None or step=1") return value + @override def process_literal_param(self, value, dialect): return self.process_bind_param(value, dialect) @@ -418,6 +427,7 @@ def __init__(self, item_type): super().__init__() self.item_type = item_type + @override def bind_processor(self, dialect): processor = _ArrayValueProcessor(self.item_type, dialect) @@ -429,9 +439,11 @@ def process(value): return process + @override def literal_processor(self, dialect): return _ArrayValueProcessor(self.item_type, dialect).literal + @override def bind_expression(self, bindvalue): expression = self.item_type.bind_expression(bindvalue) return bindvalue if expression is None else expression @@ -443,11 +455,13 @@ class _ArrayWriteIndexType(types.TypeDecorator[int]): impl = types.Integer cache_ok = True + @override def process_bind_param(self, value, dialect): if type(value) is not int: raise ValueError("ARRAY write indices must be non-NULL integers") return value + @override def process_literal_param(self, value, dialect): return self.process_bind_param(value, dialect) @@ -478,6 +492,7 @@ def __init__(self, column, path, value, value_type): ) @property + @override def _from_objects(self): return self.column._from_objects + self.value._from_objects diff --git a/pyathena/sqlalchemy/arrow.py b/pyathena/sqlalchemy/arrow.py index d5e8a05f7..ea31b0733 100644 --- a/pyathena/sqlalchemy/arrow.py +++ b/pyathena/sqlalchemy/arrow.py @@ -8,7 +8,7 @@ from typing import TYPE_CHECKING from pyathena.sqlalchemy.base import AthenaDialect -from pyathena.util import strtobool +from pyathena.util import override, strtobool if TYPE_CHECKING: from types import ModuleType @@ -47,6 +47,7 @@ class AthenaArrowDialect(AthenaDialect): driver = "arrow" supports_statement_cache = True + @override def create_connect_args(self, url): from pyathena.arrow.cursor import ArrowCursor @@ -60,5 +61,6 @@ def create_connect_args(self, url): return [[], opts] @classmethod + @override def import_dbapi(cls) -> "ModuleType": return super().import_dbapi() diff --git a/pyathena/sqlalchemy/base.py b/pyathena/sqlalchemy/base.py index 5295d08e8..876c0098c 100644 --- a/pyathena/sqlalchemy/base.py +++ b/pyathena/sqlalchemy/base.py @@ -46,6 +46,7 @@ RetryConfig, _get_error_code, _without_retries, + override, strtobool, ) @@ -236,10 +237,12 @@ def __init__(self, json_deserializer=None, json_serializer=None, **kwargs): ) @classmethod + @override def import_dbapi(cls) -> ModuleType: return pyathena @classmethod + @override def dbapi(cls) -> ModuleType: # type: ignore[override] return pyathena @@ -248,6 +251,7 @@ def _raw_connection(self, connection: Engine | Connection) -> PoolProxiedConnect return connection.raw_connection() return connection.connection + @override def create_connect_args(self, url: URL) -> tuple[tuple[str], MutableMapping[str, Any]]: # Connection string format: # awsathena+rest:// @@ -588,10 +592,12 @@ def _get_tables(self, connection, schema: str | None = None, **kw): info_cache.setdefault(("pyathena_table_metadata", catalog, schema, name), metadata) return tables + @override def get_schema_names(self, connection, **kw): schemas = self._get_schemas(connection, **kw) return [s.name for s in schemas] + @override def get_table_names(self, connection: Connection, schema: str | None = None, **kw): # Tables created by Athena are always classified as `EXTERNAL_TABLE`, # but Athena can also query tables classified as `MANAGED_TABLE`, `EXTERNAL`, or `customer`. @@ -606,10 +612,12 @@ def get_table_names(self, connection: Connection, schema: str | None = None, **k if t.table_type in ["EXTERNAL_TABLE", "MANAGED_TABLE", "EXTERNAL", "customer"] ] + @override def get_view_names(self, connection: Connection, schema: str | None = None, **kw): tables = self._get_tables(connection, schema, **kw) return [t.name for t in tables if t.table_type == "VIRTUAL_VIEW"] + @override def get_table_comment( self, connection: Connection, table_name: str, schema: str | None = None, **kw ): @@ -617,6 +625,7 @@ def get_table_comment( # An empty comment is no comment here too; the DDL compiler skips one. return {"text": metadata.comment or None} + @override def get_table_options( self, connection: Connection, table_name: str, schema: str | None = None, **kw ): @@ -631,6 +640,7 @@ def get_table_options( "awsathena_tblproperties": _HashableDict(metadata.table_properties), } + @override @reflection.cache def has_table(self, connection: Connection, table_name: str, schema: str | None = None, **kw): try: @@ -638,6 +648,7 @@ def has_table(self, connection: Connection, table_name: str, schema: str | None except exc.NoSuchTableError: return False + @override @reflection.cache def get_view_definition( self, connection: Connection, view_name: str, schema: str | None = None, **kw @@ -665,6 +676,7 @@ def get_view_definition( # empty values, which are part of the definition. return "\n".join(row[0] or "" for row in rows) + @override @reflection.cache def get_columns(self, connection: Connection, table_name: str, schema: str | None = None, **kw): return self._get_columns(connection, table_name, schema=schema, **kw) @@ -728,24 +740,28 @@ def _get_column_type(self, type_: str, _nested: bool = False): return col_type(*args) + @override def get_foreign_keys( self, connection: Connection, table_name: str, schema: str | None = None, **kw ) -> list[ReflectedForeignKeyConstraint]: # Athena has no support for foreign keys. return [] # pragma: no cover + @override def get_pk_constraint( self, connection: Connection, table_name: str, schema: str | None = None, **kw ) -> ReflectedPrimaryKeyConstraint: # Athena has no support for primary keys. return {"name": None, "constrained_columns": []} # pragma: no cover + @override def get_indexes( self, connection: Connection, table_name: str, schema: str | None = None, **kw ) -> list[ReflectedIndex]: # Athena has no support for indexes. return [] # pragma: no cover + @override def do_execute(self, cursor, statement, parameters, context=None): """Execute a statement with the DB API cursor. @@ -781,6 +797,7 @@ def do_execute(self, cursor, statement, parameters, context=None): else: context._rowcount = total + count if total >= 0 and count >= 0 else -1 + @override def do_rollback(self, dbapi_connection: PoolProxiedConnection) -> None: # No transactions for Athena pass # pragma: no cover diff --git a/pyathena/sqlalchemy/compiler.py b/pyathena/sqlalchemy/compiler.py index 8687224e0..d40cff29a 100644 --- a/pyathena/sqlalchemy/compiler.py +++ b/pyathena/sqlalchemy/compiler.py @@ -47,6 +47,7 @@ AthenaTimestamp, ) from pyathena.sqlalchemy.util import _split_type_arguments +from pyathena.util import override if TYPE_CHECKING: from sqlalchemy import ( @@ -98,21 +99,27 @@ class AthenaTypeCompiler(GenericTypeCompiler): https://docs.aws.amazon.com/athena/latest/ug/data-types.html """ + @override def visit_FLOAT(self, type_: types.Float[Any], **kw: Any) -> str: return self.visit_REAL(type_, **kw) # type: ignore[arg-type] + @override def visit_REAL(self, type_: types.REAL[Any], **kw: Any) -> str: return "FLOAT" + @override def visit_DOUBLE(self, type_, **kw) -> str: return "DOUBLE" + @override def visit_DOUBLE_PRECISION(self, type_, **kw) -> str: return "DOUBLE" + @override def visit_NUMERIC(self, type_: types.Numeric[Any], **kw: Any) -> str: return self.visit_DECIMAL(type_, **kw) # type: ignore[arg-type] + @override def visit_DECIMAL(self, type_: types.DECIMAL[Any], **kw: Any) -> str: if type_.precision is None: return "DECIMAL" @@ -123,82 +130,105 @@ def visit_DECIMAL(self, type_: types.DECIMAL[Any], **kw: Any) -> str: def visit_TINYINT(self, type_: types.Integer, **kw: Any) -> str: return "TINYINT" + @override def visit_INTEGER(self, type_: types.Integer, **kw: Any) -> str: return "INT" if kw.get("_athena_hive_ddl") else "INTEGER" + @override def visit_SMALLINT(self, type_: types.SmallInteger, **kw: Any) -> str: return "SMALLINT" + @override def visit_BIGINT(self, type_: types.BigInteger, **kw: Any) -> str: return "BIGINT" + @override def visit_TIMESTAMP(self, type_: types.TIMESTAMP, **kw: Any) -> str: return "TIMESTAMP" + @override def visit_DATETIME(self, type_: types.DateTime, **kw: Any) -> str: return self.visit_TIMESTAMP(type_, **kw) # type: ignore[arg-type] + @override def visit_DATE(self, type_: types.Date, **kw: Any) -> str: return "DATE" + @override def visit_TIME(self, type_: types.Time, **kw: Any) -> str: raise exc.CompileError(f"Data type `{type_}` is not supported") + @override def visit_CLOB(self, type_: types.CLOB, **kw: Any) -> str: return self.visit_BINARY(type_, **kw) # type: ignore[arg-type] + @override def visit_NCLOB(self, type_: types.Text, **kw: Any) -> str: return self.visit_BINARY(type_, **kw) # type: ignore[arg-type] + @override def visit_CHAR(self, type_: types.CHAR, **kw: Any) -> str: if type_.length: return self._render_string_type("CHAR", type_.length, type_.collation) return "STRING" + @override def visit_NCHAR(self, type_: types.NCHAR, **kw: Any) -> str: return self.visit_CHAR(type_, **kw) # type: ignore[arg-type] + @override def visit_VARCHAR(self, type_: types.String, **kw: Any) -> str: if type_.length: return self._render_string_type("VARCHAR", type_.length, type_.collation) return "STRING" + @override def visit_NVARCHAR(self, type_: types.NVARCHAR, **kw: Any) -> str: return self.visit_VARCHAR(type_, **kw) # type: ignore[arg-type] + @override def visit_TEXT(self, type_: types.Text, **kw: Any) -> str: return "STRING" + @override def visit_BLOB(self, type_: types.LargeBinary, **kw: Any) -> str: return self.visit_BINARY(type_, **kw) # type: ignore[arg-type] + @override def visit_BINARY(self, type_: types.BINARY, **kw: Any) -> str: return "BINARY" + @override def visit_VARBINARY(self, type_: types.VARBINARY, **kw: Any) -> str: return self.visit_BINARY(type_, **kw) # type: ignore[arg-type] + @override def visit_BOOLEAN(self, type_: types.Boolean, **kw: Any) -> str: return "BOOLEAN" def visit_JSON(self, type_: types.JSON, **kw: Any) -> str: return "JSON" + @override def visit_string(self, type_, **kw): return "STRING" + @override def visit_unicode(self, type_, **kw): return "STRING" + @override def visit_unicode_text(self, type_, **kw): return "STRING" + @override def visit_null(self, type_, **kw): return "NULL" def visit_tinyint(self, type_, **kw): return self.visit_TINYINT(type_, **kw) + @override def visit_enum(self, type_, **kw): return self.visit_string(type_, **kw) @@ -300,6 +330,7 @@ def _original_froms(elements): element = element._is_clone_of yield element + @override def visit_update(self, update_stmt, visiting_cte=None, **kw): """Rewrite partial array assignments into one native Athena UPDATE.""" return super().visit_update( @@ -323,6 +354,7 @@ def _array_lambda_name(self): self._array_lambda_index = index + 1 return f"_pyathena_element_{index}" + @override def visit_binary( self, binary, @@ -435,6 +467,7 @@ def _array_slice_step(self, sql, step, array_type, **kw): ) return f"IF({step_sql} = 1, {sql}, slice({empty}, {failure}, 0))" + @override def translate_select_structure(self, select_stmt, **kw): """Keep DISTINCT and ordering on native arrays before result serialization.""" if ( @@ -447,6 +480,7 @@ def translate_select_structure(self, select_stmt, **kw): return self._array_result_select(select_stmt) return select_stmt + @override def visit_compound_select(self, cs, asfrom=False, compound_index=None, **kw): if ( not self.stack @@ -624,6 +658,7 @@ def visit_filter_func(self, fn: Function[Any], **kw: Any) -> str: return f"filter({array_sql}, {lambda_sql})" + @override def visit_truediv_binary(self, binary, operator, **kw): """Render true division with explicit Athena numeric coercions.""" left_type = binary.left.type @@ -656,6 +691,7 @@ def visit_truediv_binary(self, binary, operator, **kw): return super().visit_truediv_binary(binary, operator, **kw) + @override def visit_cast(self, cast: Cast[Any], **kwargs): """Render a CAST with the Athena DML name of the target type. @@ -828,6 +864,7 @@ def _array_json(self, value, type_, depth=0): return f"CAST(to_hex({value}) AS JSON)" return f"CAST(CAST({value} AS VARCHAR) AS JSON)" + @override def limit_clause(self, select: GenerativeSelect, **kw): text = [] if select._offset_clause is not None: @@ -836,9 +873,11 @@ def limit_clause(self, select: GenerativeSelect, **kw): text.append(" LIMIT " + self.process(select._limit_clause, **kw)) return "\n".join(text) + @override def get_from_hint_text(self, table, text): return text + @override def format_from_hint_text(self, sqltext, table, hint, iscrud): hint_upper = hint.upper() if ( @@ -899,7 +938,8 @@ class AthenaDDLCompiler(DDLCompiler): https://docs.aws.amazon.com/athena/latest/ug/create-table.html """ - @property + @property # type: ignore[explicit-override] # python/mypy#15900 + @override def preparer(self) -> IdentifierPreparer: return self._preparer @@ -1174,6 +1214,7 @@ def _get_table_properties_specification( text.append(")") return "\n".join(text) + @override def get_column_specification(self, column: Column[Any], **kwargs) -> str: if type(column.type) in [types.Integer, types.INTEGER, types.INT]: # https://docs.aws.amazon.com/athena/latest/ug/create-table.html @@ -1188,18 +1229,23 @@ def get_column_specification(self, column: Column[Any], **kwargs) -> str: text.append(f"{self._get_comment_specification(column.comment)}") return " ".join(text) + @override def visit_check_constraint(self, constraint: CheckConstraint, **kw: Any) -> str: return "" + @override def visit_column_check_constraint(self, constraint: CheckConstraint, **kw: Any) -> str: return "" + @override def visit_foreign_key_constraint(self, constraint: ForeignKeyConstraint, **kw: Any) -> str: return "" + @override def visit_primary_key_constraint(self, constraint: PrimaryKeyConstraint, **kw: Any) -> str: return "" + @override def visit_unique_constraint(self, constraint: UniqueConstraint, **kw: Any) -> str: return "" @@ -1290,6 +1336,7 @@ def _prepared_columns( ) from e return columns, partitions, buckets + @override def visit_create_table(self, create: CreateTable, **kwargs) -> str: table = create.element dialect_opts = table.dialect_options["awsathena"] @@ -1330,6 +1377,7 @@ def visit_create_table(self, create: CreateTable, **kwargs) -> str: text.append(f"{self.post_create_table(table)}\n") return "\n".join(text) + @override def post_create_table(self, table: Table) -> str: dialect_opts: _DialectArgDict = table.dialect_options["awsathena"] dialect = cast("AthenaDialect", self.dialect) diff --git a/pyathena/sqlalchemy/map.py b/pyathena/sqlalchemy/map.py index 61d292b18..e2e828b00 100644 --- a/pyathena/sqlalchemy/map.py +++ b/pyathena/sqlalchemy/map.py @@ -14,6 +14,8 @@ from sqlalchemy.sql import sqltypes from sqlalchemy.sql.type_api import TypeEngine +from pyathena.util import override + class AthenaMap(TypeEngine[dict[str, Any]]): """SQLAlchemy type for Athena MAP complex type. @@ -58,6 +60,7 @@ def __init__(self, key_type: Any = None, value_type: Any = None) -> None: self.value_type = value_type() @property + @override def python_type(self) -> type: return dict diff --git a/pyathena/sqlalchemy/pandas.py b/pyathena/sqlalchemy/pandas.py index a697cbfc9..0d9759344 100644 --- a/pyathena/sqlalchemy/pandas.py +++ b/pyathena/sqlalchemy/pandas.py @@ -8,7 +8,7 @@ from typing import TYPE_CHECKING from pyathena.sqlalchemy.base import AthenaDialect -from pyathena.util import strtobool +from pyathena.util import override, strtobool if TYPE_CHECKING: from types import ModuleType @@ -48,6 +48,7 @@ class AthenaPandasDialect(AthenaDialect): driver = "pandas" supports_statement_cache = True + @override def create_connect_args(self, url): from pyathena.pandas.cursor import PandasCursor @@ -65,5 +66,6 @@ def create_connect_args(self, url): return [[], opts] @classmethod + @override def import_dbapi(cls) -> "ModuleType": return super().import_dbapi() diff --git a/pyathena/sqlalchemy/polars.py b/pyathena/sqlalchemy/polars.py index 84dc5c99c..a715bc8de 100644 --- a/pyathena/sqlalchemy/polars.py +++ b/pyathena/sqlalchemy/polars.py @@ -8,7 +8,7 @@ from typing import TYPE_CHECKING from pyathena.sqlalchemy.base import AthenaDialect -from pyathena.util import strtobool +from pyathena.util import override, strtobool if TYPE_CHECKING: from types import ModuleType @@ -47,6 +47,7 @@ class AthenaPolarsDialect(AthenaDialect): driver = "polars" supports_statement_cache = True + @override def create_connect_args(self, url): from pyathena.polars.cursor import PolarsCursor @@ -60,5 +61,6 @@ def create_connect_args(self, url): return [[], opts] @classmethod + @override def import_dbapi(cls) -> "ModuleType": return super().import_dbapi() diff --git a/pyathena/sqlalchemy/requirements.py b/pyathena/sqlalchemy/requirements.py index 901370daf..048e92294 100644 --- a/pyathena/sqlalchemy/requirements.py +++ b/pyathena/sqlalchemy/requirements.py @@ -8,105 +8,130 @@ from sqlalchemy.testing import exclusions from sqlalchemy.testing.requirements import SuiteRequirements +from pyathena.util import override + supported = exclusions.open unsupported = exclusions.closed class Requirements(SuiteRequirements): @property + @override def comment_reflection(self): # The upstream requirement also needs COMMENT ON TABLE. Athena only # reflects table comments from Hive tables, not Iceberg tables. return unsupported() @property + @override def reflect_table_options(self): return supported() @property + @override def array_type(self): return supported() @property + @override def uuid_data_type(self): return unsupported() @property + @override def foreign_keys(self): return unsupported() @property + @override def on_update_cascade(self): return unsupported() @property + @override def self_referential_foreign_keys(self): return unsupported() @property + @override def foreign_key_ddl(self): return unsupported() @property + @override def autoincrement_insert(self): return unsupported() @property + @override def primary_key_constraint_reflection(self): return unsupported() @property + @override def foreign_key_constraint_reflection(self): return unsupported() @property + @override def temp_table_reflection(self): return unsupported() @property + @override def temporary_tables(self): return unsupported() @property + @override def index_reflection(self): return unsupported() @property + @override def indexes_with_ascdesc(self): return unsupported() @property + @override def reflect_indexes_with_ascdesc(self): return unsupported() @property + @override def unique_constraint_reflection(self): return unsupported() @property + @override def duplicate_key_raises_integrity_error(self): return unsupported() @property + @override def update_where_target_in_subquery(self): # Verified with Iceberg tables on Athena engine version 3. return supported() @property + @override def recursive_fk_cascade(self): return unsupported() @property + @override def datetime_literals(self): return supported() @property + @override def timestamp_microseconds(self): # Iceberg tables store microseconds; Hive tables store milliseconds. # The compliance suite creates Iceberg tables. return supported() @property + @override def precision_generic_float_type(self): return exclusions.skip_if( lambda _: True, @@ -115,72 +140,88 @@ def precision_generic_float_type(self): ) @property + @override def precision_numerics_many_significant_digits(self): return supported() @property + @override def precision_numerics_retains_significant_digits(self): return supported() @property + @override def window_functions(self): return supported() @property + @override def ctes(self): # Recursive CTEs require Athena engine version 3 and have a maximum depth of 10. return supported() @property + @override def ctes_with_values(self): return supported() @property + @override def ctes_with_update_delete(self): return exclusions.skip_if( lambda _: True, "Athena does not support WITH preceding UPDATE or DELETE." ) @property + @override def ctes_on_dml(self): return exclusions.skip_if( lambda _: True, "Athena does not support INSERT, UPDATE, or DELETE inside a CTE." ) @property + @override def update_from(self): return exclusions.skip_if(lambda _: True, "Athena does not support UPDATE ... FROM.") @property + @override def delete_from(self): return exclusions.skip_if( lambda _: True, "Athena does not support DELETE ... USING or multi-table DELETE." ) @property + @override def views(self): return supported() @property + @override def schemas(self): return supported() @property + @override def implicit_default_schema(self): return supported() @property + @override def datetime_historic(self): return supported() @property + @override def date_historic(self): return supported() @property + @override def precision_numerics_enotation_small(self): return supported() @property + @override def order_by_label_with_expression(self): return supported() diff --git a/pyathena/sqlalchemy/rest.py b/pyathena/sqlalchemy/rest.py index dd4d94079..eeea5e12f 100644 --- a/pyathena/sqlalchemy/rest.py +++ b/pyathena/sqlalchemy/rest.py @@ -8,6 +8,7 @@ from typing import TYPE_CHECKING from pyathena.sqlalchemy.base import AthenaDialect +from pyathena.util import override if TYPE_CHECKING: from types import ModuleType @@ -43,5 +44,6 @@ class AthenaRestDialect(AthenaDialect): supports_statement_cache = True @classmethod + @override def import_dbapi(cls) -> "ModuleType": return super().import_dbapi() diff --git a/pyathena/sqlalchemy/s3fs.py b/pyathena/sqlalchemy/s3fs.py index a9df5d78f..293855f5d 100644 --- a/pyathena/sqlalchemy/s3fs.py +++ b/pyathena/sqlalchemy/s3fs.py @@ -8,6 +8,7 @@ from typing import TYPE_CHECKING from pyathena.sqlalchemy.base import AthenaDialect +from pyathena.util import override if TYPE_CHECKING: from types import ModuleType @@ -30,6 +31,7 @@ class AthenaS3FSDialect(AthenaDialect): driver = "s3fs" supports_statement_cache = True + @override def create_connect_args(self, url): from pyathena.s3fs.cursor import S3FSCursor @@ -38,5 +40,6 @@ def create_connect_args(self, url): return [[], opts] @classmethod + @override def import_dbapi(cls) -> "ModuleType": return super().import_dbapi() diff --git a/pyathena/sqlalchemy/struct.py b/pyathena/sqlalchemy/struct.py index 63ef0914f..865be0830 100644 --- a/pyathena/sqlalchemy/struct.py +++ b/pyathena/sqlalchemy/struct.py @@ -14,6 +14,8 @@ from sqlalchemy.sql import sqltypes from sqlalchemy.sql.type_api import TypeEngine +from pyathena.util import override + class AthenaStruct(TypeEngine[dict[str, Any]]): """SQLAlchemy type for Athena STRUCT/ROW complex type. @@ -66,6 +68,7 @@ def __getitem__(self, key: str) -> TypeEngine[Any]: return self.fields[key] @property + @override def _static_cache_key(self): return ( type(self), @@ -73,6 +76,7 @@ def _static_cache_key(self): ) @property + @override def python_type(self) -> type: return dict diff --git a/pyathena/sqlalchemy/temporal.py b/pyathena/sqlalchemy/temporal.py index 7e5c00f0d..2ab6db8d3 100644 --- a/pyathena/sqlalchemy/temporal.py +++ b/pyathena/sqlalchemy/temporal.py @@ -11,6 +11,7 @@ from sqlalchemy.sql.type_api import TypeEngine from pyathena.formatter import _date_literal, _escape_trino, _timestamp_literal +from pyathena.util import override if TYPE_CHECKING: from sqlalchemy import Dialect @@ -61,6 +62,7 @@ def __init__(self, precision: int | None = None) -> None: self.precision = precision @property + @override def python_type(self) -> type[datetime]: """The Python type of TIMESTAMP values. @@ -69,6 +71,7 @@ def python_type(self) -> type[datetime]: """ return datetime + @override def bind_processor(self, dialect: Dialect) -> _BindProcessorType[datetime] | None: """Return a processor truncating bound datetimes to the precision. @@ -89,6 +92,7 @@ def process(value: datetime | Any | None) -> datetime | Any | None: return process + @override def coerce_compared_value(self, op: OperatorType | None, value: Any) -> TypeEngine[Any]: """Keep this type for a datetime compared with a column of it. @@ -124,6 +128,7 @@ def process( return _timestamp_literal(value, precision) return f"TIMESTAMP {quote(str(value))}" + @override def literal_processor(self, dialect: Dialect) -> _LiteralProcessorType[datetime] | None: """Return the literal renderer for the dialect. @@ -155,6 +160,7 @@ class AthenaDate(TypeEngine[date]): __visit_name__ = "DATE" @property + @override def python_type(self) -> type[date]: """The Python type of DATE values. @@ -180,6 +186,7 @@ def process(value: date | Any, quote: Callable[[str], str] = _escape_trino) -> s return _date_literal(value) return f"DATE {quote(str(value))}" + @override def literal_processor(self, dialect: Dialect) -> _LiteralProcessorType[date] | None: """Return the literal renderer for the dialect. diff --git a/pyathena/sqlalchemy/types.py b/pyathena/sqlalchemy/types.py index 226e5b62d..07c8e69d8 100644 --- a/pyathena/sqlalchemy/types.py +++ b/pyathena/sqlalchemy/types.py @@ -15,6 +15,7 @@ from pyathena.sqlalchemy.map import MAP, AthenaMap from pyathena.sqlalchemy.struct import STRUCT, AthenaStruct from pyathena.sqlalchemy.temporal import AthenaDate, AthenaTimestamp +from pyathena.util import override if TYPE_CHECKING: from sqlalchemy import Dialect @@ -38,6 +39,7 @@ class AthenaBinary(types.LargeBinary): """SQLAlchemy binary type with Athena hexadecimal literals.""" + @override def literal_processor(self, dialect: Dialect) -> _LiteralProcessorType[bytes]: def process(value: bytes) -> str: return f"X'{value.hex()}'" diff --git a/pyathena/sqlalchemy/util.py b/pyathena/sqlalchemy/util.py index 4206ef772..467c2e141 100644 --- a/pyathena/sqlalchemy/util.py +++ b/pyathena/sqlalchemy/util.py @@ -7,6 +7,8 @@ """Utility classes for PyAthena SQLAlchemy dialect.""" +from pyathena.util import override + def _split_type_arguments(value: str) -> list[str]: """Split type arguments without splitting nested types or quoted field names.""" @@ -48,5 +50,6 @@ class _HashableDict(dict): # type: ignore[type-arg] making them hashable through tuple conversion. """ - def __hash__(self): # type: ignore[override] + @override + def __hash__(self): return hash(tuple(sorted(self.items()))) diff --git a/pyproject.toml b/pyproject.toml index 360573dae..414ff78ff 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -188,6 +188,8 @@ warn_no_return = true warn_return_any = true warn_unreachable = true warn_unused_configs = true +# Overrides are marked with pyathena.util.override (PEP 698). +enable_error_code = ["explicit-override"] exclude = [ "benchmarks.*", "tests.*",