From c9b55624c6e37597feaaf8e60711d644e3956ca6 Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Fri, 2 Oct 2026 01:06:48 +0900 Subject: [PATCH 1/8] Re-raise and stop interrupted queries in the SQL cursors With kill_on_interrupt enabled, the SQL cursors returned the final execution after cancelling an interrupted query instead of re-raising the KeyboardInterrupt or asyncio.CancelledError, and an interrupt while StartQueryExecution was in flight left the started query running without an ID on the cursor. Move the Spark cursor's interrupt handling into BaseCursor (_poll, _cancel_and_wait, _start_execution, _wait_for_start) so that the SQL and synchronous Spark cursors share it, and give AioBaseCursor the asyncio counterparts. Cursors record the ID of a started execution through _set_interrupted_execution_id(), which WithResultSet and SparkBaseCursor override. The Spark cursors also clear the previous calculation before starting a new one. Closes #840 Co-Authored-By: Claude Opus 5.5 --- docs/aio.md | 31 +++ docs/usage.md | 33 +++ pyathena/aio/common.py | 123 +++++++++-- pyathena/aio/spark/cursor.py | 3 + pyathena/common.py | 171 +++++++++++--- pyathena/result_set.py | 8 + pyathena/spark/common.py | 90 +------- pyathena/spark/cursor.py | 3 + tests/pyathena/aio/spark/test_cursor.py | 19 +- tests/pyathena/aio/test_cursor.py | 258 +++++++++++++++++++++- tests/pyathena/spark/test_common.py | 43 +--- tests/pyathena/spark/test_spark_cursor.py | 7 +- tests/pyathena/test_cursor.py | 213 +++++++++++++++++- tests/pyathena/util.py | 33 +++ 14 files changed, 860 insertions(+), 175 deletions(-) diff --git a/docs/aio.md b/docs/aio.md index 060bc1905..b429af071 100644 --- a/docs/aio.md +++ b/docs/aio.md @@ -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 diff --git a/docs/usage.md b/docs/usage.md index 227ad3be0..6896439ce 100644 --- a/docs/usage.md +++ b/docs/usage.md @@ -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`. +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`. + +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 diff --git a/pyathena/aio/common.py b/pyathena/aio/common.py index 6ad79f121..c56a3a668 100644 --- a/pyathena/aio/common.py +++ b/pyathena/aio/common.py @@ -2,7 +2,7 @@ import asyncio import logging -from collections.abc import Awaitable, Callable +from collections.abc import Awaitable, Callable, Coroutine from typing import Any, NoReturn, TypeVar, cast from botocore.exceptions import BotoCoreError, ClientError @@ -64,6 +64,8 @@ async def _execute( # type: ignore[override] The query execution ID. Raises: + asyncio.CancelledError: If the task is cancelled while starting the + query; see ``_start_execution()``. ProgrammingError: If the formatter rejects the query or its parameters. DatabaseError: If the ``StartQueryExecution`` request fails. """ @@ -88,18 +90,74 @@ async def _execute( # type: ignore[override] cache_expiration_time=options.cache_expiration_time, ) if query_id is None: + query_id = await self._start_execution(self._start_query_execution(request)) + return query_id + + @override + async def _start_query_execution(self, request: dict[str, Any]) -> str: # type: ignore[override] + """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")) + + @override + async def _start_execution( # type: ignore[override] + self, start: Coroutine[Any, Any, str] + ) -> str: + """Send a start request so that task cancellation stops the execution it starts. + + With ``kill_on_interrupt`` enabled, the request is shielded from task + cancellation. On cancellation, 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 + ``asyncio.CancelledError``. Another cancellation during that wait + propagates at once. + + Args: + start: Sends the start request and returns the execution ID. + + Returns: + The execution ID. + + Raises: + asyncio.CancelledError: If the task is cancelled 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 await start + + task = asyncio.ensure_future(start) + try: + return await asyncio.shield(task) + except asyncio.CancelledError as cancellation: + _logger.warning("Query canceled by user.") try: - response = await async_retry_api_call( - self._connection.client.start_query_execution, - config=self._retry_config, - logger=_logger, - **request, - ) - query_id = response.get("QueryExecutionId") + execution_id = await task + self._set_interrupted_execution_id(execution_id) + await self._cancel_and_wait(execution_id) except Exception as e: - _logger.exception("Failed to execute query.") - raise DatabaseError(*e.args) from e - return query_id + raise cancellation from e + raise @override async def _get_query_execution(self, query_id: str) -> AthenaQueryExecution: # type: ignore[override] @@ -154,10 +212,12 @@ async def _poll_until_terminal(self, query_id: str) -> AthenaQueryExecution: # @override async def _poll(self, query_id: str) -> AthenaQueryExecution: # type: ignore[override] - """Wait for a query execution to finish. + """Wait for a query execution to reach a terminal state. - On ``asyncio.CancelledError`` with ``kill_on_interrupt`` enabled, stops the - query and returns its final execution instead of re-raising. + On task cancellation with ``kill_on_interrupt`` enabled, requests + cancellation with ``_cancel_and_wait()`` and re-raises + ``asyncio.CancelledError``. Cancellation is a best-effort request, so the + query can still end as ``SUCCEEDED`` or ``FAILED`` instead of ``CANCELLED``. Args: query_id: The query execution ID. @@ -166,19 +226,34 @@ async def _poll(self, query_id: str) -> AthenaQueryExecution: # type: ignore[ov 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. + 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. """ try: - 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_until_terminal(query_id) - else: + return await self._poll_until_terminal(query_id) + except asyncio.CancelledError as cancellation: + if not self._kill_on_interrupt: raise - return query_execution + _logger.warning("Query canceled by user.") + try: + await self._cancel_and_wait(query_id) + except Exception as e: + raise cancellation from e + raise + + @override + async def _cancel_and_wait(self, query_id: str) -> None: # type: ignore[override] + """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_until_terminal(query_id) @override async def _cancel(self, query_id: str) -> None: # type: ignore[override] diff --git a/pyathena/aio/spark/cursor.py b/pyathena/aio/spark/cursor.py index 6c8ece8aa..f84e671ce 100644 --- a/pyathena/aio/spark/cursor.py +++ b/pyathena/aio/spark/cursor.py @@ -340,6 +340,9 @@ async def execute( 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, diff --git a/pyathena/common.py b/pyathena/common.py index 963ec4599..8797bfb9d 100644 --- a/pyathena/common.py +++ b/pyathena/common.py @@ -2,9 +2,11 @@ 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 @@ -39,6 +41,11 @@ _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. @@ -888,31 +895,126 @@ def _poll_until_terminal( time.sleep(self._poll_interval) def _poll(self, query_id: str) -> AthenaQueryExecution | AthenaCalculationExecution: - """Wait for a query execution to finish. + """Wait for an execution to reach a terminal state. - On ``KeyboardInterrupt`` with ``kill_on_interrupt`` enabled, stops the query - and returns its final execution instead of re-raising. + 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 query execution ID. + query_id: The execution ID. Returns: - The query execution in a terminal state. + The execution in a terminal state. Raises: - KeyboardInterrupt: If interrupted and ``kill_on_interrupt`` is disabled. - OperationalError: If a status or stop request fails. + KeyboardInterrupt: If interrupted while waiting. A failure to cancel or + wait for the execution becomes its ``__cause__``. + OperationalError: If a status request fails. """ try: - 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_until_terminal(query_id) - else: - raise e - return query_execution + 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 @@ -1152,6 +1254,8 @@ def _execute( 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. """ @@ -1176,18 +1280,33 @@ def _execute( cache_expiration_time=options.cache_expiration_time, ) if query_id is None: - try: - query_id = retry_api_call( - self._connection.client.start_query_execution, - config=self._retry_config, - logger=_logger, - **request, - ).get("QueryExecutionId") - except Exception as e: - _logger.exception("Failed to execute query.") - raise DatabaseError(*e.args) from e + 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")) + @abstractmethod def execute( self, diff --git a/pyathena/result_set.py b/pyathena/result_set.py index dade34f42..8b6f0e3f1 100644 --- a/pyathena/result_set.py +++ b/pyathena/result_set.py @@ -901,6 +901,14 @@ def query_id(self) -> str | None: def query_id(self, val: str | None) -> None: self._query_id = val + def _set_interrupted_execution_id(self, execution_id: str) -> None: + """Keep the ID of a query started by an interrupted start request. + + Args: + execution_id: The query execution ID. + """ + self.query_id = execution_id + @property def query(self) -> str | None: if not self.result_set: diff --git a/pyathena/spark/common.py b/pyathena/spark/common.py index d4871ef9e..53a42e9e6 100644 --- a/pyathena/spark/common.py +++ b/pyathena/spark/common.py @@ -9,11 +9,9 @@ import contextlib import logging -import threading import time import uuid from abc import ABCMeta, abstractmethod -from concurrent.futures import Future, wait from datetime import datetime from typing import Any, cast @@ -31,11 +29,6 @@ _logger = logging.getLogger(__name__) -# How often a wait for the 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 - class SparkBaseCursor(BaseCursor, metaclass=ABCMeta): """Abstract base class for Spark-enabled cursor implementations. @@ -336,38 +329,6 @@ def _poll_until_terminal( time.sleep(self._poll_interval) @override - def _poll(self, query_id: str) -> AthenaQueryExecution | AthenaCalculationExecution: - """Wait for a calculation execution to reach a terminal state. - - On ``KeyboardInterrupt`` with ``kill_on_interrupt`` enabled, requests - cancellation, waits for the calculation to reach a terminal state, stores - it as the cursor's calculation execution, and re-raises the interrupt. - Cancellation is a best-effort request, so the terminal state can be - ``COMPLETED`` or ``FAILED`` instead of ``CANCELED``. - - Args: - query_id: The calculation execution ID. - - Returns: - The calculation execution in a terminal state. - - Raises: - KeyboardInterrupt: If interrupted while waiting. A failure to cancel or - wait for the calculation 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, calculation_id: str) -> None: """Request cancellation and store the calculation's terminal state. @@ -425,34 +386,16 @@ def _calculate( description=description, client_request_token=client_request_token or str(uuid.uuid4()), ) - if not self._kill_on_interrupt: - return self._start_calculation_execution(request) - - future: Future[str] = Future() + return self._start_execution(lambda: self._start_calculation_execution(request)) - def start() -> 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(self._start_calculation_execution(request)) - except BaseException as e: - future.set_exception(e) + @override + def _set_interrupted_execution_id(self, execution_id: str) -> None: + """Keep the ID of a calculation started by an interrupted start request. - try: - threading.Thread(target=start, name="pyathena-spark-start", daemon=True).start() - return self._wait_for_calculation_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: - self._calculation_id = self._wait_for_calculation_start(future) - self._cancel_and_wait(self._calculation_id) - except Exception as e: - raise interrupt from e - raise + Args: + execution_id: The calculation execution ID. + """ + self._calculation_id = execution_id def _start_calculation_execution(self, request: dict[str, Any]) -> str: """Send a ``StartCalculationExecution`` request. @@ -478,23 +421,6 @@ def _start_calculation_execution(self, request: dict[str, Any]) -> str: raise DatabaseError(*e.args) from e return cast(str, response.get("CalculationExecutionId")) - @staticmethod - def _wait_for_calculation_start(future: Future[str]) -> str: - """Wait for the start request on a helper thread to finish. - - Args: - future: The future of the start request. - - Returns: - The calculation execution ID. - - Raises: - DatabaseError: If the request failed. - """ - while not future.done(): - wait((future,), timeout=_INTERRUPT_CHECK_INTERVAL) - return future.result() - @override def _cancel(self, query_id: str) -> None: """Stop a calculation execution with ``StopCalculationExecution``. diff --git a/pyathena/spark/cursor.py b/pyathena/spark/cursor.py index 151c24858..ad6e655da 100644 --- a/pyathena/spark/cursor.py +++ b/pyathena/spark/cursor.py @@ -109,6 +109,9 @@ def execute( work_group: str | None = None, **kwargs, ) -> SparkCursor: + # A failure below must not leave the previous calculation on the cursor. + self._calculation_id = None + self._calculation_execution = None self._calculation_id = self._calculate( session_id=session_id if session_id else self._session_id, code_block=operation, diff --git a/tests/pyathena/aio/spark/test_cursor.py b/tests/pyathena/aio/spark/test_cursor.py index 92053ef57..ef5feeb20 100644 --- a/tests/pyathena/aio/spark/test_cursor.py +++ b/tests/pyathena/aio/spark/test_cursor.py @@ -280,6 +280,11 @@ async def test_execute_kill_on_interrupt_failure(self, failing): kill_on_interrupt=True, final_state=AthenaCalculationExecutionStatus.STATE_COMPLETED, ) + # Left by a previous calculation on the same cursor. + cursor._calculation_id = "previous_calculation_id" + cursor._calculation_execution = MagicMock( + state=AthenaCalculationExecutionStatus.STATE_COMPLETED + ) # Raise the cancellation from the first status request directly, so that the # test receives the re-raised exception itself rather than one made by a task. cursor._get_calculation_execution_status = AsyncMock( @@ -292,6 +297,7 @@ async def test_execute_kill_on_interrupt_failure(self, failing): assert exc_info.value.__cause__ is error cancel.assert_awaited_once_with("calculation_id") + assert cursor.calculation_id == "calculation_id" assert cursor.calculation_execution is None async def test_execute_cancellation_without_kill_on_interrupt(self): @@ -391,14 +397,17 @@ async def test_execute_cancelled_while_starting(self): async def test_execute_timeout_while_starting(self): cursor, cancel, started, release = _starting_cursor() - timer = threading.Timer(0.2, release.set) - timer.start() + task = asyncio.create_task(asyncio.wait_for(cursor.execute("code"), timeout=0.05)) try: - with pytest.raises(asyncio.TimeoutError): - await asyncio.wait_for(cursor.execute("code"), timeout=0.05) + assert await asyncio.to_thread(started.wait, _TIMEOUT) + # The timeout is due before this sleep ends, so the event loop handles it + # while the start request is still blocked. + await asyncio.sleep(0.1) + assert not task.done() finally: - timer.cancel() release.set() + with pytest.raises(asyncio.TimeoutError): + await task assert started.is_set() cancel.assert_awaited_once_with("calculation_id") diff --git a/tests/pyathena/aio/test_cursor.py b/tests/pyathena/aio/test_cursor.py index 9679854ae..7bece9167 100644 --- a/tests/pyathena/aio/test_cursor.py +++ b/tests/pyathena/aio/test_cursor.py @@ -5,6 +5,7 @@ from unittest.mock import AsyncMock, MagicMock, call, patch import pytest +from botocore.exceptions import ClientError from pyathena import BINARY, Binary, ExecuteOptions from pyathena.aio.cursor import AioCursor @@ -15,7 +16,93 @@ from pyathena.util import RetryConfig from tests import ENV from tests.pyathena.aio.conftest import _aio_connect -from tests.pyathena.util import succeeded_query_execution, throttle_metadata_api +from tests.pyathena.util import ( + EVENT_TIMEOUT, + succeeded_query_execution, + throttle_metadata_api, +) + + +def _offline_cursor(kill_on_interrupt, final_state): + """An AioCursor whose first status request blocks until the task is cancelled. + + Later status requests report ``RUNNING`` once, then ``final_state``. + + Args: + kill_on_interrupt: Whether the cursor cancels the query on cancellation. + final_state: The state of the query after cancellation. + + Returns: + The cursor, the mock of its cancellation request, and an event set when the + first status request starts. + """ + polling = asyncio.Event() + states = iter([AthenaQueryExecution.STATE_RUNNING, final_state]) + + async def get_query_execution(query_id): + if not polling.is_set(): + polling.set() + await asyncio.Event().wait() + return MagicMock(state=next(states)) + + cursor = AioCursor.__new__(AioCursor) # bypass __init__ to avoid AWS calls + cursor._rowcount = -1 + cursor._result_set = None + cursor._poll_interval = 0 + cursor._kill_on_interrupt = kill_on_interrupt + cursor._on_poll = None + cursor._on_start_query_execution = None + cursor._execute = AsyncMock(return_value="query_id") + cursor._get_query_execution = get_query_execution + # A successful query builds a result set from these. + cursor._connection = MagicMock() + cursor._converter = MagicMock() + cursor._arraysize = 1 + cursor._retry_config = RetryConfig() + cursor._result_set_class = MagicMock(create=AsyncMock()) + cancel = cursor._cancel = AsyncMock() + return cursor, cancel, polling + + +def _starting_cursor(kill_on_interrupt=True, response=None): + """An AioCursor whose StartQueryExecution request blocks in a thread until released. + + Status requests report ``RUNNING`` once, then ``CANCELLED``. + + Args: + kill_on_interrupt: Whether the cursor cancels the query on cancellation. + response: The exception the start request raises once released; the + default returns ``query_id``. + + Returns: + The cursor, the mock of its cancellation request, an event set when the + start request starts, and an event that releases it. + """ + started = threading.Event() + release = threading.Event() + + def start_query_execution(**kwargs): + started.set() + assert release.wait(EVENT_TIMEOUT) + if response: + raise response + return {"QueryExecutionId": "query_id"} + + cursor, cancel, _ = _offline_cursor(kill_on_interrupt, AthenaQueryExecution.STATE_CANCELLED) + del cursor._execute # use the real _execute() + cursor._query_id = None + cursor._prepare_query = MagicMock(return_value=("SELECT 1", None)) + cursor._build_start_query_execution_request = MagicMock(return_value={}) + cursor._find_previous_query_id = AsyncMock(return_value=None) + cursor._connection.client.start_query_execution.side_effect = start_query_execution + cursor._retry_config = RetryConfig(attempt=2, multiplier=0) + cursor._get_query_execution = AsyncMock( + side_effect=[ + MagicMock(state=AthenaQueryExecution.STATE_RUNNING), + MagicMock(state=AthenaQueryExecution.STATE_CANCELLED), + ] + ) + return cursor, cancel, started, release class TestAioCursor: @@ -115,6 +202,7 @@ async def test_execute_internal_legacy_kwargs_passthrough(self): "QueryExecutionId": "test_query_id" } cursor._retry_config = RetryConfig() + cursor._kill_on_interrupt = True with ( patch.object( @@ -152,6 +240,174 @@ async def test_execute_internal_legacy_kwargs_passthrough(self): cache_expiration_time=100, ) + @pytest.mark.parametrize( + "final_state", + [AthenaQueryExecution.STATE_CANCELLED, AthenaQueryExecution.STATE_SUCCEEDED], + ) + async def test_execute_kill_on_interrupt(self, final_state): + """Task cancellation cancels the query, waits for it, and is re-raised (no AWS).""" + polled = [] + cursor, cancel, polling = _offline_cursor(kill_on_interrupt=True, final_state=final_state) + cursor._on_poll = polled.append + task = asyncio.create_task(cursor.execute("SELECT 1")) + await polling.wait() + task.cancel() + with pytest.raises(asyncio.CancelledError): + await task + + assert task.cancelled() + cancel.assert_awaited_once_with("query_id") + # The cancellation is re-raised only after the query reaches a terminal state. + assert [execution.state for execution in polled] == [ + AthenaQueryExecution.STATE_RUNNING, + final_state, + ] + assert cursor.query_id == "query_id" + assert cursor.result_set is None + + async def test_execute_kill_on_interrupt_timeout(self): + """A timeout cancels the query and raises TimeoutError (no AWS).""" + cursor, cancel, _ = _offline_cursor( + kill_on_interrupt=True, final_state=AthenaQueryExecution.STATE_CANCELLED + ) + with pytest.raises(asyncio.TimeoutError): + await asyncio.wait_for(cursor.execute("SELECT 1"), timeout=0.01) + + cancel.assert_awaited_once_with("query_id") + + @pytest.mark.parametrize("failing", ["cancel", "wait"]) + async def test_execute_kill_on_interrupt_failure(self, failing): + """A failure to cancel or wait becomes the cause of the cancellation (no AWS).""" + error = OperationalError("failed") + cursor, cancel, _ = _offline_cursor( + kill_on_interrupt=True, final_state=AthenaQueryExecution.STATE_SUCCEEDED + ) + # Raise the cancellation from the first status request directly, so that the + # test receives the re-raised exception itself rather than one made by a task. + cursor._get_query_execution = AsyncMock(side_effect=[asyncio.CancelledError(), error]) + if failing == "cancel": + cancel.side_effect = error + with pytest.raises(asyncio.CancelledError) as exc_info: + await cursor.execute("SELECT 1") + + assert exc_info.value.__cause__ is error + cancel.assert_awaited_once_with("query_id") + + async def test_execute_without_kill_on_interrupt(self): + """Without kill_on_interrupt, cancellation propagates at once (no AWS).""" + cursor, cancel, polling = _offline_cursor( + kill_on_interrupt=False, final_state=AthenaQueryExecution.STATE_SUCCEEDED + ) + task = asyncio.create_task(cursor.execute("SELECT 1")) + await polling.wait() + task.cancel() + with pytest.raises(asyncio.CancelledError): + await task + + cancel.assert_not_awaited() + assert cursor.query_id == "query_id" + + async def test_execute_cancelled_while_starting(self): + """Cancellation during the start request stops the query it starts (no AWS).""" + polled = [] + cursor, cancel, started, release = _starting_cursor() + cursor._on_poll = polled.append + task = asyncio.create_task(cursor.execute("SELECT 1")) + assert await asyncio.to_thread(started.wait, EVENT_TIMEOUT) + task.cancel() + # Let the task handle the cancellation before the start request finishes. + for _ in range(5): + await asyncio.sleep(0) + assert not task.done() + release.set() + with pytest.raises(asyncio.CancelledError): + await task + + assert task.cancelled() + cursor._connection.client.start_query_execution.assert_called_once() + cancel.assert_awaited_once_with("query_id") + assert [execution.state for execution in polled] == [ + AthenaQueryExecution.STATE_RUNNING, + AthenaQueryExecution.STATE_CANCELLED, + ] + assert cursor.query_id == "query_id" + assert cursor.result_set is None + + async def test_execute_timeout_while_starting(self): + """A timeout during the start request stops the query and raises TimeoutError (no AWS).""" + cursor, cancel, started, release = _starting_cursor() + task = asyncio.create_task(asyncio.wait_for(cursor.execute("SELECT 1"), timeout=0.05)) + try: + assert await asyncio.to_thread(started.wait, EVENT_TIMEOUT) + # The timeout is due before this sleep ends, so the event loop handles it + # while the start request is still blocked. + await asyncio.sleep(0.1) + assert not task.done() + finally: + release.set() + with pytest.raises(asyncio.TimeoutError): + await task + + assert started.is_set() + cancel.assert_awaited_once_with("query_id") + assert cursor.query_id == "query_id" + + @pytest.mark.parametrize("failing", ["start", "cancel"]) + async def test_execute_cancelled_while_starting_failure(self, failing): + """A failure to start or cancel becomes the cause of the cancellation (no AWS).""" + error = OperationalError("failed") + cursor, cancel, started, release = _starting_cursor( + response=ClientError( + {"Error": {"Code": "InvalidRequestException", "Message": "failed"}}, + "StartQueryExecution", + ) + if failing == "start" + else None + ) + if failing == "cancel": + cancel.side_effect = error + raised = [] + + async def execute(): + try: + await cursor.execute("SELECT 1") + except asyncio.CancelledError as e: + raised.append(e) + raise + + task = asyncio.create_task(execute()) + assert await asyncio.to_thread(started.wait, EVENT_TIMEOUT) + task.cancel() + release.set() + with pytest.raises(asyncio.CancelledError): + await task + + assert task.cancelled() + if failing == "start": + assert isinstance(raised[0].__cause__, DatabaseError) + cancel.assert_not_awaited() + assert cursor.query_id is None + else: + assert raised[0].__cause__ is error + cancel.assert_awaited_once_with("query_id") + assert cursor.query_id == "query_id" + + async def test_execute_cancelled_while_starting_without_kill_on_interrupt(self): + """Without kill_on_interrupt, cancellation during the start propagates at once (no AWS).""" + cursor, cancel, started, release = _starting_cursor(kill_on_interrupt=False) + task = asyncio.create_task(cursor.execute("SELECT 1")) + assert await asyncio.to_thread(started.wait, EVENT_TIMEOUT) + try: + task.cancel() + with pytest.raises(asyncio.CancelledError): + await task + finally: + release.set() + + assert task.cancelled() + cancel.assert_not_awaited() + assert cursor.query_id is None + async def test_cache_size_different_schema(self): """A cached result is only reused when it ran against the same schema (#739). diff --git a/tests/pyathena/spark/test_common.py b/tests/pyathena/spark/test_common.py index 35ed2efff..5d16726e3 100644 --- a/tests/pyathena/spark/test_common.py +++ b/tests/pyathena/spark/test_common.py @@ -9,7 +9,6 @@ import logging import threading import uuid -from concurrent.futures import wait from unittest.mock import MagicMock, patch import pytest @@ -22,6 +21,7 @@ from pyathena.spark.common import SparkBaseCursor from pyathena.spark.cursor import SparkCursor from pyathena.util import RetryConfig +from tests.pyathena.util import interrupt_start_waits SPARK_CURSOR_CLASSES = [SparkCursor, AsyncSparkCursor, AioSparkCursor] SYNC_SPARK_CURSOR_CLASSES = [SparkCursor, AsyncSparkCursor] @@ -110,33 +110,6 @@ def start_calculation_execution(**kwargs): return started, release -def _interrupt_waits(started, release, interrupts=1): - """Patch the wait for the start request to raise KeyboardInterrupt. - - The first ``interrupts`` waits raise ``KeyboardInterrupt`` once the request - has started; later waits release the request and wait for it. - - Args: - started: Set when the start request starts. - release: Releases the start request. - interrupts: How many waits raise. - - Returns: - The patcher, and the list of raised interrupts. - """ - raised = [] - - def interrupting_wait(futures, timeout=None): - if len(raised) < interrupts: - assert started.wait(_TIMEOUT) - raised.append(KeyboardInterrupt()) - raise raised[-1] - release.set() - return wait(futures, timeout) - - return patch("pyathena.spark.common.wait", side_effect=interrupting_wait), raised - - def _init_cursor(cursor_class, connection, **kwargs): return cursor_class( connection=connection, @@ -261,7 +234,7 @@ def test_calculate_reuses_generated_token_on_retry(self, cursor_class, kill_on_i assert tokens[0] == tokens[1] uuid.UUID(tokens[0]) for thread in threading.enumerate(): - if thread.name == "pyathena-spark-start": + if thread.name == "pyathena-start": thread.join(_TIMEOUT) assert not thread.is_alive() @@ -302,7 +275,7 @@ def test_calculate_interrupted_while_starting(self, cursor_class, final_state): final_execution = MagicMock(state=final_state) cursor._get_calculation_execution.return_value = final_execution started, release = _block_start(cursor) - waits, raised = _interrupt_waits(started, release) + waits, raised = interrupt_start_waits(started, release) with waits, pytest.raises(KeyboardInterrupt) as exc_info: cursor._calculate(session_id="session_id", code_block="code") @@ -332,7 +305,7 @@ def test_calculate_interrupted_while_starting_failure(self, cursor_class, failin cursor._cancel.side_effect = error if failing == "wait": cursor._get_calculation_execution_status.side_effect = error - waits, raised = _interrupt_waits(started, release) + waits, raised = interrupt_start_waits(started, release) with waits, pytest.raises(KeyboardInterrupt) as exc_info: cursor._calculate(session_id="session_id", code_block="code") @@ -352,7 +325,7 @@ def test_calculate_interrupted_while_starting_failure(self, cursor_class, failin def test_calculate_second_interrupt_while_starting(self, cursor_class): cursor = _calculation_cursor(cursor_class) started, release = _block_start(cursor) - waits, raised = _interrupt_waits(started, release, interrupts=2) + waits, raised = interrupt_start_waits(started, release, interrupts=2) try: with waits, pytest.raises(KeyboardInterrupt) as exc_info: @@ -380,7 +353,7 @@ def start(self): raise KeyboardInterrupt with ( - patch("pyathena.spark.common.threading.Thread", InterruptedThread), + patch("pyathena.common.threading.Thread", InterruptedThread), pytest.raises(KeyboardInterrupt) as exc_info, ): cursor._calculate(session_id="session_id", code_block="code") @@ -399,7 +372,7 @@ def test_calculate_interrupt_without_kill_on_interrupt(self, cursor_class): cursor._connection.client.start_calculation_execution.side_effect = KeyboardInterrupt() with ( - patch("pyathena.spark.common.threading.Thread") as thread, + patch("pyathena.common.threading.Thread") as thread, pytest.raises(KeyboardInterrupt), ): cursor._calculate(session_id="session_id", code_block="code") @@ -411,7 +384,7 @@ def test_calculate_interrupt_without_kill_on_interrupt(self, cursor_class): def test_execute_interrupted_while_starting(self): cursor = _calculation_cursor(SparkCursor) started, release = _block_start(cursor) - waits, _ = _interrupt_waits(started, release) + waits, _ = interrupt_start_waits(started, release) with waits, pytest.raises(KeyboardInterrupt): cursor.execute("code") diff --git a/tests/pyathena/spark/test_spark_cursor.py b/tests/pyathena/spark/test_spark_cursor.py index 152678545..b8e6433ba 100644 --- a/tests/pyathena/spark/test_spark_cursor.py +++ b/tests/pyathena/spark/test_spark_cursor.py @@ -221,7 +221,11 @@ def test_execute_kill_on_interrupt_failure(self, failing): cursor._poll_interval = 0 cursor._kill_on_interrupt = True cursor._on_poll = None - cursor._calculation_execution = None + # Left by a previous calculation on the same cursor. + cursor._calculation_id = "previous_calculation_id" + cursor._calculation_execution = MagicMock( + state=AthenaCalculationExecutionStatus.STATE_COMPLETED + ) with ( patch.object(SparkCursor, "_calculate", return_value="calculation_id"), @@ -239,6 +243,7 @@ def test_execute_kill_on_interrupt_failure(self, failing): assert exc_info.value.__cause__ is error cancel.assert_called_once_with("calculation_id") + assert cursor.calculation_id == "calculation_id" assert cursor.calculation_execution is None def test_execute_interrupt_without_kill_on_interrupt(self): diff --git a/tests/pyathena/test_cursor.py b/tests/pyathena/test_cursor.py index 5533efcbe..720a5b230 100644 --- a/tests/pyathena/test_cursor.py +++ b/tests/pyathena/test_cursor.py @@ -29,6 +29,7 @@ Binary, ExecuteOptions, ) +from pyathena.async_cursor import AsyncCursor from pyathena.converter import _to_array, _to_map, _to_struct from pyathena.cursor import Cursor from pyathena.error import DatabaseError, NotSupportedError, OperationalError, ProgrammingError @@ -37,11 +38,74 @@ from tests import ENV from tests.pyathena.conftest import connect from tests.pyathena.tables import TABLES, VIEWS -from tests.pyathena.util import succeeded_query_execution, throttle_metadata_api, unreachable_glue +from tests.pyathena.util import ( + EVENT_TIMEOUT, + interrupt_start_waits, + succeeded_query_execution, + throttle_metadata_api, + unreachable_glue, +) _logger = logging.getLogger(__name__) +def _offline_cursor(kill_on_interrupt, cursor_class=Cursor): + """A cursor whose query requests go to a mocked Athena client. + + Args: + kill_on_interrupt: Whether the cursor cancels the query on interrupt. + cursor_class: The cursor class. + + Returns: + The cursor and the mock of its cancellation request. + """ + cursor = cursor_class.__new__(cursor_class) # bypass __init__ to avoid AWS calls + cursor._rowcount = -1 + cursor._result_set = None + cursor._query_id = None + cursor._poll_interval = 0 + cursor._kill_on_interrupt = kill_on_interrupt + cursor._on_poll = None + cursor._on_start_query_execution = None + cursor._prepare_query = MagicMock(return_value=("SELECT 1", None)) + cursor._build_start_query_execution_request = MagicMock(return_value={}) + cursor._find_previous_query_id = MagicMock(return_value=None) + cursor._connection = MagicMock() + cursor._connection.client.start_query_execution.return_value = {"QueryExecutionId": "query_id"} + cursor._retry_config = RetryConfig(attempt=2, multiplier=0) + # A successful query builds a result set from these. + cursor._converter = MagicMock() + cursor._arraysize = 1 + cursor._result_set_class = MagicMock() + cancel = cursor._cancel = MagicMock() + return cursor, cancel + + +def _block_start(cursor, response=None): + """Make the cursor's StartQueryExecution request block until released. + + Args: + cursor: The cursor from ``_offline_cursor``. + response: The exception to raise once released; the default returns + ``query_id``. + + Returns: + An event set when the request starts, and an event that releases it. + """ + started = threading.Event() + release = threading.Event() + + def start_query_execution(**kwargs): + started.set() + assert release.wait(EVENT_TIMEOUT) + if response: + raise response + return {"QueryExecutionId": "query_id"} + + cursor._connection.client.start_query_execution.side_effect = start_query_execution + return started, release + + class TestCursor: def test_fetchone(self, cursor): cursor.execute("SELECT * FROM one_row") @@ -1201,6 +1265,7 @@ def test_execute_internal_legacy_kwargs_passthrough(self): "QueryExecutionId": "test_query_id" } cursor._retry_config = RetryConfig() + cursor._kill_on_interrupt = True with ( patch.object( @@ -1354,6 +1419,152 @@ def test_on_poll_none_is_noop(self): assert result is execution + @pytest.mark.parametrize( + "final_state", + [AthenaQueryExecution.STATE_CANCELLED, AthenaQueryExecution.STATE_SUCCEEDED], + ) + def test_execute_kill_on_interrupt(self, final_state): + """An interrupt cancels the query, waits for it, and is re-raised (no AWS).""" + polled = [] + cursor, cancel = _offline_cursor(kill_on_interrupt=True) + cursor._on_poll = polled.append + cursor._get_query_execution = MagicMock( + side_effect=[ + KeyboardInterrupt(), + MagicMock(state=AthenaQueryExecution.STATE_RUNNING), + MagicMock(state=final_state), + ] + ) + with pytest.raises(KeyboardInterrupt): + cursor.execute("SELECT 1") + + cancel.assert_called_once_with("query_id") + # The interrupt is re-raised only after the query reaches a terminal state. + assert [execution.state for execution in polled] == [ + AthenaQueryExecution.STATE_RUNNING, + final_state, + ] + assert cursor.query_id == "query_id" + assert cursor.result_set is None + + @pytest.mark.parametrize("failing", ["cancel", "wait"]) + def test_execute_kill_on_interrupt_failure(self, failing): + """A failure to cancel or wait becomes the cause of the interrupt (no AWS).""" + error = OperationalError("failed") + cursor, cancel = _offline_cursor(kill_on_interrupt=True) + cursor._get_query_execution = MagicMock(side_effect=[KeyboardInterrupt(), error]) + if failing == "cancel": + cancel.side_effect = error + with pytest.raises(KeyboardInterrupt) as exc_info: + cursor.execute("SELECT 1") + + assert exc_info.value.__cause__ is error + cancel.assert_called_once_with("query_id") + + def test_execute_without_kill_on_interrupt(self): + """Without kill_on_interrupt, an interrupt propagates without cancellation (no AWS).""" + cursor, cancel = _offline_cursor(kill_on_interrupt=False) + cursor._get_query_execution = MagicMock(side_effect=[KeyboardInterrupt()]) + with pytest.raises(KeyboardInterrupt): + cursor.execute("SELECT 1") + + cancel.assert_not_called() + assert cursor.query_id == "query_id" + + @pytest.mark.parametrize( + "final_state", + [AthenaQueryExecution.STATE_CANCELLED, AthenaQueryExecution.STATE_SUCCEEDED], + ) + def test_execute_interrupted_while_starting(self, final_state): + """An interrupt during the start request stops the query it starts (no AWS).""" + polled = [] + cursor, cancel = _offline_cursor(kill_on_interrupt=True) + cursor._on_poll = polled.append + cursor._get_query_execution = MagicMock( + side_effect=[ + MagicMock(state=AthenaQueryExecution.STATE_RUNNING), + MagicMock(state=final_state), + ] + ) + started, release = _block_start(cursor) + waits, raised = interrupt_start_waits(started, release) + + with waits, pytest.raises(KeyboardInterrupt) as exc_info: + cursor.execute("SELECT 1") + + assert exc_info.value is raised[0] + assert exc_info.value.__cause__ is None + cursor._connection.client.start_query_execution.assert_called_once() + cancel.assert_called_once_with("query_id") + # The interrupt is re-raised only after the query reaches a terminal state. + assert [execution.state for execution in polled] == [ + AthenaQueryExecution.STATE_RUNNING, + final_state, + ] + assert cursor.query_id == "query_id" + assert cursor.result_set is None + + @pytest.mark.parametrize("failing", ["start", "cancel"]) + def test_execute_interrupted_while_starting_failure(self, failing): + """A failure to start or cancel becomes the cause of the interrupt (no AWS).""" + error = OperationalError("failed") + cursor, cancel = _offline_cursor(kill_on_interrupt=True) + started, release = _block_start( + cursor, + response=ClientError( + {"Error": {"Code": "InvalidRequestException", "Message": "failed"}}, + "StartQueryExecution", + ) + if failing == "start" + else None, + ) + if failing == "cancel": + cancel.side_effect = error + waits, raised = interrupt_start_waits(started, release) + + with waits, pytest.raises(KeyboardInterrupt) as exc_info: + cursor.execute("SELECT 1") + + assert exc_info.value is raised[0] + if failing == "start": + assert isinstance(exc_info.value.__cause__, DatabaseError) + cancel.assert_not_called() + assert cursor.query_id is None + else: + assert exc_info.value.__cause__ is error + cancel.assert_called_once_with("query_id") + assert cursor.query_id == "query_id" + + def test_execute_interrupted_while_starting_without_kill_on_interrupt(self): + """Without kill_on_interrupt, the request runs on the caller's thread (no AWS).""" + cursor, cancel = _offline_cursor(kill_on_interrupt=False) + cursor._connection.client.start_query_execution.side_effect = KeyboardInterrupt() + + with ( + patch("pyathena.common.threading.Thread") as thread, + pytest.raises(KeyboardInterrupt), + ): + cursor.execute("SELECT 1") + + thread.assert_not_called() + cancel.assert_not_called() + assert cursor.query_id is None + + def test_async_cursor_execute_interrupted_while_starting(self): + """AsyncCursor starts queries on the caller's thread, so it stops them too (no AWS).""" + cursor, cancel = _offline_cursor(kill_on_interrupt=True, cursor_class=AsyncCursor) + cursor._get_query_execution = MagicMock( + return_value=MagicMock(state=AthenaQueryExecution.STATE_CANCELLED) + ) + started, release = _block_start(cursor) + waits, raised = interrupt_start_waits(started, release) + + with waits, pytest.raises(KeyboardInterrupt) as exc_info: + cursor.execute("SELECT 1") + + assert exc_info.value is raised[0] + cancel.assert_called_once_with("query_id") + def test_on_poll_connection_level(self): """Connection-level on_poll fires during query execution.""" states = [] diff --git a/tests/pyathena/util.py b/tests/pyathena/util.py index bfe4dd3cc..4c6a9bf4d 100644 --- a/tests/pyathena/util.py +++ b/tests/pyathena/util.py @@ -6,7 +6,9 @@ # SPDX-License-Identifier: MIT import time +from concurrent.futures import wait from pathlib import Path +from unittest.mock import patch from botocore.config import Config from botocore.exceptions import ClientError @@ -161,3 +163,34 @@ def wait_for_spark_session_state(client, session_id, state, timeout=120): return time.sleep(1) raise AssertionError(f"Session {session_id} did not become {state} in {timeout} seconds.") + + +# Seconds a test waits for an event from a helper thread before failing. +EVENT_TIMEOUT = 10 + + +def interrupt_start_waits(started, release, interrupts=1): + """Patch the wait for a start request on a helper thread to raise KeyboardInterrupt. + + The first ``interrupts`` waits raise ``KeyboardInterrupt`` once the request + has started; later waits release the request and wait for it. + + Args: + started: Set when the start request starts. + release: Releases the start request. + interrupts: How many waits raise. + + Returns: + The patcher, and the list of raised interrupts. + """ + raised = [] + + def interrupting_wait(futures, timeout=None): + if len(raised) < interrupts: + assert started.wait(EVENT_TIMEOUT) + raised.append(KeyboardInterrupt()) + raise raised[-1] + release.set() + return wait(futures, timeout) + + return patch("pyathena.common.wait", side_effect=interrupting_wait), raised From 7a4f4ae36b151241f36a1f32f3c69d65258e2b38 Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Fri, 2 Oct 2026 01:17:53 +0900 Subject: [PATCH 2/8] Send no start request after an asyncio cancellation that preceded it The asyncio cursors sent StartQueryExecution or StartCalculationExecution and then stopped the execution when the task was cancelled after execute() had scheduled the shielded start task but before that task began. Skip the request when the caller has been cancelled since, as the synchronous cursors already do, so that it is never sent. Co-Authored-By: Claude Opus 5.5 --- pyathena/aio/common.py | 35 ++++++++++++++++++------- pyathena/aio/spark/cursor.py | 29 +++++++++++++++----- tests/pyathena/aio/spark/test_cursor.py | 17 ++++++++++++ tests/pyathena/aio/test_cursor.py | 17 ++++++++++++ 4 files changed, 81 insertions(+), 17 deletions(-) diff --git a/pyathena/aio/common.py b/pyathena/aio/common.py index c56a3a668..d143bcd0a 100644 --- a/pyathena/aio/common.py +++ b/pyathena/aio/common.py @@ -90,7 +90,7 @@ async def _execute( # type: ignore[override] cache_expiration_time=options.cache_expiration_time, ) if query_id is None: - query_id = await self._start_execution(self._start_query_execution(request)) + query_id = await self._start_execution(lambda: self._start_query_execution(request)) return query_id @override @@ -120,19 +120,22 @@ async def _start_query_execution(self, request: dict[str, Any]) -> str: # type: @override async def _start_execution( # type: ignore[override] - self, start: Coroutine[Any, Any, str] + self, start: Callable[[], Coroutine[Any, Any, str]] ) -> str: """Send a start request so that task cancellation stops the execution it starts. - With ``kill_on_interrupt`` enabled, the request is shielded from task - cancellation. On cancellation, the cursor waits for the request to finish, - records the execution ID with ``_set_interrupted_execution_id()``, requests + With ``kill_on_interrupt`` enabled, the request runs in a task shielded + from task cancellation. On cancellation, the request is abandoned if that + task has not begun it by then; it is never sent, and the cancellation + 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 ``asyncio.CancelledError``. Another cancellation during that wait propagates at once. Args: - start: Sends the start request and returns the execution ID. + start: Returns a coroutine that sends the start request and returns the + execution ID. Returns: The execution ID. @@ -144,15 +147,27 @@ async def _start_execution( # type: ignore[override] DatabaseError: If the request fails. """ if not self._kill_on_interrupt: - return await start + return await start() - task = asyncio.ensure_future(start) + caller = asyncio.current_task() + cancel_requests = caller.cancelling() if caller else 0 + + async def run() -> str | None: + # Begin the request only if the caller has not been cancelled since. + if caller and caller.cancelling() > cancel_requests: + return None + return await start() + + task = asyncio.ensure_future(run()) try: - return await asyncio.shield(task) + return cast(str, await asyncio.shield(task)) except asyncio.CancelledError as cancellation: - _logger.warning("Query canceled by user.") try: execution_id = await task + if execution_id is None: + # The task did not begin the request, so it was never sent. + raise cancellation + _logger.warning("Query canceled by user.") self._set_interrupted_execution_id(execution_id) await self._cancel_and_wait(execution_id) except Exception as e: diff --git a/pyathena/aio/spark/cursor.py b/pyathena/aio/spark/cursor.py index f84e671ce..b6091666a 100644 --- a/pyathena/aio/spark/cursor.py +++ b/pyathena/aio/spark/cursor.py @@ -106,8 +106,10 @@ async def _calculate( # type: ignore[override] retried request returns the calculation an earlier attempt started instead of starting another one. - With ``kill_on_interrupt`` enabled, the request is shielded from task - cancellation. On cancellation, waits for the request to finish, requests + With ``kill_on_interrupt`` enabled, the request runs in a task shielded + from task cancellation. On cancellation, the request is abandoned if that + task has not begun it by then; it is never sent, and the cancellation + propagates. Otherwise the cursor waits for the request to finish, requests cancellation of the calculation it started, waits for a terminal state, stores the calculation ID and execution on the cursor, and re-raises ``asyncio.CancelledError``. Another cancellation during that wait @@ -137,14 +139,27 @@ async def _calculate( # type: ignore[override] if not self._kill_on_interrupt: return await self._start_calculation_execution(request) - start = asyncio.ensure_future(self._start_calculation_execution(request)) + caller = asyncio.current_task() + cancel_requests = caller.cancelling() if caller else 0 + + async def run() -> str | None: + # Begin the request only if the caller has not been cancelled since. + if caller and caller.cancelling() > cancel_requests: + return None + return await self._start_calculation_execution(request) + + start = asyncio.ensure_future(run()) try: - return await asyncio.shield(start) + return cast(str, 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) + calculation_id = await start + if calculation_id is None: + # The task did not begin the request, so it was never sent. + raise cancellation + _logger.warning("Query canceled by user.") + self._calculation_id = calculation_id + await self._cancel_and_wait(calculation_id) except Exception as e: raise cancellation from e raise diff --git a/tests/pyathena/aio/spark/test_cursor.py b/tests/pyathena/aio/spark/test_cursor.py index ef5feeb20..228e07d83 100644 --- a/tests/pyathena/aio/spark/test_cursor.py +++ b/tests/pyathena/aio/spark/test_cursor.py @@ -395,6 +395,23 @@ async def test_execute_cancelled_while_starting(self): assert cursor.calculation_id == "calculation_id" assert cursor.state == AthenaCalculationExecutionStatus.STATE_CANCELED + async def test_execute_cancelled_before_request_is_sent(self): + """Cancellation before the start task begins sends no request (no AWS).""" + cursor, cancel, started, release = _starting_cursor() + release.set() + task = asyncio.create_task(cursor.execute("code")) + # The task runs until it awaits the start task, which has not begun yet. + await asyncio.sleep(0) + task.cancel() + with pytest.raises(asyncio.CancelledError): + await task + + assert task.cancelled() + assert not started.is_set() + cursor._connection.client.start_calculation_execution.assert_not_called() + cancel.assert_not_awaited() + assert cursor.calculation_id is None + async def test_execute_timeout_while_starting(self): cursor, cancel, started, release = _starting_cursor() task = asyncio.create_task(asyncio.wait_for(cursor.execute("code"), timeout=0.05)) diff --git a/tests/pyathena/aio/test_cursor.py b/tests/pyathena/aio/test_cursor.py index 7bece9167..568449568 100644 --- a/tests/pyathena/aio/test_cursor.py +++ b/tests/pyathena/aio/test_cursor.py @@ -333,6 +333,23 @@ async def test_execute_cancelled_while_starting(self): assert cursor.query_id == "query_id" assert cursor.result_set is None + async def test_execute_cancelled_before_request_is_sent(self): + """Cancellation before the start task begins sends no request (no AWS).""" + cursor, cancel, started, release = _starting_cursor() + release.set() + task = asyncio.create_task(cursor.execute("SELECT 1")) + # The task runs until it awaits the start task, which has not begun yet. + await asyncio.sleep(0) + task.cancel() + with pytest.raises(asyncio.CancelledError): + await task + + assert task.cancelled() + assert not started.is_set() + cursor._connection.client.start_query_execution.assert_not_called() + cancel.assert_not_awaited() + assert cursor.query_id is None + async def test_execute_timeout_while_starting(self): """A timeout during the start request stops the query and raises TimeoutError (no AWS).""" cursor, cancel, started, release = _starting_cursor() From ab0ddaf8b1977a915871c40e141f9dbdc58d29e2 Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Fri, 2 Oct 2026 01:19:05 +0900 Subject: [PATCH 3/8] Document the abandoned asyncio start request and AsyncCursor query IDs Co-Authored-By: Claude Opus 5.5 --- docs/aio.md | 1 + docs/spark.md | 1 + docs/usage.md | 3 ++- 3 files changed, 4 insertions(+), 1 deletion(-) diff --git a/docs/aio.md b/docs/aio.md index b429af071..18f4fe299 100644 --- a/docs/aio.md +++ b/docs/aio.md @@ -141,6 +141,7 @@ 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. +If the request has not been sent yet when the cancellation is handled, it is never sent. 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. diff --git a/docs/spark.md b/docs/spark.md index c55218a55..60e6f2100 100644 --- a/docs/spark.md +++ b/docs/spark.md @@ -512,4 +512,5 @@ With `kill_on_interrupt` enabled, which is the default, cancelling the task whil requests cancellation of the calculation, waits until it reaches a terminal state, and then raises `asyncio.CancelledError`. Cancelling the task while `execute()` is still starting the calculation first waits for the start request to finish, and then cancels the calculation it started in the same way. +If the request has not been sent yet when the cancellation is handled, it is never sent. Cancelling the task again during this wait raises `asyncio.CancelledError` at once without cancelling the calculation. diff --git a/docs/usage.md b/docs/usage.md index 6896439ce..d2c224c37 100644 --- a/docs/usage.md +++ b/docs/usage.md @@ -514,7 +514,8 @@ A `KeyboardInterrupt` while `execute()` is still starting the query first waits 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()`. +`AsyncCursor` and its variants also stop a query whose start is interrupted in `execute()`, +but they have no `query_id` property, so the ID of that query is not available. They wait for queries on worker threads, which do not receive `KeyboardInterrupt`. A second `KeyboardInterrupt` during the cancellation request or these waits propagates immediately, and the query can keep running. From ed50a797e3335897f7d14b7dc37b11c1fa326962 Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Fri, 2 Oct 2026 01:20:38 +0900 Subject: [PATCH 4/8] Say when an interrupt abandons the start request Co-Authored-By: Claude Opus 5.5 --- docs/aio.md | 2 +- docs/spark.md | 2 +- docs/usage.md | 2 +- 3 files changed, 3 insertions(+), 3 deletions(-) diff --git a/docs/aio.md b/docs/aio.md index 18f4fe299..5f0f85b64 100644 --- a/docs/aio.md +++ b/docs/aio.md @@ -141,7 +141,7 @@ 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. -If the request has not been sent yet when the cancellation is handled, it is never sent. +If the task is cancelled before `execute()` begins the request, the request is never sent. 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. diff --git a/docs/spark.md b/docs/spark.md index 60e6f2100..956a929d5 100644 --- a/docs/spark.md +++ b/docs/spark.md @@ -512,5 +512,5 @@ With `kill_on_interrupt` enabled, which is the default, cancelling the task whil requests cancellation of the calculation, waits until it reaches a terminal state, and then raises `asyncio.CancelledError`. Cancelling the task while `execute()` is still starting the calculation first waits for the start request to finish, and then cancels the calculation it started in the same way. -If the request has not been sent yet when the cancellation is handled, it is never sent. +If the task is cancelled before `execute()` begins the request, the request is never sent. Cancelling the task again during this wait raises `asyncio.CancelledError` at once without cancelling the calculation. diff --git a/docs/usage.md b/docs/usage.md index d2c224c37..dc5a43173 100644 --- a/docs/usage.md +++ b/docs/usage.md @@ -513,7 +513,7 @@ A `KeyboardInterrupt` while `execute()` is still starting the query first waits [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. +If `execute()` has not begun the request when the interrupt is handled, the request is never sent. `AsyncCursor` and its variants also stop a query whose start is interrupted in `execute()`, but they have no `query_id` property, so the ID of that query is not available. They wait for queries on worker threads, which do not receive `KeyboardInterrupt`. From 5f55851e61c02d6ea62cc3e08fc4960caa9d33f8 Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Fri, 2 Oct 2026 01:24:11 +0900 Subject: [PATCH 5/8] Tie the unset asyncio query_id to when execute() begins the request Co-Authored-By: Claude Opus 5.5 --- docs/aio.md | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/docs/aio.md b/docs/aio.md index 5f0f85b64..3b6f38818 100644 --- a/docs/aio.md +++ b/docs/aio.md @@ -146,7 +146,7 @@ Cancelling the task again during the cancellation request or these waits raises 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. +`query_id` is `None` if the timeout expires before `execute()` begins the start request, for example while looking up a cached result. ```python import asyncio From b309307118aac222893aeb85dc7a68bacde3d8f5 Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Fri, 2 Oct 2026 01:26:34 +0900 Subject: [PATCH 6/8] State the only asyncio timeout window that leaves query_id unset Co-Authored-By: Claude Opus 5.5 --- docs/aio.md | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/docs/aio.md b/docs/aio.md index 3b6f38818..a460c75ff 100644 --- a/docs/aio.md +++ b/docs/aio.md @@ -146,7 +146,7 @@ Cancelling the task again during the cancellation request or these waits raises 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 `execute()` begins the start request, for example while looking up a cached result. +If the timeout expires while `execute()` is still looking up a cached result, no query is started and `query_id` is `None`. ```python import asyncio From 7ae628c14f155b3dcfcb0a400a51ccecc5daf378 Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Fri, 2 Oct 2026 01:29:18 +0900 Subject: [PATCH 7/8] Scope the asyncio timeout and kill_on_interrupt=False notes to started queries Co-Authored-By: Claude Opus 5.5 --- docs/aio.md | 4 ++-- docs/usage.md | 2 +- 2 files changed, 3 insertions(+), 3 deletions(-) diff --git a/docs/aio.md b/docs/aio.md index a460c75ff..171ecbd5d 100644 --- a/docs/aio.md +++ b/docs/aio.md @@ -143,9 +143,9 @@ Cancelling the task while `execute()` is still starting the query first waits fo and then cancels the query it started in the same way. If the task is cancelled before `execute()` begins the request, the request is never sent. 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. +With `kill_on_interrupt=False`, `asyncio.CancelledError` is raised immediately, and a query that has already started keeps running. -A timeout from `asyncio.wait_for()` therefore cancels the query and raises `asyncio.TimeoutError`. +With `kill_on_interrupt` enabled, a timeout from `asyncio.wait_for()` therefore cancels a query that has started and raises `asyncio.TimeoutError`. If the timeout expires while `execute()` is still looking up a cached result, no query is started and `query_id` is `None`. ```python diff --git a/docs/usage.md b/docs/usage.md index dc5a43173..0638fc489 100644 --- a/docs/usage.md +++ b/docs/usage.md @@ -519,7 +519,7 @@ but they have no `query_id` property, so the ID of that query is not available. They wait for queries on worker threads, which do not receive `KeyboardInterrupt`. 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. +With `kill_on_interrupt=False`, the `KeyboardInterrupt` propagates immediately, and a query that has already started keeps running. ```python from pyathena import connect From c6357d07e5c0e0508f46f535bcbfd9d596ba2fe8 Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Fri, 2 Oct 2026 01:31:40 +0900 Subject: [PATCH 8/8] Limit the asyncio timeout note to starting and waiting for the query Co-Authored-By: Claude Opus 5.5 --- docs/aio.md | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/docs/aio.md b/docs/aio.md index 171ecbd5d..d48895a69 100644 --- a/docs/aio.md +++ b/docs/aio.md @@ -145,7 +145,7 @@ If the task is cancelled before `execute()` begins the request, the request is n 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 a query that has already started keeps running. -With `kill_on_interrupt` enabled, a timeout from `asyncio.wait_for()` therefore cancels a query that has started and raises `asyncio.TimeoutError`. +With `kill_on_interrupt` enabled, a timeout from `asyncio.wait_for()` while `execute()` starts or waits for the query therefore cancels it and raises `asyncio.TimeoutError`. If the timeout expires while `execute()` is still looking up a cached result, no query is started and `query_id` is `None`. ```python