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
2 changes: 1 addition & 1 deletion code_sandboxes/__version__.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,4 +3,4 @@

"""Code Sandboxes."""

__version__ = "0.0.19"
__version__ = "0.0.21"
46 changes: 40 additions & 6 deletions code_sandboxes/jupyter_sandbox.py
Original file line number Diff line number Diff line change
Expand Up @@ -60,6 +60,10 @@ def __init__(
port: int = DEFAULT_PORT,
python_executable: str | None = None,
separate_process: bool = True,
kernel_id: str | None = None,
kernel_path: str | None = None,
client_kwargs: dict | None = None,
reuse_kernel: bool = True,
**kwargs,
):
super().__init__(config)
Expand Down Expand Up @@ -90,6 +94,14 @@ def __init__(
self._workdir_tmp: str | None = None
self._extra_kwargs = kwargs
self._owns_server = server_url is None
# Explicit kernel to connect to. When ``kernel_id`` is None and
# ``reuse_kernel`` is True, the sandbox attempts to reuse a pre-warmed
# kernel on the server; when ``kernel_id`` is None and
# ``reuse_kernel`` is False, a brand-new kernel is created.
self._kernel_id = kernel_id
self._kernel_path = kernel_path
self._client_kwargs = client_kwargs
self._reuse_kernel = reuse_kernel

@classmethod
def list_environments(cls) -> list[SandboxEnvironment]:
Expand Down Expand Up @@ -339,15 +351,26 @@ def start(self) -> None:

self._wait_for_server(timeout=self.config.timeout or DEFAULT_STARTUP_TIMEOUT)

# Try to reuse an existing pre-warmed kernel instead of creating a new one.
# Jupyter runtimes pre-warm a kernel at startup; connecting to it avoids
# unnecessary kernel proliferation (3 kernels → 2, or 2 → 1).
kernel_id = self._find_existing_kernel()
# Decide which kernel to connect to:
# - an explicit kernel_id always wins;
# - otherwise, when reuse is enabled, reuse a pre-warmed kernel (Jupyter
# runtimes pre-warm one at startup) to avoid kernel proliferation;
# - otherwise connect with no id so the client starts a brand-new kernel.
if self._kernel_id is not None:
kernel_id = self._kernel_id
elif self._reuse_kernel:
kernel_id = self._find_existing_kernel()
else:
kernel_id = None
Comment thread
echarles marked this conversation as resolved.

self._client = KernelClient(
server_url=self._server_url, token=self._token, kernel_id=kernel_id
server_url=self._server_url,
token=self._token,
kernel_id=kernel_id,
client_kwargs=self._client_kwargs or None,
)
self._client.start()

self._client.start(path=self._kernel_path)
Comment thread
echarles marked this conversation as resolved.

self._default_context = self.create_context("default")
self._info = SandboxInfo(
Expand All @@ -361,6 +384,17 @@ def start(self) -> None:
)
self._started = True

@property
def kernel_client(self):
"""The underlying ``jupyter_kernel_client.KernelClient``.

Exposed so callers that need the full low-level kernel API (for
example streaming execution via ``execute_interactive``) can delegate
to the same client the sandbox uses internally. ``None`` until
:meth:`start` has been called.
"""
return self._client

def _setup_tool_caller(self) -> None:
"""Keep tool calling on the client side for Jupyter sandboxes."""
return
Expand Down
132 changes: 132 additions & 0 deletions tests/test_jupyter.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,8 @@
"""jupyter sandbox tests."""

import os
import sys
import types
from pathlib import Path

import pytest
Expand All @@ -13,6 +15,136 @@
from code_sandboxes.models import SandboxConfig


def test_explicit_kernel_id_wins_over_reuse(monkeypatch):
"""Explicit kernel_id takes precedence even when reuse is enabled."""

captured: dict[str, object] = {}

class _KernelClientStub:
def __init__(self, server_url, token, kernel_id, client_kwargs=None):
captured["server_url"] = server_url
captured["token"] = token
captured["kernel_id"] = kernel_id
captured["client_kwargs"] = client_kwargs

def start(self, path=None):
captured["path"] = path

def stop(self):
return None

monkeypatch.setitem(
sys.modules,
"jupyter_kernel_client",
types.SimpleNamespace(KernelClient=_KernelClientStub),
)

sandbox = JupyterSandbox(
server_url="http://localhost:8888",
kernel_id="explicit-kernel",
reuse_kernel=True,
)

monkeypatch.setattr(sandbox, "_wait_for_server", lambda timeout=None: None)

def _should_not_be_called():
raise AssertionError("_find_existing_kernel should not be called with explicit kernel_id")

monkeypatch.setattr(sandbox, "_find_existing_kernel", _should_not_be_called)

sandbox.start()
try:
assert captured["kernel_id"] == "explicit-kernel"
finally:
sandbox.stop()


def test_reuse_kernel_false_forces_new_kernel(monkeypatch):
"""When reuse_kernel is False and no kernel_id is provided, connect with kernel_id=None."""

captured: dict[str, object] = {}

class _KernelClientStub:
def __init__(self, server_url, token, kernel_id, client_kwargs=None):
captured["kernel_id"] = kernel_id

def start(self, path=None):
return None

def stop(self):
return None

monkeypatch.setitem(
sys.modules,
"jupyter_kernel_client",
types.SimpleNamespace(KernelClient=_KernelClientStub),
)

sandbox = JupyterSandbox(
server_url="http://localhost:8888",
kernel_id=None,
reuse_kernel=False,
)

monkeypatch.setattr(sandbox, "_wait_for_server", lambda timeout=None: None)

def _should_not_be_called():
raise AssertionError("_find_existing_kernel should not be called when reuse_kernel=False")

monkeypatch.setattr(sandbox, "_find_existing_kernel", _should_not_be_called)

sandbox.start()
try:
assert captured["kernel_id"] is None
finally:
sandbox.stop()


def test_kernel_client_forwards_client_kwargs(monkeypatch, tmp_path: Path):
"""JupyterSandbox forwards client_kwargs to KernelClient."""

captured: dict[str, object] = {}

class _KernelClientStub:
def __init__(self, server_url, token, kernel_id, client_kwargs=None):
captured["server_url"] = server_url
captured["token"] = token
captured["kernel_id"] = kernel_id
captured["client_kwargs"] = client_kwargs

def start(self, path=None):
captured["start_path"] = path

def stop(self):
return None

monkeypatch.setitem(
sys.modules,
"jupyter_kernel_client",
types.SimpleNamespace(KernelClient=_KernelClientStub),
)

notebook_path = str(tmp_path / "notebook.ipynb")

sandbox = JupyterSandbox(
server_url="http://localhost:8888",
kernel_id="kernel-1",
kernel_path=notebook_path,
client_kwargs={"reconnect_interval": 5},
reuse_kernel=False,
)

monkeypatch.setattr(sandbox, "_wait_for_server", lambda timeout=None: None)

sandbox.start()
try:
assert captured["kernel_id"] == "kernel-1"
assert captured.get("client_kwargs") == {"reconnect_interval": 5}
assert captured.get("start_path") == notebook_path
finally:
sandbox.stop()


class TestJupyterSandbox:
"""Tests for JupyterSandbox."""

Expand Down