Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions changelog.md
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,7 @@ Upcoming (TBD)
Bug Fixes
--------
* Fix exceptions when interrupting "streaming" unbuffered queries.
* Fix exceptions when interrupting "executing" unbuffered queries.


Internal
Expand Down
30 changes: 24 additions & 6 deletions mycli/packages/execution/background_runner.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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

Expand All @@ -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)
Expand All @@ -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

Expand Down
131 changes: 130 additions & 1 deletion mycli_test/pytests/test_background_runner.py
Original file line number Diff line number Diff line change
@@ -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
Expand All @@ -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

Expand Down Expand Up @@ -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('<H', code) + b'#70100Query execution was interrupted'
connection._rfile = BytesIO(struct.pack('<I', len(payload)) + payload)
result = MySQLResult(connection)
result.unbuffered_active = True
connection._result = result
cursor = BackgroundSSCursor(connection)
cursor._result = result
return cursor


def test_cancelled_execute_defers_interruption_to_cleanup(
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
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('<I', len(payload))[:3] + b'\x01' + payload)
next_cursor = connection.cursor()
assert next_cursor.execute('SET @value = 1') == 0
assert next_cursor._cancelled_result is None
next_cursor.close()
finally:
connection._force_close()


def test_marker_for_old_result_does_not_suppress_interruption() -> 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)
Expand Down
Loading