diff --git a/docs/aio.md b/docs/aio.md index 060bc190..b429af07 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 227ad3be..6896439c 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 4592ef29..a2422d31 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 @@ -63,6 +63,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. """ @@ -87,18 +89,72 @@ 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 + + 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")) + + 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 async def _get_query_execution(self, query_id: str) -> AthenaQueryExecution: # type: ignore[override] """Get a query execution with ``GetQueryExecution``. @@ -150,10 +206,12 @@ async def _poll_until_terminal(self, query_id: str) -> AthenaQueryExecution: # await asyncio.sleep(self._poll_interval) 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. @@ -162,19 +220,33 @@ 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 + + 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) async def _cancel(self, query_id: str) -> None: # type: ignore[override] """Stop a query execution with ``StopQueryExecution``. diff --git a/pyathena/aio/spark/cursor.py b/pyathena/aio/spark/cursor.py index 5c82c536..c3356cc9 100644 --- a/pyathena/aio/spark/cursor.py +++ b/pyathena/aio/spark/cursor.py @@ -328,6 +328,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, diff --git a/pyathena/common.py b/pyathena/common.py index 963ec459..8797bfb9 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 8bfe5692..f30ba44d 100644 --- a/pyathena/result_set.py +++ b/pyathena/result_set.py @@ -897,6 +897,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 01637313..05aa5c0e 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. @@ -334,38 +327,6 @@ def _poll_until_terminal( return self._get_calculation_execution(query_id) time.sleep(self._poll_interval) - 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. @@ -423,34 +384,15 @@ 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) + return self._start_execution(lambda: self._start_calculation_execution(request)) - future: Future[str] = Future() + def _set_interrupted_execution_id(self, execution_id: str) -> None: + """Keep the ID of a calculation started by an interrupted start 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) - - 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. @@ -476,23 +418,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() - 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 7ca17442..4d2cb63f 100644 --- a/pyathena/spark/cursor.py +++ b/pyathena/spark/cursor.py @@ -106,6 +106,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 92053ef5..ef5feeb2 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 9679854a..7bece916 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 35ed2eff..5d16726e 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 15267854..b8e6433b 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 eb64f494..18d32109 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 @@ -36,11 +37,74 @@ from pyathena.util import RetryConfig from tests import ENV from tests.pyathena.conftest import connect -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") @@ -1200,6 +1264,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( @@ -1353,6 +1418,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 bfe4dd3c..4c6a9bf4 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