diff --git a/pyathena/aio/common.py b/pyathena/aio/common.py index 1de17ab4..8510f155 100644 --- a/pyathena/aio/common.py +++ b/pyathena/aio/common.py @@ -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): + """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 diff --git a/pyathena/arrow/cursor.py b/pyathena/arrow/cursor.py index 5ad7cd60..6b2dabda 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 WithFetch +from pyathena.result_set import WithResultSet if TYPE_CHECKING: import polars as pl @@ -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 diff --git a/pyathena/common.py b/pyathena/common.py index 7f5c98cf..e556b0cd 100644 --- a/pyathena/common.py +++ b/pyathena/common.py @@ -914,6 +914,76 @@ def _poll(self, query_id: str) -> AthenaQueryExecution | AthenaCalculationExecut raise e return query_execution + def _cache_search_limits( + self, cache_size: int, cache_expiration_time: int + ) -> tuple[int, datetime | None]: + """Resolve how far the result cache search looks back. No I/O. + + Args: + 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 number of executions to search, unbounded when only + ``cache_expiration_time`` is set, and the oldest completion time + to accept, or None for no limit. + """ + 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) + return cache_size, expiration_time + return cache_size, None + + def _match_previous_query( + self, + query: str, + query_executions: list[AthenaQueryExecution], + expiration_time: datetime | None, + ) -> tuple[str | None, bool]: + """Find the latest reusable execution of a query in one page. No I/O. + + Reusable executions are succeeded DML queries with the same query + string, schema, and catalog (case-insensitive). Executions are checked + from the latest completion; the check stops at the first one completed + before ``expiration_time``. + + Args: + query: The query string. + query_executions: One page of the work group's executions. + expiration_time: The oldest completion time to accept, or None for + no limit. + + Returns: + The matching query ID or None, and whether the check reached an + expired execution. + """ + for execution in sorted( + ( + e + for e in query_executions + if e.state == AthenaQueryExecution.STATE_SUCCEEDED + and e.statement_type == AthenaQueryExecution.STATEMENT_TYPE_DML + ), + # https://github.com/python/mypy/issues/9656 + key=lambda e: e.completion_date_time, # type: ignore[arg-type, return-value] + reverse=True, + ): + if ( + expiration_time + and execution.completion_date_time + and execution.completion_date_time.astimezone(UTC) < expiration_time + ): + return None, True + if ( + execution.query == query + and execution.database == self._schema_name + and (execution.catalog or "").lower() == (self._catalog_name or "").lower() + ): + return execution.query_id, False + return None, False + def _find_previous_query_id( self, query: str, @@ -921,13 +991,23 @@ def _find_previous_query_id( 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: @@ -936,32 +1016,10 @@ def _find_previous_query_id( next_token, query_executions = 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 - ), - # https://github.com/python/mypy/issues/9656 - 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) @@ -1030,6 +1088,36 @@ def _call_on_start_query_execution(self, query_id: str, options: ExecuteOptions) if options.on_start_query_execution: options.on_start_query_execution(query_id) + def _build_execute_request( + self, + operation: str, + parameters: dict[str, Any] | list[str] | None, + options: ExecuteOptions, + ) -> tuple[str, dict[str, Any]]: + """Format a query and build its ``StartQueryExecution`` request. No I/O. + + Args: + operation: SQL query string. + parameters: Query parameters. + options: The resolved execution options. + + Returns: + Tuple of (formatted_query, request). + + Raises: + ProgrammingError: If the formatter rejects the query or its parameters. + """ + 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, + ) + return query, request + def _execute( self, operation: str, @@ -1043,6 +1131,30 @@ def _execute( 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 (e.g. dbt-athena <= 1.10.x # calls _execute() with work_group/s3_staging_dir/cache_* keywords). @@ -1056,16 +1168,7 @@ def _execute( 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 = self._find_previous_query_id( query, options.work_group, diff --git a/pyathena/cursor.py b/pyathena/cursor.py index eb302b30..cc365ce0 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, WithFetch +from pyathena.result_set import AthenaDictResultSet, AthenaResultSet, WithResultSet _logger = logging.getLogger(__name__) -class Cursor(WithFetch): +class Cursor(WithResultSet): """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 883172c0..38b66c98 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 WithFetch +from pyathena.result_set import WithResultSet if TYPE_CHECKING: from pandas import DataFrame @@ -26,7 +26,7 @@ _logger = logging.getLogger(__name__) -class PandasCursor(WithFetch): +class PandasCursor(WithResultSet): """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 a92ed580..4bd8c78e 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 WithFetch +from pyathena.result_set import WithResultSet if TYPE_CHECKING: import polars as pl @@ -27,7 +27,7 @@ _logger = logging.getLogger(__name__) -class PolarsCursor(WithFetch): +class PolarsCursor(WithResultSet): """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 d75326c3..fffbb43a 100644 --- a/pyathena/result_set.py +++ b/pyathena/result_set.py @@ -2,7 +2,6 @@ import collections import logging -from abc import abstractmethod from datetime import datetime from typing import ( TYPE_CHECKING, @@ -795,9 +794,25 @@ def _get_rows( ] -class WithResultSet: - def __init__(self): - super().__init__() +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. + """ + + def __init__(self, **kwargs) -> None: + """Initialize the cursor with no query ID and no result set. + + Args: + **kwargs: Arguments passed to ``BaseCursor.__init__``. + """ + super().__init__(**kwargs) + self._query_id: str | None = None + self._result_set: AthenaResultSet | None = None def _reset_state(self) -> None: self._rowcount = -1 @@ -807,14 +822,17 @@ def _reset_state(self) -> None: self.result_set = None @property - @abstractmethod def result_set(self) -> AthenaResultSet | None: - raise NotImplementedError # pragma: no cover + """The result set of the last executed query. + + Returns: + The result set, or None before a query succeeds or after a reset. + """ + return self._result_set @result_set.setter - @abstractmethod def result_set(self, val: AthenaResultSet | None) -> None: - raise NotImplementedError # pragma: no cover + self._result_set = val @property def has_result_set(self) -> bool: @@ -841,14 +859,21 @@ def catalog(self) -> str | None: return self.result_set.catalog @property - @abstractmethod def query_id(self) -> str | None: - raise NotImplementedError # pragma: no cover + """The query execution ID of the last execution. + + With ``cache_size`` or ``cache_expiration_time``, this can be the ID of + a previous execution whose result is reused. + + Returns: + The query execution ID, or None if there is none since the last + reset. + """ + return self._query_id @query_id.setter - @abstractmethod def query_id(self, val: str | None) -> None: - raise NotImplementedError # pragma: no cover + self._query_id = val @property def query(self) -> str | None: @@ -1044,26 +1069,17 @@ def rowcount(self) -> int: """ return self.result_set.rowcount if self.result_set else self._rowcount - -class WithFetch(BaseCursor, CursorIterator, WithResultSet): - """Mixin providing shared properties, fetch, lifecycle, and sync iteration for 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 sync iteration protocol. - - 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: + """The default number of rows per ``fetchmany()`` call. + + ``execute()`` passes it to the new result set, so a change applies to + the result sets of later executions. Setting it to zero or a negative + value raises ``ProgrammingError``. + + Returns: + The default number of rows per ``fetchmany()`` call. + """ return self._arraysize @arraysize.setter @@ -1073,28 +1089,65 @@ def arraysize(self, value: int) -> None: self._arraysize = value @property - def result_set(self) -> AthenaResultSet | None: - return self._result_set + def rownumber(self) -> int | None: + """The zero-based index of the next row in the result set. - @result_set.setter - def result_set(self, val) -> None: - self._result_set = val + Returns: + The row index, or None if there is no result set or the index is + unknown. + """ + return self.result_set.rownumber if self.result_set else None - @property - def query_id(self) -> str | None: - return self._query_id + def fetchone( + self, + ) -> tuple[Any | None, ...] | dict[Any, Any | None] | None: + """Fetch the next row of the result set. - @query_id.setter - def query_id(self, val) -> None: - self._query_id = val + Returns: + A tuple representing the next row, or None if no more rows. - @property - def rownumber(self) -> int | None: - return self.result_set.rownumber if self.result_set else None + 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() - @property - def rowcount(self) -> int: - return self.result_set.rowcount if self.result_set else self._rowcount + 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 close(self) -> None: """Close the cursor and release associated resources.""" @@ -1148,54 +1201,3 @@ def cancel(self) -> None: if not self.query_id: raise ProgrammingError("QueryExecutionId is none or empty.") 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() diff --git a/pyathena/s3fs/cursor.py b/pyathena/s3fs/cursor.py index 6052b66e..27299ff5 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 WithFetch +from pyathena.result_set import WithResultSet from pyathena.s3fs.converter import DefaultS3FSTypeConverter from pyathena.s3fs.result_set import AthenaS3FSResultSet, CSVReaderType _logger = logging.getLogger(__name__) -class S3FSCursor(WithFetch): +class S3FSCursor(WithResultSet): """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/aio/test_cursor.py b/tests/pyathena/aio/test_cursor.py index 1a35025c..9679854a 100644 --- a/tests/pyathena/aio/test_cursor.py +++ b/tests/pyathena/aio/test_cursor.py @@ -1,8 +1,8 @@ import asyncio import re import threading -from datetime import UTC, datetime -from unittest.mock import AsyncMock, MagicMock, patch +from datetime import UTC, datetime, timedelta +from unittest.mock import AsyncMock, MagicMock, call, patch import pytest @@ -15,7 +15,7 @@ from pyathena.util import RetryConfig from tests import ENV from tests.pyathena.aio.conftest import _aio_connect -from tests.pyathena.util import throttle_metadata_api +from tests.pyathena.util import succeeded_query_execution, throttle_metadata_api class TestAioCursor: @@ -233,6 +233,74 @@ def execution(catalog): == "query_id_awsdatacatalog" ) + async def test_cache_search_stops_at_expired_execution(self, caplog): + query = "SELECT * FROM one_row" + now = datetime.now(UTC) + cursor = AioCursor.__new__(AioCursor) + cursor._schema_name = "this_schema" + cursor._catalog_name = None + page = [ + succeeded_query_execution("expired", query, now - timedelta(hours=2)), + succeeded_query_execution("other", "SELECT 1", now), + ] + + # A page read after the expired execution would return this match. + next_page = (None, [succeeded_query_execution("next_page", query, now)]) + + with patch.object( + AioCursor, + "_list_query_executions", + new_callable=AsyncMock, + side_effect=[("next_token", page), next_page], + ) as list_mock: + # Without cache_size, cache_expiration_time alone bounds the search. + assert ( + await cursor._find_previous_query_id(query, None, cache_expiration_time=3600) + is None + ) + list_mock.assert_awaited_once_with(None, next_token=None, max_results=50) + # A failed search also returns None; the expiry must stop it without an error. + assert "Failed to check the cache" not in caplog.text + + async def test_cache_search_reads_pages_up_to_cache_size(self): + query = "SELECT * FROM one_row" + now = datetime.now(UTC) + cursor = AioCursor.__new__(AioCursor) + cursor._schema_name = "this_schema" + cursor._catalog_name = None + pages = [ + ("next_token", [succeeded_query_execution("other", "SELECT 1", now)]), + ("last_token", [succeeded_query_execution("match", query, now)]), + ] + + with patch.object( + AioCursor, "_list_query_executions", new_callable=AsyncMock, side_effect=pages + ) as list_mock: + assert await cursor._find_previous_query_id(query, "wg", cache_size=60) == "match" + assert list_mock.await_args_list == [ + call("wg", next_token=None, max_results=50), + call("wg", next_token="next_token", max_results=10), + ] + + async def test_cache_search_prefers_latest_execution(self): + query = "SELECT * FROM one_row" + now = datetime.now(UTC) + cursor = AioCursor.__new__(AioCursor) + cursor._schema_name = "this_schema" + cursor._catalog_name = None + page = [ + succeeded_query_execution("older", query, now - timedelta(minutes=1)), + succeeded_query_execution("latest", query, now), + ] + + with patch.object( + AioCursor, + "_list_query_executions", + new_callable=AsyncMock, + return_value=(None, page), + ): + assert await cursor._find_previous_query_id(query, None, cache_size=10) == "latest" + async def test_no_result_set_raises(self, aio_cursor): with pytest.raises(ProgrammingError): await aio_cursor.fetchone() diff --git a/tests/pyathena/test_cursor.py b/tests/pyathena/test_cursor.py index 02db8ef8..eb64f494 100644 --- a/tests/pyathena/test_cursor.py +++ b/tests/pyathena/test_cursor.py @@ -9,10 +9,10 @@ import uuid from concurrent import futures from concurrent.futures.thread import ThreadPoolExecutor -from datetime import UTC, date, datetime +from datetime import UTC, date, datetime, timedelta from decimal import Decimal from random import randint -from unittest.mock import MagicMock, patch +from unittest.mock import MagicMock, call, patch import pytest from botocore.exceptions import ClientError @@ -36,7 +36,7 @@ from pyathena.util import RetryConfig from tests import ENV from tests.pyathena.conftest import connect -from tests.pyathena.util import throttle_metadata_api, unreachable_glue +from tests.pyathena.util import succeeded_query_execution, throttle_metadata_api, unreachable_glue _logger = logging.getLogger(__name__) @@ -277,6 +277,61 @@ def execution(catalog): == "query_id_awsdatacatalog" ) + def test_cache_search_stops_at_expired_execution(self, caplog): + query = "SELECT * FROM one_row" + now = datetime.now(UTC) + cursor = Cursor.__new__(Cursor) + cursor._schema_name = "this_schema" + cursor._catalog_name = None + page = [ + succeeded_query_execution("expired", query, now - timedelta(hours=2)), + succeeded_query_execution("other", "SELECT 1", now), + ] + + # A page read after the expired execution would return this match. + next_page = (None, [succeeded_query_execution("next_page", query, now)]) + + with patch.object( + Cursor, "_list_query_executions", side_effect=[("next_token", page), next_page] + ) as list_mock: + # Without cache_size, cache_expiration_time alone bounds the search. + assert cursor._find_previous_query_id(query, None, cache_expiration_time=3600) is None + list_mock.assert_called_once_with(None, next_token=None, max_results=50) + # A failed search also returns None; the expiry must stop it without an error. + assert "Failed to check the cache" not in caplog.text + + def test_cache_search_reads_pages_up_to_cache_size(self): + query = "SELECT * FROM one_row" + now = datetime.now(UTC) + cursor = Cursor.__new__(Cursor) + cursor._schema_name = "this_schema" + cursor._catalog_name = None + pages = [ + ("next_token", [succeeded_query_execution("other", "SELECT 1", now)]), + ("last_token", [succeeded_query_execution("match", query, now)]), + ] + + with patch.object(Cursor, "_list_query_executions", side_effect=pages) as list_mock: + assert cursor._find_previous_query_id(query, "wg", cache_size=60) == "match" + assert list_mock.call_args_list == [ + call("wg", next_token=None, max_results=50), + call("wg", next_token="next_token", max_results=10), + ] + + def test_cache_search_prefers_latest_execution(self): + query = "SELECT * FROM one_row" + now = datetime.now(UTC) + cursor = Cursor.__new__(Cursor) + cursor._schema_name = "this_schema" + cursor._catalog_name = None + page = [ + succeeded_query_execution("older", query, now - timedelta(minutes=1)), + succeeded_query_execution("latest", query, now), + ] + + with patch.object(Cursor, "_list_query_executions", return_value=(None, page)): + assert cursor._find_previous_query_id(query, None, cache_size=10) == "latest" + @pytest.mark.parametrize( "cursor", [{"work_group": ENV.work_group, "result_reuse_enable": True, "result_reuse_minutes": 5}], diff --git a/tests/pyathena/util.py b/tests/pyathena/util.py index f5ac673d..bfe4dd3c 100644 --- a/tests/pyathena/util.py +++ b/tests/pyathena/util.py @@ -14,7 +14,7 @@ from sqlalchemy import types from pyathena.glue import GlueMetadataClient -from pyathena.model import AthenaCalculationExecutionStatus +from pyathena.model import AthenaCalculationExecutionStatus, AthenaQueryExecution _queries = Environment( loader=FileSystemLoader(Path(__file__).parents[1].resolve() / "resources" / "queries") @@ -51,6 +51,34 @@ def fail(**kwargs): return calls +def succeeded_query_execution(query_id, query, completion_date_time, schema="this_schema"): + """Build a succeeded DML query execution, as the result cache search lists it. + + Args: + query_id: The query execution ID. + query: The query string. + completion_date_time: The completion time. + schema: The database the query ran against. + + Returns: + The query execution. + """ + return AthenaQueryExecution( + { + "QueryExecution": { + "QueryExecutionId": query_id, + "Query": query, + "StatementType": AthenaQueryExecution.STATEMENT_TYPE_DML, + "QueryExecutionContext": {"Database": schema}, + "Status": { + "State": AthenaQueryExecution.STATE_SUCCEEDED, + "CompletionDateTime": completion_date_time, + }, + } + } + ) + + def unreachable_glue(connection): """A Glue client for the connection that sends requests to a closed proxy port.""" return GlueMetadataClient(