"""Asyncio result sets that fetch Athena query results with ``GetQueryResults``."""
from __future__ import annotations
import logging
from typing import (
TYPE_CHECKING,
Any,
NoReturn,
cast,
)
from pyathena.aio.util import async_retry_api_call
from pyathena.converter import Converter
from pyathena.error import OperationalError, ProgrammingError
from pyathena.model import AthenaQueryExecution
from pyathena.result_set import AthenaDictResultSet, AthenaResultSet
from pyathena.util import RetryConfig, override
if TYPE_CHECKING:
from pyathena.connection import Connection
_logger = logging.getLogger(__name__)
[docs]
class AthenaAioResultSet(AthenaResultSet):
"""Async result set that provides async fetch methods.
Skips the synchronous ``_pre_fetch`` by passing ``_pre_fetch=False`` to
the parent ``__init__`` and provides an ``async create()`` classmethod
factory instead. Synchronous iteration raises ``TypeError``; use
``async for`` instead.
"""
[docs]
def __init__(
self,
connection: Connection[Any],
converter: Converter,
query_execution: AthenaQueryExecution,
arraysize: int,
retry_config: RetryConfig,
result_set_type_hints: dict[str | int, str] | None = None,
) -> None:
"""Initialize the result set without fetching rows; ``create()`` fetches the first page.
Args:
connection: The connection that ran the query.
converter: The converter for result values.
query_execution: The query execution whose results to read.
arraysize: The number of rows per ``GetQueryResults`` page and the default
``fetchmany()`` size.
retry_config: The retry configuration for API calls.
result_set_type_hints: Athena type signatures for complex-type columns,
keyed by column name (case-insensitive) or zero-based column index.
Raises:
ProgrammingError: If ``query_execution`` is not given.
"""
super().__init__(
connection=connection,
converter=converter,
query_execution=query_execution,
arraysize=arraysize,
retry_config=retry_config,
_pre_fetch=False,
result_set_type_hints=result_set_type_hints,
)
[docs]
@classmethod
async def create(
cls,
connection: Connection[Any],
converter: Converter,
query_execution: AthenaQueryExecution,
arraysize: int,
retry_config: RetryConfig,
result_set_type_hints: dict[str | int, str] | None = None,
**kwargs: Any,
) -> AthenaAioResultSet:
"""Async factory method.
Creates an ``AthenaAioResultSet`` and awaits the initial data fetch.
Args:
connection: The database connection.
converter: Type converter for result values.
query_execution: Query execution metadata.
arraysize: Number of rows to fetch per request.
retry_config: Retry configuration for API calls.
result_set_type_hints: Athena type signatures for complex-type columns,
keyed by column name (case-insensitive) or zero-based column index.
**kwargs: Additional arguments passed to the constructor of ``cls``,
such as ``dict_type`` for ``AthenaAioDictResultSet``.
Returns:
A fully initialized ``AthenaAioResultSet``.
"""
result_set = cls(
connection,
converter,
query_execution,
arraysize,
retry_config,
result_set_type_hints=result_set_type_hints,
**kwargs,
)
if result_set.state == AthenaQueryExecution.STATE_SUCCEEDED:
await result_set._async_pre_fetch()
return result_set
async def _async_get_query_results(
self, max_results: int, next_token: str | None = None
) -> dict[str, Any]:
"""Get a page of query results with ``GetQueryResults``.
Args:
max_results: The maximum number of rows in the page.
next_token: The token of the page to get; the first page if None.
Returns:
The ``GetQueryResults`` response.
Raises:
ProgrammingError: If the query ID is missing, the query has not
succeeded, or the result set is closed.
OperationalError: If the request fails.
"""
request = self._build_get_query_results_request(max_results, next_token)
if self.is_closed:
raise ProgrammingError("AthenaAioResultSet is closed.")
try:
response = await async_retry_api_call(
self.connection.client.get_query_results,
config=self._retry_config,
logger=_logger,
**request,
)
except Exception as e:
_logger.exception("Failed to fetch result set.")
raise OperationalError(*e.args) from e
else:
return cast(dict[str, Any], response)
async def _async_fetch(self) -> None:
"""Fetch the next page of rows into the result set.
Raises:
ProgrammingError: If there is no next page.
OperationalError: If the request fails.
"""
if not self._next_token:
raise ProgrammingError("NextToken is none or empty.")
response = await self._async_get_query_results(self._arraysize, self._next_token)
rows, self._next_token = self._parse_result_rows(response)
self._process_rows(rows)
async def _async_pre_fetch(self) -> None:
"""Fetch the first page of rows along with the result metadata.
Raises:
ProgrammingError: If the query ID is missing, the query has not
succeeded, or the result set is closed.
OperationalError: If the request fails.
"""
response = await self._async_get_query_results(self._arraysize)
self._process_metadata(response)
self._process_update_count(response)
rows, self._next_token = self._parse_result_rows(response)
offset = 1 if rows and self._is_first_row_column_labels(rows) else 0
self._process_rows(rows, offset)
[docs]
@override
async def fetchone( # type: ignore[override]
self,
) -> tuple[Any | None, ...] | dict[Any, Any | None] | None:
"""Fetch the next row of the result set.
Automatically fetches the next page from Athena when the current
page is exhausted and more pages are available.
Returns:
The next row (a tuple, or a dict for ``AthenaAioDictResultSet``), or
None if no more rows.
"""
if not self._rows and self._next_token:
await self._async_fetch()
if not self._rows:
return None
if self._rownumber is None:
self._rownumber = 0
self._rownumber += 1
return self._rows.popleft()
[docs]
@override
async def fetchmany( # type: ignore[override]
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. If None or not positive,
``arraysize`` is used.
Returns:
The rows, fewer than ``size`` when the result is exhausted.
"""
if not size or size <= 0:
size = self._arraysize
rows = []
for _ in range(size):
row = await self.fetchone()
if row:
rows.append(row)
else:
break
return rows
[docs]
@override
async def fetchall( # type: ignore[override]
self,
) -> list[tuple[Any | None, ...] | dict[Any, Any | None]]:
"""Fetch all remaining rows from the result set.
Returns:
The remaining rows.
"""
rows = []
while True:
row = await self.fetchone()
if row:
rows.append(row)
else:
break
return rows
[docs]
@override
def __iter__(self) -> NoReturn:
"""Reject synchronous iteration; use ``async for`` instead.
Raises:
TypeError: Always, because the fetch methods are coroutines.
"""
raise TypeError(f"'{type(self).__name__}' object is not iterable; use 'async for' instead.")
def __aiter__(self):
return self
async def __anext__(self):
row = await self.fetchone()
if row is None:
raise StopAsyncIteration
return row
[docs]
class AthenaAioDictResultSet(AthenaDictResultSet, AthenaAioResultSet):
"""Async result set that returns rows as dictionaries.
Inherits ``_get_rows`` from ``AthenaDictResultSet`` and async fetch
methods from ``AthenaAioResultSet`` via multiple inheritance.
"""