From 168c45bfe2c40862e209b45f1a03f7bc43c56ea8 Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Sat, 26 Sep 2026 16:58:56 +0900 Subject: [PATCH 1/6] Re-raise the interrupt after kill_on_interrupt cancellation With kill_on_interrupt enabled, the shared _poll() of the SQL cursors requested StopQueryExecution after a KeyboardInterrupt or task cancellation, waited for a terminal state, and then returned the final execution instead of re-raising. Callers saw an OperationalError for a CANCELLED or FAILED query, or a normal return when the query SUCCEEDED first, and asyncio.wait_for()/asyncio.timeout() could not turn the cancellation into a timeout. Re-raise the original interrupt after the stop request and the wait, as #833 does for the Spark cursors, with a failure to stop or wait as its cause. Closes #840 Co-Authored-By: Claude Opus 5.5 --- docs/aio.md | 27 +++++++++ docs/usage.md | 23 ++++++++ pyathena/aio/common.py | 37 +++++++++--- pyathena/common.py | 37 +++++++++--- tests/pyathena/aio/test_cursor.py | 98 +++++++++++++++++++++++++++++++ tests/pyathena/test_cursor.py | 69 ++++++++++++++++++++++ 6 files changed, 275 insertions(+), 16 deletions(-) diff --git a/docs/aio.md b/docs/aio.md index 060bc190..6835b488 100644 --- a/docs/aio.md +++ b/docs/aio.md @@ -130,6 +130,33 @@ 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. +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`: + +```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 {cursor.query_id} timed out") +``` + (aio-dict-cursor)= ## AioDictCursor diff --git a/docs/usage.md b/docs/usage.md index 227ad3be..ae624015 100644 --- a/docs/usage.md +++ b/docs/usage.md @@ -501,6 +501,29 @@ 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. +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 16547d1e..bd2f835f 100644 --- a/pyathena/aio/common.py +++ b/pyathena/aio/common.py @@ -127,16 +127,37 @@ async def __poll(self, query_id: str) -> AthenaQueryExecution: await asyncio.sleep(self._poll_interval) async def _poll(self, query_id: str) -> AthenaQueryExecution: # type: ignore[override] + """Wait for a query execution to reach a terminal state. + + On task cancellation with ``kill_on_interrupt`` enabled, requests + cancellation, waits for the query to reach a terminal state, and re-raises + ``asyncio.CancelledError``. + Cancellation is a best-effort request, so the terminal state can be + ``SUCCEEDED`` or ``FAILED`` instead of ``CANCELLED``. + + Args: + query_id: The query execution ID. + + Returns: + The query execution in a terminal state. + + Raises: + asyncio.CancelledError: If the task is cancelled while waiting. A failure + to cancel or wait for the query becomes its ``__cause__``. + OperationalError: If a status request fails. + """ try: - query_execution = await self.__poll(query_id) - except asyncio.CancelledError: - if self._kill_on_interrupt: - _logger.warning("Query canceled by user.") - await self._cancel(query_id) - query_execution = await self.__poll(query_id) - else: + return await self.__poll(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(query_id) + await self.__poll(query_id) + except Exception as e: + raise cancellation from e + raise async def _cancel(self, query_id: str) -> None: # type: ignore[override] """Stop a query execution with ``StopQueryExecution``. diff --git a/pyathena/common.py b/pyathena/common.py index cfb7d09f..58534cea 100644 --- a/pyathena/common.py +++ b/pyathena/common.py @@ -876,16 +876,37 @@ def __poll(self, query_id: str) -> AthenaQueryExecution | AthenaCalculationExecu time.sleep(self._poll_interval) def _poll(self, query_id: str) -> AthenaQueryExecution | AthenaCalculationExecution: + """Wait for a query execution to reach a terminal state. + + On ``KeyboardInterrupt`` with ``kill_on_interrupt`` enabled, requests + cancellation, waits for the query to reach a terminal state, and re-raises + the interrupt. + Cancellation is a best-effort request, so the terminal state can be + ``SUCCEEDED`` or ``FAILED`` instead of ``CANCELLED``. + + Args: + query_id: The query execution ID. + + Returns: + The query execution in a terminal state. + + Raises: + KeyboardInterrupt: If interrupted while waiting. A failure to cancel or + wait for the query becomes its ``__cause__``. + OperationalError: If a status request fails. + """ try: - query_execution = self.__poll(query_id) - except KeyboardInterrupt as e: - if self._kill_on_interrupt: - _logger.warning("Query canceled by user.") + return self.__poll(query_id) + except KeyboardInterrupt as interrupt: + if not self._kill_on_interrupt: + raise + _logger.warning("Query canceled by user.") + try: self._cancel(query_id) - query_execution = self.__poll(query_id) - else: - raise e - return query_execution + self.__poll(query_id) + except Exception as e: + raise interrupt from e + raise def _find_previous_query_id( self, diff --git a/tests/pyathena/aio/test_cursor.py b/tests/pyathena/aio/test_cursor.py index 0ca9a66e..dc9e5908 100644 --- a/tests/pyathena/aio/test_cursor.py +++ b/tests/pyathena/aio/test_cursor.py @@ -18,6 +18,44 @@ from tests.pyathena.util import throttle_metadata_api +def _offline_cursor(kill_on_interrupt, final_state): + """An AioCursor whose first status request blocks until the task is cancelled. + + 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() + + async def get_query_execution(query_id): + if not polling.is_set(): + polling.set() + await asyncio.Event().wait() + return MagicMock(state=final_state) + + 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 + + class TestAioCursor: @pytest.mark.parametrize( ("value", "expected"), @@ -152,6 +190,66 @@ 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).""" + cursor, cancel, polling = _offline_cursor(kill_on_interrupt=True, final_state=final_state) + 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") + 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_cache_size_different_schema(self): """A cached result is only reused when it ran against the same schema (#739). diff --git a/tests/pyathena/test_cursor.py b/tests/pyathena/test_cursor.py index f142287b..817876c6 100644 --- a/tests/pyathena/test_cursor.py +++ b/tests/pyathena/test_cursor.py @@ -40,6 +40,33 @@ _logger = logging.getLogger(__name__) +def _offline_cursor(kill_on_interrupt): + """A Cursor that starts a query and polls it without AWS. + + Args: + kill_on_interrupt: Whether the cursor cancels the query on interrupt. + + Returns: + The cursor and the mock of its cancellation request. + """ + cursor = Cursor.__new__(Cursor) # 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 = MagicMock(return_value="query_id") + # 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() + cancel = cursor._cancel = MagicMock() + return cursor, cancel + + class TestCursor: def test_fetchone(self, cursor): cursor.execute("SELECT * FROM one_row") @@ -1285,6 +1312,48 @@ 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).""" + cursor, cancel = _offline_cursor(kill_on_interrupt=True) + cursor._get_query_execution = MagicMock( + side_effect=[KeyboardInterrupt(), MagicMock(state=final_state)] + ) + with pytest.raises(KeyboardInterrupt): + cursor.execute("SELECT 1") + + cancel.assert_called_once_with("query_id") + assert cursor._get_query_execution.call_count == 2 + 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" + def test_on_poll_connection_level(self): """Connection-level on_poll fires during query execution.""" states = [] From da57283dc81c056882fe87600d55d48e83d1b0ab Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Sat, 26 Sep 2026 17:09:08 +0900 Subject: [PATCH 2/6] Address independent review of the interrupt re-raise Limit the documented asyncio timeout behavior to a timeout while the query is being waited for, note that a repeated interrupt or task cancellation skips the wait, and make the tests assert that the interrupt is re-raised only after a non-terminal poll reaches a terminal state. Co-Authored-By: Claude Opus 5.5 --- docs/aio.md | 7 +++++-- docs/usage.md | 1 + tests/pyathena/aio/test_cursor.py | 12 +++++++++++- tests/pyathena/test_cursor.py | 14 ++++++++++++-- 4 files changed, 29 insertions(+), 5 deletions(-) diff --git a/docs/aio.md b/docs/aio.md index 6835b488..965b7c71 100644 --- a/docs/aio.md +++ b/docs/aio.md @@ -139,9 +139,12 @@ requests cancellation of the query, waits until it reaches a terminal state, and 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 again during that wait raises `asyncio.CancelledError` without waiting for the terminal state. 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`: +A timeout from `asyncio.wait_for()` that expires while `execute()` waits for the query therefore cancels the query +and raises `asyncio.TimeoutError`. +If it expires while the query is being started, `query_id` is `None`, and the pending start request can still start the query. ```python import asyncio @@ -154,7 +157,7 @@ async with await aio_connect(s3_staging_dir="s3://YOUR_S3_BUCKET/path/to/", try: await asyncio.wait_for(cursor.execute("SELECT * FROM many_rows"), timeout=60) except asyncio.TimeoutError: - print(f"Query {cursor.query_id} timed out") + print(f"Query timed out: {cursor.query_id}") ``` (aio-dict-cursor)= diff --git a/docs/usage.md b/docs/usage.md index ae624015..28c87767 100644 --- a/docs/usage.md +++ b/docs/usage.md @@ -508,6 +508,7 @@ requests cancellation, waits until the query reaches a terminal state, and then 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 second `KeyboardInterrupt` during that wait propagates without waiting for the terminal state. With `kill_on_interrupt=False`, the `KeyboardInterrupt` propagates immediately and the query keeps running. ```python diff --git a/tests/pyathena/aio/test_cursor.py b/tests/pyathena/aio/test_cursor.py index dc9e5908..c35193f2 100644 --- a/tests/pyathena/aio/test_cursor.py +++ b/tests/pyathena/aio/test_cursor.py @@ -21,6 +21,8 @@ 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. @@ -30,12 +32,13 @@ def _offline_cursor(kill_on_interrupt, final_state): 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=final_state) + return MagicMock(state=next(states)) cursor = AioCursor.__new__(AioCursor) # bypass __init__ to avoid AWS calls cursor._rowcount = -1 @@ -196,7 +199,9 @@ async def test_execute_internal_legacy_kwargs_passthrough(self): ) 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() @@ -205,6 +210,11 @@ async def test_execute_kill_on_interrupt(self, final_state): 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 diff --git a/tests/pyathena/test_cursor.py b/tests/pyathena/test_cursor.py index 817876c6..104bf4cb 100644 --- a/tests/pyathena/test_cursor.py +++ b/tests/pyathena/test_cursor.py @@ -1318,15 +1318,25 @@ def test_on_poll_none_is_noop(self): ) 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=final_state)] + 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") - assert cursor._get_query_execution.call_count == 2 + # 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 From ef17decf45fbbb6acbada3515325d782d5a9607e Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Sat, 26 Sep 2026 17:15:33 +0900 Subject: [PATCH 3/6] Cover the cancellation request in the repeated-interrupt note A second interrupt or task cancellation also escapes while the cancellation request is in progress, not only during the follow-up wait. Co-Authored-By: Claude Opus 5.5 --- docs/aio.md | 2 +- docs/usage.md | 2 +- 2 files changed, 2 insertions(+), 2 deletions(-) diff --git a/docs/aio.md b/docs/aio.md index 965b7c71..e00413b6 100644 --- a/docs/aio.md +++ b/docs/aio.md @@ -139,7 +139,7 @@ requests cancellation of the query, waits until it reaches a terminal state, and 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 again during that wait raises `asyncio.CancelledError` without waiting for the terminal state. +Cancelling the task again during the cancellation request or that wait 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()` that expires while `execute()` waits for the query therefore cancels the query diff --git a/docs/usage.md b/docs/usage.md index 28c87767..40bb4501 100644 --- a/docs/usage.md +++ b/docs/usage.md @@ -508,7 +508,7 @@ requests cancellation, waits until the query reaches a terminal state, and then 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 second `KeyboardInterrupt` during that wait propagates without waiting for the terminal state. +A second `KeyboardInterrupt` during the cancellation request or that wait propagates immediately, and the query can keep running. With `kill_on_interrupt=False`, the `KeyboardInterrupt` propagates immediately and the query keeps running. ```python From 3338a1e6419c707781bdb8bdf60340715e02fc50 Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Mon, 28 Sep 2026 15:47:21 +0900 Subject: [PATCH 4/6] Stop a query interrupted while it is being started Move the Spark start-phase interrupt handling from #861 into shared helpers and use them for the SQL cursors too. With kill_on_interrupt, StartQueryExecution now runs on a helper thread (sync) or is shielded from task cancellation (asyncio); an interrupt waits for the request, records the query ID on cursors that expose query_id, stops the query, waits for a terminal state, and re-raises. A request the helper has not begun is abandoned and never sent. The polling-phase recovery of the SQL and Spark cursors now goes through the same helpers, so the Spark cursors keep their behavior with less duplicated code. Co-Authored-By: Claude Opus 5.5 --- docs/aio.md | 9 +- docs/usage.md | 11 +- pyathena/aio/common.py | 181 ++++++++++++++++++++--- pyathena/aio/spark/cursor.py | 38 +++-- pyathena/common.py | 222 +++++++++++++++++++++++++--- pyathena/result_set.py | 8 + pyathena/spark/common.py | 75 +++------- tests/pyathena/aio/test_cursor.py | 143 +++++++++++++++++- tests/pyathena/spark/test_common.py | 41 +---- tests/pyathena/test_cursor.py | 145 +++++++++++++++++- tests/pyathena/util.py | 33 +++++ 11 files changed, 737 insertions(+), 169 deletions(-) diff --git a/docs/aio.md b/docs/aio.md index e00413b6..cd7562f0 100644 --- a/docs/aio.md +++ b/docs/aio.md @@ -139,12 +139,13 @@ requests cancellation of the query, waits until it reaches a terminal state, and 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 again during the cancellation request or that wait raises `asyncio.CancelledError` immediately, and the query can keep running. +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()` that expires while `execute()` waits for the query therefore cancels the query -and raises `asyncio.TimeoutError`. -If it expires while the query is being started, `query_id` is `None`, and the pending start request can still start the query. +A timeout from `asyncio.wait_for()` therefore cancels the query and raises `asyncio.TimeoutError`. +`query_id` is `None` only if the timeout expires before the query is started, for example while looking up a cached result. ```python import asyncio diff --git a/docs/usage.md b/docs/usage.md index 40bb4501..6896439c 100644 --- a/docs/usage.md +++ b/docs/usage.md @@ -508,7 +508,16 @@ requests cancellation, waits until the query reaches a terminal state, and then 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 second `KeyboardInterrupt` during the cancellation request or that wait propagates immediately, and the query can keep running. + +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 diff --git a/pyathena/aio/common.py b/pyathena/aio/common.py index bd2f835f..759f072f 100644 --- a/pyathena/aio/common.py +++ b/pyathena/aio/common.py @@ -3,7 +3,7 @@ import asyncio import logging import sys -from collections.abc import Awaitable, Callable +from collections.abc import Awaitable, Callable, Coroutine from datetime import datetime, timedelta, timezone from typing import Any, TypeVar, cast @@ -23,6 +23,71 @@ _T = TypeVar("_T") +async def _start_interruptibly( + start: Coroutine[Any, Any, str], stop: Callable[[str], Awaitable[None]] +) -> str: + """Send a start request shielded from task cancellation, stopping what it starts. + + On ``asyncio.CancelledError``, waits for the request to finish, stops the + execution it started, and re-raises the cancellation. Another cancellation + during that wait propagates at once. + + Args: + start: Sends the start request and returns the execution ID. + stop: Requests cancellation of an execution and waits for a terminal state. + + Returns: + The execution ID. + + Raises: + asyncio.CancelledError: If the task is cancelled while starting. A failure + to start or stop the execution becomes its ``__cause__``. + """ + task = asyncio.ensure_future(start) + try: + return await asyncio.shield(task) + except asyncio.CancelledError as cancellation: + _logger.warning("Query canceled by user.") + try: + await stop(await task) + except Exception as e: + raise cancellation from e + raise + + +async def _poll_interruptibly( + execution_id: str, + poll: Callable[[str], Awaitable[_T]], + stop: Callable[[str], Awaitable[None]], +) -> _T: + """Wait for an execution to reach a terminal state, stopping it on cancellation. + + On ``asyncio.CancelledError``, stops the execution and re-raises the + cancellation. Another cancellation while stopping propagates at once. + + Args: + execution_id: The execution ID. + poll: Waits for the execution to reach a terminal state. + stop: Requests cancellation of the execution and waits for a terminal state. + + Returns: + The result of ``poll``. + + Raises: + asyncio.CancelledError: If the task is cancelled while waiting. A failure + to stop the execution becomes its ``__cause__``. + """ + try: + return await poll(execution_id) + except asyncio.CancelledError as cancellation: + _logger.warning("Query canceled by user.") + try: + await stop(execution_id) + except Exception as e: + raise cancellation from e + raise + + class AioBaseCursor(BaseCursor): """Async base cursor that overrides I/O methods with async equivalents. @@ -44,6 +109,36 @@ async def _execute( # type: ignore[override] paramstyle: str | None = None, options: ExecuteOptions | None = None, ) -> str: + """Start a query execution, or find a cached one, and return its ID. + + With ``kill_on_interrupt`` enabled, the ``StartQueryExecution`` request is + shielded from task cancellation. On cancellation, the cursor waits for the + request to finish, records the query ID with ``_set_interrupted_query_id()``, + requests cancellation of the query, waits for a terminal state, and + re-raises ``asyncio.CancelledError``. + + Args: + operation: SQL query string to execute. + parameters: Query parameters. + work_group: Athena workgroup. + s3_staging_dir: S3 location for query results. + cache_size: Number of recent executions to search for a cached result. + cache_expiration_time: Maximum age in seconds of a cached result. + result_reuse_enable: Whether to enable Athena's result reuse. + result_reuse_minutes: Maximum age in minutes of a reused result. + paramstyle: Parameter style of the query. + options: Shared execution options. Individual keyword arguments take + precedence over its fields. + + Returns: + The query execution ID. + + Raises: + asyncio.CancelledError: If the task is cancelled while starting the + query. A failure to start or cancel the query becomes its + ``__cause__``. + DatabaseError: If the request fails. + """ # The individual keyword arguments are retained for backward compatibility # with external callers that predate ExecuteOptions, mirroring # BaseCursor._execute(). @@ -74,19 +169,50 @@ async def _execute( # type: ignore[override] cache_expiration_time=options.cache_expiration_time, ) if query_id is None: - try: - response = await async_retry_api_call( - self._connection.client.start_query_execution, - config=self._retry_config, - logger=_logger, - **request, + if self._kill_on_interrupt: + query_id = await _start_interruptibly( + self.__start_query_execution(request), self.__stop_started_query ) - query_id = response.get("QueryExecutionId") - except Exception as e: - _logger.exception("Failed to execute query.") - raise DatabaseError(*e.args) from e + else: + query_id = await self.__start_query_execution(request) return query_id + async def __start_query_execution(self, request: dict[str, Any]) -> str: + """Send a ``StartQueryExecution`` request. + + Args: + request: The request parameters. + + Returns: + The query execution ID. + + Raises: + DatabaseError: If the request fails. + """ + try: + response = await async_retry_api_call( + self._connection.client.start_query_execution, + config=self._retry_config, + logger=_logger, + **request, + ) + except Exception as e: + _logger.exception("Failed to execute query.") + raise DatabaseError(*e.args) from e + return cast(str, response.get("QueryExecutionId")) + + async def __stop_started_query(self, query_id: str) -> None: + """Record a query that a cancelled start request started, and stop it. + + Args: + query_id: The query execution ID. + + Raises: + OperationalError: If the cancellation or a status request fails. + """ + self._set_interrupted_query_id(query_id) + await self.__cancel_and_wait(query_id) + async def _get_query_execution(self, query_id: str) -> AthenaQueryExecution: # type: ignore[override] """Get a query execution with ``GetQueryExecution``. @@ -146,18 +272,21 @@ async def _poll(self, query_id: str) -> AthenaQueryExecution: # type: ignore[ov to cancel or wait for the query becomes its ``__cause__``. OperationalError: If a status request fails. """ - try: + if not self._kill_on_interrupt: return await self.__poll(query_id) - except asyncio.CancelledError as cancellation: - if not self._kill_on_interrupt: - raise - _logger.warning("Query canceled by user.") - try: - await self._cancel(query_id) - await self.__poll(query_id) - except Exception as e: - raise cancellation from e - raise + return await _poll_interruptibly(query_id, self.__poll, self.__cancel_and_wait) + + async def __cancel_and_wait(self, query_id: str) -> None: + """Request cancellation of a query and wait for a terminal state. + + Args: + query_id: The query execution ID. + + Raises: + OperationalError: If the cancellation or a status request fails. + """ + await self._cancel(query_id) + await self.__poll(query_id) async def _cancel(self, query_id: str) -> None: # type: ignore[override] """Stop a query execution with ``StopQueryExecution``. @@ -655,6 +784,14 @@ def query_id(self) -> str | None: def query_id(self, val) -> None: self._query_id = val + def _set_interrupted_query_id(self, query_id: str) -> None: + """Expose the ID of a query started by a cancelled start request. + + Args: + query_id: The query execution ID. + """ + self.query_id = query_id + @property def rownumber(self) -> int | None: return self.result_set.rownumber if self.result_set else None diff --git a/pyathena/aio/spark/cursor.py b/pyathena/aio/spark/cursor.py index 71cebdd0..bc5d6501 100644 --- a/pyathena/aio/spark/cursor.py +++ b/pyathena/aio/spark/cursor.py @@ -12,6 +12,7 @@ import uuid from typing import Any, cast +from pyathena.aio.common import _poll_interruptibly, _start_interruptibly from pyathena.aio.util import async_retry_api_call from pyathena.error import DatabaseError, NotSupportedError, OperationalError, ProgrammingError from pyathena.model import ( @@ -170,18 +171,21 @@ async def _calculate( # type: ignore[override] ) if not self._kill_on_interrupt: return await self.__start_calculation(request) + return await _start_interruptibly( + self.__start_calculation(request), self.__stop_started_calculation + ) - start = asyncio.ensure_future(self.__start_calculation(request)) - try: - return await asyncio.shield(start) - except asyncio.CancelledError as cancellation: - _logger.warning("Query canceled by user.") - try: - self._calculation_id = await start - await self.__cancel_and_wait(self._calculation_id) - except Exception as e: - raise cancellation from e - raise + async def __stop_started_calculation(self, calculation_id: str) -> None: + """Store a calculation that a cancelled start request started, and stop it. + + Args: + calculation_id: The calculation execution ID. + + Raises: + OperationalError: If the cancellation or a status request fails. + """ + self._calculation_id = calculation_id + await self.__cancel_and_wait(calculation_id) async def __start_calculation(self, request: dict[str, Any]) -> str: """Send a ``StartCalculationExecution`` request. @@ -243,17 +247,9 @@ async def _poll( # type: ignore[override] to cancel or wait for the calculation becomes its ``__cause__``. OperationalError: If a status request fails. """ - try: + if not self._kill_on_interrupt: return await self.__poll(query_id) - except asyncio.CancelledError as cancellation: - if not self._kill_on_interrupt: - raise - _logger.warning("Query canceled by user.") - try: - await self.__cancel_and_wait(query_id) - except Exception as e: - raise cancellation from e - raise + return await _poll_interruptibly(query_id, self.__poll, self.__cancel_and_wait) async def __cancel_and_wait(self, calculation_id: str) -> None: """Request cancellation and store the calculation's terminal state. diff --git a/pyathena/common.py b/pyathena/common.py index 58534cea..de0a5728 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 datetime, timedelta, timezone from typing import TYPE_CHECKING, Any, TypeVar, cast @@ -47,6 +49,108 @@ :class:`~pyathena.model.AthenaCalculationExecutionStatus` for Spark calculations. """ +# 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 + + +def _start_interruptibly( + start: Callable[[], str], stop: Callable[[str], None], thread_name: str +) -> str: + """Send a start request on a helper thread so that an interrupt can stop what it starts. + + On ``KeyboardInterrupt``, the request is abandoned if the helper has not begun + it by then; the helper then never sends it, and the interrupt propagates. + Otherwise this waits for the request to finish, stops the execution it started, + and re-raises the interrupt. Another ``KeyboardInterrupt`` during that wait + propagates at once. + + Args: + start: Sends the start request and returns the execution ID. + stop: Requests cancellation of an execution and waits for a terminal state. + thread_name: The name of the helper thread. + + Returns: + The execution ID. + + Raises: + KeyboardInterrupt: If interrupted while starting. A failure to start or + stop the execution becomes its ``__cause__``. + """ + 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=thread_name, daemon=True).start() + return _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: + stop(_wait_for_start(future)) + except Exception as e: + raise interrupt from e + raise + + +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 _poll_interruptibly( + execution_id: str, poll: Callable[[str], _T], stop: Callable[[str], None] +) -> _T: + """Wait for an execution to reach a terminal state, stopping it on an interrupt. + + On ``KeyboardInterrupt``, stops the execution and re-raises the interrupt. + Another ``KeyboardInterrupt`` while stopping propagates at once. + + Args: + execution_id: The execution ID. + poll: Waits for the execution to reach a terminal state. + stop: Requests cancellation of the execution and waits for a terminal state. + + Returns: + The result of ``poll``. + + Raises: + KeyboardInterrupt: If interrupted while waiting. A failure to stop the + execution becomes its ``__cause__``. + """ + try: + return poll(execution_id) + except KeyboardInterrupt as interrupt: + _logger.warning("Query canceled by user.") + try: + stop(execution_id) + except Exception as e: + raise interrupt from e + raise + class CursorIterator(metaclass=ABCMeta): """Abstract base class providing iteration and result fetching capabilities for cursors. @@ -895,18 +999,30 @@ def _poll(self, query_id: str) -> AthenaQueryExecution | AthenaCalculationExecut wait for the query becomes its ``__cause__``. OperationalError: If a status request fails. """ - try: + if not self._kill_on_interrupt: return self.__poll(query_id) - except KeyboardInterrupt as interrupt: - if not self._kill_on_interrupt: - raise - _logger.warning("Query canceled by user.") - try: - self._cancel(query_id) - self.__poll(query_id) - except Exception as e: - raise interrupt from e - raise + return _poll_interruptibly(query_id, self.__poll, self.__cancel_and_wait) + + def __cancel_and_wait(self, query_id: str) -> None: + """Request cancellation of a query and wait for a terminal state. + + Args: + query_id: The query execution ID. + + Raises: + OperationalError: If the cancellation or a status request fails. + """ + self._cancel(query_id) + self.__poll(query_id) + + def _set_interrupted_query_id(self, query_id: str) -> None: # noqa: B027 + """Record the ID of a query started by a start request that was interrupted. + + Does nothing by default; cursors that expose ``query_id`` override this. + + Args: + query_id: The query execution ID. + """ def _find_previous_query_id( self, @@ -1038,6 +1154,36 @@ def _execute( paramstyle: str | None = None, options: ExecuteOptions | None = None, ) -> str: + """Start a query execution, or find a cached one, and return its ID. + + With ``kill_on_interrupt`` enabled, the ``StartQueryExecution`` request + runs on a helper thread. On ``KeyboardInterrupt``, the request is abandoned + if the helper has not begun it by then. Otherwise the cursor waits for the + request to finish, records the query ID with ``_set_interrupted_query_id()``, + requests cancellation of the query, waits for a terminal state, and + re-raises the interrupt. + + Args: + operation: SQL query string to execute. + parameters: Query parameters. + work_group: Athena workgroup. + s3_staging_dir: S3 location for query results. + cache_size: Number of recent executions to search for a cached result. + cache_expiration_time: Maximum age in seconds of a cached result. + result_reuse_enable: Whether to enable Athena's result reuse. + result_reuse_minutes: Maximum age in minutes of a reused result. + paramstyle: Parameter style of the query. + options: Shared execution options. Individual keyword arguments take + precedence over its fields. + + Returns: + The query execution ID. + + Raises: + KeyboardInterrupt: If interrupted while starting the query. A failure + to start or cancel the query becomes its ``__cause__``. + DatabaseError: If the request fails. + """ # The individual keyword arguments are retained for backward compatibility # with external callers that predate ExecuteOptions (e.g. dbt-athena <= 1.10.x # calls _execute() with work_group/s3_staging_dir/cache_* keywords). @@ -1068,18 +1214,52 @@ 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 + if self._kill_on_interrupt: + query_id = _start_interruptibly( + lambda: self.__start_query_execution(request), + self.__stop_started_query, + "pyathena-query-start", + ) + else: + query_id = 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")) + + def __stop_started_query(self, query_id: str) -> None: + """Record a query that an interrupted start request started, and stop it. + + Args: + query_id: The query execution ID. + + Raises: + OperationalError: If the cancellation or a status request fails. + """ + self._set_interrupted_query_id(query_id) + self.__cancel_and_wait(query_id) + def _calculate( self, session_id: str, diff --git a/pyathena/result_set.py b/pyathena/result_set.py index f6fe6017..ee734847 100644 --- a/pyathena/result_set.py +++ b/pyathena/result_set.py @@ -1047,6 +1047,14 @@ def query_id(self) -> str | None: def query_id(self, val) -> None: self._query_id = val + def _set_interrupted_query_id(self, query_id: str) -> None: + """Expose the ID of a query started by an interrupted start request. + + Args: + query_id: The query execution ID. + """ + self.query_id = query_id + @property def rownumber(self) -> int | None: return self.result_set.rownumber if self.result_set else None diff --git a/pyathena/spark/common.py b/pyathena/spark/common.py index 8114dc9b..8a752961 100644 --- a/pyathena/spark/common.py +++ b/pyathena/spark/common.py @@ -9,18 +9,16 @@ 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 DatabaseError, NotSupportedError, OperationalError -from pyathena.common import BaseCursor +from pyathena.common import BaseCursor, _poll_interruptibly, _start_interruptibly from pyathena.model import ( AthenaCalculationExecution, AthenaCalculationExecutionStatus, @@ -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. @@ -341,17 +334,9 @@ def _poll(self, query_id: str) -> AthenaQueryExecution | AthenaCalculationExecut wait for the calculation becomes its ``__cause__``. OperationalError: If a status request fails. """ - try: + if not self._kill_on_interrupt: return self.__poll(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 + return _poll_interruptibly(query_id, self.__poll, self.__cancel_and_wait) def __cancel_and_wait(self, calculation_id: str) -> None: """Request cancellation and store the calculation's terminal state. @@ -410,32 +395,23 @@ def _calculate( ) if not self._kill_on_interrupt: return self.__start_calculation(request) + return _start_interruptibly( + lambda: self.__start_calculation(request), + self.__stop_started_calculation, + "pyathena-spark-start", + ) - future: Future[str] = Future() + def __stop_started_calculation(self, calculation_id: str) -> None: + """Store a calculation that an interrupted start request started, and stop it. - 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) + Args: + calculation_id: The calculation execution ID. - 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 + Raises: + OperationalError: If the cancellation or a status request fails. + """ + self._calculation_id = calculation_id + self.__cancel_and_wait(calculation_id) def __start_calculation(self, request: dict[str, Any]) -> str: """Send a ``StartCalculationExecution`` request. @@ -461,23 +437,6 @@ def __start_calculation(self, request: dict[str, Any]) -> str: 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/test_cursor.py b/tests/pyathena/aio/test_cursor.py index c35193f2..93c54c38 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, patch import pytest +from botocore.exceptions import ClientError from pyathena import BINARY, Binary, ExecuteOptions from pyathena.aio.cursor import AioCursor @@ -15,7 +16,7 @@ from pyathena.util import RetryConfig from tests import ENV from tests.pyathena.aio.conftest import _aio_connect -from tests.pyathena.util import throttle_metadata_api +from tests.pyathena.util import EVENT_TIMEOUT, throttle_metadata_api def _offline_cursor(kill_on_interrupt, final_state): @@ -59,6 +60,47 @@ async def get_query_execution(query_id): 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: @pytest.mark.parametrize( ("value", "expected"), @@ -156,6 +198,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( @@ -260,6 +303,104 @@ async def test_execute_without_kill_on_interrupt(self): 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() + timer = threading.Timer(0.2, release.set) + timer.start() + try: + with pytest.raises(asyncio.TimeoutError): + await asyncio.wait_for(cursor.execute("SELECT 1"), timeout=0.05) + finally: + timer.cancel() + release.set() + + 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..935cf12d 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, @@ -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/test_cursor.py b/tests/pyathena/test_cursor.py index 104bf4cb..ebe3960a 100644 --- a/tests/pyathena/test_cursor.py +++ b/tests/pyathena/test_cursor.py @@ -28,6 +28,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 @@ -35,38 +36,73 @@ from pyathena.util import RetryConfig from tests import ENV from tests.pyathena.conftest import connect -from tests.pyathena.util import throttle_metadata_api, unreachable_glue +from tests.pyathena.util import ( + EVENT_TIMEOUT, + interrupt_start_waits, + throttle_metadata_api, + unreachable_glue, +) _logger = logging.getLogger(__name__) -def _offline_cursor(kill_on_interrupt): - """A Cursor that starts a query and polls it without AWS. +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.__new__(Cursor) # bypass __init__ to avoid AWS calls + 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._execute = MagicMock(return_value="query_id") - # A successful query builds a result set from these. + 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._retry_config = RetryConfig() 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") @@ -1159,6 +1195,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( @@ -1364,6 +1401,100 @@ def test_execute_without_kill_on_interrupt(self): 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 f5ac673d..530aa732 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 @@ -133,3 +135,34 @@ def wait_for_spark_session_state(client, session_id, state, timeout=120): return time.sleep(1) raise AssertionError(f"Session {session_id} did not become {state} in {timeout} seconds.") + + +# Seconds a test waits for an event from a helper thread before failing. +EVENT_TIMEOUT = 10 + + +def interrupt_start_waits(started, release, interrupts=1): + """Patch the wait for a start request on a helper thread to raise KeyboardInterrupt. + + The first ``interrupts`` waits raise ``KeyboardInterrupt`` once the request + has started; later waits release the request and wait for it. + + Args: + started: Set when the start request starts. + release: Releases the start request. + interrupts: How many waits raise. + + Returns: + The patcher, and the list of raised interrupts. + """ + raised = [] + + def interrupting_wait(futures, timeout=None): + if len(raised) < interrupts: + assert started.wait(EVENT_TIMEOUT) + raised.append(KeyboardInterrupt()) + raise raised[-1] + release.set() + return wait(futures, timeout) + + return patch("pyathena.common.wait", side_effect=interrupting_wait), raised From d66c6994034894313f6b68f23f9f5d421ac13632 Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Mon, 28 Sep 2026 15:49:54 +0900 Subject: [PATCH 5/6] State when an asyncio timeout leaves query_id unset Co-Authored-By: Claude Opus 5.5 --- docs/aio.md | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/docs/aio.md b/docs/aio.md index cd7562f0..4d6bc5ff 100644 --- a/docs/aio.md +++ b/docs/aio.md @@ -145,7 +145,7 @@ Cancelling the task again during the cancellation request or these waits raises With `kill_on_interrupt=False`, `asyncio.CancelledError` is raised immediately and the query keeps running. A timeout from `asyncio.wait_for()` therefore cancels the query and raises `asyncio.TimeoutError`. -`query_id` is `None` only if the timeout expires before the query is started, for example while looking up a cached result. +`query_id` is `None` only if no query was started, for example when the timeout expires while looking up a cached result. ```python import asyncio From 10f7a832ebbac10440706167d1e370f90b23c2ab Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Mon, 28 Sep 2026 17:08:41 +0900 Subject: [PATCH 6/6] Address independent review of the shared interrupt handling Clear the previous calculation when a Spark cursor starts a new one, so that a failed execution no longer leaves the earlier calculation's state next to the new calculation ID. Make the asyncio start-phase timeout tests wait for the timeout before releasing the start request, and state when an asyncio timeout leaves query_id unset without implying that no query started. Co-Authored-By: Claude Opus 5.5 --- docs/aio.md | 2 +- pyathena/aio/spark/cursor.py | 3 +++ pyathena/spark/cursor.py | 3 +++ tests/pyathena/aio/spark/test_cursor.py | 19 ++++++++++++++----- tests/pyathena/aio/test_cursor.py | 13 ++++++++----- tests/pyathena/spark/test_spark_cursor.py | 7 ++++++- 6 files changed, 35 insertions(+), 12 deletions(-) diff --git a/docs/aio.md b/docs/aio.md index 4d6bc5ff..b429af07 100644 --- a/docs/aio.md +++ b/docs/aio.md @@ -145,7 +145,7 @@ Cancelling the task again during the cancellation request or these waits raises With `kill_on_interrupt=False`, `asyncio.CancelledError` is raised immediately and the query keeps running. A timeout from `asyncio.wait_for()` therefore cancels the query and raises `asyncio.TimeoutError`. -`query_id` is `None` only if no query was started, for example when the timeout expires while looking up a cached result. +`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 diff --git a/pyathena/aio/spark/cursor.py b/pyathena/aio/spark/cursor.py index bc5d6501..7f72bc3e 100644 --- a/pyathena/aio/spark/cursor.py +++ b/pyathena/aio/spark/cursor.py @@ -348,6 +348,9 @@ async def execute( # type: ignore[override] Returns: Self reference for method chaining. """ + # A failure below must not leave the previous calculation on the cursor. + self._calculation_id = None + self._calculation_execution = None self._calculation_id = await self._calculate( session_id=session_id if session_id else self._session_id, code_block=operation, diff --git a/pyathena/spark/cursor.py b/pyathena/spark/cursor.py index f0f25241..61de3151 100644 --- a/pyathena/spark/cursor.py +++ b/pyathena/spark/cursor.py @@ -144,6 +144,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 93c54c38..efda5852 100644 --- a/tests/pyathena/aio/test_cursor.py +++ b/tests/pyathena/aio/test_cursor.py @@ -332,14 +332,17 @@ async def test_execute_cancelled_while_starting(self): 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() - timer = threading.Timer(0.2, release.set) - timer.start() + task = asyncio.create_task(asyncio.wait_for(cursor.execute("SELECT 1"), timeout=0.05)) try: - with pytest.raises(asyncio.TimeoutError): - await asyncio.wait_for(cursor.execute("SELECT 1"), timeout=0.05) + 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: - timer.cancel() release.set() + with pytest.raises(asyncio.TimeoutError): + await task assert started.is_set() cancel.assert_awaited_once_with("query_id") 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):