diff --git a/solara/server/kernel.py b/solara/server/kernel.py index 09420a9cc..76f047b80 100644 --- a/solara/server/kernel.py +++ b/solara/server/kernel.py @@ -250,7 +250,7 @@ def close(self): except: # noqa pass - def send( + def _wire_message( self, stream, msg_or_type, @@ -263,32 +263,87 @@ def send( metadata=None, ): if stream is None: - return # can happen when the kernel is closed but someone was still trying to send a message + return None # can happen when the kernel is closed but someone was still trying to send a message + if isinstance(msg_or_type, dict): + msg = msg_or_type + else: + msg = self.msg( + msg_or_type, + content=content, + parent=parent, + header=header, + metadata=metadata, + ) + _fix_msg(msg) + msg["channel"] = stream.channel try: - if isinstance(msg_or_type, dict): - msg = msg_or_type - else: - msg = self.msg( - msg_or_type, - content=content, - parent=parent, - header=header, - metadata=metadata, - ) - _fix_msg(msg) - msg["channel"] = stream.channel - # not using pdb guard for performance reasons - try: - if buffers: - msg["buffers"] = [memoryview(k).cast("b") for k in buffers] - wire_message = serialize_binary_message(msg) - else: - wire_message = json_dumps(msg) - except Exception: - logger.exception("Could not serialize message: %r", msg) - if settings.main.use_pdb: - pdb.post_mortem() - raise + if buffers: + msg["buffers"] = [memoryview(k).cast("b") for k in buffers] + return serialize_binary_message(msg) + return json_dumps(msg) + except Exception: + logger.exception("Could not serialize message: %r", msg) + if settings.main.use_pdb: + pdb.post_mortem() + raise + + def send_to( + self, + websocket_target: websocket.WebsocketWrapper, + stream, + msg_or_type, + content=None, + parent=None, + ident=None, + buffers=None, + track=False, + header=None, + metadata=None, + ): + try: + wire_message = self._wire_message( + stream, + msg_or_type, + content=content, + parent=parent, + ident=ident, + buffers=buffers, + track=track, + header=header, + metadata=metadata, + ) + if wire_message is None: + return + send_websockets({websocket_target}, wire_message) + except Exception as e: + logger.exception("Error sending message: %s", e) + + def send( + self, + stream, + msg_or_type, + content=None, + parent=None, + ident=None, + buffers=None, + track=False, + header=None, + metadata=None, + ): + try: + wire_message = self._wire_message( + stream, + msg_or_type, + content=content, + parent=parent, + ident=ident, + buffers=buffers, + track=track, + header=header, + metadata=metadata, + ) + if wire_message is None: + return send_websockets(self.websockets, wire_message) except Exception as e: logger.exception("Error sending message: %s", e) diff --git a/solara/server/kernel_context.py b/solara/server/kernel_context.py index a5a7a90b2..98d0dedbd 100644 --- a/solara/server/kernel_context.py +++ b/solara/server/kernel_context.py @@ -475,7 +475,18 @@ def initialize_virtual_kernel(session_id: str, kernel_id: str, websocket: websoc context = contexts[kernel_id] if context.session_id != session_id: logger.critical("Session id mismatch when reusing kernel (hack attempt?): %s != %s", context.session_id, session_id) - websocket.send_text("Session id mismatch when reusing kernel (hack attempt?)") + context.kernel.session.send_to( + websocket, + context.kernel.iopub_socket, + "solara_kernel_terminated", + { + "reason": "session_id_mismatch", + "message": "Rejected websocket connection for existing virtual kernel because the session id did not match.", + "kernel_id": kernel_id, + "retryable": False, + }, + ) + websocket.close() # to avoid very fast reconnects (we are in a thread anyway) time.sleep(0.5) raise ValueError("Session id mismatch") diff --git a/solara/server/static/main-vuetify.js b/solara/server/static/main-vuetify.js index d4bce7148..f0ce6c86f 100644 --- a/solara/server/static/main-vuetify.js +++ b/solara/server/static/main-vuetify.js @@ -133,6 +133,15 @@ async function solaraInit(mountId, appName) { } const close_url = `${solara.rootPath}/_solara/api/close/${kernelId}?session_id=${kernel.clientId}`; let skipReconnectedCheck = true; + kernel.iopubMessage.connect((_, msg) => { + if (msg.header.msg_type !== 'solara_kernel_terminated') { + return; + } + console.error('Solara kernel terminated:', msg.content); + window.dispatchEvent(new CustomEvent('solara.kernelTerminated', { + detail: msg.content, + })); + }); kernel.statusChanged.connect(() => { app.$data.kernelBusy = kernel.status == 'busy'; }); diff --git a/tests/unit/lifecycle_test.py b/tests/unit/lifecycle_test.py index 4a31e88ea..9f1e2e54d 100644 --- a/tests/unit/lifecycle_test.py +++ b/tests/unit/lifecycle_test.py @@ -1,4 +1,5 @@ import asyncio +import json import sys import time from unittest.mock import Mock @@ -114,3 +115,26 @@ async def test_kernel_lifecycle_close_while_disconnected(close_first, short_cull assert not context.closed_event.is_set() await cull_task_2 assert context.closed_event.is_set() + + +def test_kernel_lifecycle_session_id_mismatch_sends_solara_message(): + websocket = Mock() + kernel_context.initialize_virtual_kernel("session-id-1", "kernel-id-1", websocket) + + mismatched_websocket = Mock() + with pytest.raises(ValueError, match="Session id mismatch"): + kernel_context.initialize_virtual_kernel("session-id-2", "kernel-id-1", mismatched_websocket) + + mismatched_websocket.send.assert_called_once() + raw_message = mismatched_websocket.send.call_args.args[0] + msg = json.loads(raw_message) + + assert msg["channel"] == "iopub" + assert msg["header"]["msg_type"] == "solara_kernel_terminated" + assert msg["content"] == { + "reason": "session_id_mismatch", + "message": "Rejected websocket connection for existing virtual kernel because the session id did not match.", + "kernel_id": "kernel-id-1", + "retryable": False, + } + mismatched_websocket.close.assert_called_once()