Skip to content
Open
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
107 changes: 81 additions & 26 deletions solara/server/kernel.py
Original file line number Diff line number Diff line change
Expand Up @@ -250,7 +250,7 @@ def close(self):
except: # noqa
pass

def send(
def _wire_message(
self,
stream,
msg_or_type,
Expand All @@ -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)
Expand Down
13 changes: 12 additions & 1 deletion solara/server/kernel_context.py
Original file line number Diff line number Diff line change
Expand Up @@ -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")
Expand Down
9 changes: 9 additions & 0 deletions solara/server/static/main-vuetify.js
Original file line number Diff line number Diff line change
Expand Up @@ -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';
});
Expand Down
24 changes: 24 additions & 0 deletions tests/unit/lifecycle_test.py
Original file line number Diff line number Diff line change
@@ -1,4 +1,5 @@
import asyncio
import json
import sys
import time
from unittest.mock import Mock
Expand Down Expand Up @@ -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()
Loading