diff --git a/pyathena/aio/common.py b/pyathena/aio/common.py index 56563cca6..11f8c4c26 100644 --- a/pyathena/aio/common.py +++ b/pyathena/aio/common.py @@ -8,7 +8,7 @@ from botocore.exceptions import BotoCoreError, ClientError from pyathena.aio.util import async_retry_api_call -from pyathena.common import BaseCursor +from pyathena.common import BaseCursor, CursorIterator from pyathena.error import DatabaseError, OperationalError, ProgrammingError from pyathena.glue import GlueMetadataClient from pyathena.model import AthenaDatabase, AthenaQueryExecution, AthenaTableMetadata @@ -617,13 +617,13 @@ async def athena_request( ) -class WithAsyncFetch(AioBaseCursor, WithResultSet): +class WithAsyncFetch(WithResultSet, AioBaseCursor, CursorIterator): """Base class of the asyncio SQL cursors. - Overrides ``executemany`` and ``cancel`` of ``WithResultSet`` with async - versions and adds async iteration and the async context manager protocol. - Synchronous iteration raises ``TypeError``. Subclasses override the fetch - methods with async versions. + Combines ``WithResultSet`` with ``AioBaseCursor`` and ``CursorIterator``, + and provides async ``executemany`` and ``cancel``, async iteration, and the + async context manager protocol. Synchronous iteration raises + ``TypeError``. Subclasses implement the fetch methods as coroutines. Subclasses override ``execute()`` and optionally ``__init__`` and format-specific helpers. diff --git a/pyathena/arrow/cursor.py b/pyathena/arrow/cursor.py index 6b2dabda6..5ad7cd603 100644 --- a/pyathena/arrow/cursor.py +++ b/pyathena/arrow/cursor.py @@ -13,7 +13,7 @@ from pyathena.error import OperationalError, ProgrammingError from pyathena.model import AthenaQueryExecution from pyathena.options import ExecuteOptions -from pyathena.result_set import WithResultSet +from pyathena.result_set import WithFetch if TYPE_CHECKING: import polars as pl @@ -22,7 +22,7 @@ _logger = logging.getLogger(__name__) -class ArrowCursor(WithResultSet): +class ArrowCursor(WithFetch): """Cursor for handling Apache Arrow Table results from Athena queries. This cursor returns query results as Apache Arrow Tables, which provide diff --git a/pyathena/cursor.py b/pyathena/cursor.py index cc365ce00..eb302b306 100644 --- a/pyathena/cursor.py +++ b/pyathena/cursor.py @@ -8,12 +8,12 @@ from pyathena.error import OperationalError, ProgrammingError from pyathena.model import AthenaQueryExecution from pyathena.options import ExecuteOptions -from pyathena.result_set import AthenaDictResultSet, AthenaResultSet, WithResultSet +from pyathena.result_set import AthenaDictResultSet, AthenaResultSet, WithFetch _logger = logging.getLogger(__name__) -class Cursor(WithResultSet): +class Cursor(WithFetch): """A DB API 2.0 compliant cursor for executing SQL queries on Amazon Athena. The Cursor class provides methods for executing SQL queries against Amazon Athena diff --git a/pyathena/pandas/cursor.py b/pyathena/pandas/cursor.py index 38b66c98a..883172c05 100644 --- a/pyathena/pandas/cursor.py +++ b/pyathena/pandas/cursor.py @@ -18,7 +18,7 @@ DefaultPandasUnloadTypeConverter, ) from pyathena.pandas.result_set import AthenaPandasResultSet, PandasDataFrameIterator -from pyathena.result_set import WithResultSet +from pyathena.result_set import WithFetch if TYPE_CHECKING: from pandas import DataFrame @@ -26,7 +26,7 @@ _logger = logging.getLogger(__name__) -class PandasCursor(WithResultSet): +class PandasCursor(WithFetch): """Cursor for handling pandas DataFrame results from Athena queries. This cursor returns query results as pandas DataFrames with memory-efficient diff --git a/pyathena/polars/cursor.py b/pyathena/polars/cursor.py index 4bd8c78e8..a92ed580a 100644 --- a/pyathena/polars/cursor.py +++ b/pyathena/polars/cursor.py @@ -18,7 +18,7 @@ DefaultPolarsUnloadTypeConverter, ) from pyathena.polars.result_set import AthenaPolarsResultSet -from pyathena.result_set import WithResultSet +from pyathena.result_set import WithFetch if TYPE_CHECKING: import polars as pl @@ -27,7 +27,7 @@ _logger = logging.getLogger(__name__) -class PolarsCursor(WithResultSet): +class PolarsCursor(WithFetch): """Cursor for handling Polars DataFrame results from Athena queries. This cursor returns query results as Polars DataFrames using Polars' native diff --git a/pyathena/result_set.py b/pyathena/result_set.py index c292ee892..388c6ade5 100644 --- a/pyathena/result_set.py +++ b/pyathena/result_set.py @@ -794,14 +794,13 @@ def _get_rows( ] -class WithResultSet(BaseCursor, CursorIterator): - """Base class of the SQL cursors that keep a result set. - - Provides the result set and its properties, fetch, ``close``, - ``executemany``, ``cancel``, and sync iteration. The sync SQL cursors - subclass it directly. For the asyncio cursors, ``WithAsyncFetch`` - overrides ``executemany`` and ``cancel`` with async versions, and its - subclasses override the fetch methods. +class WithResultSet: + """Mixin that keeps a cursor's query ID and result set. + + Provides the query ID, the result set and its properties, ``arraysize``, + ``rownumber``, ``rowcount``, and ``close``. ``WithFetch`` and + ``WithAsyncFetch`` list it before ``BaseCursor`` / ``AioBaseCursor`` and + ``CursorIterator``, so that these members take precedence over theirs. """ def __init__(self, arraysize: int | None = None, **kwargs) -> None: @@ -811,7 +810,7 @@ def __init__(self, arraysize: int | None = None, **kwargs) -> None: arraysize: Default number of rows per ``fetchmany()`` call, validated by the ``arraysize`` setter. If None, ``DEFAULT_FETCH_SIZE`` is used. - **kwargs: Arguments passed to ``BaseCursor.__init__``. + **kwargs: Arguments passed to the next ``__init__`` in the MRO. Raises: ProgrammingError: If ``arraysize`` is outside the range the @@ -1107,6 +1106,23 @@ def rownumber(self) -> int | None: """ return self.result_set.rownumber if self.result_set else None + def close(self) -> None: + """Close the cursor and release associated resources.""" + self._rowcount = -1 + if self.result_set and not self.result_set.is_closed: + self.result_set.close() + + +class WithFetch(WithResultSet, BaseCursor, CursorIterator): + """Base class of the sync SQL cursors. + + Combines ``WithResultSet`` with ``BaseCursor`` and ``CursorIterator``, and + provides sync fetch, ``executemany``, ``cancel``, and sync iteration. + + Subclasses override ``execute()`` and optionally ``__init__`` and + format-specific helpers. + """ + def fetchone( self, ) -> tuple[Any | None, ...] | dict[Any, Any | None] | None: @@ -1158,12 +1174,6 @@ def fetchall( result_set = cast(AthenaResultSet, self.result_set) return result_set.fetchall() - def close(self) -> None: - """Close the cursor and release associated resources.""" - self._rowcount = -1 - if self.result_set and not self.result_set.is_closed: - self.result_set.close() - def executemany( self, operation: str, diff --git a/pyathena/s3fs/cursor.py b/pyathena/s3fs/cursor.py index 27299ff5d..6052b66e1 100644 --- a/pyathena/s3fs/cursor.py +++ b/pyathena/s3fs/cursor.py @@ -8,14 +8,14 @@ from pyathena.error import OperationalError from pyathena.model import AthenaQueryExecution from pyathena.options import ExecuteOptions -from pyathena.result_set import WithResultSet +from pyathena.result_set import WithFetch from pyathena.s3fs.converter import DefaultS3FSTypeConverter from pyathena.s3fs.result_set import AthenaS3FSResultSet, CSVReaderType _logger = logging.getLogger(__name__) -class S3FSCursor(WithResultSet): +class S3FSCursor(WithFetch): """Cursor for reading CSV results via S3FileSystem without pandas/pyarrow. This cursor uses Python's standard csv module and PyAthena's S3FileSystem diff --git a/tests/pyathena/test_result_set.py b/tests/pyathena/test_result_set.py new file mode 100644 index 000000000..26635f24d --- /dev/null +++ b/tests/pyathena/test_result_set.py @@ -0,0 +1,24 @@ +# Copyright 2026 The PyAthena authors +# +# Licensed under the MIT License. +# See LICENSE or https://opensource.org/licenses/MIT. +# +# SPDX-License-Identifier: MIT +import pytest + +from pyathena.aio.common import AioBaseCursor, WithAsyncFetch +from pyathena.common import BaseCursor, CursorIterator +from pyathena.result_set import WithFetch, WithResultSet + + +class TestWithResultSet: + def test_is_mixin(self): + assert WithResultSet.__bases__ == (object,) + + @pytest.mark.parametrize( + ("cursor_base", "base"), + [(WithFetch, BaseCursor), (WithAsyncFetch, AioBaseCursor)], + ) + def test_precedes_cursor_bases(self, cursor_base, base): + # Listed first, so that its members take precedence over the cursor bases'. + assert cursor_base.__bases__ == (WithResultSet, base, CursorIterator)