diff --git a/changelog.md b/changelog.md index 2d7a5acb..8ba34562 100644 --- a/changelog.md +++ b/changelog.md @@ -4,6 +4,7 @@ Upcoming (TBD) Bug Fixes -------- * Fix exceptions when interrupting "streaming" unbuffered queries. +* Fix exceptions when interrupting "executing" unbuffered queries. Internal diff --git a/mycli/packages/execution/background_runner.py b/mycli/packages/execution/background_runner.py index 76a64362..3f9cfb1a 100644 --- a/mycli/packages/execution/background_runner.py +++ b/mycli/packages/execution/background_runner.py @@ -17,7 +17,9 @@ from typing import Any, Callable, Iterator, TypeVar from prompt_toolkit.utils import get_cwidth +from pymysql import OperationalError from pymysql.connections import Connection, MySQLResult +from pymysql.constants import ER from pymysql.cursors import Cursor, SSCursor from mycli.constants import DEFAULT_WIDTH, TTY_ERASE_LINE, QueryState @@ -53,7 +55,12 @@ def call(cursor: Cursor, *args: Any, **kwargs: Any) -> Any: runner = getattr(cursor.connection, '_mycli_background_runner', None) if runner is None: return method(cursor, *args, **kwargs) - return runner.call(lambda: method(cursor, *args, **kwargs), new_statement=method.__name__ == 'execute') + try: + return runner.call(lambda: method(cursor, *args, **kwargs), new_statement=method.__name__ == 'execute') + except QueryCancelled: + if isinstance(cursor, BackgroundSSCursor): + cursor._cancelled_result = cursor._result + raise return call @@ -67,6 +74,7 @@ class BackgroundCursor(Cursor): class BackgroundSSCursor(SSCursor): connection: Connection | None _result: MySQLResult | None + _cancelled_result: MySQLResult | None = None execute = background(SSCursor.execute) nextset = background(SSCursor.nextset) @@ -82,14 +90,24 @@ def close(self) -> None: try: if connection.open: super().close() + except OperationalError as exc: + if exc.args[0] != ER.QUERY_INTERRUPTED or self._result is None or self._result is not self._cancelled_result: + raise + self._discard_result() finally: if not connection.open: # A disconnected stream cannot be drained, even during destruction. - if self._result is not None: - self._result.unbuffered_active = False - self._result.connection = None - self._result = None - self.connection = None + self._discard_result() + if self.connection is None: + self._cancelled_result = None + + def _discard_result(self) -> None: + if self._result is not None: + self._result.unbuffered_active = False + self._result.has_next = False + self._result.connection = None + self._result = None + self.connection = None __del__ = close diff --git a/mycli_test/pytests/test_background_runner.py b/mycli_test/pytests/test_background_runner.py index 7bf96dd7..874f6189 100644 --- a/mycli_test/pytests/test_background_runner.py +++ b/mycli_test/pytests/test_background_runner.py @@ -1,8 +1,9 @@ from collections.abc import Iterator from concurrent.futures import Future import gc -from io import StringIO +from io import BytesIO, StringIO import signal +import struct import sys from threading import Event, get_ident from time import monotonic @@ -13,6 +14,7 @@ from pymysql import OperationalError from pymysql.connections import Connection, MySQLResult +from pymysql.constants import ER from pymysql.cursors import Cursor, SSCursor import pytest @@ -493,6 +495,133 @@ def drain() -> None: assert streaming_cursor.connection is None +def _cursor_with_pending_error(code: int = ER.QUERY_INTERRUPTED) -> BackgroundSSCursor: + connection = Connection(defer_connect=True, cursorclass=SSCursor) + connection.server_thread_id = (42,) + connection._sock = Mock() + connection._current_timeout = connection._read_timeout + connection._next_seq_id = 0 + payload = b'\xff' + struct.pack(' None: + cursor = _cursor_with_pending_error() + connection = cursor.connection + result = cursor._result + assert connection is not None and result is not None + runner.attach(connection, Mock()) + original_call = runner._call + interrupts = iter([True]) + try: + with monkeypatch.context() as patch: + patch.setattr(cursor, '_query', lambda sql: 1) + patch.setattr(runner, '_call', lambda operation, interrupted: original_call(operation, lambda: next(interrupts, False))) + with pytest.raises(QueryCancelled): + cursor.execute('SELECT 1') + + assert cursor._cancelled_result is result + assert result is not None + cursor.close() + cursor.close() + + assert not result.unbuffered_active + assert result.connection is None + assert not result.has_next + assert cursor.connection is None + assert cursor._cancelled_result is None + assert connection.open + finally: + result.unbuffered_active = False + cursor.connection = None + connection._force_close() + + +def test_cancelled_streaming_cursor_destruction_consumes_interruption( + runner: BackgroundRunner, + monkeypatch: pytest.MonkeyPatch, +) -> None: + cursor = _cursor_with_pending_error() + connection = cursor.connection + result = cursor._result + assert connection is not None and result is not None + cursor._cancelled_result = result + runner.attach(connection, Mock()) + reference = weakref.ref(cursor) + unraisable = Mock() + monkeypatch.setattr(sys, 'unraisablehook', unraisable) + try: + del cursor + gc.collect() + + assert reference() is None + unraisable.assert_not_called() + assert not result.unbuffered_active + assert connection.open + finally: + result.unbuffered_active = False + connection._force_close() + + +@pytest.mark.parametrize('marked, code', [(False, ER.QUERY_INTERRUPTED), (True, 2013)]) +def test_streaming_cleanup_does_not_hide_unexpected_error(marked: bool, code: int) -> None: + cursor = _cursor_with_pending_error(code) + connection = cursor.connection + result = cursor._result + assert connection is not None and result is not None + if marked: + cursor._cancelled_result = result + try: + with pytest.raises(OperationalError) as raised: + cursor.close() + assert raised.value.args[0] == code + finally: + result.unbuffered_active = False + cursor.connection = None + connection._force_close() + + +def test_cancelled_result_does_not_mark_successor(runner: BackgroundRunner) -> None: + cursor = _cursor_with_pending_error() + connection = cursor.connection + assert connection is not None + cursor._cancelled_result = cursor._result + runner.attach(connection, Mock()) + try: + cursor.close() + # An OK packet proves a subsequent command can use the same connection. + payload = b'\x00\x00\x00\x02\x00\x00\x00' + connection._rfile = BytesIO(struct.pack(' None: + cursor = _cursor_with_pending_error() + connection = cursor.connection + assert connection is not None + cursor._cancelled_result = MySQLResult(connection) + try: + with pytest.raises(OperationalError, match='Query execution was interrupted'): + cursor.close() + finally: + cursor.connection = None + connection._force_close() + + def test_monitor_queries_separate_connection(runner: BackgroundRunner) -> None: cursor = Mock() cursor.__enter__ = Mock(return_value=cursor)