From b558849f204c5e8fc489e5d14fc06035e4e0facb Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Sun, 27 Sep 2026 11:41:18 +0900 Subject: [PATCH 1/7] Cancel a Spark calculation interrupted while it is being started An interrupt that arrived while execute() was sending StartCalculationExecution left the calculation running with no ID on the cursor, and the session rejected new calculations until it finished. With kill_on_interrupt, the sync cursors now send the start request from a helper thread and the asyncio cursor shields it from cancellation; on an interrupt they wait for the request to finish, cancel the calculation it started, wait for a terminal state, and re-raise the interrupt, as the polling phase already does since #833. The Spark cursors also send a generated ClientRequestToken when the caller passes none. botocore does not generate one for StartCalculationExecution, so a retried start could fail on the busy session and leave the first calculation running unseen; with the token, the retry returns that calculation instead. Co-Authored-By: Claude Opus 5.5 --- docs/spark.md | 13 ++ pyathena/aio/spark/cursor.py | 80 +++++++- pyathena/spark/common.py | 130 ++++++++++++- tests/pyathena/aio/spark/test_cursor.py | 196 ++++++++++++++++++- tests/pyathena/spark/test_common.py | 246 +++++++++++++++++++++++- 5 files changed, 650 insertions(+), 15 deletions(-) diff --git a/docs/spark.md b/docs/spark.md index 720702b1..71ecab97 100644 --- a/docs/spark.md +++ b/docs/spark.md @@ -285,6 +285,16 @@ 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. + (async-spark-cursor)= ## AsyncSparkCursor @@ -498,3 +508,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..0e256439 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,10 @@ _logger = logging.getLogger(__name__) +# How often a wait for the start request returns to the interpreter, so that +# a KeyboardInterrupt is raised promptly on every platform. +_START_WAIT_INTERVAL = 0.1 + class SparkBaseCursor(BaseCursor, metaclass=ABCMeta): """Abstract base class for Spark-enabled cursor implementations. @@ -340,14 +347,127 @@ 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``, 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: + try: + future.set_result(self.__start_calculation(request)) + except BaseException as e: + future.set_exception(e) + + threading.Thread(target=start, name="pyathena-spark-start", daemon=True).start() + try: + return self.__wait_for_start(future) + except KeyboardInterrupt as interrupt: + _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=_START_WAIT_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..e66df295 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,158 @@ 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) + 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() From e2968c3483f9007770241197eb1681e4c2fbab58 Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Sun, 27 Sep 2026 11:42:52 +0900 Subject: [PATCH 2/7] Note that a caller's calculation token must be unique Co-Authored-By: Claude Opus 5.5 --- docs/spark.md | 2 ++ 1 file changed, 2 insertions(+) diff --git a/docs/spark.md b/docs/spark.md index 71ecab97..c55218a5 100644 --- a/docs/spark.md +++ b/docs/spark.md @@ -294,6 +294,8 @@ A cancellation request sent right after a calculation starts can occasionally ha 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)= From 25bd87e2073973f0d098c3df8bfbabae98e5e382 Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Sun, 27 Sep 2026 11:47:09 +0900 Subject: [PATCH 3/7] Send no Spark start request after an interrupt that preceded it An interrupt raised while the helper thread was being started escaped before the recovery, but the helper could still send the request and leave the calculation running unseen. The helper now claims the future before sending, and the recovery cancels an unclaimed future and re-raises at once, so a request is either never sent or waited for. Co-Authored-By: Claude Opus 5.5 --- pyathena/spark/common.py | 18 +++++++++++++----- tests/pyathena/spark/test_common.py | 29 +++++++++++++++++++++++++++++ 2 files changed, 42 insertions(+), 5 deletions(-) diff --git a/pyathena/spark/common.py b/pyathena/spark/common.py index 0e256439..022be231 100644 --- a/pyathena/spark/common.py +++ b/pyathena/spark/common.py @@ -378,10 +378,12 @@ def _calculate( of starting another one. With ``kill_on_interrupt`` enabled, the request runs on a helper thread. - On ``KeyboardInterrupt``, 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. + A ``KeyboardInterrupt`` before the helper sends the request propagates, and + the request is not sent. A later one 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. @@ -410,15 +412,21 @@ def _calculate( future: Future[str] = Future() def start() -> None: + # Send nothing if the interrupt already gave up on this request. + 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) - threading.Thread(target=start, name="pyathena-spark-start", daemon=True).start() 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 request has not been sent and never will be. + raise _logger.warning("Query canceled by user.") try: self._calculation_id = self.__wait_for_start(future) diff --git a/tests/pyathena/spark/test_common.py b/tests/pyathena/spark/test_common.py index e66df295..35ed2eff 100644 --- a/tests/pyathena/spark/test_common.py +++ b/tests/pyathena/spark/test_common.py @@ -364,6 +364,35 @@ def test_calculate_second_interrupt_while_starting(self, cursor_class): 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) From 2ca18129314af888728d576abc0c208f7c9017e6 Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Sun, 27 Sep 2026 11:49:32 +0900 Subject: [PATCH 4/7] Describe the Spark start interrupt by when the request begins Co-Authored-By: Claude Opus 5.5 --- pyathena/spark/common.py | 16 ++++++++-------- 1 file changed, 8 insertions(+), 8 deletions(-) diff --git a/pyathena/spark/common.py b/pyathena/spark/common.py index 022be231..9f6a30db 100644 --- a/pyathena/spark/common.py +++ b/pyathena/spark/common.py @@ -378,12 +378,12 @@ def _calculate( of starting another one. With ``kill_on_interrupt`` enabled, the request runs on a helper thread. - A ``KeyboardInterrupt`` before the helper sends the request propagates, and - the request is not sent. A later one 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. + A ``KeyboardInterrupt`` before the helper begins the request propagates, + and the helper does not send it. Once the helper has begun, an interrupt + 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. @@ -412,7 +412,7 @@ def _calculate( future: Future[str] = Future() def start() -> None: - # Send nothing if the interrupt already gave up on this request. + # Begin the request only if no interrupt has given up on it yet. if not future.set_running_or_notify_cancel(): return try: @@ -425,7 +425,7 @@ def start() -> None: return self.__wait_for_start(future) except KeyboardInterrupt as interrupt: if future.cancel(): - # The request has not been sent and never will be. + # The helper has not begun the request and never will. raise _logger.warning("Query canceled by user.") try: From 7cb6e58bba4e07b075a0aa947c6c5204b7d1cb00 Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Sun, 27 Sep 2026 11:50:54 +0900 Subject: [PATCH 5/7] Say when an interrupt abandons the Spark start request Co-Authored-By: Claude Opus 5.5 --- pyathena/spark/common.py | 13 +++++++------ 1 file changed, 7 insertions(+), 6 deletions(-) diff --git a/pyathena/spark/common.py b/pyathena/spark/common.py index 9f6a30db..8ec98d82 100644 --- a/pyathena/spark/common.py +++ b/pyathena/spark/common.py @@ -378,12 +378,13 @@ def _calculate( of starting another one. With ``kill_on_interrupt`` enabled, the request runs on a helper thread. - A ``KeyboardInterrupt`` before the helper begins the request propagates, - and the helper does not send it. Once the helper has begun, an interrupt - 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. + When a ``KeyboardInterrupt`` is handled, a request the helper has not yet + begun is abandoned: the helper never sends it and the interrupt + propagates. If the helper has already begun the request, the cursor waits + for it 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. From de2cf809bb34bb2b764e9479410a0c48ba5f0085 Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Sun, 27 Sep 2026 11:51:57 +0900 Subject: [PATCH 6/7] Tie abandoning the Spark start request to the handler's attempt Co-Authored-By: Claude Opus 5.5 --- pyathena/spark/common.py | 14 +++++++------- 1 file changed, 7 insertions(+), 7 deletions(-) diff --git a/pyathena/spark/common.py b/pyathena/spark/common.py index 8ec98d82..7f5492b3 100644 --- a/pyathena/spark/common.py +++ b/pyathena/spark/common.py @@ -378,13 +378,13 @@ def _calculate( of starting another one. With ``kill_on_interrupt`` enabled, the request runs on a helper thread. - When a ``KeyboardInterrupt`` is handled, a request the helper has not yet - begun is abandoned: the helper never sends it and the interrupt - propagates. If the helper has already begun the request, the cursor waits - for it 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. + 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. From 1466e7b5bfd477b943708e1e95d575cfd90b8715 Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Mon, 28 Sep 2026 00:34:52 +0900 Subject: [PATCH 7/7] Name the Spark start wait interval after its purpose The interval exists so that the wait for the start request wakes up to check for Ctrl-C: before Python 3.14, an untimed lock wait on Windows cannot be interrupted by signals. Co-Authored-By: Claude Opus 5.5 --- pyathena/spark/common.py | 9 +++++---- 1 file changed, 5 insertions(+), 4 deletions(-) diff --git a/pyathena/spark/common.py b/pyathena/spark/common.py index 7f5492b3..8114dc9b 100644 --- a/pyathena/spark/common.py +++ b/pyathena/spark/common.py @@ -31,9 +31,10 @@ _logger = logging.getLogger(__name__) -# How often a wait for the start request returns to the interpreter, so that -# a KeyboardInterrupt is raised promptly on every platform. -_START_WAIT_INTERVAL = 0.1 +# 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): @@ -474,7 +475,7 @@ def __wait_for_start(future: Future[str]) -> str: DatabaseError: If the request failed. """ while not future.done(): - wait((future,), timeout=_START_WAIT_INTERVAL) + wait((future,), timeout=_INTERRUPT_CHECK_INTERVAL) return future.result() def _cancel(self, query_id: str) -> None: