Skip to content
Merged
203 changes: 54 additions & 149 deletions pyathena/aio/common.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,20 +2,18 @@

import asyncio
import logging
import sys
from collections.abc import Awaitable, Callable
from datetime import UTC, datetime, timedelta
from typing import Any, TypeVar, cast
from typing import Any, TypeVar

from botocore.exceptions import BotoCoreError, ClientError

from pyathena.aio.util import async_retry_api_call
from pyathena.common import BaseCursor, CursorIterator
from pyathena.common import BaseCursor
from pyathena.error import DatabaseError, OperationalError, ProgrammingError
from pyathena.glue import GlueMetadataClient
from pyathena.model import AthenaDatabase, AthenaQueryExecution, AthenaTableMetadata
from pyathena.options import ExecuteOptions
from pyathena.result_set import AthenaResultSet, WithResultSet
from pyathena.result_set import WithResultSet
from pyathena.util import _is_throttling_error

_logger = logging.getLogger(__name__)
Expand Down Expand Up @@ -44,6 +42,30 @@ async def _execute( # type: ignore[override]
paramstyle: str | None = None,
options: ExecuteOptions | None = None,
) -> str:
"""Start a query execution, or find a previous one to reuse.

The individual keyword arguments override the ``options`` field of the
same name unless None.

Args:
operation: SQL query string.
parameters: Query parameters.
work_group: Athena work group.
s3_staging_dir: S3 location for query results.
cache_size: Number of recent executions to search for a reusable result.
cache_expiration_time: Maximum age of a reusable result in seconds.
result_reuse_enable: Whether to enable Athena result reuse.
result_reuse_minutes: Maximum age of an Athena-reused result in minutes.
paramstyle: Parameter style ('qmark' or 'pyformat').
options: The execution options.

Returns:
The query execution ID.

Raises:
ProgrammingError: If the formatter rejects the query or its parameters.

Copy link
Copy Markdown
Member Author

Choose a reason for hiding this comment

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

Self-review round two (claims, callers, operations): FINDINGS (3, corrected)

Range: git diff a86a180ebbe5e18920cd75802ea5350ce8d33f10..a4a5d61099f10182071e15e7cab1081e8966d270 (84a46f5 → a4a5d61 changes only docstrings)

Claims checked:

  • PR body, MRO: "Cursor / AioCursor keep the capped arraysize, format cursors the uncapped one, as before". Checked against the owning class of each member for all 12 SQL cursors on this head and on a86a180. Holds.
  • PR body, compatibility: subclasses of WithFetch / WithAsyncFetch resolve every member as before. The names shared by WithResultSet / CursorIterator were all overridden by the old mixins. Holds; added to the body.
  • PR body, AWS operator: the ListQueryExecutions / StartQueryExecution request count and order are unchanged. The paging test pins max_results 50 then 10 for cache_size=60; the matching and paging code moved without changes. Holds.
  • PR body, pinning evidence: "passed on unchanged master" predated the 84a46f5 test change. Re-ran the current tests against a86a180's common.py / aio/common.py / result_set.py: 10 passed. The body now names each commit's evidence.
  • PR body, duplication table: re-ran the script on a4a5d61, still 18 pairs / 284 lines. The list of remaining pairs had left out the I/O parts of _find_previous_query_id and _execute; corrected. Also noted that the script differs from the counting behind the issue's estimate.
  • PR body: "Following the issue discussion" misattributed the parse-helper decision, which the maintainer made for this PR; corrected.
  • Docstrings: see the inline findings.

Existing callers: signatures of _execute / _find_previous_query_id are unchanged (dbt-athena 1.x, #734). WithFetch.__init__ is gone, but super().__init__(**kwargs) from subclasses now reaches BaseCursor.__init__ with the same arguments.

Not measured: live AWS request counts; the equivalence above is static plus the mocked paging test.

DatabaseError: If the ``StartQueryExecution`` request fails.
"""
# The individual keyword arguments are retained for backward compatibility
# with external callers that predate ExecuteOptions, mirroring
# BaseCursor._execute().
Expand All @@ -57,16 +79,7 @@ async def _execute( # type: ignore[override]
result_reuse_minutes=result_reuse_minutes,
paramstyle=paramstyle,
)
query, execution_parameters = self._prepare_query(operation, parameters, options.paramstyle)

