From 85f0baa6b093adfc3b262d0f9a29dbdcd71ccb78 Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Thu, 1 Oct 2026 09:22:43 +0900 Subject: [PATCH] Restore WithResultSet as a mixin with sync and async sibling bases #883 folded WithFetch into WithResultSet and made the mixin a subclass of BaseCursor and CursorIterator, with WithAsyncFetch overriding its sync methods. Restore the mixin design: WithResultSet has no base class again, WithFetch is restored for the sync fetch, executemany, and cancel, and WithAsyncFetch is its asyncio sibling. The members shared by both bases stay in WithResultSet, which both list first so that these members take precedence over those of the cursor bases and CursorIterator. Co-Authored-By: Claude Opus 5.5 --- pyathena/aio/common.py | 12 +++++----- pyathena/arrow/cursor.py | 4 ++-- pyathena/cursor.py | 4 ++-- pyathena/pandas/cursor.py | 4 ++-- pyathena/polars/cursor.py | 4 ++-- pyathena/result_set.py | 40 +++++++++++++++++++------------ pyathena/s3fs/cursor.py | 4 ++-- tests/pyathena/test_result_set.py | 24 +++++++++++++++++++ 8 files changed, 65 insertions(+), 31 deletions(-) create mode 100644 tests/pyathena/test_result_set.py 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)