-
Notifications
You must be signed in to change notification settings - Fork 114
Share the pure cursor-base logic between the sync and asyncio cursors #883
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
Merged
Merged
Changes from all commits
Commits
Show all changes
9 commits
Select commit
Hold shift + click to select a range
f78d06c
Share the pure cursor-base logic between the sync and asyncio cursors
laughingman7743 84a46f5
Fail rather than page on in the cache expiry tests
laughingman7743 a4a5d61
Correct the fetch and error descriptions in the cursor base docstrings
laughingman7743 ae77fbc
Catch a failing expiry check in the cache search tests
laughingman7743 7d4562a
Fold WithFetch into WithResultSet
laughingman7743 1b594b1
Describe what WithAsyncFetch overrides in the WithResultSet docstring
laughingman7743 a874aed
Document the result set properties of WithResultSet
laughingman7743 3e376d2
Correct the query_id and arraysize property docstrings
laughingman7743 d4938df
Merge remote-tracking branch 'origin/master' into refactor/880-cursor…
laughingman7743 File filter
Filter by extension
Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
There are no files selected for viewing
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
| Original file line number | Diff line number | Diff line change |
|---|---|---|
|
|
@@ -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__) | ||
|
|
@@ -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. | ||
| DatabaseError: If the ``StartQueryExecution`` request fails. | ||
| """ | ||
| # The individual keyword arguments are retained for backward compatibility | ||
| # with external callers that predate ExecuteOptions, mirroring | ||
| # BaseCursor._execute(). | ||
|
|
@@ -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, | ||
|
|
@@ -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: | ||
|
|
@@ -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) | ||
|
|
@@ -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): | ||
|
Member
Author
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. Self-review round two (claims, callers) of the Claims checked:
|
||
| """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, | ||
|
|
@@ -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: | ||
|
|
@@ -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 | ||
|
|
||
|
|
||
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Oops, something went wrong.
Oops, something went wrong.
Add this suggestion to a batch that can be applied as a single commit.
This suggestion is invalid because no changes were made to the code.
Suggestions cannot be applied while the pull request is closed.
Suggestions cannot be applied while viewing a subset of changes.
Only one suggestion per line can be applied in a batch.
Add this suggestion to a batch that can be applied as a single commit.
Applying suggestions on deleted lines is not supported.
You must change the existing code in this line in order to create a valid suggestion.
Outdated suggestions cannot be applied.
This suggestion has been applied or marked resolved.
Suggestions cannot be applied from pending reviews.
Suggestions cannot be applied on multi-line comments.
Suggestions cannot be applied while the pull request is queued to merge.
Suggestion cannot be applied right now. Please check back later.
There was a problem hiding this comment.
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:
Cursor/AioCursorkeep the cappedarraysize, 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.WithFetch/WithAsyncFetchresolve every member as before. The names shared byWithResultSet/CursorIteratorwere all overridden by the old mixins. Holds; added to the body.ListQueryExecutions/StartQueryExecutionrequest count and order are unchanged. The paging test pinsmax_results50 then 10 forcache_size=60; the matching and paging code moved without changes. Holds.master" predated the 84a46f5 test change. Re-ran the current tests against a86a180'scommon.py/aio/common.py/result_set.py: 10 passed. The body now names each commit's evidence._find_previous_query_idand_execute; corrected. Also noted that the script differs from the counting behind the issue's estimate.Existing callers: signatures of
_execute/_find_previous_query_idare unchanged (dbt-athena 1.x, #734).WithFetch.__init__is gone, butsuper().__init__(**kwargs)from subclasses now reachesBaseCursor.__init__with the same arguments.Not measured: live AWS request counts; the equivalence above is static plus the mocked paging test.