Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
41 changes: 33 additions & 8 deletions pyathena/aio/common.py
Original file line number Diff line number Diff line change
Expand Up @@ -113,27 +113,52 @@ async def _get_query_execution(self, query_id: str) -> AthenaQueryExecution: #
else:
return AthenaQueryExecution(response)

async def __poll(self, query_id: str) -> AthenaQueryExecution:
async def _poll_until_terminal(self, query_id: str) -> AthenaQueryExecution: # type: ignore[override]
"""Poll a query execution until it reaches a terminal state.

Calls ``on_poll`` with every status and awaits ``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 = await self._get_query_execution(query_id)
if self._on_poll:
self._on_poll(query_execution)
if query_execution.state in [
AthenaQueryExecution.STATE_SUCCEEDED,
AthenaQueryExecution.STATE_FAILED,
AthenaQueryExecution.STATE_CANCELLED,
]:
if query_execution.state in AthenaQueryExecution.TERMINAL_STATES:
return query_execution
await asyncio.sleep(self._poll_interval)

async def _poll(self, query_id: str) -> AthenaQueryExecution: # type: ignore[override]
"""Wait for a query execution to finish.

On ``asyncio.CancelledError`` with ``kill_on_interrupt`` enabled, stops the
query and returns its final execution instead of re-raising.

Args:
query_id: The query execution ID.

Returns:
The query execution in a terminal state.

