Skip to content
Draft
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
31 changes: 31 additions & 0 deletions docs/aio.md
Original file line number Diff line number Diff line change
Expand Up @@ -130,6 +130,37 @@ async with await aio_connect(s3_staging_dir="s3://YOUR_S3_BUCKET/path/to/",
await cursor.cancel()
```

(aio-task-cancellation)=

### Task cancellation

With `kill_on_interrupt` enabled, which is the default, cancelling the task while `execute()` waits for the query
requests cancellation of the query, waits until it reaches a terminal state, and then raises `asyncio.CancelledError`.
Cancellation is a best-effort request, so the query can still end as `SUCCEEDED` or `FAILED`.
The `query_id` property keeps the ID of the cancelled query.
If the cancellation request fails, `asyncio.CancelledError` is raised with the error as its cause.
Cancelling the task while `execute()` is still starting the query first waits for the start request to finish,
and then cancels the query it started in the same way.
Cancelling the task again during the cancellation request or these waits raises `asyncio.CancelledError` immediately, and the query can keep running.
With `kill_on_interrupt=False`, `asyncio.CancelledError` is raised immediately and the query keeps running.

A timeout from `asyncio.wait_for()` therefore cancels the query and raises `asyncio.TimeoutError`.
`query_id` is `None` if the timeout expires before the start request is sent, for example while looking up a cached result.

```python
import asyncio

from pyathena import aio_connect

async with await aio_connect(s3_staging_dir="s3://YOUR_S3_BUCKET/path/to/",
region_name="us-west-2") as conn:
async with conn.cursor() as cursor:
try:
await asyncio.wait_for(cursor.execute("SELECT * FROM many_rows"), timeout=60)
except asyncio.TimeoutError:
print(f"Query timed out: {cursor.query_id}")
```

(aio-dict-cursor)=

## AioDictCursor
Expand Down
33 changes: 33 additions & 0 deletions docs/usage.md
Original file line number Diff line number Diff line change
Expand Up @@ -501,6 +501,39 @@ The `on_start_query_execution` callback is supported by the following cursor typ
Note: `AsyncCursor` and its variants do not support this callback as they already
return the query ID immediately through their different execution model.

## Query cancellation on interrupt

With `kill_on_interrupt` enabled, which is the default, a `KeyboardInterrupt` while `execute()` waits for the query
requests cancellation, waits until the query reaches a terminal state, and then propagates.
Cancellation is a best-effort request, so the query can still end as `SUCCEEDED` or `FAILED`.

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 2: claims, callers, and operations. Result: FINDINGS (PR description only, repaired)

Scope: c01c56f73c7dbf973fe73b52b6093f0a32952321..2de91e768196dc061251f7609ae0d9b9be2905cd, all claims in the PR description, commit message, _poll() docstrings, docs/usage.md, and docs/aio.md.

Claims checked:

  • Best-effort stop can still end SUCCEEDED: measured with one SELECT 1 on the CI account. StopQueryExecution on an already SUCCEEDED query returns HTTP 200 and the state stays SUCCEEDED. So the race re-raises the interrupt without a cause, and this sentence and the _poll() docstrings hold.
  • "If the cancellation request fails, ... as its cause": covered by the failure[cancel] tests on both bases.
  • kill_on_interrupt=False keeps the query running: no stop request is made (test_execute_without_kill_on_interrupt).
  • Timeouts: the commit message claim holds. On the original code, asyncio.wait_for() surfaces OperationalError (the new timeout test fails this way), and CPython's Timeout.__aexit__ converts only CancelledError into TimeoutError.
  • Existing callers: dbt-athena's current connection manager uses boto3 directly, and its legacy PyAthena cursor overrides _poll() (dbt-athena/src/dbt/adapters/athena/connections_legacy.py:167), so it is unaffected. The kill_on_interrupt parameter descriptions in pyathena/connection.py:228 and the cursor docstrings remain accurate.

Findings, repaired in the PR description:

  1. The WHY claimed that TaskGroup could not handle the cancellation. A TaskGroup still raises its ExceptionGroup when a sibling fails, so this is narrowed to the verified asyncio.wait_for() effect.
  2. The description omitted the executemany() consequence: it now stops at the interrupted execution, where an interrupted execution that ended SUCCEEDED used to let the loop continue. This is added, along with the dbt-athena compatibility note and the live stop measurement.

Deferred (pre-existing, out of scope): the thread-pool Async* cursor docstrings (e.g. pyathena/arrow/async_cursor.py:91) say kill_on_interrupt cancels on keyboard interrupt, but their polling runs in executor threads, which never receive KeyboardInterrupt. This PR does not change that path.

The `query_id` property keeps the ID of the interrupted query.
If the cancellation request fails, the `KeyboardInterrupt` propagates with the error as its cause.

A `KeyboardInterrupt` while `execute()` is still starting the query first waits for the
[StartQueryExecution](https://docs.aws.amazon.com/athena/latest/APIReference/API_StartQueryExecution.html)
request to finish, and then cancels the query it started in the same way.
The `query_id` property returns that query's ID.
If the request has not been sent yet when the interrupt is handled, it is never sent.
`AsyncCursor` and its variants also stop a query whose start is interrupted in `execute()`.
They wait for queries on worker threads, which do not receive `KeyboardInterrupt`.

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 2 (expanded scope): claims, callers, and operations. Result: CLEAN

Scope: 659676c07e2c09397b7cbc5740cc10bfe6fe41cb..d66c6994034894313f6b68f23f9f5d421ac13632, the rewritten PR description, docstrings, and docs/usage.md / docs/aio.md.

Claims checked:

  • No SQL token generation needed: botocore's Athena model marks StartQueryExecution.ClientRequestToken as idempotencyToken (auto-generated per call) and StartCalculationExecution's not. RetryConfig defaults to THROTTLING_ERROR_CODES, which do not start a query.
  • "AsyncCursor and its variants also stop a query whose start is interrupted": pyathena/async_cursor.py:225 and the pandas/arrow/polars/s3fs async cursors call self._execute() on the caller's thread. Polling runs in self._executor.
  • "If the request has not been sent yet ..., it is never sent": this is the sync abandon path, covered by the Cancel a Spark calculation interrupted while it is being started #861 tests test_calculate_interrupted_before_request_is_sent[False/True], which now exercise the shared helper.
  • Real signal: a SIGINT via os.kill while Cursor.execute() was blocked in a mocked StartQueryExecution produced one start request, a stop with the returned ID, query_id set, and KeyboardInterrupt without a cause. Checked on 3.13.1 and 3.10.16.
  • Operational: one short-lived daemon thread per StartQueryExecution with kill_on_interrupt, and no extra AWS calls on the normal path. The start still goes through retry_api_call with the same config.
  • Evidence scope: local results are offline only (3.13.1 and 3.10.16). The AWS suites were not run locally for this revision. An accidental local run of all of tests/pyathena with --noconftest sent some queries from tests that call connect() directly; its missing-schema failures are not used as evidence, and the PR description says so.

No corrections were needed beyond the round 1 repair.


A second `KeyboardInterrupt` during the cancellation request or these waits propagates immediately, and the query can keep running.
With `kill_on_interrupt=False`, the `KeyboardInterrupt` propagates immediately and the query keeps running.

```python
from pyathena import connect

cursor = connect(s3_staging_dir="s3://YOUR_S3_BUCKET/path/to/",
region_name="us-west-2").cursor()
try:
cursor.execute("SELECT * FROM many_rows")
except KeyboardInterrupt:
print(f"Query {cursor.query_id} was interrupted")
raise
```

For the native asyncio cursors, see {ref}`aio-task-cancellation`.

## Query polling callback

PyAthena provides an `on_poll` callback that is invoked once per poll iteration with the
Expand Down
200 changes: 179 additions & 21 deletions pyathena/aio/common.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,7 +3,7 @@
import asyncio
import logging
import sys
from collections.abc import Awaitable, Callable
from collections.abc import Awaitable, Callable, Coroutine
from datetime import datetime, timedelta, timezone
from typing import Any, TypeVar, cast

Expand All @@ -23,6 +23,71 @@
_T = TypeVar("_T")


async def _start_interruptibly(
start: Coroutine[Any, Any, str], stop: Callable[[str], Awaitable[None]]
) -> str:
"""Send a start request shielded from task cancellation, stopping what it starts.

On ``asyncio.CancelledError``, waits for the request to finish, stops the
execution it started, and re-raises the cancellation. Another cancellation
during that wait propagates at once.

Args:
start: Sends the start request and returns the execution ID.
stop: Requests cancellation of an execution and waits for a terminal state.

Returns:
The execution ID.

Raises:
asyncio.CancelledError: If the task is cancelled while starting. A failure
to start or stop the execution becomes its ``__cause__``.
"""
task = asyncio.ensure_future(start)
try:
return await asyncio.shield(task)
except asyncio.CancelledError as cancellation:
_logger.warning("Query canceled by user.")
try:
await stop(await task)
except Exception as e:
raise cancellation from e
raise


async def _poll_interruptibly(
execution_id: str,
poll: Callable[[str], Awaitable[_T]],
stop: Callable[[str], Awaitable[None]],
) -> _T:
"""Wait for an execution to reach a terminal state, stopping it on cancellation.

On ``asyncio.CancelledError``, stops the execution and re-raises the
cancellation. Another cancellation while stopping propagates at once.

Args:
execution_id: The execution ID.
poll: Waits for the execution to reach a terminal state.
stop: Requests cancellation of the execution and waits for a terminal state.

Returns:
The result of ``poll``.

Raises:
asyncio.CancelledError: If the task is cancelled while waiting. A failure
to stop the execution becomes its ``__cause__``.
"""
try:
return await poll(execution_id)
except asyncio.CancelledError as cancellation:
_logger.warning("Query canceled by user.")
try:
await stop(execution_id)
except Exception as e:
raise cancellation from e
raise


class AioBaseCursor(BaseCursor):
"""Async base cursor that overrides I/O methods with async equivalents.

Expand All @@ -44,6 +109,36 @@ async def _execute( # type: ignore[override]
paramstyle: str | None = None,
options: ExecuteOptions | None = None,
) -> str:
"""Start a query execution, or find a cached one, and return its ID.

With ``kill_on_interrupt`` enabled, the ``StartQueryExecution`` request is
shielded from task cancellation. On cancellation, the cursor waits for the
request to finish, records the query ID with ``_set_interrupted_query_id()``,
requests cancellation of the query, waits for a terminal state, and
re-raises ``asyncio.CancelledError``.

Args:
operation: SQL query string to execute.
parameters: Query parameters.
work_group: Athena workgroup.
s3_staging_dir: S3 location for query results.
cache_size: Number of recent executions to search for a cached result.
cache_expiration_time: Maximum age in seconds of a cached result.
result_reuse_enable: Whether to enable Athena's result reuse.
result_reuse_minutes: Maximum age in minutes of a reused result.
paramstyle: Parameter style of the query.
options: Shared execution options. Individual keyword arguments take
precedence over its fields.

Returns:
The query execution ID.

Raises:
asyncio.CancelledError: If the task is cancelled while starting the
query. A failure to start or cancel the query becomes its
``__cause__``.
DatabaseError: If the request fails.
"""
# The individual keyword arguments are retained for backward compatibility
# with external callers that predate ExecuteOptions, mirroring
# BaseCursor._execute().
Expand Down Expand Up @@ -74,19 +169,50 @@ async def _execute( # type: ignore[override]
cache_expiration_time=options.cache_expiration_time,
)
if query_id is None:
try:
response = await async_retry_api_call(
self._connection.client.start_query_execution,
config=self._retry_config,
logger=_logger,
**request,
if self._kill_on_interrupt:
query_id = await _start_interruptibly(
self.__start_query_execution(request), self.__stop_started_query
)
query_id = response.get("QueryExecutionId")
except Exception as e:
_logger.exception("Failed to execute query.")
raise DatabaseError(*e.args) from e
else:
query_id = await self.__start_query_execution(request)
return query_id

async 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 = await async_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"))

async def __stop_started_query(self, query_id: str) -> None:
"""Record a query that a cancelled start request started, and stop it.

Args:
query_id: The query execution ID.

Raises:
OperationalError: If the cancellation or a status request fails.
"""
self._set_interrupted_query_id(query_id)
await self.__cancel_and_wait(query_id)

async def _get_query_execution(self, query_id: str) -> AthenaQueryExecution: # type: ignore[override]
"""Get a query execution with ``GetQueryExecution``.

Expand Down Expand Up @@ -127,16 +253,40 @@ async def __poll(self, query_id: str) -> AthenaQueryExecution:
await asyncio.sleep(self._poll_interval)

async def _poll(self, query_id: str) -> AthenaQueryExecution: # type: ignore[override]
try:
query_execution = await self.__poll(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)
else:
raise
return query_execution
"""Wait for a query execution to reach a terminal state.

On task cancellation with ``kill_on_interrupt`` enabled, requests
cancellation, waits for the query to reach a terminal state, and re-raises
``asyncio.CancelledError``.
Cancellation is a best-effort request, so the terminal state can be
``SUCCEEDED`` or ``FAILED`` instead of ``CANCELLED``.

Args:
query_id: The query execution ID.

Returns:
The query execution in a terminal state.

Raises:
asyncio.CancelledError: If the task is cancelled while waiting. A failure
to cancel or wait for the query becomes its ``__cause__``.
OperationalError: If a status request fails.
"""
if not self._kill_on_interrupt:
return await self.__poll(query_id)
return await _poll_interruptibly(query_id, self.__poll, self.__cancel_and_wait)

async def __cancel_and_wait(self, query_id: str) -> None:
"""Request cancellation of a query and wait for a terminal state.

Args:
query_id: The query execution ID.

Raises:
OperationalError: If the cancellation or a status request fails.
"""
await self._cancel(query_id)
await self.__poll(query_id)

async def _cancel(self, query_id: str) -> None: # type: ignore[override]
"""Stop a query execution with ``StopQueryExecution``.
Expand Down Expand Up @@ -634,6 +784,14 @@ def query_id(self) -> str | None:
def query_id(self, val) -> None:
self._query_id = val

def _set_interrupted_query_id(self, query_id: str) -> None:
"""Expose the ID of a query started by a cancelled start request.

Args:
query_id: The query execution ID.
"""
self.query_id = query_id

@property
def rownumber(self) -> int | None:
return self.result_set.rownumber if self.result_set else None
Expand Down
41 changes: 20 additions & 21 deletions pyathena/aio/spark/cursor.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,6 +12,7 @@
import uuid
from typing import Any, cast

from pyathena.aio.common import _poll_interruptibly, _start_interruptibly
from pyathena.aio.util import async_retry_api_call
from pyathena.error import DatabaseError, NotSupportedError, OperationalError, ProgrammingError
from pyathena.model import (
Expand Down Expand Up @@ -170,18 +171,21 @@ async def _calculate( # type: ignore[override]
)
if not self._kill_on_interrupt:
return await self.__start_calculation(request)
return await _start_interruptibly(
self.__start_calculation(request), self.__stop_started_calculation
)

start = asyncio.ensure_future(self.__start_calculation(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)
except Exception as e:
raise cancellation from e
raise
async def __stop_started_calculation(self, calculation_id: str) -> None:
"""Store a calculation that a cancelled start request started, and stop it.

Args:
calculation_id: The calculation execution ID.

Raises:
OperationalError: If the cancellation or a status request fails.
"""
self._calculation_id = calculation_id
await self.__cancel_and_wait(calculation_id)

async def __start_calculation(self, request: dict[str, Any]) -> str:
"""Send a ``StartCalculationExecution`` request.
Expand Down Expand Up @@ -243,17 +247,9 @@ async def _poll( # type: ignore[override]
to cancel or wait for the calculation becomes its ``__cause__``.
OperationalError: If a status request fails.
"""
try:
if not self._kill_on_interrupt:
return await self.__poll(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)
except Exception as e:
raise cancellation from e
raise
return await _poll_interruptibly(query_id, self.__poll, self.__cancel_and_wait)

async def __cancel_and_wait(self, calculation_id: str) -> None:
"""Request cancellation and store the calculation's terminal state.
Expand Down Expand Up @@ -352,6 +348,9 @@ async def execute( # type: ignore[override]
Returns:
Self reference for method chaining.
"""
# A failure below must not leave the previous calculation on the cursor.
self._calculation_id = None
self._calculation_execution = None
self._calculation_id = await self._calculate(
session_id=session_id if session_id else self._session_id,
code_block=operation,
Expand Down
Loading
Loading