From e4c0c2692764cdcf80d1142bb2997c84d36068a3 Mon Sep 17 00:00:00 2001 From: Alex Date: Thu, 9 Jul 2026 19:56:13 +0100 Subject: [PATCH] fix: more sandbox protections to terminate plus max pages --- application/agents/tools/code_executor.py | 6 +- .../agents/workflows/workflow_engine.py | 6 +- application/parser/document_reader.py | 93 ++++++- application/sandbox/base.py | 3 + application/sandbox/jupyter_gateway.py | 133 ++++++++-- application/sandbox/manager.py | 46 ++-- tests/agents/test_workflow_code_node.py | 26 +- tests/parser/test_document_reader.py | 65 +++++ .../sandbox/test_jupyter_gateway_isolation.py | 242 ++++++++++++++++++ tests/sandbox/test_sandbox_manager.py | 75 ++++++ tests/test_code_executor_tool.py | 21 ++ 11 files changed, 661 insertions(+), 55 deletions(-) diff --git a/application/agents/tools/code_executor.py b/application/agents/tools/code_executor.py index 0bcc7910..074460ff 100644 --- a/application/agents/tools/code_executor.py +++ b/application/agents/tools/code_executor.py @@ -234,10 +234,10 @@ class CodeExecutorTool(Tool): logger.exception("code_executor: exec raised") return {"status": "error", "error": f"execution failed: {type(exc).__name__}: {exc}"} - # Capture even on error/timeout so partial outputs aren't lost; a - # capture failure must never mask the run's real status. + # Capture even on error/timeout while the runtime remains reachable + # so partial outputs aren't lost; capture never masks the run status. artifacts: List[Dict[str, Any]] = [] - if should_capture: + if should_capture and not result.runtime_invalidated: try: artifacts = self._capture_artifacts(manager, session_id, pre_signatures, outputs) except Exception: diff --git a/application/agents/workflows/workflow_engine.py b/application/agents/workflows/workflow_engine.py index f4f49c6a..87f14e01 100644 --- a/application/agents/workflows/workflow_engine.py +++ b/application/agents/workflows/workflow_engine.py @@ -499,7 +499,11 @@ class WorkflowEngine: "status": "completed" if result.ok else "error", } ] - if self.run_persisted: + if result.runtime_invalidated: + # A hard timeout destroyed the workspace runtime, so there is + # nothing reachable to capture and the timeout below stays primary. + artifacts = [] + elif self.run_persisted: artifacts = capture_artifacts( manager, session_id, diff --git a/application/parser/document_reader.py b/application/parser/document_reader.py index c9483759..3edbfb6a 100644 --- a/application/parser/document_reader.py +++ b/application/parser/document_reader.py @@ -15,7 +15,7 @@ import os import tempfile import zipfile from pathlib import Path -from typing import Any, Dict, List, Optional +from typing import Any, Dict, Iterator, List, Optional from application.core.settings import settings from application.parser.file.bulk import get_default_file_extractor @@ -35,6 +35,8 @@ _MAX_CELL_CHARS = 200 # Caps applied to the bounded view that rides back through the Redis result # backend (the full result still lives in the persisted artifact). _MAX_CHUNKS_RETURNED = 50 +_MAX_PAGE_SELECTIONS = 10_000 +_MAX_PAGE_SELECTOR_TOKENS = 20_000 _VALID_OUTPUTS = ("markdown", "text", "structured", "chunks") _VALID_OCR = ("auto", "on", "off") @@ -256,35 +258,100 @@ def _structure_summary(structured: Any) -> Dict[str, int]: def _apply_pages(text: str, pages: Any) -> str: """Best-effort page-range slice on a page-delimited markdown blob (``\\f`` separated).""" - if not pages: + if not pages or "\f" not in text: return text - parts = text.split("\f") - if len(parts) <= 1: - return text - selected = _selected_page_indices(pages, len(parts)) + total = text.count("\f") + 1 + selected = _selected_page_indices(pages, total) if not selected: return text - return "\f".join(parts[i] for i in selected if 0 <= i < len(parts)) + + # Splitting a hostile form-feed blob could allocate millions of strings. + # Walk boundaries and retain only the bounded set of requested pages. + wanted = set(selected) + found: Dict[int, tuple[int, int]] = {} + page_index = 0 + start = 0 + while wanted: + end = text.find("\f", start) + if end < 0: + end = len(text) + if page_index in wanted: + found[page_index] = (start, end) + wanted.remove(page_index) + if end == len(text): + break + page_index += 1 + start = end + 1 + + # Preserve requested order and ordinary duplicates, but never let repeated + # selectors amplify the output beyond the source text's resident size. + output = io.StringIO() + output_chars = 0 + wrote_page = False + for index in selected: + span = found.get(index) + if span is None: + continue + page_start, page_end = span + added = (page_end - page_start) + (1 if wrote_page else 0) + if output_chars + added > len(text): + break + if wrote_page: + output.write("\f") + output.write(text[page_start:page_end]) + output_chars += added + wrote_page = True + return output.getvalue() if wrote_page else text + + +def _iter_page_tokens(pages: Any) -> Iterator[Any]: + """Yield bounded selector tokens without materializing a comma-split list.""" + if isinstance(pages, list): + for position, token in enumerate(pages): + if position >= _MAX_PAGE_SELECTOR_TOKENS: + break + yield token + return + + raw = str(pages) + start = 0 + emitted = 0 + while emitted < _MAX_PAGE_SELECTOR_TOKENS: + end = raw.find(",", start) + if end < 0: + yield raw[start:] + return + yield raw[start:end] + emitted += 1 + start = end + 1 def _selected_page_indices(pages: Any, total: int) -> List[int]: - """Parse ``pages`` ("1-3", "2", [1,2]) into 0-based indices bounded by ``total``.""" + """Parse ``pages`` into a bounded list of valid 0-based page occurrences.""" + if total <= 0: + return [] + indices: List[int] = [] - tokens = pages if isinstance(pages, list) else str(pages).split(",") - for token in tokens: + for token in _iter_page_tokens(pages): + if len(indices) >= _MAX_PAGE_SELECTIONS: + break token = str(token).strip() if "-" in token: try: lo, hi = (int(p) for p in token.split("-", 1)) except ValueError: continue - indices.extend(range(lo - 1, hi)) + start = max(lo - 1, 0) + stop = min(hi, total, start + (_MAX_PAGE_SELECTIONS - len(indices))) + indices.extend(range(start, stop)) else: try: - indices.append(int(token) - 1) + index = int(token) - 1 except ValueError: continue - return [i for i in indices if 0 <= i < total] + if 0 <= index < total: + indices.append(index) + return indices def _to_chunks(text: str, max_chars: Optional[int]) -> List[str]: diff --git a/application/sandbox/base.py b/application/sandbox/base.py index 5c3dceec..3fa892ce 100644 --- a/application/sandbox/base.py +++ b/application/sandbox/base.py @@ -37,6 +37,9 @@ class ExecResult: display_data: List[DisplayData] = field(default_factory=list) plots: List[Plot] = field(default_factory=list) truncated: bool = False # output exceeded the budget and was cut; status stays "ok" + # The backend invalidated the runtime while producing this result. Managers + # must discard their cached handle so the next open performs a cold start. + runtime_invalidated: bool = False @property def ok(self) -> bool: diff --git a/application/sandbox/jupyter_gateway.py b/application/sandbox/jupyter_gateway.py index fe917e6b..424bda2f 100644 --- a/application/sandbox/jupyter_gateway.py +++ b/application/sandbox/jupyter_gateway.py @@ -53,9 +53,12 @@ _CONTAINMENT_SNIPPET = ( class _Kernel: """Tracks one gateway kernel plus the per-session workspace it executes in.""" - def __init__(self, kernel_id: str, workspace: str) -> None: + def __init__( + self, kernel_id: str, workspace: str, session_id: Optional[str] = None + ) -> None: self.kernel_id = kernel_id self.workspace = workspace + self.session_id = session_id self.initialized = False @@ -86,6 +89,9 @@ class JupyterKernelGatewaySandbox(CodeSandbox): self._max_output_bytes = max_output_bytes self._max_file_bytes = max_file_bytes self._kernels: Dict[str, _Kernel] = {} + # Timed-out kernels whose DELETE could not yet be confirmed. Immutable + # ids stay retryable here and are never reused as active runtimes. + self._quarantined_kernels: Dict[str, Optional[str]] = {} self._lock = threading.Lock() # Session ids with a create in flight; a second open() for the same id # waits on this CV and reuses the result instead of double-creating a @@ -110,6 +116,9 @@ class JupyterKernelGatewaySandbox(CodeSandbox): return urlunparse((scheme, parsed.netloc, path, "", "", "")) def _get_kernel(self, session_id: str) -> _Kernel: + # A manager may cache a replacement handle, so normal traffic is also + # an opportunity to retire an older quarantined kernel. + self._retry_quarantined_kernels(session_id) with self._lock: kernel = self._kernels.get(session_id) if kernel is None: @@ -136,10 +145,24 @@ class JupyterKernelGatewaySandbox(CodeSandbox): while session_id in self._creating: self._create_cv.wait() existing = self._kernels.get(session_id) - if existing is not None: + has_quarantine = any( + owner == session_id for owner in self._quarantined_kernels.values() + ) + if existing is not None and not has_quarantine: return existing.kernel_id self._creating.add(session_id) try: + unresolved = self._retry_quarantined_kernels(session_id) + with self._lock: + # A replacement may already exist. It is safe to use even when + # deletion of an older exact id remains pending. + existing = self._kernels.get(session_id) + if existing is not None: + return existing.kernel_id + if unresolved: + raise RuntimeError( + f"Previous timed-out kernel for {session_id!r} could not be terminated" + ) resp = requests.post( f"{self._base_url}/api/kernels", headers=self._headers(), @@ -149,7 +172,7 @@ class JupyterKernelGatewaySandbox(CodeSandbox): resp.raise_for_status() kernel_id = resp.json()["id"] workspace = f"{_WORKSPACE_ROOT}/{session_id}" - kernel = _Kernel(kernel_id, workspace) + kernel = _Kernel(kernel_id, workspace, session_id) with self._lock: self._kernels[session_id] = kernel self._prime(kernel) @@ -162,6 +185,7 @@ class JupyterKernelGatewaySandbox(CodeSandbox): def attach(self, session_id: str) -> str: """Reattach to a still-running kernel for ``session_id``; open a cold one if gone.""" self._validate_session_id(session_id) + unresolved = self._retry_quarantined_kernels(session_id) with self._lock: existing = self._kernels.get(session_id) if existing is not None and self._kernel_alive(existing.kernel_id): @@ -169,6 +193,10 @@ class JupyterKernelGatewaySandbox(CodeSandbox): if existing is not None: with self._lock: self._kernels.pop(session_id, None) + if unresolved: + raise RuntimeError( + f"Previous timed-out kernel for {session_id!r} could not be terminated" + ) logger.warning("Re-attaching session %s to a cold kernel; previous state is lost", session_id) return self.open(session_id) @@ -205,16 +233,25 @@ class JupyterKernelGatewaySandbox(CodeSandbox): self._kernels.pop(session_id, None) self._delete_kernel(kernel_id) - def _delete_kernel(self, kernel_id: str) -> None: - """Best-effort DELETE of a gateway kernel by id (teardown never raises).""" - try: - requests.delete( - f"{self._base_url}/api/kernels/{kernel_id}", - headers=self._headers(), - timeout=self._http_timeout, - ) - except requests.RequestException as exc: # teardown is best-effort - logger.warning("Failed to delete kernel %s: %s", kernel_id, exc) + def _delete_kernel(self, kernel_id: str) -> bool: + """Best-effort DELETE of a gateway kernel, retrying once and never raising.""" + last_error = "unknown error" + for _attempt in range(2): + try: + resp = requests.delete( + f"{self._base_url}/api/kernels/{kernel_id}", + headers=self._headers(), + timeout=self._http_timeout, + ) + except requests.RequestException as exc: # teardown is best-effort + last_error = str(exc) + continue + # A missing kernel is already in the desired state. + if 200 <= resp.status_code < 300 or resp.status_code == 404: + return True + last_error = f"HTTP {resp.status_code}" + logger.warning("Failed to delete kernel %s after retry: %s", kernel_id, last_error) + return False def _kernel_alive(self, kernel_id: str) -> bool: try: @@ -238,8 +275,8 @@ class JupyterKernelGatewaySandbox(CodeSandbox): except requests.RequestException as exc: logger.warning("Failed to interrupt kernel %s: %s", kernel_id, exc) - def _interrupt_and_drain(self, ws: websocket.WebSocket, msg_id: str, kernel_id: str) -> None: - """Interrupt the kernel then drain frames until it idles, leaving the session reusable.""" + def _interrupt_and_drain(self, ws: websocket.WebSocket, msg_id: str, kernel_id: str) -> bool: + """Interrupt and drain the request, returning whether matching idle was observed.""" self._interrupt(kernel_id) drain_deadline = time.monotonic() + self._http_timeout while time.monotonic() < drain_deadline: @@ -247,7 +284,7 @@ class JupyterKernelGatewaySandbox(CodeSandbox): ws.settimeout(max(0.05, drain_deadline - time.monotonic())) raw = ws.recv() except (websocket.WebSocketTimeoutException, websocket.WebSocketConnectionClosedException): - return + return False if not raw: continue try: @@ -258,7 +295,51 @@ class JupyterKernelGatewaySandbox(CodeSandbox): continue msg_type = msg.get("msg_type") or msg.get("header", {}).get("msg_type") if msg_type == "status" and msg.get("content", {}).get("execution_state") == "idle": - return + return True + return False + + def _retry_quarantined_kernels(self, session_id: str) -> List[str]: + """Retry pending exact-id deletes for this session; return unresolved ids.""" + with self._lock: + pending = [ + (kernel_id, owner) + for kernel_id, owner in self._quarantined_kernels.items() + if owner == session_id + ] + unresolved: List[str] = [] + for kernel_id, owner in pending: + if self._delete_kernel(kernel_id): + with self._lock: + if ( + kernel_id in self._quarantined_kernels + and self._quarantined_kernels[kernel_id] == owner + ): + self._quarantined_kernels.pop(kernel_id, None) + else: + unresolved.append(kernel_id) + return unresolved + + def _invalidate_kernel( + self, kernel_id: str, session_id: Optional[str] = None + ) -> bool: + """Quarantine and hard-delete one exact kernel without touching a replacement.""" + with self._lock: + stale_sessions = [ + session_id + for session_id, kernel in self._kernels.items() + if kernel.kernel_id == kernel_id + ] + owner = session_id or (stale_sessions[0] if len(stale_sessions) == 1 else None) + # Register before eviction so a failed DELETE never loses the only + # retryable reference to this runaway kernel id. + self._quarantined_kernels[kernel_id] = owner + for session_id in stale_sessions: + self._kernels.pop(session_id, None) + deleted = self._delete_kernel(kernel_id) + if deleted: + with self._lock: + self._quarantined_kernels.pop(kernel_id, None) + return deleted def _prime(self, kernel: _Kernel) -> None: """Create the per-session workspace (mode 0700) and chdir the kernel into it.""" @@ -314,7 +395,14 @@ class JupyterKernelGatewaySandbox(CodeSandbox): try: msg_id = uuid.uuid4().hex ws.send(json.dumps(self._execute_request(msg_id, code))) - return self._collect(ws, msg_id, timeout, kernel.kernel_id, max_output_bytes) + return self._collect( + ws, + msg_id, + timeout, + kernel.kernel_id, + max_output_bytes, + session_id=kernel.session_id, + ) finally: try: ws.close() @@ -357,6 +445,7 @@ class JupyterKernelGatewaySandbox(CodeSandbox): timeout: float, kernel_id: str, max_output_bytes: Optional[int] = None, + session_id: Optional[str] = None, ) -> ExecResult: """Read iopub/shell frames until ``execute_reply``/idle, a wall-clock deadline, or a closed socket.""" effective = max_output_bytes if max_output_bytes is not None else self._max_output_bytes @@ -374,14 +463,18 @@ class JupyterKernelGatewaySandbox(CodeSandbox): remaining = deadline - now if remaining <= 0: self._fail(result, "TimeoutError", f"execution exceeded {timeout}s") - self._interrupt_and_drain(ws, msg_id, kernel_id) + if not self._interrupt_and_drain(ws, msg_id, kernel_id): + self._invalidate_kernel(kernel_id, session_id) + result.runtime_invalidated = True break try: ws.settimeout(remaining) raw = ws.recv() except websocket.WebSocketTimeoutException: self._fail(result, "TimeoutError", f"execution exceeded {timeout}s") - self._interrupt_and_drain(ws, msg_id, kernel_id) + if not self._interrupt_and_drain(ws, msg_id, kernel_id): + self._invalidate_kernel(kernel_id, session_id) + result.runtime_invalidated = True break except websocket.WebSocketConnectionClosedException: self._fail(result, "KernelDiedError", "kernel channel closed before completion") diff --git a/application/sandbox/manager.py b/application/sandbox/manager.py index c471eb6e..9007e2f5 100644 --- a/application/sandbox/manager.py +++ b/application/sandbox/manager.py @@ -194,40 +194,43 @@ class SandboxManager: def exec(self, session_id: str, code: str, timeout: Optional[float] = None) -> ExecResult: """Execute ``code`` in the bound session, holding it in-use so a reap/evict can't pull it.""" - self._enter(session_id) + session = self._enter(session_id) try: - return self._backend.exec(session_id, code, timeout) + result = self._backend.exec(session_id, code, timeout) + if result.runtime_invalidated: + self._drop_invalidated_session(session_id, session) + return result finally: - self._leave(session_id) + self._leave(session_id, expected=session) def put_file(self, session_id: str, dest_path: str, data: bytes) -> None: """Write ``data`` into the bound session's workspace.""" - self._enter(session_id) + session = self._enter(session_id) try: self._backend.put_file(session_id, dest_path, data) finally: - self._leave(session_id) + self._leave(session_id, expected=session) def get_file(self, session_id: str, path: str) -> bytes: """Read ``path`` from the bound session's workspace.""" - self._enter(session_id) + session = self._enter(session_id) try: return self._backend.get_file(session_id, path) finally: - self._leave(session_id) + self._leave(session_id, expected=session) def list_files(self, session_id: str) -> List[str]: """List files in the bound session's workspace.""" - self._enter(session_id) + session = self._enter(session_id) try: return self._backend.list_files(session_id) finally: - self._leave(session_id) + self._leave(session_id, expected=session) def remove_path(self, session_id: str, path: str) -> None: """Best-effort delete a workspace-relative path; never raises (cleanup must not fail an op).""" try: - self._enter(session_id) + session = self._enter(session_id) except KeyError: return try: @@ -239,7 +242,7 @@ class SandboxManager: except Exception: logger.exception("SandboxManager: best-effort remove_path failed for %r", path) finally: - self._leave(session_id) + self._leave(session_id, expected=session) def _remove_via_exec(self, session_id: str, path: str) -> None: """Fallback workspace cleanup: run a contained shutil.rmtree of the relative path.""" @@ -300,7 +303,7 @@ class SandboxManager: session = self._sessions.get(session_id) return session.ttl if session else None - def _enter(self, session_id: str) -> None: + def _enter(self, session_id: str) -> _Session: """Touch the idle clock and mark the session in-use so a concurrent reap/evict skips it.""" with self._lock: session = self._sessions.get(session_id) @@ -308,19 +311,28 @@ class SandboxManager: raise KeyError(f"No sandbox session bound for {session_id!r}") session.last_access = time.monotonic() session.in_use += 1 + return session - def _leave(self, session_id: str) -> None: + def _drop_invalidated_session(self, session_id: str, expected: _Session) -> None: + """Drop the exact manager entry whose backend runtime was already destroyed.""" + with self._lock: + if self._sessions.get(session_id) is expected: + self._sessions.pop(session_id, None) + + def _leave(self, session_id: str, expected: Optional[_Session] = None) -> None: """Release an in-use hold taken by ``_enter``; run a close deferred by ``close`` on the last release. - Idempotent if the session was already closed. When the final hold is released and - a ``close`` was deferred (``pending_close``), the session is popped here and its - backend torn down OUTSIDE the lock, keyed by the captured handle. + Idempotent if the session was already closed. If ``expected`` is supplied, + a newer entry under the same id is never modified (generation/ABA guard). + When the final hold is released and a ``close`` was deferred + (``pending_close``), the session is popped here and its backend torn down + OUTSIDE the lock, keyed by the captured handle. """ handle: Optional[str] = None do_close = False with self._lock: session = self._sessions.get(session_id) - if session is not None and session.in_use > 0: + if session is not None and (expected is None or session is expected) and session.in_use > 0: session.in_use -= 1 session.last_access = time.monotonic() if session.in_use == 0 and session.pending_close: diff --git a/tests/agents/test_workflow_code_node.py b/tests/agents/test_workflow_code_node.py index a1961ba7..e76de7f0 100644 --- a/tests/agents/test_workflow_code_node.py +++ b/tests/agents/test_workflow_code_node.py @@ -50,12 +50,20 @@ def _code_node(node_id="code_1", **config) -> WorkflowNode: class _Result: - def __init__(self, ok=True, stdout="", error_name=None, error_value=None): + def __init__( + self, + ok=True, + stdout="", + error_name=None, + error_value=None, + runtime_invalidated=False, + ): self.status = "ok" if ok else "error" self.stdout = stdout self.stderr = "" self.error_name = error_name self.error_value = error_value + self.runtime_invalidated = runtime_invalidated @property def ok(self): @@ -275,6 +283,22 @@ def test_code_node_failure_raises(patch_sandbox): list(engine._execute_code_node(node)) +def test_code_node_invalidated_runtime_skips_capture_and_preserves_timeout(patch_sandbox): + engine = _engine() + patch_sandbox["result"] = _Result( + ok=False, + error_name="TimeoutError", + error_value="execution exceeded 30s", + runtime_invalidated=True, + ) + node = _code_node(code="while True: pass") + + with pytest.raises(ValueError, match="failed: TimeoutError: execution exceeded 30s"): + list(engine._execute_code_node(node)) + + assert patch_sandbox["capture_calls"] == 0 + + def test_code_node_empty_code_raises(patch_sandbox): engine = _engine() node = _code_node(code=" ") diff --git a/tests/parser/test_document_reader.py b/tests/parser/test_document_reader.py index dbf464c0..e43e7547 100644 --- a/tests/parser/test_document_reader.py +++ b/tests/parser/test_document_reader.py @@ -260,6 +260,71 @@ def test_pages_slices_form_feed_blob(monkeypatch): assert out["content"] == "page2" +@pytest.mark.unit +def test_page_ranges_are_bounded_before_materializing(monkeypatch): + """Clamp hostile ranges to real pages while preserving selection semantics.""" + real_range = range + range_calls: List[tuple[int, int]] = [] + + def _guarded_range(start: int, stop: int) -> range: + range_calls.append((start, stop)) + if stop - start > 100: + raise AssertionError("attempted to materialize an unbounded page range") + return real_range(start, stop) + + monkeypatch.setattr(dr, "range", _guarded_range, raising=False) + + selected = dr._selected_page_indices("1-1000000000,1-1000000000", total=3) + + assert selected == [0, 1, 2, 0, 1, 2] + assert range_calls == [(0, 3), (0, 3)] + + +@pytest.mark.unit +def test_page_range_expansion_has_an_absolute_cap(monkeypatch): + """An attacker-controlled page count cannot turn a range into millions of ints.""" + real_range = range + range_calls: List[tuple[int, int]] = [] + + def _guarded_range(start: int, stop: int) -> range: + range_calls.append((start, stop)) + if stop - start > dr._MAX_PAGE_SELECTIONS: + raise AssertionError("attempted to expand beyond the page-selection cap") + return real_range(start, stop) + + monkeypatch.setattr(dr, "range", _guarded_range, raising=False) + + selected = dr._selected_page_indices("1-1000000000", total=25_000_000) + + assert len(selected) == dr._MAX_PAGE_SELECTIONS + assert selected[0] == 0 + assert selected[-1] == dr._MAX_PAGE_SELECTIONS - 1 + assert range_calls == [(0, dr._MAX_PAGE_SELECTIONS)] + + +@pytest.mark.unit +def test_page_slicing_never_split_materializes_all_pages(): + """Form-feed slicing walks boundaries without allocating one string per page.""" + + class _NoSplitText(str): + def split(self, *_args, **_kwargs): + raise AssertionError("page slicing must not call str.split") + + text = _NoSplitText("page1\fpage2\fpage3") + + assert dr._apply_pages(text, "3,1,3") == "page3\fpage1\fpage3" + + +@pytest.mark.unit +def test_duplicate_page_selection_cannot_amplify_source_text(): + """Repeated selectors remain bounded to the source text's resident size.""" + text = ("x" * 100) + "\fy" + + selected = dr._apply_pages(text, [1] * (dr._MAX_PAGE_SELECTIONS * 2)) + + assert len(selected) <= len(text) + + # --------------------------------------------------------------------------- # structured output (Docling stubbed) # --------------------------------------------------------------------------- diff --git a/tests/sandbox/test_jupyter_gateway_isolation.py b/tests/sandbox/test_jupyter_gateway_isolation.py index a23a0607..f87338df 100644 --- a/tests/sandbox/test_jupyter_gateway_isolation.py +++ b/tests/sandbox/test_jupyter_gateway_isolation.py @@ -175,6 +175,248 @@ class _FakeWS: raise websocket.WebSocketConnectionClosedException() +class _TimeoutWS: + """A channel that never yields another frame, even after an interrupt.""" + + def settimeout(self, _t): + pass + + def recv(self): + import websocket + + raise websocket.WebSocketTimeoutException() + + +class _TimeoutThenIdleWS: + """Reach the exec deadline, then report matching idle after the interrupt.""" + + def __init__(self, msg_id): + self._msg_id = msg_id + self._calls = 0 + + def settimeout(self, _t): + pass + + def recv(self): + import websocket + + self._calls += 1 + if self._calls == 1: + raise websocket.WebSocketTimeoutException() + return _frame(self._msg_id, "status", {"execution_state": "idle"}) + + +def test_timeout_deletes_and_invalidates_kernel_when_interrupt_never_idles(monkeypatch): + """A kernel that catches the interrupt cannot outlive the execution deadline.""" + sb = JupyterKernelGatewaySandbox(gateway_url="http://unused") + kernel = _Kernel("stuck-kernel", "/tmp/stuck") + sb._kernels["session-stuck"] = kernel + interrupted = [] + deleted = [] + monkeypatch.setattr(sb, "_interrupt", lambda kernel_id: interrupted.append(kernel_id)) + monkeypatch.setattr(sb, "_delete_kernel", lambda kernel_id: deleted.append(kernel_id) or True) + + result = sb._collect( + _TimeoutWS(), + "timed-out-message", + timeout=5, + kernel_id=kernel.kernel_id, + ) + + assert result.error_name == "TimeoutError" + assert result.runtime_invalidated is True + assert interrupted == [kernel.kernel_id] + assert deleted == [kernel.kernel_id] + assert "session-stuck" not in sb._kernels + + +def test_timeout_keeps_kernel_when_interrupt_reaches_matching_idle(monkeypatch): + """A confirmed-idle interrupted kernel remains reusable and is not deleted.""" + sb = JupyterKernelGatewaySandbox(gateway_url="http://unused") + kernel = _Kernel("reusable-kernel", "/tmp/reusable") + sb._kernels["session-reusable"] = kernel + interrupted = [] + deleted = [] + monkeypatch.setattr(sb, "_interrupt", lambda kernel_id: interrupted.append(kernel_id)) + monkeypatch.setattr(sb, "_delete_kernel", lambda kernel_id: deleted.append(kernel_id) or True) + msg_id = "timed-out-message" + + result = sb._collect( + _TimeoutThenIdleWS(msg_id), + msg_id, + timeout=5, + kernel_id=kernel.kernel_id, + ) + + assert result.error_name == "TimeoutError" + assert result.runtime_invalidated is False + assert interrupted == [kernel.kernel_id] + assert deleted == [] + assert sb._kernels["session-reusable"] is kernel + + +def test_timeout_invalidation_does_not_remove_replacement_kernel(monkeypatch): + """Late cleanup of an old timeout must leave a newly opened generation registered.""" + sb = JupyterKernelGatewaySandbox(gateway_url="http://unused") + old = _Kernel("old-kernel", "/tmp/old") + replacement = _Kernel("new-kernel", "/tmp/new") + sb._kernels["session-race"] = old + deleted = [] + + def replace_during_drain(*_args): + sb._kernels["session-race"] = replacement + return False + + monkeypatch.setattr(sb, "_interrupt_and_drain", replace_during_drain) + monkeypatch.setattr(sb, "_delete_kernel", lambda kernel_id: deleted.append(kernel_id) or True) + + result = sb._collect( + _TimeoutWS(), + "timed-out-message", + timeout=5, + kernel_id=old.kernel_id, + ) + + assert result.runtime_invalidated is True + assert deleted == [old.kernel_id] + assert sb._kernels["session-race"] is replacement + + +def test_failed_replacement_race_quarantine_does_not_block_unrelated_open(monkeypatch): + """A failed old-id DELETE remains scoped to its original session.""" + sb = JupyterKernelGatewaySandbox(gateway_url="http://unused") + old = _Kernel("old-kernel", "/tmp/old", "session-race") + replacement = _Kernel("new-kernel", "/tmp/new", "session-race") + + def replace_during_drain(*_args): + sb._kernels["session-race"] = replacement + return False + + class _Response: + status_code = 201 + + def raise_for_status(self): + pass + + def json(self): + return {"id": "unrelated-kernel"} + + monkeypatch.setattr(sb, "_interrupt_and_drain", replace_during_drain) + monkeypatch.setattr(sb, "_delete_kernel", lambda _kernel_id: False) + monkeypatch.setattr(jupyter_gateway.requests, "post", lambda *args, **kwargs: _Response()) + monkeypatch.setattr(sb, "_prime", lambda _kernel: None) + + result = sb._collect( + _TimeoutWS(), + "timed-out-message", + timeout=5, + kernel_id=old.kernel_id, + session_id=old.session_id, + ) + + assert result.runtime_invalidated is True + assert sb._quarantined_kernels == {old.kernel_id: "session-race"} + assert sb.open("unrelated") == "unrelated-kernel" + assert sb._kernels["session-race"] is replacement + assert sb._quarantined_kernels == {old.kernel_id: "session-race"} + + +def test_delete_kernel_retries_transient_http_failure(monkeypatch): + """A transient gateway failure gets one more termination attempt.""" + sb = JupyterKernelGatewaySandbox(gateway_url="http://unused") + statuses = iter((500, 204)) + calls = [] + + class _Response: + def __init__(self, status_code): + self.status_code = status_code + + def fake_delete(url, **_kwargs): + calls.append(url) + return _Response(next(statuses)) + + monkeypatch.setattr(jupyter_gateway.requests, "delete", fake_delete) + + assert sb._delete_kernel("stuck-kernel") is True + assert calls == [ + "http://unused/api/kernels/stuck-kernel", + "http://unused/api/kernels/stuck-kernel", + ] + + +def test_delete_kernel_reports_failure_after_both_attempts(monkeypatch, caplog): + """Two non-success responses leave termination explicitly unconfirmed.""" + sb = JupyterKernelGatewaySandbox(gateway_url="http://unused") + statuses = iter((500, 503)) + calls = [] + + class _Response: + def __init__(self, status_code): + self.status_code = status_code + + def fake_delete(url, **_kwargs): + calls.append(url) + return _Response(next(statuses)) + + monkeypatch.setattr(jupyter_gateway.requests, "delete", fake_delete) + + assert sb._delete_kernel("stuck-kernel") is False + assert len(calls) == 2 + assert "Failed to delete kernel stuck-kernel after retry: HTTP 503" in caplog.text + + +def test_failed_timeout_delete_is_quarantined_and_cold_open_fails_closed(monkeypatch): + """An undeletable runaway remains retryable and blocks a duplicate cold kernel.""" + sb = JupyterKernelGatewaySandbox(gateway_url="http://unused") + kernel = _Kernel("stuck-kernel", "/tmp/stuck") + sb._kernels["session-stuck"] = kernel + delete_calls = [] + post_calls = [] + monkeypatch.setattr(sb, "_interrupt", lambda _kernel_id: None) + monkeypatch.setattr( + sb, + "_delete_kernel", + lambda kernel_id: (delete_calls.append(kernel_id), False)[1], + ) + monkeypatch.setattr( + jupyter_gateway.requests, + "post", + lambda *args, **kwargs: post_calls.append((args, kwargs)), + ) + + result = sb._collect( + _TimeoutWS(), "timed-out-message", timeout=5, kernel_id=kernel.kernel_id + ) + + assert result.runtime_invalidated is True + assert "session-stuck" not in sb._kernels + assert sb._quarantined_kernels == {kernel.kernel_id: "session-stuck"} + with pytest.raises(RuntimeError, match="could not be terminated"): + sb.open("session-stuck") + assert delete_calls == [kernel.kernel_id, kernel.kernel_id] + assert post_calls == [] + assert sb._quarantined_kernels == {kernel.kernel_id: "session-stuck"} + + +def test_cached_replacement_retries_old_quarantine_without_touching_replacement(monkeypatch): + """Cached-session traffic retires only the old immutable kernel id.""" + sb = JupyterKernelGatewaySandbox(gateway_url="http://unused") + replacement = _Kernel("replacement-kernel", "/tmp/replacement") + sb._kernels["session-race"] = replacement + sb._quarantined_kernels["old-kernel"] = "session-race" + deleted = [] + monkeypatch.setattr( + sb, + "_delete_kernel", + lambda kernel_id: (deleted.append(kernel_id), True)[1], + ) + + assert sb._get_kernel("session-race") is replacement + assert deleted == ["old-kernel"] + assert sb._quarantined_kernels == {} + assert sb._kernels["session-race"] is replacement + + def test_collect_caps_oversize_rich_output(monkeypatch): # A huge execute_result must NOT be buffered: once it would exceed the byte # budget the bundle is dropped and the result is marked truncated. diff --git a/tests/sandbox/test_sandbox_manager.py b/tests/sandbox/test_sandbox_manager.py index 02b9abcb..60533bba 100644 --- a/tests/sandbox/test_sandbox_manager.py +++ b/tests/sandbox/test_sandbox_manager.py @@ -174,6 +174,81 @@ def test_exec_and_file_roundtrip_through_manager(backend): assert mgr.list_files("conv-1") == ["a.txt"] +def test_runtime_invalidation_drops_cached_handle_and_next_open_is_cold(): + """A backend-killed runtime must not remain reusable in the manager cache.""" + + class _InvalidatingBackend(FakeBackend): + def exec(self, session_id, code, timeout=None): + self._handles.pop(session_id, None) + return ExecResult( + status="error", + error_name="TimeoutError", + error_value="execution exceeded its deadline", + exit_code=-1, + runtime_invalidated=True, + ) + + backend = _InvalidatingBackend() + mgr = SandboxManager(backend, max_ttl=600) + old_handle = mgr.open("conv-1") + + result = mgr.exec("conv-1", "while True: pass", timeout=1) + + assert result.runtime_invalidated is True + assert not mgr.has_session("conv-1") + new_handle = mgr.open("conv-1") + assert new_handle != old_handle + assert backend.open_calls == ["conv-1", "conv-1"] + assert backend.closed_handles == [] # the backend already destroyed the invalid runtime + + +def test_old_file_operation_cannot_release_reopened_session_after_invalidation(): + """A stale operation's leave must not decrement a replacement session generation.""" + + class _BlockedReadInvalidatingBackend(FakeBackend): + def __init__(self): + super().__init__() + self.read_started = threading.Event() + self.release_read = threading.Event() + + def get_file(self, session_id, path): + self.read_started.set() + assert self.release_read.wait(timeout=5) + return b"old-generation" + + def exec(self, session_id, code, timeout=None): + self._handles.pop(session_id, None) + return ExecResult(status="error", error_name="TimeoutError", runtime_invalidated=True) + + backend = _BlockedReadInvalidatingBackend() + mgr = SandboxManager(backend, max_ttl=600) + mgr.open("conv-1") + read_result: Dict[str, bytes] = {} + reader = threading.Thread( + target=lambda: read_result.__setitem__("data", mgr.get_file("conv-1", "old.txt")) + ) + reader.start() + assert backend.read_started.wait(timeout=5) + + mgr.exec("conv-1", "while True: pass", timeout=1) + assert not mgr.has_session("conv-1") + mgr.open("conv-1") + replacement = mgr._enter("conv-1") + + backend.release_read.set() + reader.join(timeout=5) + assert not reader.is_alive() + assert read_result == {"data": b"old-generation"} + assert replacement.in_use == 1 + + # Closing is still deferred for the replacement's genuine hold. A stale + # unguarded leave would have decremented it to zero and closed immediately. + mgr.close("conv-1") + assert mgr.has_session("conv-1") + mgr._leave("conv-1", expected=replacement) + assert not mgr.has_session("conv-1") + + def test_file_ops_require_open_session(backend): mgr = SandboxManager(backend, max_ttl=600) with pytest.raises(KeyError): diff --git a/tests/test_code_executor_tool.py b/tests/test_code_executor_tool.py index 496af9f1..df5e441a 100644 --- a/tests/test_code_executor_tool.py +++ b/tests/test_code_executor_tool.py @@ -26,6 +26,7 @@ class _FakeManager: self._result = result self.closed: list = [] self.opened: list = [] + self.list_calls = 0 def open(self, session_id, ttl=None): self.opened.append((session_id, ttl)) @@ -35,6 +36,7 @@ class _FakeManager: return self._result def list_files(self, session_id): + self.list_calls += 1 return [] def close(self, session_id): @@ -302,6 +304,25 @@ def test_session_kept_alive_on_positive_ttl(monkeypatch): assert manager.closed == [] +def test_invalidated_runtime_skips_post_exec_artifact_capture(monkeypatch): + manager = _FakeManager( + ExecResult( + status="error", + error_name="TimeoutError", + error_value="execution exceeded 60s", + runtime_invalidated=True, + ) + ) + + payload = _run_with_fake_manager(monkeypatch, manager, code="while True: pass", persist=True) + + assert payload["status"] == "error" + assert "timed out" in payload["error"].lower() + # The pre-exec signature snapshot lists files once; a post-exec capture + # would issue a second list against the now-destroyed runtime. + assert manager.list_calls == 1 + + # --------------------------------------------------------------------------- # Input materialization: short-ref + uuid resolution (no live sandbox/DB) # ---------------------------------------------------------------------------