Raises:
asyncio.CancelledError: If cancelled and ``kill_on_interrupt`` is disabled.
OperationalError: If a status or stop request fails.
"""
try:
query_execution = await self.__poll(query_id)
query_execution = await self._poll_until_terminal(query_id)
except asyncio.CancelledError:
if self._kill_on_interrupt:
_logger.warning("Query canceled by user.")
await self._cancel(query_id)
query_execution = await self.__poll(query_id)
query_execution = await self._poll_until_terminal(query_id)
else:
raise
return query_execution
Expand Down
47 changes: 31 additions & 16 deletions pyathena/aio/result_set.py
Original file line number Diff line number Diff line change
Expand Up @@ -85,21 +85,26 @@ async def create(
await result_set._async_pre_fetch()
return result_set

async def __async_get_query_results(
async def _async_get_query_results(
self, max_results: int, next_token: str | None = None
) -> dict[str, Any]:
if not self.query_id:
raise ProgrammingError("QueryExecutionId is none or empty.")
if self.state != AthenaQueryExecution.STATE_SUCCEEDED:
raise ProgrammingError("QueryExecutionState is not SUCCEEDED.")
"""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.")
request: dict[str, Any] = {
"QueryExecutionId": self.query_id,
"MaxResults": max_results,
}
if next_token:
request["NextToken"] = next_token
try:
response = await async_retry_api_call(
self.connection.client.get_query_results,
Expand All @@ -113,18 +118,28 @@ async def __async_get_query_results(
else:
return cast(dict[str, Any], response)

async def __async_fetch(self, next_token: str | None = None) -> dict[str, Any]:
return await self.__async_get_query_results(self._arraysize, next_token)

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_fetch(self._next_token)
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:
response = await self.__async_fetch()
"""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)
Expand Down
44 changes: 29 additions & 15 deletions pyathena/aio/spark/cursor.py
Original file line number Diff line number Diff line change
Expand Up @@ -169,21 +169,23 @@ async def _calculate( # type: ignore[override]
client_request_token=client_request_token or str(uuid.uuid4()),
)
if not self._kill_on_interrupt:
return await self.__start_calculation(request)
return await self._start_calculation_execution(request)

start = asyncio.ensure_future(self.__start_calculation(request))
start = asyncio.ensure_future(self._start_calculation_execution(request))
try:
return await asyncio.shield(start)
except asyncio.CancelledError as cancellation:
_logger.warning("Query canceled by user.")
try:
self._calculation_id = await start
await self.__cancel_and_wait(self._calculation_id)
await self._cancel_and_wait(self._calculation_id)
except Exception as e:
raise cancellation from e
raise

async def __start_calculation(self, request: dict[str, Any]) -> str:
async def _start_calculation_execution( # type: ignore[override]
self, request: dict[str, Any]
) -> str:
"""Send a ``StartCalculationExecution`` request.

Args:
Expand All @@ -207,16 +209,28 @@ async def __start_calculation(self, request: dict[str, Any]) -> str:
raise DatabaseError(*e.args) from e
return cast(str, response.get("CalculationExecutionId"))

async def __poll(self, query_id: str) -> AthenaQueryExecution | AthenaCalculationExecution:
async def _poll_until_terminal( # type: ignore[override]
self, query_id: str
) -> AthenaQueryExecution | AthenaCalculationExecution:
"""Poll a calculation execution until it reaches a terminal state.

Calls ``on_poll`` with every status and awaits ``poll_interval`` seconds
between requests.

Args:
query_id: The calculation execution ID.

Returns:
The calculation execution in a terminal state.

Raises:
OperationalError: If a status request fails.
"""
while True:
calculation_status = await self._get_calculation_execution_status(query_id)
if self._on_poll:
self._on_poll(calculation_status)
if calculation_status.state in [
AthenaCalculationExecutionStatus.STATE_COMPLETED,
AthenaCalculationExecutionStatus.STATE_FAILED,
AthenaCalculationExecutionStatus.STATE_CANCELED,
]:
if calculation_status.state in AthenaCalculationExecutionStatus.TERMINAL_STATES:
return await self._get_calculation_execution(query_id)
await asyncio.sleep(self._poll_interval)

Expand Down Expand Up @@ -244,18 +258,18 @@ async def _poll( # type: ignore[override]
OperationalError: If a status request fails.
"""
try:
return await self.__poll(query_id)
return await self._poll_until_terminal(query_id)
except asyncio.CancelledError as cancellation:
if not self._kill_on_interrupt:
raise
_logger.warning("Query canceled by user.")
try:
await self.__cancel_and_wait(query_id)
await self._cancel_and_wait(query_id)
except Exception as e:
raise cancellation from e
raise

async def __cancel_and_wait(self, calculation_id: str) -> None:
async def _cancel_and_wait(self, calculation_id: str) -> None: # type: ignore[override]
"""Request cancellation and store the calculation's terminal state.

Args:
Expand All @@ -266,7 +280,7 @@ async def __cancel_and_wait(self, calculation_id: str) -> None:
"""
await self._cancel(calculation_id)
self._calculation_execution = cast(
AthenaCalculationExecution, await self.__poll(calculation_id)
AthenaCalculationExecution, await self._poll_until_terminal(calculation_id)
)

async def _cancel(self, query_id: str) -> None: # type: ignore[override]
Expand All @@ -292,7 +306,7 @@ async def _terminate_session(self) -> None: # type: ignore[override]
**request,
)
except Exception as e:
_logger.exception("Failed to terminate session.")
_logger.exception(f"Failed to terminate session: {self._session_id}.")
raise OperationalError(*e.args) from e

async def _read_s3_file_as_text(self, uri) -> str: # type: ignore[override]
Expand Down
9 changes: 7 additions & 2 deletions pyathena/arrow/result_set.py
Original file line number Diff line number Diff line change
Expand Up @@ -110,7 +110,7 @@ def __init__(
self._connect_timeout = connect_timeout
self._request_timeout = request_timeout
self._kwargs = kwargs
self._fs = self.__s3_file_system()
self._fs = self._create_s3_file_system()
if self.state == AthenaQueryExecution.STATE_SUCCEEDED and self.output_location:
self._table = self._as_arrow()
elif self.state == AthenaQueryExecution.STATE_SUCCEEDED:
Expand All @@ -121,7 +121,12 @@ def __init__(
self._table = pa.Table.from_pydict({})
self._batches = iter(self._table.to_batches(arraysize))

def __s3_file_system(self):
def _create_s3_file_system(self):
"""Create a pyarrow ``S3FileSystem`` from the connection settings.

Returns:
The pyarrow S3 filesystem for reading the query results.
"""
from pyarrow import fs

connection = self.connection
Expand Down
43 changes: 35 additions & 8 deletions pyathena/common.py
Original file line number Diff line number Diff line change
Expand Up @@ -862,27 +862,54 @@ def _list_query_executions(
return next_token, []
return next_token, self._batch_get_query_execution(query_ids)

def __poll(self, query_id: str) -> AthenaQueryExecution | AthenaCalculationExecution:
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.STATE_SUCCEEDED,
AthenaQueryExecution.STATE_FAILED,
AthenaQueryExecution.STATE_CANCELLED,
]:
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 a query execution to finish.

On ``KeyboardInterrupt`` with ``kill_on_interrupt`` enabled, stops the query
and returns its final execution instead of re-raising.

Args:
query_id: The query execution ID.

Returns:
The query execution in a terminal state.

Raises:
KeyboardInterrupt: If interrupted and ``kill_on_interrupt`` is disabled.
OperationalError: If a status or stop request fails.
"""
try:
query_execution = self.__poll(query_id)
query_execution = self._poll_until_terminal(query_id)
except KeyboardInterrupt as e:
if self._kill_on_interrupt:
_logger.warning("Query canceled by user.")
self._cancel(query_id)
query_execution = self.__poll(query_id)
query_execution = self._poll_until_terminal(query_id)
else:
raise e
return query_execution
Expand Down
2 changes: 2 additions & 0 deletions pyathena/model.py
Original file line number Diff line number Diff line change
Expand Up @@ -49,6 +49,7 @@ class AthenaQueryExecution:
STATE_SUCCEEDED: str = "SUCCEEDED"
STATE_FAILED: str = "FAILED"
STATE_CANCELLED: str = "CANCELLED"
TERMINAL_STATES: tuple[str, ...] = (STATE_SUCCEEDED, STATE_FAILED, STATE_CANCELLED)

Copy link
Copy Markdown
Member Author

Choose a reason for hiding this comment

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

Self-review round two: claims, callers, and operational behavior. Result: FINDINGS in the PR description only; corrected. No code changes.

Scope: git diff 7b77a536342b00694a9e76c4edf9b4811d47bca0..f4cc848b8502cb4292ca932b8341cf898728454d, plus the PR body, commit message, and changed docstrings.

Claims checked:

  • "16 non-dunder name-mangled helpers", "16 → 0". git grep -nE "def __[a-z0-9_]*[a-z0-9]\(" <sha> -- pyathena/ returns 16 at the base and 0 at the head.
  • Async overrides "as _poll / _cancel / _get_query_execution already do". aio/common.py:90, :139, and :166 are async def with # type: ignore[override].
  • Signatures of _execute, _poll, _cancel, _get_query_execution, _find_previous_query_id, and _terminate_session unchanged (Backwards incompatible changes in 3.35 breaks dbt projects #734). The diff adds or removes no def line for any of them.
  • _create_s3_file_system "the name AthenaS3FSResultSet already uses". See s3fs/result_set.py:107.
  • _terminate_session_by_id docstring: "session startup calls this synchronously in every cursor variant". AioSparkCursor also starts its session in SparkBaseCursor.__init__, which runs sync, and it does not override _terminate_session_by_id.
  • "Checks run in the same order as before". Query ID, then state, then closed. The request dict is built before the closed check, which does no I/O.

Corrections to the PR body:

  1. "Mangled helpers cannot be reused or overridden" overstated it: a subclass can still call self._Class__name. Reworded to "cannot override mangled helpers and can reuse them only through the mangled name".
  2. The body did not say that TERMINAL_STATES is a public addition. It now does: "The only public addition is the TERMINAL_STATES class constant", noting that AthenaCalculationExecution inherits it.

Existing and downstream callers:

  • The new single-underscore names become visible to subclasses.
  • GitHub code search of dbt-labs/dbt-adapters (dbt-athena), run on 2026-09-28, found no definitions or uses of _poll_until_terminal, _get_query_results, _async_get_query_results, _build_get_query_results_request, _cancel_and_wait, _start_calculation_execution, _wait_for_calculation_start, _terminate_session_by_id, or _create_s3_file_system.
  • The same batch found TERMINAL_STATES only as dbt-spark's and dbt-bigquery's own constants, which shows the search was working.
  • Limitation: code search covers only the default branch. Other downstream subclasses were not surveyed.

AWS operator: the same APIs are called the same number of times. Polling cadence, retries, and interrupt and cancel paths are unchanged, and no quota or cost change is claimed.

Documentation: no docs page references the removed names; git grep over docs/ finds none.

Evidence:

  • The AWS results in the PR body are local runs and are labelled by tree.
  • The rename-only run used an uncommitted working tree, before the shared changes.
  • On f4cc848 only the SQL cursor and aio cursor tests ran.
  • The Spark, format-cursor, and SQLAlchemy suites on the published head come from the Ready CI.


STATEMENT_TYPE_DDL: str = "DDL"
STATEMENT_TYPE_DML: str = "DML"
Expand Down Expand Up @@ -353,6 +354,7 @@ class AthenaCalculationExecutionStatus:
STATE_CANCELED: str = "CANCELED"
STATE_COMPLETED: str = "COMPLETED"
STATE_FAILED: str = "FAILED"
TERMINAL_STATES: tuple[str, ...] = (STATE_COMPLETED, STATE_FAILED, STATE_CANCELED)

def __init__(self, response: dict[str, Any]) -> None:
status = response.get("Status")
Expand Down
9 changes: 7 additions & 2 deletions pyathena/pandas/result_set.py
Original file line number Diff line number Diff line change
Expand Up @@ -304,7 +304,7 @@ def __init__(
self._auto_optimize_chunksize = auto_optimize_chunksize
self._data_manifest: list[str] = []
self._kwargs = kwargs
self._fs = self.__s3_file_system()
self._fs = self._create_s3_file_system()
self._csv_stream: TextIOWrapper | None = None

# Cache time column names for efficient _trunc_date processing
Expand Down Expand Up @@ -427,7 +427,12 @@ def _auto_determine_chunksize(self, file_size_bytes: int) -> int | None:
return self.AUTO_CHUNK_SIZE_MEDIUM
return None

def __s3_file_system(self):
def _create_s3_file_system(self):
"""Create PyAthena's ``S3FileSystem`` from the connection settings.

Returns:
The S3 filesystem for reading the query results.
"""
from pyathena.filesystem.s3 import S3FileSystem

return S3FileSystem(
Expand Down
Loading
Loading