"""Base classes and callback types shared by PyAthena cursors."""
from __future__ import annotations
import logging
import sys
import threading
import time
from abc import ABCMeta, abstractmethod
from collections.abc import Callable
from concurrent.futures import Future, wait
from datetime import UTC, datetime, timedelta
from typing import TYPE_CHECKING, Any, TypeVar, cast
from botocore.exceptions import BotoCoreError, ClientError
import pyathena
from pyathena.converter import Converter, DefaultTypeConverter
from pyathena.error import DatabaseError, OperationalError, ProgrammingError
from pyathena.formatter import Formatter
from pyathena.glue import GlueMetadataClient
from pyathena.model import (
AthenaCalculationExecution,
AthenaCalculationExecutionStatus,
AthenaCompression,
AthenaDatabase,
AthenaFileFormat,
AthenaQueryExecution,
AthenaTableMetadata,
)
from pyathena.options import ExecuteOptions
from pyathena.util import (
RetryConfig,
_get_error_code,
_is_throttling_error,
retry_api_call,
)
if TYPE_CHECKING:
from pyathena.connection import Connection
_logger = logging.getLogger(__name__)
_T = TypeVar("_T")
# How often a wait for a start request wakes up to check for Ctrl-C, so that
# a KeyboardInterrupt is raised promptly where an untimed lock wait cannot be
# interrupted by signals (Windows before Python 3.14).
_INTERRUPT_CHECK_INTERVAL = 0.1
OnPollCallback = Callable[[AthenaQueryExecution | AthenaCalculationExecutionStatus], None]
"""Type of the optional ``on_poll`` callback.
Invoked once per poll iteration with the current execution object: an
:class:`~pyathena.model.AthenaQueryExecution` for SQL queries, or an
:class:`~pyathena.model.AthenaCalculationExecutionStatus` for Spark calculations.
"""
[docs]
class CursorIterator(metaclass=ABCMeta):
"""Abstract base class providing iteration and result fetching capabilities for cursors.
This mixin class provides common functionality for iterating through query results
and managing cursor state. It implements the iterator protocol and provides
standard fetch methods that conform to the DB API 2.0 specification.
Attributes:
DEFAULT_FETCH_SIZE: Default number of rows to fetch per request (1000).
DEFAULT_RESULT_REUSE_MINUTES: Default minutes for Athena result reuse (60).
arraysize: Number of rows to fetch with fetchmany() if size not specified.
Note:
This is an abstract base class used by concrete cursor implementations.
It should not be instantiated directly.
"""
# https://docs.aws.amazon.com/athena/latest/APIReference/API_GetQueryResults.html
# Valid Range: Minimum value of 1. Maximum value of 1000.
DEFAULT_FETCH_SIZE: int = 1000
# https://docs.aws.amazon.com/athena/latest/APIReference/API_ResultReuseByAgeConfiguration.html
# Specifies, in minutes, the maximum age of a previous query result
# that Athena should consider for reuse. The default is 60.
DEFAULT_RESULT_REUSE_MINUTES = 60
[docs]
def __init__(self, **kwargs) -> None:
"""Initialize the iterator with no current row and an unknown row count.
Args:
**kwargs: Keyword arguments, of which only ``arraysize`` is used. If it is
absent, ``DEFAULT_FETCH_SIZE`` is used.
Raises:
ProgrammingError: If ``arraysize`` is outside the range the
``arraysize`` setter accepts.
"""
super().__init__()
self.arraysize: int = kwargs.get("arraysize", self.DEFAULT_FETCH_SIZE)
self._rownumber: int | None = None
self._rowcount: int = -1 # By default, return -1 to indicate that this is not supported.
@property
def arraysize(self) -> int:
"""The default number of rows per ``fetchmany()`` call."""
return self._arraysize
@arraysize.setter
def arraysize(self, value: int) -> None:
if value <= 0 or value > self.DEFAULT_FETCH_SIZE:
raise ProgrammingError(
f"MaxResults is more than maximum allowed length {self.DEFAULT_FETCH_SIZE}."
)
self._arraysize = value
@property
def rownumber(self) -> int | None:
"""The zero-based index of the next row, or None if it is unknown."""
return self._rownumber
@property
def rowcount(self) -> int:
"""The number of rows affected by the last operation, or -1 if it is unknown."""
return self._rowcount
[docs]
@abstractmethod
def fetchone(self):
"""Fetch the next row of the result."""
raise NotImplementedError # pragma: no cover
[docs]
@abstractmethod
def fetchmany(self):
"""Fetch the next set of rows of the result."""
raise NotImplementedError # pragma: no cover
[docs]
@abstractmethod
def fetchall(self):
"""Fetch all remaining rows of the result."""
raise NotImplementedError # pragma: no cover
def __next__(self):
row = self.fetchone()
if row is None:
raise StopIteration
return row
def __iter__(self):
return self
[docs]
class BaseCursor(metaclass=ABCMeta):
"""Abstract base class for all PyAthena cursor implementations.
This class provides the foundational functionality for executing SQL queries
and calculations on Amazon Athena. It handles AWS API interactions, query
execution management, metadata operations, and result polling.
All concrete cursor implementations (Cursor, DictCursor, PandasCursor,
ArrowCursor, SparkCursor, AsyncCursor) inherit from this base class and
implement the abstract methods according to their specific use cases.
Attributes:
LIST_QUERY_EXECUTIONS_MAX_RESULTS: Maximum results per query listing API call (50).
LIST_TABLE_METADATA_MAX_RESULTS: Maximum results per table metadata API call (50).
LIST_DATABASES_MAX_RESULTS: Maximum results per database listing API call (50).
Key Features:
- Query execution and polling with configurable retry logic
- Table and database metadata operations
- Result caching and reuse capabilities
- Encryption and security configuration support
- Workgroup and catalog management
- Query cancellation and interruption handling
Example:
This is an abstract base class and should not be instantiated directly.
Use concrete implementations like Cursor or PandasCursor instead:
>>> cursor = connection.cursor() # Creates default Cursor
>>> cursor.execute("SELECT * FROM my_table")
>>> results = cursor.fetchall()
Note:
This class contains AWS service quotas as constants. These limits
are enforced by the AWS Athena service and should not be modified.
"""
# https://docs.aws.amazon.com/athena/latest/APIReference/API_ListQueryExecutions.html
# Valid Range: Minimum value of 0. Maximum value of 50.
LIST_QUERY_EXECUTIONS_MAX_RESULTS = 50
# https://docs.aws.amazon.com/athena/latest/APIReference/API_ListTableMetadata.html
# Valid Range: Minimum value of 1. Maximum value of 50.
LIST_TABLE_METADATA_MAX_RESULTS = 50
# https://docs.aws.amazon.com/athena/latest/APIReference/API_ListDatabases.html
# Valid Range: Minimum value of 1. Maximum value of 50.
LIST_DATABASES_MAX_RESULTS = 50
[docs]
def __init__(
self,
connection: Connection[Any],
converter: Converter,
formatter: Formatter,
retry_config: RetryConfig,
s3_staging_dir: str | None,
schema_name: str | None,
catalog_name: str | None,
work_group: str | None,
poll_interval: float,
encryption_option: str | None,
kms_key: str | None,
kill_on_interrupt: bool,
result_reuse_enable: bool,
result_reuse_minutes: int,
on_start_query_execution: Callable[[str], None] | None = None,
on_poll: OnPollCallback | None = None,
**kwargs,
) -> None:
"""Initialize the cursor with the settings it uses to run queries.
Args:
connection: The connection that created the cursor.
converter: Converter for result values.
formatter: Formatter for query parameters.
retry_config: Retry configuration for API calls.
s3_staging_dir: S3 location for query results.
schema_name: Default schema name.
catalog_name: Default catalog name.
work_group: Athena workgroup name.
poll_interval: Query status polling interval in seconds.
encryption_option: S3 encryption option (SSE_S3, SSE_KMS, CSE_KMS).
kms_key: KMS key for encryption.
kill_on_interrupt: Cancel the execution when a ``KeyboardInterrupt`` interrupts
starting it or waiting for it.
result_reuse_enable: Enable Athena query result reuse.
result_reuse_minutes: Maximum age in minutes of a reused result.
on_start_query_execution: Callback invoked with each query ID before the cursor
waits for the query, by cursors whose ``execute()`` supports it.
on_poll: Callback invoked once per poll iteration with the current
execution object.
**kwargs: Ignored.
"""
super().__init__()
self._connection = connection
self._converter = converter
self._formatter = formatter
self._retry_config = retry_config
self._s3_staging_dir = s3_staging_dir
self._schema_name = schema_name
self._catalog_name = catalog_name
self._work_group = work_group
self._poll_interval = poll_interval
self._encryption_option = encryption_option
self._kms_key = kms_key
self._kill_on_interrupt = kill_on_interrupt
self._result_reuse_enable = result_reuse_enable
self._result_reuse_minutes = result_reuse_minutes
# ``on_start_query_execution`` is invoked by cursors whose ``execute()``
# supports it (the synchronous and aio cursors). Async/Spark cursors return
# the query id immediately through their execution model and do not invoke it.
self._on_start_query_execution = on_start_query_execution
self._on_poll = on_poll
[docs]
@staticmethod
def get_default_converter(unload: bool = False) -> DefaultTypeConverter | Any:
"""Get the default type converter for this cursor class.
Args:
unload: Whether the converter is for UNLOAD operations. Some cursor
types may return different converters for UNLOAD operations.
Returns:
The default type converter instance for this cursor type.
"""
return DefaultTypeConverter()
@property
def connection(self) -> Connection[Any]:
"""The connection that created this cursor."""
return self._connection
def _build_start_query_execution_request(
self,
query: str,
work_group: str | None = None,
s3_staging_dir: str | None = None,
result_reuse_enable: bool | None = None,
result_reuse_minutes: int | None = None,
execution_parameters: list[str] | None = None,
) -> dict[str, Any]:
request: dict[str, Any] = {
"QueryString": query,
"QueryExecutionContext": {},
}
if self._schema_name:
request["QueryExecutionContext"].update({"Database": self._schema_name})
if self._catalog_name:
request["QueryExecutionContext"].update({"Catalog": self._catalog_name})
result_configuration: dict[str, Any] = {}
if self._s3_staging_dir or s3_staging_dir:
result_configuration["OutputLocation"] = (
s3_staging_dir if s3_staging_dir else self._s3_staging_dir
)
if self._work_group or work_group:
request.update({"WorkGroup": work_group if work_group else self._work_group})
if self._encryption_option:
enc_conf = {
"EncryptionOption": self._encryption_option,
}
if self._kms_key:
enc_conf.update({"KmsKey": self._kms_key})
result_configuration["EncryptionConfiguration"] = enc_conf
if result_configuration:
request["ResultConfiguration"] = result_configuration
if self._result_reuse_enable or result_reuse_enable:
reuse_conf = {
"Enabled": result_reuse_enable
if result_reuse_enable is not None
else self._result_reuse_enable,
"MaxAgeInMinutes": result_reuse_minutes
if result_reuse_minutes is not None
else self._result_reuse_minutes,
}
request["ResultReuseConfiguration"] = {"ResultReuseByAgeConfiguration": reuse_conf}
if execution_parameters:
request["ExecutionParameters"] = execution_parameters
return request
def _build_start_calculation_execution_request(
self,
session_id: str,
code_block: str,
description: str | None = None,
client_request_token: str | None = None,
):
request: dict[str, Any] = {
"SessionId": session_id,
"CodeBlock": code_block,
}
if description:
request.update({"Description": description})
if client_request_token:
request.update({"ClientRequestToken": client_request_token})
return request
def _build_list_query_executions_request(
self,
work_group: str | None,
next_token: str | None = None,
max_results: int | None = None,
) -> dict[str, Any]:
request: dict[str, Any] = {
"MaxResults": max_results if max_results else self.LIST_QUERY_EXECUTIONS_MAX_RESULTS
}
if self._work_group or work_group:
request.update({"WorkGroup": work_group if work_group else self._work_group})
if next_token:
request.update({"NextToken": next_token})
return request
def _build_list_table_metadata_request(
self,
catalog_name: str | None,
schema_name: str | None,
expression: str | None = None,
next_token: str | None = None,
max_results: int | None = None,
) -> dict[str, Any]:
request: dict[str, Any] = {
"CatalogName": catalog_name if catalog_name else self._catalog_name,
"DatabaseName": schema_name if schema_name else self._schema_name,
"MaxResults": max_results if max_results else self.LIST_TABLE_METADATA_MAX_RESULTS,
}
if expression:
request.update({"Expression": expression})
if next_token:
request.update({"NextToken": next_token})
if self._work_group:
request.update({"WorkGroup": self._work_group})
return request
def _glue_catalog_name(self, catalog_name: str | None) -> str | None:
"""The catalog to read from Glue when Athena throttles.
Args:
catalog_name: The requested catalog, or None for the cursor's catalog.
Returns:
The catalog name if the fallback is on, the catalog is Glue-backed and
Glue is still reachable; otherwise None.
"""
catalog = catalog_name if catalog_name else self._catalog_name
connection = self._connection
if connection.glue_metadata_fallback and connection._glue.usable_for(catalog):
return catalog
return None
def _glue_request_failed(
self, e: BotoCoreError | ClientError, description: str, absence_is_final: bool
) -> None:
"""Raise Glue's answer that a table is absent; otherwise log the failure.
After a request that could not reach Glue at all, the connection's
``GlueMetadataClient`` stops using Glue, so later throttled requests
do not wait for it again.
Args:
e: The failed Glue request's exception.
description: What the request does, for the log message.
absence_is_final: Whether Glue's ``EntityNotFoundException`` answers
the request.
Raises:
OperationalError: For Glue's ``EntityNotFoundException`` when
``absence_is_final`` is set.
"""
if (
absence_is_final
and isinstance(e, ClientError)
and _get_error_code(e) == "EntityNotFoundException"
):
raise OperationalError(*e.args) from e
unreachable = isinstance(e, GlueMetadataClient.UNREACHABLE_ERRORS)
suffix = " and not using Glue again on this connection" if unreachable else ""
_logger.warning(
f"Glue request to {description} failed: {e}; retrying the Athena request{suffix}."
)
def _with_glue_fallback(
self,
catalog_name: str | None,
athena_request: Callable[[Callable[[BaseException], bool] | None, bool], _T],
glue_request: Callable[[GlueMetadataClient, str], _T],
description: str,
logging_: bool = True,
absence_is_final: bool = False,
) -> _T:
"""Run a metadata request, answering its throttling from Glue.
Athena rate-limits its metadata API per account, separately from Glue.
In a Glue-backed catalog the first attempt stops at a throttled
response and the request goes to Glue instead of through the retry
policy; other errors keep the policy. When the Glue request fails, the
Athena request continues with the policy. With ``absence_is_final``,
Glue's answer that the table does not exist is raised instead.
Args:
catalog_name: The requested catalog, or None for the cursor's catalog.
athena_request: Sends the Athena request; receives the predicate
that stops its retries (or None) and whether to log a failure.
glue_request: Sends the Glue request; receives the connection's
``GlueMetadataClient`` and the catalog name.
description: What the request does, for log messages.
logging_: Whether to log a failed request.
absence_is_final: Whether Glue's ``EntityNotFoundException`` answers
the request.
Returns:
The result of the Athena or the Glue request.
Raises:
OperationalError: If the request fails.
"""
glue_catalog = self._glue_catalog_name(catalog_name)
if glue_catalog is None:
return athena_request(None, logging_)
try:
return athena_request(_is_throttling_error, False)
except OperationalError as e:
if not _is_throttling_error(e.__cause__ or e):
if logging_:
_logger.exception(f"Failed to {description}.")
raise
_logger.warning(f"Request to {description} was throttled; reading it from Glue.")
try:
return glue_request(self._connection._glue, glue_catalog)
except (BotoCoreError, ClientError) as e:
self._glue_request_failed(e, description, absence_is_final)
return athena_request(None, logging_)
def _build_list_databases_request(
self,
catalog_name: str | None,
next_token: str | None = None,
max_results: int | None = None,
):
request: dict[str, Any] = {
"CatalogName": catalog_name if catalog_name else self._catalog_name,
"MaxResults": max_results if max_results else self.LIST_DATABASES_MAX_RESULTS,
}
if next_token:
request.update({"NextToken": next_token})
if self._work_group:
request.update({"WorkGroup": self._work_group})
return request
def _list_databases(
self,
catalog_name: str | None,
next_token: str | None = None,
max_results: int | None = None,
logging_: bool = True,
stop_on: Callable[[BaseException], bool] | None = None,
) -> tuple[str | None, list[AthenaDatabase]]:
"""List one page of the catalog's databases with ``ListDatabases``.
Args:
catalog_name: The catalog, or None for the cursor's catalog.
next_token: The token of the page to read.
max_results: The page size.
logging_: Whether to log a failed request.
stop_on: Stops the retries at an exception it accepts; used for
the first attempt of the Glue fallback.
Returns:
The next page's token, or None, and the page's databases.
Raises:
OperationalError: If the request fails.
"""
request = self._build_list_databases_request(
catalog_name=catalog_name,
next_token=next_token,
max_results=max_results,
)
try:
response = retry_api_call(
self.connection._client.list_databases,
config=self._retry_config,
logger=_logger,
stop_on=stop_on,
**request,
)
except Exception as e:
if logging_:
_logger.exception("Failed to list databases.")
raise OperationalError(*e.args) from e
else:
return response.get("NextToken"), [
AthenaDatabase({"Database": r}) for r in response.get("DatabaseList", [])
]
[docs]
def list_databases(
self,
catalog_name: str | None,
max_results: int | None = None,
) -> list[AthenaDatabase]:
# Pages already read are kept, so a retried request resumes after them.
"""List the catalog's databases.
In ``AwsDataCatalog`` and S3 Tables catalogs, a throttled request is
answered from the AWS Glue Data Catalog; see ``glue_metadata_fallback``.
Args:
catalog_name: The catalog, or None for the cursor's catalog.
max_results: The page size of each request.
Returns:
The catalog's databases.
Raises:
OperationalError: If the request fails.
"""
databases: list[AthenaDatabase] = []
next_token = None
def athena_request(
stop_on: Callable[[BaseException], bool] | None, logging_: bool
) -> list[AthenaDatabase]:
nonlocal next_token
while True:
next_token, response = self._list_databases(
catalog_name=catalog_name,
next_token=next_token,
max_results=max_results,
logging_=logging_,
stop_on=stop_on,
)
databases.extend(response)
if not next_token:
return databases
return self._with_glue_fallback(
catalog_name, athena_request, GlueMetadataClient.list_databases, "list databases"
)
def _build_get_table_metadata_request(
self,
table_name: str,
catalog_name: str | None = None,
schema_name: str | None = None,
) -> dict[str, Any]:
request: dict[str, Any] = {
"CatalogName": catalog_name if catalog_name else self._catalog_name,
"DatabaseName": schema_name if schema_name else self._schema_name,
"TableName": table_name,
}
if self._work_group:
request.update({"WorkGroup": self._work_group})
return request
def _get_table_metadata(
self,
table_name: str,
catalog_name: str | None = None,
schema_name: str | None = None,
logging_: bool = True,
stop_on: Callable[[BaseException], bool] | None = None,
) -> AthenaTableMetadata:
"""Get one table's metadata with ``GetTableMetadata``.
Args:
table_name: The table name.
catalog_name: The catalog, or None for the cursor's catalog.
schema_name: The database, or None for the cursor's schema.
logging_: Whether to log a failed request.
stop_on: Stops the retries at an exception it accepts; used for
the first attempt of the Glue fallback.
Returns:
The table's metadata.
Raises:
OperationalError: If the request fails.
"""
request = self._build_get_table_metadata_request(
table_name=table_name,
catalog_name=catalog_name,
schema_name=schema_name,
)
try:
response = retry_api_call(
self._connection.client.get_table_metadata,
config=self._retry_config,
logger=_logger,
stop_on=stop_on,
**request,
)
except Exception as e:
if logging_:
_logger.exception("Failed to get table metadata.")
raise OperationalError(*e.args) from e
else:
return AthenaTableMetadata(response)
def _list_table_metadata(
self,
catalog_name: str | None = None,
schema_name: str | None = None,
expression: str | None = None,
next_token: str | None = None,
max_results: int | None = None,
logging_: bool = True,
stop_on: Callable[[BaseException], bool] | None = None,
) -> tuple[str | None, list[AthenaTableMetadata]]:
"""List one page of a database's table metadata with ``ListTableMetadata``.
Args:
catalog_name: The catalog, or None for the cursor's catalog.
schema_name: The database, or None for the cursor's schema.
expression: A table name pattern.
next_token: The token of the page to read.
max_results: The page size.
logging_: Whether to log a failed request.
stop_on: Stops the retries at an exception it accepts; used for
the first attempt of the Glue fallback.
Returns:
The next page's token, or None, and the page's table metadata.
Raises:
OperationalError: If the request fails.
"""
request = self._build_list_table_metadata_request(
catalog_name=catalog_name,
schema_name=schema_name,
expression=expression,
next_token=next_token,
max_results=max_results,
)
try:
response = retry_api_call(
self.connection._client.list_table_metadata,
config=self._retry_config,
logger=_logger,
stop_on=stop_on,
**request,
)
except Exception as e:
if logging_:
_logger.exception("Failed to list table metadata.")
raise OperationalError(*e.args) from e
else:
return response.get("NextToken"), [
AthenaTableMetadata({"TableMetadata": r})
for r in response.get("TableMetadataList", [])
]
def _get_query_execution(self, query_id: str) -> AthenaQueryExecution:
"""Get a query execution with ``GetQueryExecution``.
Args:
query_id: The query execution ID.
Returns:
The query execution.
Raises:
OperationalError: If the request fails.
"""
request: dict[str, Any] = {"QueryExecutionId": query_id}
try:
response = retry_api_call(
self._connection.client.get_query_execution,
config=self._retry_config,
logger=_logger,
**request,
)
except Exception as e:
_logger.exception("Failed to get query execution.")
raise OperationalError(*e.args) from e
else:
return AthenaQueryExecution(response)
def _get_calculation_execution_status(self, query_id: str) -> AthenaCalculationExecutionStatus:
"""Get a calculation's status with ``GetCalculationExecutionStatus``.
Args:
query_id: The calculation execution ID.
Returns:
The calculation's status.
Raises:
OperationalError: If the request fails.
"""
request: dict[str, Any] = {"CalculationExecutionId": query_id}
try:
response = retry_api_call(
self._connection.client.get_calculation_execution_status,
config=self._retry_config,
logger=_logger,
**request,
)
except Exception as e:
_logger.exception("Failed to get calculation execution status.")
raise OperationalError(*e.args) from e
else:
return AthenaCalculationExecutionStatus(response)
def _get_calculation_execution(self, query_id: str) -> AthenaCalculationExecution:
"""Get a calculation execution with ``GetCalculationExecution``.
Args:
query_id: The calculation execution ID.
Returns:
The calculation execution.
Raises:
OperationalError: If the request fails.
"""
request: dict[str, Any] = {"CalculationExecutionId": query_id}
try:
response = retry_api_call(
self._connection.client.get_calculation_execution,
config=self._retry_config,
logger=_logger,
**request,
)
except Exception as e:
_logger.exception("Failed to get calculation execution.")
raise OperationalError(*e.args) from e
else:
return AthenaCalculationExecution(response)
def _batch_get_query_execution(self, query_ids: list[str]) -> list[AthenaQueryExecution]:
try:
response = retry_api_call(
self.connection._client.batch_get_query_execution,
config=self._retry_config,
logger=_logger,
QueryExecutionIds=query_ids,
)
except Exception as e:
_logger.exception("Failed to batch get query execution.")
raise OperationalError(*e.args) from e
else:
return [
AthenaQueryExecution({"QueryExecution": r})
for r in response.get("QueryExecutions", [])
]
def _list_query_executions(
self,
work_group: str | None = None,
next_token: str | None = None,
max_results: int | None = None,
) -> tuple[str | None, list[AthenaQueryExecution]]:
request = self._build_list_query_executions_request(
work_group=work_group, next_token=next_token, max_results=max_results
)
try:
response = retry_api_call(
self.connection._client.list_query_executions,
config=self._retry_config,
logger=_logger,
**request,
)
except Exception as e:
_logger.exception("Failed to list query executions.")
raise OperationalError(*e.args) from e
else:
next_token = response.get("NextToken")
query_ids = response.get("QueryExecutionIds")
if not query_ids:
return next_token, []
return next_token, self._batch_get_query_execution(query_ids)
def _poll_until_terminal(
self, query_id: str
) -> AthenaQueryExecution | AthenaCalculationExecution:
"""Poll a query execution until it reaches a terminal state.
Calls ``on_poll`` with every status and sleeps ``poll_interval`` seconds
between requests.
Args:
query_id: The query execution ID.
Returns:
The query execution in a terminal state.
Raises:
OperationalError: If a status request fails.
"""
while True:
query_execution = self._get_query_execution(query_id)
if self._on_poll:
self._on_poll(query_execution)
if query_execution.state in AthenaQueryExecution.TERMINAL_STATES:
return query_execution
time.sleep(self._poll_interval)
def _poll(self, query_id: str) -> AthenaQueryExecution | AthenaCalculationExecution:
"""Wait for an execution to reach a terminal state.
On ``KeyboardInterrupt`` with ``kill_on_interrupt`` enabled, requests
cancellation with ``_cancel_and_wait()`` and re-raises the interrupt.
Cancellation is a best-effort request, so the execution can still end in
another terminal state.
Args:
query_id: The execution ID.
Returns:
The execution in a terminal state.
Raises:
KeyboardInterrupt: If interrupted while waiting. A failure to cancel or
wait for the execution becomes its ``__cause__``.
OperationalError: If a status request fails.
"""
try:
return self._poll_until_terminal(query_id)
except KeyboardInterrupt as interrupt:
if not self._kill_on_interrupt:
raise
_logger.warning("Query canceled by user.")
try:
self._cancel_and_wait(query_id)
except Exception as e:
raise interrupt from e
raise
def _cancel_and_wait(self, query_id: str) -> None:
"""Request cancellation of an execution and wait for a terminal state.
Args:
query_id: The execution ID.
Raises:
OperationalError: If the cancellation or a status request fails.
"""
self._cancel(query_id)
self._poll_until_terminal(query_id)
def _start_execution(self, start: Callable[[], str]) -> str:
"""Send a start request so that an interrupt stops the execution it starts.
With ``kill_on_interrupt`` enabled, the request runs on a helper thread.
On ``KeyboardInterrupt``, the cursor first tries to abandon the request.
This succeeds only if the helper has not begun the request by then; the
helper then never sends it, and the interrupt propagates. Otherwise the
cursor waits for the request to finish, records the execution ID with
``_set_interrupted_execution_id()``, requests cancellation with
``_cancel_and_wait()``, and re-raises the interrupt. Another
``KeyboardInterrupt`` during that wait propagates at once.
Args:
start: Sends the start request and returns the execution ID.
Returns:
The execution ID.
Raises:
KeyboardInterrupt: If interrupted while starting. A failure to start,
cancel, or wait for the execution becomes its ``__cause__``.
DatabaseError: If the request fails.
"""
if not self._kill_on_interrupt:
return start()
future: Future[str] = Future()
def run() -> None:
# Begin the request only if no interrupt has given up on it yet.
if not future.set_running_or_notify_cancel():
return
try:
future.set_result(start())
except BaseException as e:
future.set_exception(e)
try:
threading.Thread(target=run, name="pyathena-start", daemon=True).start()
return self._wait_for_start(future)
except KeyboardInterrupt as interrupt:
if future.cancel():
# The helper has not begun the request and never will.
raise
_logger.warning("Query canceled by user.")
try:
execution_id = self._wait_for_start(future)
self._set_interrupted_execution_id(execution_id)
self._cancel_and_wait(execution_id)
except Exception as e:
raise interrupt from e
raise
@staticmethod
def _wait_for_start(future: Future[str]) -> str:
"""Wait for a start request on a helper thread to finish.
Args:
future: The future of the start request.
Returns:
The execution ID.
Raises:
DatabaseError: If the request failed.
"""
while not future.done():
wait((future,), timeout=_INTERRUPT_CHECK_INTERVAL)
return future.result()
def _set_interrupted_execution_id(self, execution_id: str) -> None: # noqa: B027
"""Record the ID of an execution started by an interrupted start request.
Does nothing by default; cursors that expose the execution ID override this.
Args:
execution_id: The execution ID.
"""
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,
work_group: str | None,
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
try:
next_token = None
while cache_size > 0:
max_results = min(cache_size, self.LIST_QUERY_EXECUTIONS_MAX_RESULTS)
cache_size -= max_results
next_token, query_executions = self._list_query_executions(
work_group, next_token=next_token, max_results=max_results
)
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)
return query_id
def _prepare_query(
self,
operation: str,
parameters: dict[str, Any] | list[str] | None = None,
paramstyle: str | None = None,
) -> tuple[str, list[str] | None]:
"""Format query and build execution parameters. No I/O.
Args:
operation: SQL query string.
parameters: Query parameters.
paramstyle: Parameter style override.
Returns:
Tuple of (formatted_query, execution_parameters).
"""
if pyathena.paramstyle == "qmark" or paramstyle == "qmark":
query = operation
execution_parameters = cast(list[str] | None, parameters)
else:
query = self._formatter.format(operation, cast(dict[str, Any] | None, parameters))
execution_parameters = None
_logger.debug(query)
return query, execution_parameters
def _prepare_unload(
self,
operation: str,
s3_staging_dir: str | None,
) -> tuple[str, str | None]:
"""Wrap operation with UNLOAD if enabled.
Args:
operation: SQL query string.
s3_staging_dir: S3 location for query results.
Returns:
Tuple of (possibly-wrapped operation, unload_location or None).
"""
if not getattr(self, "_unload", False):
return operation, None
s3_staging_dir = s3_staging_dir if s3_staging_dir else self._s3_staging_dir
if not s3_staging_dir:
raise ProgrammingError("If the unload option is used, s3_staging_dir is required.")
return self._formatter.wrap_unload(
operation,
s3_staging_dir=s3_staging_dir,
format_=AthenaFileFormat.FILE_FORMAT_PARQUET,
compression=AthenaCompression.COMPRESSION_SNAPPY,
)
def _call_on_start_query_execution(self, query_id: str, options: ExecuteOptions) -> None:
"""Invoke the connection-level and execute-level query-start callbacks.
Both callbacks are invoked if set. Called by cursors whose execution
model supports early access to the query ID (the synchronous and aio
cursors) once ``_execute()`` returns it: after the StartQueryExecution
API call, or with a reusable query ID found through ``cache_size``.
"""
if self._on_start_query_execution:
self._on_start_query_execution(query_id)
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,
parameters: dict[str, Any] | list[str] | None = None,
work_group: str | None = None,
s3_staging_dir: str | None = None,
cache_size: int | None = None,
cache_expiration_time: int | None = None,
result_reuse_enable: bool | None = None,
result_reuse_minutes: int | None = None,
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. A query with execution parameters (``qmark``)
always starts a new execution.
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:
KeyboardInterrupt: If interrupted while starting the query; see
``_start_execution()``.
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).
options = ExecuteOptions.resolve(
options,
work_group=work_group,
s3_staging_dir=s3_staging_dir,
cache_size=cache_size,
cache_expiration_time=cache_expiration_time,
result_reuse_enable=result_reuse_enable,
result_reuse_minutes=result_reuse_minutes,
paramstyle=paramstyle,
)
query, request = self._build_execute_request(operation, parameters, options)
query_id = None
# Athena does not return the ExecutionParameters of earlier executions,
# so the cache cannot tell which parameters an execution ran with (#941).
if not request.get("ExecutionParameters"):
query_id = self._find_previous_query_id(
query,
options.work_group,
cache_size=options.cache_size,
cache_expiration_time=options.cache_expiration_time,
)
if query_id is None:
query_id = self._start_execution(lambda: self._start_query_execution(request))
return query_id
def _start_query_execution(self, request: dict[str, Any]) -> str:
"""Send a ``StartQueryExecution`` request.
Args:
request: The request parameters.
Returns:
The query execution ID.
Raises:
DatabaseError: If the request fails.
"""
try:
response = retry_api_call(
self._connection.client.start_query_execution,
config=self._retry_config,
logger=_logger,
**request,
)
except Exception as e:
_logger.exception("Failed to execute query.")
raise DatabaseError(*e.args) from e
return cast(str, response.get("QueryExecutionId"))
[docs]
@abstractmethod
def execute(
self,
operation: str,
parameters: dict[str, Any] | list[str] | None = None,
**kwargs,
):
"""Execute a SQL query.
Args:
operation: SQL query string.
parameters: Query parameters.
**kwargs: Execution options defined by the cursor implementation.
"""
raise NotImplementedError # pragma: no cover
[docs]
@abstractmethod
def executemany(
self,
operation: str,
seq_of_parameters: list[dict[str, Any] | list[str] | None],
**kwargs,
) -> None:
"""Execute a SQL query once for each set of parameters.
Args:
operation: SQL query string.
seq_of_parameters: Sequence of parameter sets.
**kwargs: Execution options defined by the cursor implementation.
"""
raise NotImplementedError # pragma: no cover
[docs]
@abstractmethod
def close(self) -> None:
"""Close the cursor."""
raise NotImplementedError # pragma: no cover
def _cancel(self, query_id: str) -> None:
"""Stop a query execution with ``StopQueryExecution``.
Args:
query_id: The query execution ID.
Raises:
OperationalError: If the request fails.
"""
request: dict[str, Any] = {"QueryExecutionId": query_id}
try:
retry_api_call(
self._connection.client.stop_query_execution,
config=self._retry_config,
logger=_logger,
**request,
)
except Exception as e:
_logger.exception("Failed to cancel query.")
raise OperationalError(*e.args) from e
[docs]
def setoutputsize(self, size, column=None): # noqa: B027
"""Accept a column buffer size as DB API 2.0 requires, and ignore it.
Args:
size: Buffer size for large columns.
column: Index of the column the size applies to, or None for all
large columns.
"""
def __enter__(self):
return self
def __exit__(self, exc_type, exc_val, exc_tb):
self.close()