request = self._build_start_query_execution_request(
query=query,
work_group=options.work_group,
s3_staging_dir=options.s3_staging_dir,
result_reuse_enable=options.result_reuse_enable,
result_reuse_minutes=options.result_reuse_minutes,
execution_parameters=execution_parameters,
)
query, request = self._build_execute_request(operation, parameters, options)
query_id = await self._find_previous_query_id(
query,
options.work_group,
Expand Down Expand Up @@ -236,13 +249,23 @@ async def _find_previous_query_id( # type: ignore[override]
cache_size: int = 0,
cache_expiration_time: int = 0,
) -> str | None:
"""Find a previous execution of a query whose result can be reused.

Searches the work group's recent executions page by page. A failed
search is logged and treated as a cache miss.

Args:
query: The query string.
work_group: The work group to search, or None for the cursor's.
cache_size: The number of recent executions to search, or 0.
cache_expiration_time: The maximum age of a reused result in
seconds, or 0 for no limit.

Returns:
The query ID of the latest reusable execution, or None.
"""
cache_size, expiration_time = self._cache_search_limits(cache_size, cache_expiration_time)
query_id = None
if cache_size == 0 and cache_expiration_time > 0:
cache_size = sys.maxsize
if cache_expiration_time > 0:
expiration_time = datetime.now(UTC) - timedelta(seconds=cache_expiration_time)
else:
expiration_time = datetime.now(UTC)
try:
next_token = None
while cache_size > 0:
Expand All @@ -251,31 +274,10 @@ async def _find_previous_query_id( # type: ignore[override]
next_token, query_executions = await self._list_query_executions(
work_group, next_token=next_token, max_results=max_results
)
for execution in sorted(
(
e
for e in query_executions
if e.state == AthenaQueryExecution.STATE_SUCCEEDED
and e.statement_type == AthenaQueryExecution.STATEMENT_TYPE_DML
),
key=lambda e: e.completion_date_time, # type: ignore[arg-type, return-value]
reverse=True,
):
if (
cache_expiration_time > 0
and execution.completion_date_time
and execution.completion_date_time.astimezone(UTC) < expiration_time
):
next_token = None
break
if (
execution.query == query
and execution.database == self._schema_name
and (execution.catalog or "").lower() == (self._catalog_name or "").lower()
):
query_id = execution.query_id
break
if query_id or next_token is None:
query_id, expired = self._match_previous_query(
query, query_executions, expiration_time
)
if query_id or expired or next_token is None:
break
except Exception:
_logger.warning("Failed to check the cache. Moving on without cache.", exc_info=True)
Expand Down Expand Up @@ -615,63 +617,17 @@ async def athena_request(
)


class WithAsyncFetch(AioBaseCursor, CursorIterator, WithResultSet):
"""Mixin providing shared fetch, lifecycle, and async protocol for SQL cursors.
class WithAsyncFetch(AioBaseCursor, WithResultSet):

Copy link
Copy Markdown
Member Author

Choose a reason for hiding this comment

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

Self-review round two (claims, callers) of the WithFetch fold: FINDINGS (1, corrected)

Claims checked:

  • PR body, breaking change: "a class that lists BaseCursor or CursorIterator before WithResultSet fails with an MRO TypeError". Verified with type() for (BaseCursor, CursorIterator, WithResultSet), (BaseCursor, WithResultSet), and (CursorIterator, WithResultSet), which all raise TypeError, and (WithResultSet,) and (WithResultSet, BaseCursor, CursorIterator), which work. The body's advice now says "first or alone".
  • PR body: WithFetch "was added in v3.29.0 and is not in the API docs". Checked: git tag --contains b753495 starts at v3.29.0, and docs/ has no WithFetch. WithResultSet is in docs/api/connection.rst (:members: only, under Result Sets; its placement is left as is).
  • Existing callers: from pyathena.result_set import WithFetch now raises ImportError. The maintainer chose no alias, for the 4.0.0 major release.
  • Finding: the new WithResultSet docstring (7d4562a) said WithAsyncFetch "overrides the fetch and lifecycle methods". It overrides only executemany / cancel; its subclasses override fetch; close is shared. Corrected in 1b594b1 (docstring only).

"""Base class of the asyncio SQL cursors.

Provides properties (``arraysize``, ``result_set``, ``query_id``,
``rownumber``, ``rowcount``), lifecycle methods (``close``, ``executemany``,
``cancel``), default sync fetch (for cursors whose result sets load all
data eagerly in ``__init__``), and the async iteration protocol.
Overrides ``executemany`` and ``cancel`` of ``WithResultSet`` with async
versions and adds async iteration and the async context manager protocol.
Subclasses override the fetch methods with async versions.

Subclasses override ``execute()`` and optionally ``__init__`` and
format-specific helpers.
"""

def __init__(self, **kwargs) -> None:
super().__init__(**kwargs)
self._query_id: str | None = None
self._result_set: AthenaResultSet | None = None

@property
def arraysize(self) -> int:
return self._arraysize

@arraysize.setter
def arraysize(self, value: int) -> None:
if value <= 0:
raise ProgrammingError("arraysize must be a positive integer value.")
self._arraysize = value

@property # type: ignore[override]
def result_set(self) -> AthenaResultSet | None:
return self._result_set

@result_set.setter
def result_set(self, val) -> None:
self._result_set = val

@property
def query_id(self) -> str | None:
return self._query_id

@query_id.setter
def query_id(self, val) -> None:
self._query_id = val

@property
def rownumber(self) -> int | None:
return self.result_set.rownumber if self.result_set else None

@property
def rowcount(self) -> int:
return self.result_set.rowcount if self.result_set else self._rowcount

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()

async def executemany( # type: ignore[override]
self,
operation: str,
Expand Down Expand Up @@ -709,7 +665,7 @@ async def executemany( # type: ignore[override]
self._reset_state()
self._rowcount = rowcount

async def cancel(self) -> None:
async def cancel(self) -> None: # type: ignore[override]
"""Cancel the currently executing query.

Raises:
Expand All @@ -719,57 +675,6 @@ async def cancel(self) -> None:
raise ProgrammingError("QueryExecutionId is none or empty.")
await self._cancel(self.query_id)

def fetchone(
self,
) -> tuple[Any | None, ...] | dict[Any, Any | None] | None:
"""Fetch the next row of the result set.

Returns:
A tuple representing the next row, or None if no more rows.

Raises:
ProgrammingError: If no result set is available.
"""
if not self.has_result_set:
raise ProgrammingError("No result set.")
result_set = cast(AthenaResultSet, self.result_set)
return result_set.fetchone()

def fetchmany(
self, size: int | None = None
) -> list[tuple[Any | None, ...] | dict[Any, Any | None]]:
"""Fetch multiple rows from the result set.

Args:
size: Maximum number of rows to fetch. Defaults to arraysize.

Returns:
List of tuples representing the fetched rows.

Raises:
ProgrammingError: If no result set is available.
"""
if not self.has_result_set:
raise ProgrammingError("No result set.")
result_set = cast(AthenaResultSet, self.result_set)
return result_set.fetchmany(size)

def fetchall(
self,
) -> list[tuple[Any | None, ...] | dict[Any, Any | None]]:
"""Fetch all remaining rows from the result set.

Returns:
List of tuples representing all remaining rows.

Raises:
ProgrammingError: If no result set is available.
"""
if not self.has_result_set:
raise ProgrammingError("No result set.")
result_set = cast(AthenaResultSet, self.result_set)
return result_set.fetchall()

def __aiter__(self):
return self

Expand Down
4 changes: 2 additions & 2 deletions pyathena/arrow/cursor.py
Original file line number Diff line number Diff line change
Expand Up @@ -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 WithFetch
from pyathena.result_set import WithResultSet

if TYPE_CHECKING:
import polars as pl
Expand All @@ -22,7 +22,7 @@
_logger = logging.getLogger(__name__)


class ArrowCursor(WithFetch):
class ArrowCursor(WithResultSet):
"""Cursor for handling Apache Arrow Table results from Athena queries.

This cursor returns query results as Apache Arrow Tables, which provide
Expand Down
Loading
Loading