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 @@ -6,6 +6,7 @@ Bug Fixes
* Fix exceptions when interrupting "streaming" unbuffered queries.
* Fix exceptions when interrupting "executing" unbuffered queries.
* Don't try to kill a remote thread when a query is already in the "rendering" state.
* Don't try to kill a remote thread when a query is already in the "transforming" state.


Documentation
Expand Down
2 changes: 1 addition & 1 deletion mycli/packages/execution/background_runner.py
Original file line number Diff line number Diff line change
Expand Up @@ -154,7 +154,7 @@ def rendering(self, state: QueryState = QueryState.RENDERING) -> Iterator[None]:
try:
yield
except KeyboardInterrupt:
if self._local_state() == QueryState.RENDERING and not self._busy:
if self._local_state() in (QueryState.RENDERING, QueryState.TRANSFORMING) and not self._busy:
raise QueryCancelled(False) from None
raise
finally:
Expand Down
18 changes: 10 additions & 8 deletions mycli_test/pytests/test_background_runner.py
Original file line number Diff line number Diff line change
Expand Up @@ -914,7 +914,7 @@ def test_new_statement_resets_update_deadline(runner: BackgroundRunner, monkeypa
@pytest.mark.parametrize('state', [QueryState.RENDERING, QueryState.TRANSFORMING])
def test_rendering_error_stops_ticker(runner: BackgroundRunner, error: BaseException, state: QueryState) -> None:
runner.show_state = True
expected = QueryCancelled if isinstance(error, KeyboardInterrupt) and state == QueryState.RENDERING else type(error)
expected = QueryCancelled if isinstance(error, KeyboardInterrupt) else type(error)
with pytest.raises(expected):
with runner.rendering(state):
thread = runner._render_thread
Expand Down Expand Up @@ -990,23 +990,25 @@ def test_disabled_rendering_starts_no_thread(runner: BackgroundRunner, state: Qu
assert runner._render_thread is None


def test_nested_rendering_interrupt_is_local(runner: BackgroundRunner) -> None:
with runner.rendering(QueryState.TRANSFORMING):
@pytest.mark.parametrize('state', [QueryState.RENDERING, QueryState.TRANSFORMING])
def test_nested_rendering_interrupt_is_local(runner: BackgroundRunner, state: QueryState) -> None:
outer_state = QueryState.TRANSFORMING if state == QueryState.RENDERING else QueryState.RENDERING
with runner.rendering(outer_state):
with pytest.raises(QueryCancelled) as raised:
with runner.rendering():
with runner.rendering(state):
raise KeyboardInterrupt
assert not raised.value.disconnected
assert runner._render_state == QueryState.TRANSFORMING
assert runner._render_state == outer_state
assert runner._render_depth == 1
assert runner._render_depth == 0


@pytest.mark.parametrize('phase', ['streaming', 'transforming', 'database'])
@pytest.mark.parametrize('phase', ['streaming', 'transforming_database', 'database'])
def test_non_rendering_interrupt_keeps_existing_handling(runner: BackgroundRunner, phase: str) -> None:
state = QueryState.TRANSFORMING if phase == 'transforming' else QueryState.RENDERING
state = QueryState.TRANSFORMING if phase == 'transforming_database' else QueryState.RENDERING
if phase == 'streaming':
runner.attach(Connection(defer_connect=True, cursorclass=SSCursor), Mock())
runner._busy = phase == 'database'
runner._busy = phase != 'streaming'
try:
with pytest.raises(KeyboardInterrupt):
with runner.rendering(state):
Expand Down
63 changes: 63 additions & 0 deletions mycli_test/pytests/test_main_modes_repl.py
Original file line number Diff line number Diff line change
Expand Up @@ -2409,6 +2409,69 @@ def run(self, text: str) -> Iterator[SQLResult]:
assert cli.query_history[-1].successful is False


@pytest.mark.parametrize('show_state', [False, True])
@pytest.mark.parametrize('unbuffered', [False, True])
@pytest.mark.parametrize('save', [False, True])
def test_transform_interrupt_does_not_cancel_remote_query(
monkeypatch: pytest.MonkeyPatch,
show_state: bool,
unbuffered: bool,
save: bool,
) -> None:
patch_repl_runtime_defaults(monkeypatch)
runner = BackgroundRunner(0)
runner.show_state = show_state
runner.interval = 60
connection = pymysql.Connection(
defer_connect=True,
cursorclass=pymysql.cursors.SSCursor if unbuffered else pymysql.cursors.Cursor,
)
connection.server_thread_id = (42,)
connection._sock = Mock()
runner.attach(connection, Mock())
sql = SimpleNamespace(dbname='db', connection_id=42, conn=connection, background_runner=runner, connect=Mock())
cli = make_repl_cli(sql)
cli.reconnect = Mock()
sql.run = Mock(side_effect=lambda text: iter([SQLResult(header=['id'], rows=[(1,)])]))
remote_cancel = Mock()
hook = Mock()
monkeypatch.setattr(runner, '_control', remote_cancel)
monkeypatch.setattr(repl_mode.special_commands, 'run_post_redirect_hook', hook)
monkeypatch.setattr(repl_mode, 'prepare_polars_transform', lambda *args: object())

def interrupt_transform(transform: Any, results: Iterator[SQLResult], *args: Any, **kwargs: Any) -> SQLResult:
list(results)
raise KeyboardInterrupt

monkeypatch.setattr(repl_mode, 'run_polars_transform', interrupt_transform)
command = 'SELECT 1 .| df.head()' + (' .> result.parquet' if save else '')
try:
repl_mode._one_iteration(cli, repl_mode.ReplState(), command)

sql.run.assert_called_once_with('SELECT 1')
remote_cancel.assert_not_called()
sql.connect.assert_not_called()
cli.reconnect.assert_not_called()
hook.assert_not_called()
assert cli.output_calls == []
assert 'Query cancelled.' in cli.echo_calls
assert cli.query_history[-1].successful is False
assert cli.query_history[-1].query == command
assert runner._render_depth == 0
assert runner._render_thread is None
assert not runner.visible
assert sql.conn is connection and connection.open

repl_mode._one_iteration(cli, repl_mode.ReplState(), 'SELECT 2')
assert bool(cli.query_history[-1].successful)
remote_cancel.assert_not_called()
sql.connect.assert_not_called()
cli.reconnect.assert_not_called()
finally:
runner.close()
connection._force_close()


@pytest.mark.parametrize('show_state', [False, True])
@pytest.mark.parametrize('lazy', [False, True])
def test_rendering_interrupt_does_not_cancel_remote_query(
Expand Down
Loading