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.*",