diff --git a/docs/spark.md b/docs/spark.md index 720702b1..c55218a5 100644 --- a/docs/spark.md +++ b/docs/spark.md @@ -285,6 +285,18 @@ requests cancellation, waits until the calculation reaches a terminal state, and The `state` property returns that terminal state. If the cancellation request fails, the `KeyboardInterrupt` propagates with the error as its cause. +A `KeyboardInterrupt` while `execute()` is still starting the calculation first waits for the +[StartCalculationExecution](https://docs.aws.amazon.com/athena/latest/APIReference/API_StartCalculationExecution.html) +request to finish, and then cancels the calculation it started in the same way. +The `calculation_id` property returns that calculation's ID. +A second `KeyboardInterrupt` during this wait propagates at once without cancelling the calculation. +A cancellation request sent right after a calculation starts can occasionally have no effect, so the calculation can still end in the `COMPLETED` state. + +Unless `client_request_token` is passed to `execute()`, the cursor sends a generated `ClientRequestToken` with each calculation. +A retried start request then returns the calculation that an earlier attempt started instead of starting another one. +A token passed to `execute()` must be unique for each calculation: +Athena returns the earlier calculation for a reused token, even when the code differs. + (async-spark-cursor)= ## AsyncSparkCursor @@ -498,3 +510,6 @@ async with await aio_connect(work_group="YOUR_SPARK_WORKGROUP", With `kill_on_interrupt` enabled, which is the default, cancelling the task while `execute()` waits for the calculation 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. +Cancelling the task again during this wait raises `asyncio.CancelledError` at once without cancelling the calculation. diff --git a/pyathena/aio/spark/cursor.py b/pyathena/aio/spark/cursor.py index ed9849b5..71cebdd0 100644 --- a/pyathena/aio/spark/cursor.py +++ b/pyathena/aio/spark/cursor.py @@ -9,6 +9,7 @@ import asyncio import logging +import uuid from typing import Any, cast from pyathena.aio.util import async_retry_api_call @@ -133,12 +134,67 @@ async def _calculate( # type: ignore[override] description: str | None = None, client_request_token: str | None = None, ) -> str: + """Start a calculation execution with ``StartCalculationExecution``. + + Without ``client_request_token``, a generated token is sent, so that a + 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 + 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 + propagates at once. + + Args: + session_id: The session ID. + code_block: The code to run. + description: The calculation description. + client_request_token: The idempotency token of the request. + + Returns: + The calculation execution ID. + + Raises: + asyncio.CancelledError: If the task is cancelled while starting the + calculation. A failure to start, cancel, or wait for the + calculation becomes its ``__cause__``. + DatabaseError: If the request fails. + """ request = self._build_start_calculation_execution_request( session_id=session_id, code_block=code_block, description=description, - client_request_token=client_request_token, + client_request_token=client_request_token or str(uuid.uuid4()), ) + if not self._kill_on_interrupt: + return await self.__start_calculation(request) + + start = asyncio.ensure_future(self.__start_calculation(request)) + try: + return await asyncio.shield(start) + except asyncio.CancelledError as cancellation: + _logger.warning("Query canceled by user.") + try: + self._calculation_id = await start + await self.__cancel_and_wait(self._calculation_id) + except Exception as e: + raise cancellation from e + raise + + async def __start_calculation(self, request: dict[str, Any]) -> str: + """Send a ``StartCalculationExecution`` request. + + Args: + request: The request parameters. + + Returns: + The calculation execution ID. + + Raises: + DatabaseError: If the request fails. + """ try: response = await async_retry_api_call( self._connection.client.start_calculation_execution, @@ -146,11 +202,10 @@ async def _calculate( # type: ignore[override] logger=_logger, **request, ) - calculation_id = response.get("CalculationExecutionId") except Exception as e: _logger.exception("Failed to execute calculation.") raise DatabaseError(*e.args) from e - return cast(str, calculation_id) + return cast(str, response.get("CalculationExecutionId")) async def __poll(self, query_id: str) -> AthenaQueryExecution | AthenaCalculationExecution: while True: @@ -195,14 +250,25 @@ async def _poll( # type: ignore[override] raise _logger.warning("Query canceled by user.") try: - await self._cancel(query_id) - self._calculation_execution = cast( - AthenaCalculationExecution, await self.__poll(query_id) - ) + await self.__cancel_and_wait(query_id) except Exception as e: raise cancellation from e raise + async def __cancel_and_wait(self, calculation_id: str) -> None: + """Request cancellation and store the calculation's terminal state. + + Args: + calculation_id: The calculation execution ID. + + Raises: + OperationalError: If the cancellation or a status request fails. + """ + await self._cancel(calculation_id) + self._calculation_execution = cast( + AthenaCalculationExecution, await self.__poll(calculation_id) + ) + async def _cancel(self, query_id: str) -> None: # type: ignore[override] request: dict[str, Any] = {"CalculationExecutionId": query_id} try: diff --git a/pyathena/spark/common.py b/pyathena/spark/common.py index 74105ed0..8114dc9b 100644 --- a/pyathena/spark/common.py +++ b/pyathena/spark/common.py @@ -9,14 +9,17 @@ 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 import botocore -from pyathena import NotSupportedError, OperationalError +from pyathena import DatabaseError, NotSupportedError, OperationalError from pyathena.common import BaseCursor from pyathena.model import ( AthenaCalculationExecution, @@ -28,6 +31,11 @@ _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. @@ -340,14 +348,136 @@ def _poll(self, query_id: str) -> AthenaQueryExecution | AthenaCalculationExecut raise _logger.warning("Query canceled by user.") try: - self._cancel(query_id) - self._calculation_execution = cast( - AthenaCalculationExecution, self.__poll(query_id) - ) + 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. + + Args: + calculation_id: The calculation execution ID. + + Raises: + OperationalError: If the cancellation or a status request fails. + """ + self._cancel(calculation_id) + self._calculation_execution = cast(AthenaCalculationExecution, self.__poll(calculation_id)) + + def _calculate( + self, + session_id: str, + code_block: str, + description: str | None = None, + client_request_token: str | None = None, + ) -> str: + """Start a calculation execution with ``StartCalculationExecution``. + + Without ``client_request_token``, a generated token is sent, so that a + retried request returns the calculation an earlier attempt started instead + of starting another one. + + 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, requests cancellation of the + calculation it started, waits for a terminal state, stores the calculation + ID and execution on the cursor, and re-raises the interrupt. Another + ``KeyboardInterrupt`` during that wait propagates at once. + + Args: + session_id: The session ID. + code_block: The code to run. + description: The calculation description. + client_request_token: The idempotency token of the request. + + Returns: + The calculation execution ID. + + Raises: + KeyboardInterrupt: If interrupted while starting the calculation. A + failure to start, cancel, or wait for the calculation becomes its + ``__cause__``. + DatabaseError: If the request fails. + """ + request = self._build_start_calculation_execution_request( + session_id=session_id, + code_block=code_block, + description=description, + client_request_token=client_request_token or str(uuid.uuid4()), + ) + if not self._kill_on_interrupt: + return self.__start_calculation(request) + + future: Future[str] = Future() + + 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(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_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_start(future) + self.__cancel_and_wait(self._calculation_id) + except Exception as e: + raise interrupt from e + raise + + def __start_calculation(self, request: dict[str, Any]) -> str: + """Send a ``StartCalculationExecution`` request. + + Args: + request: The request parameters. + + Returns: + The calculation execution ID. + + Raises: + DatabaseError: If the request fails. + """ + try: + response = retry_api_call( + self._connection.client.start_calculation_execution, + config=self._retry_config, + logger=_logger, + **request, + ) + except Exception as e: + _logger.exception("Failed to execute calculation.") + raise DatabaseError(*e.args) from e + return cast(str, response.get("CalculationExecutionId")) + + @staticmethod + def __wait_for_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/tests/pyathena/aio/spark/test_cursor.py b/tests/pyathena/aio/spark/test_cursor.py index 63e21361..92053ef5 100644 --- a/tests/pyathena/aio/spark/test_cursor.py +++ b/tests/pyathena/aio/spark/test_cursor.py @@ -7,13 +7,17 @@ import asyncio import textwrap +import threading +import uuid from unittest.mock import AsyncMock, MagicMock import pytest +from botocore.exceptions import ClientError from pyathena.aio.spark.cursor import AioSparkCursor -from pyathena.error import NotSupportedError, OperationalError +from pyathena.error import DatabaseError, NotSupportedError, OperationalError from pyathena.model import AthenaCalculationExecutionStatus, AthenaSessionStatus +from pyathena.util import RetryConfig from tests import ENV from tests.pyathena.aio.conftest import _aio_connect from tests.pyathena.util import ( @@ -55,6 +59,51 @@ async def get_status(query_id): return cursor, cancel, polling +_TIMEOUT = 10 + + +def _starting_cursor(kill_on_interrupt=True, response=None): + """An AioSparkCursor whose start request blocks in a thread until released. + + Args: + kill_on_interrupt: Whether the cursor cancels the calculation on cancellation. + response: The exception the start request raises once released; the + default returns ``calculation_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_calculation_execution(**kwargs): + started.set() + assert release.wait(_TIMEOUT) + if response: + raise response + return {"CalculationExecutionId": "calculation_id"} + + cursor = AioSparkCursor.__new__(AioSparkCursor) # bypass __init__ to avoid AWS calls + cursor._session_id = "session_id" + cursor._connection = MagicMock() + cursor._connection.client.start_calculation_execution.side_effect = start_calculation_execution + cursor._retry_config = RetryConfig(attempt=2, multiplier=0) + cursor._poll_interval = 0 + cursor._kill_on_interrupt = kill_on_interrupt + cursor._on_poll = None + cursor._calculation_id = None + cursor._calculation_execution = None + cursor._get_calculation_execution_status = AsyncMock( + return_value=MagicMock(state=AthenaCalculationExecutionStatus.STATE_CANCELED) + ) + cursor._get_calculation_execution = AsyncMock( + return_value=MagicMock(state=AthenaCalculationExecutionStatus.STATE_CANCELED) + ) + cancel = cursor._cancel = AsyncMock() + return cursor, cancel, started, release + + class TestAioSparkCursor: async def test_spark_dataframe(self, aio_spark_cursor): await aio_spark_cursor.execute( @@ -284,6 +333,151 @@ async def test_session_ownership(self, aio_spark_cursor): AthenaSessionStatus.STATE_TERMINATED, ) + @pytest.mark.parametrize("kill_on_interrupt", [True, False]) + async def test_calculate_reuses_generated_token_on_retry(self, kill_on_interrupt): + cursor, _, _, release = _starting_cursor(kill_on_interrupt=kill_on_interrupt) + release.set() + client = cursor._connection.client + client.start_calculation_execution.side_effect = [ + ClientError( + {"Error": {"Code": "ThrottlingException", "Message": "Rate exceeded"}}, + "StartCalculationExecution", + ), + {"CalculationExecutionId": "calculation_id"}, + ] + + assert await cursor._calculate(session_id="session_id", code_block="code") == ( + "calculation_id" + ) + + tokens = [ + c.kwargs["ClientRequestToken"] + for c in client.start_calculation_execution.call_args_list + ] + assert len(tokens) == 2 + assert tokens[0] == tokens[1] + uuid.UUID(tokens[0]) + + async def test_calculate_keeps_caller_token(self): + cursor, _, _, release = _starting_cursor() + release.set() + + await cursor._calculate( + session_id="session_id", code_block="code", client_request_token="token" + ) + + cursor._connection.client.start_calculation_execution.assert_called_once_with( + SessionId="session_id", CodeBlock="code", ClientRequestToken="token" + ) + + async def test_execute_cancelled_while_starting(self): + cursor, cancel, started, release = _starting_cursor() + task = asyncio.create_task(cursor.execute("code")) + assert await asyncio.to_thread(started.wait, _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_calculation_execution.assert_called_once() + cancel.assert_awaited_once_with("calculation_id") + assert cursor.calculation_id == "calculation_id" + assert cursor.state == AthenaCalculationExecutionStatus.STATE_CANCELED + + async def test_execute_timeout_while_starting(self): + cursor, cancel, started, release = _starting_cursor() + timer = threading.Timer(0.2, release.set) + timer.start() + try: + with pytest.raises(asyncio.TimeoutError): + await asyncio.wait_for(cursor.execute("code"), timeout=0.05) + finally: + timer.cancel() + release.set() + + assert started.is_set() + cancel.assert_awaited_once_with("calculation_id") + assert cursor.calculation_id == "calculation_id" + + @pytest.mark.parametrize("failing", ["start", "cancel", "wait"]) + async def test_execute_cancelled_while_starting_failure(self, failing): + error = OperationalError("failed") + cursor, cancel, started, release = _starting_cursor( + response=ClientError( + {"Error": {"Code": "InvalidRequestException", "Message": "failed"}}, + "StartCalculationExecution", + ) + if failing == "start" + else None + ) + if failing == "cancel": + cancel.side_effect = error + if failing == "wait": + cursor._get_calculation_execution_status.side_effect = error + raised = [] + + async def execute(): + try: + await cursor.execute("code") + except asyncio.CancelledError as e: + raised.append(e) + raise + + task = asyncio.create_task(execute()) + assert await asyncio.to_thread(started.wait, _TIMEOUT) + task.cancel() + release.set() + with pytest.raises(asyncio.CancelledError): + await task + + assert task.cancelled() + assert cursor.calculation_execution is None + if failing == "start": + assert isinstance(raised[0].__cause__, DatabaseError) + cancel.assert_not_awaited() + assert cursor.calculation_id is None + else: + assert raised[0].__cause__ is error + cancel.assert_awaited_once_with("calculation_id") + assert cursor.calculation_id == "calculation_id" + + async def test_execute_second_cancellation_while_starting(self): + cursor, cancel, started, release = _starting_cursor() + task = asyncio.create_task(cursor.execute("code")) + assert await asyncio.to_thread(started.wait, _TIMEOUT) + try: + task.cancel() + for _ in range(5): + await asyncio.sleep(0) + task.cancel() + with pytest.raises(asyncio.CancelledError): + await task + finally: + release.set() + + assert task.cancelled() + cancel.assert_not_awaited() + + async def test_execute_cancelled_while_starting_without_kill_on_interrupt(self): + cursor, cancel, started, release = _starting_cursor(kill_on_interrupt=False) + task = asyncio.create_task(cursor.execute("code")) + assert await asyncio.to_thread(started.wait, _TIMEOUT) + try: + task.cancel() + with pytest.raises(asyncio.CancelledError): + await task + finally: + release.set() + + assert task.cancelled() + cancel.assert_not_awaited() + assert cursor.calculation_id is None + async def test_executemany(self, aio_spark_cursor): with pytest.raises(NotSupportedError): await aio_spark_cursor.executemany("SELECT 1", []) diff --git a/tests/pyathena/spark/test_common.py b/tests/pyathena/spark/test_common.py index 7f6815e0..35ed2eff 100644 --- a/tests/pyathena/spark/test_common.py +++ b/tests/pyathena/spark/test_common.py @@ -7,20 +7,25 @@ import asyncio import logging +import threading +import uuid +from concurrent.futures import wait from unittest.mock import MagicMock, patch import pytest from botocore.exceptions import ClientError -from pyathena import OperationalError +from pyathena import DatabaseError, OperationalError from pyathena.aio.spark.cursor import AioSparkCursor -from pyathena.model import AthenaSessionStatus +from pyathena.model import AthenaCalculationExecutionStatus, AthenaSessionStatus from pyathena.spark.async_cursor import AsyncSparkCursor from pyathena.spark.common import SparkBaseCursor from pyathena.spark.cursor import SparkCursor from pyathena.util import RetryConfig SPARK_CURSOR_CLASSES = [SparkCursor, AsyncSparkCursor, AioSparkCursor] +SYNC_SPARK_CURSOR_CLASSES = [SparkCursor, AsyncSparkCursor] +_TIMEOUT = 10 def _session_status(state: str, reason: str | None = None) -> AthenaSessionStatus: @@ -47,6 +52,91 @@ def _connection(): return connection +def _calculation_cursor(cursor_class, kill_on_interrupt=True): + """A cursor whose calculation requests go to a mocked Athena client. + + Args: + cursor_class: The synchronous Spark cursor class. + kill_on_interrupt: Whether the cursor cancels the calculation on interrupt. + + Returns: + The cursor. Its ``_cancel`` is a mock, and status requests report + ``CANCELED``. + """ + cursor = cursor_class.__new__(cursor_class) # bypass __init__ to avoid AWS calls + cursor._session_id = "session_id" + cursor._connection = MagicMock() + cursor._connection.client.start_calculation_execution.return_value = { + "CalculationExecutionId": "calculation_id" + } + cursor._retry_config = RetryConfig(attempt=2, multiplier=0) + cursor._poll_interval = 0 + cursor._kill_on_interrupt = kill_on_interrupt + cursor._on_poll = None + cursor._calculation_id = None + cursor._calculation_execution = None + cursor._cancel = MagicMock() + cursor._get_calculation_execution_status = MagicMock( + return_value=MagicMock(state=AthenaCalculationExecutionStatus.STATE_CANCELED) + ) + cursor._get_calculation_execution = MagicMock( + return_value=MagicMock(state=AthenaCalculationExecutionStatus.STATE_CANCELED) + ) + return cursor + + +def _block_start(cursor, response=None): + """Make the cursor's start request block until released. + + Args: + cursor: The cursor from ``_calculation_cursor``. + response: The exception to raise once released; the default returns + ``calculation_id``. + + Returns: + An event set when the request starts, and an event that releases it. + """ + started = threading.Event() + release = threading.Event() + + def start_calculation_execution(**kwargs): + started.set() + assert release.wait(_TIMEOUT) + if response: + raise response + return {"CalculationExecutionId": "calculation_id"} + + cursor._connection.client.start_calculation_execution.side_effect = start_calculation_execution + 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, @@ -148,6 +238,187 @@ def test_exists_session_raises_on_failure_state(self): cursor._exists_session("session_id") cursor._connection.client.get_session.assert_called_once_with(SessionId="session_id") + @pytest.mark.parametrize("cursor_class", SYNC_SPARK_CURSOR_CLASSES) + @pytest.mark.parametrize("kill_on_interrupt", [True, False]) + def test_calculate_reuses_generated_token_on_retry(self, cursor_class, kill_on_interrupt): + cursor = _calculation_cursor(cursor_class, kill_on_interrupt=kill_on_interrupt) + client = cursor._connection.client + client.start_calculation_execution.side_effect = [ + ClientError( + {"Error": {"Code": "ThrottlingException", "Message": "Rate exceeded"}}, + "StartCalculationExecution", + ), + {"CalculationExecutionId": "calculation_id"}, + ] + + assert cursor._calculate(session_id="session_id", code_block="code") == "calculation_id" + + tokens = [ + c.kwargs["ClientRequestToken"] + for c in client.start_calculation_execution.call_args_list + ] + assert len(tokens) == 2 + assert tokens[0] == tokens[1] + uuid.UUID(tokens[0]) + for thread in threading.enumerate(): + if thread.name == "pyathena-spark-start": + thread.join(_TIMEOUT) + assert not thread.is_alive() + + @pytest.mark.parametrize("cursor_class", SYNC_SPARK_CURSOR_CLASSES) + def test_calculate_generates_token_per_call(self, cursor_class): + cursor = _calculation_cursor(cursor_class) + client = cursor._connection.client + + cursor._calculate(session_id="session_id", code_block="code") + cursor._calculate(session_id="session_id", code_block="code") + + tokens = [ + c.kwargs["ClientRequestToken"] + for c in client.start_calculation_execution.call_args_list + ] + assert tokens[0] != tokens[1] + + @pytest.mark.parametrize("cursor_class", SYNC_SPARK_CURSOR_CLASSES) + def test_calculate_keeps_caller_token(self, cursor_class): + cursor = _calculation_cursor(cursor_class) + + cursor._calculate(session_id="session_id", code_block="code", client_request_token="token") + + cursor._connection.client.start_calculation_execution.assert_called_once_with( + SessionId="session_id", CodeBlock="code", ClientRequestToken="token" + ) + + @pytest.mark.parametrize("cursor_class", SYNC_SPARK_CURSOR_CLASSES) + @pytest.mark.parametrize( + "final_state", + [ + AthenaCalculationExecutionStatus.STATE_CANCELED, + AthenaCalculationExecutionStatus.STATE_COMPLETED, + ], + ) + def test_calculate_interrupted_while_starting(self, cursor_class, final_state): + cursor = _calculation_cursor(cursor_class) + 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) + + with waits, pytest.raises(KeyboardInterrupt) as exc_info: + cursor._calculate(session_id="session_id", code_block="code") + + assert exc_info.value is raised[0] + assert exc_info.value.__cause__ is None + cursor._connection.client.start_calculation_execution.assert_called_once() + cursor._cancel.assert_called_once_with("calculation_id") + assert cursor.calculation_id == "calculation_id" + assert cursor._calculation_execution is final_execution + + @pytest.mark.parametrize("cursor_class", SYNC_SPARK_CURSOR_CLASSES) + @pytest.mark.parametrize("failing", ["start", "cancel", "wait"]) + def test_calculate_interrupted_while_starting_failure(self, cursor_class, failing): + cursor = _calculation_cursor(cursor_class) + error = OperationalError("failed") + started, release = _block_start( + cursor, + response=ClientError( + {"Error": {"Code": "InvalidRequestException", "Message": "failed"}}, + "StartCalculationExecution", + ) + if failing == "start" + else None, + ) + if failing == "cancel": + cursor._cancel.side_effect = error + if failing == "wait": + cursor._get_calculation_execution_status.side_effect = error + waits, raised = _interrupt_waits(started, release) + + with waits, pytest.raises(KeyboardInterrupt) as exc_info: + cursor._calculate(session_id="session_id", code_block="code") + + assert exc_info.value is raised[0] + assert cursor._calculation_execution is None + if failing == "start": + assert isinstance(exc_info.value.__cause__, DatabaseError) + cursor._cancel.assert_not_called() + assert cursor.calculation_id is None + else: + assert exc_info.value.__cause__ is error + cursor._cancel.assert_called_once_with("calculation_id") + assert cursor.calculation_id == "calculation_id" + + @pytest.mark.parametrize("cursor_class", SYNC_SPARK_CURSOR_CLASSES) + 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) + + try: + with waits, pytest.raises(KeyboardInterrupt) as exc_info: + cursor._calculate(session_id="session_id", code_block="code") + finally: + release.set() + + assert exc_info.value is raised[1] + assert exc_info.value.__context__ is raised[0] + cursor._cancel.assert_not_called() + + @pytest.mark.parametrize("cursor_class", SYNC_SPARK_CURSOR_CLASSES) + @pytest.mark.parametrize("helper_starts", [False, True]) + def test_calculate_interrupted_before_request_is_sent(self, cursor_class, helper_starts): + cursor = _calculation_cursor(cursor_class) + targets = [] + + class InterruptedThread: + """A thread whose start() is interrupted; the test runs its target later.""" + + def __init__(self, target, name, daemon): + targets.append(target) + + def start(self): + raise KeyboardInterrupt + + with ( + patch("pyathena.spark.common.threading.Thread", InterruptedThread), + pytest.raises(KeyboardInterrupt) as exc_info, + ): + cursor._calculate(session_id="session_id", code_block="code") + if helper_starts: + # The helper thread starts running after the interrupt was handled. + targets[0]() + + assert exc_info.value.__cause__ is None + cursor._connection.client.start_calculation_execution.assert_not_called() + cursor._cancel.assert_not_called() + assert cursor.calculation_id is None + + @pytest.mark.parametrize("cursor_class", SYNC_SPARK_CURSOR_CLASSES) + def test_calculate_interrupt_without_kill_on_interrupt(self, cursor_class): + cursor = _calculation_cursor(cursor_class, kill_on_interrupt=False) + cursor._connection.client.start_calculation_execution.side_effect = KeyboardInterrupt() + + with ( + patch("pyathena.spark.common.threading.Thread") as thread, + pytest.raises(KeyboardInterrupt), + ): + cursor._calculate(session_id="session_id", code_block="code") + + thread.assert_not_called() + cursor._connection.client.start_calculation_execution.assert_called_once() + cursor._cancel.assert_not_called() + + def test_execute_interrupted_while_starting(self): + cursor = _calculation_cursor(SparkCursor) + started, release = _block_start(cursor) + waits, _ = _interrupt_waits(started, release) + + with waits, pytest.raises(KeyboardInterrupt): + cursor.execute("code") + + assert cursor.calculation_id == "calculation_id" + assert cursor.state == AthenaCalculationExecutionStatus.STATE_CANCELED + @pytest.mark.parametrize("cursor_class", SPARK_CURSOR_CLASSES) def test_init_starts_session(self, cursor_class): connection = _connection()