From 1fe62362810a393a1079da855ee28cb2ac6fc591 Mon Sep 17 00:00:00 2001 From: Alex Date: Wed, 24 Jun 2026 12:00:09 +0100 Subject: [PATCH] Add code_executor tool to run sandboxed code and persist artifacts Add a code_executor agent tool with a run_code action that runs agent-provided code in the per-conversation sandbox session and captures produced files as artifacts. Inputs are materialized only from artifacts the caller can access (parent-scoped); produced files are stored under the user's namespace with server-computed size and sha256, and the storage write is ordered last in the transaction so a failure cannot orphan bytes. Output is a compact payload with no raw bytes, and the produced artifact lights up the existing tool artifact rail. Execution honors a wall-clock timeout and an agent-selectable session TTL clamped by the global cap, and the action can be gated behind approval. Tool-call argument logging is redacted so code bodies are not written to logs. --- application/agents/tool_executor.py | 64 ++- application/agents/tools/code_executor.py | 476 ++++++++++++++++++++ application/agents/tools/tool_manager.py | 6 +- tests/integration/test_code_executor_e2e.py | 279 ++++++++++++ tests/test_code_executor_tool.py | 268 +++++++++++ 5 files changed, 1089 insertions(+), 4 deletions(-) create mode 100644 application/agents/tools/code_executor.py create mode 100644 tests/integration/test_code_executor_e2e.py create mode 100644 tests/test_code_executor_tool.py diff --git a/application/agents/tool_executor.py b/application/agents/tool_executor.py index 3a4b107e..7cc49ef9 100644 --- a/application/agents/tool_executor.py +++ b/application/agents/tool_executor.py @@ -33,6 +33,28 @@ def _sanitize_tool_prefix(tool_name: Optional[str]) -> str: return re.sub(r"[^a-zA-Z0-9_-]+", "_", str(tool_name or "")).strip("_") +# Longest string value rendered into a debug log line; longer values (e.g. an +# LLM-authored ``code`` body or an api_tool ``body``) are truncated so the full +# program/secret is never written to logs even at DEBUG level. +_LOG_VALUE_PREVIEW_LEN = 80 + + +def _redact_args_for_log(args: Any) -> Any: + """Truncate long string values so a code/body argument never lands in logs in full.""" + if not isinstance(args, dict): + text = str(args) + return text if len(text) <= _LOG_VALUE_PREVIEW_LEN else f"{text[:_LOG_VALUE_PREVIEW_LEN]}...(truncated)" + redacted: Dict[str, Any] = {} + for key, value in args.items(): + if isinstance(value, str) and len(value) > _LOG_VALUE_PREVIEW_LEN: + redacted[key] = f"{value[:_LOG_VALUE_PREVIEW_LEN]}...(truncated, {len(value)} chars)" + elif isinstance(value, (dict, list)): + redacted[key] = f"<{type(value).__name__} omitted>" + else: + redacted[key] = value + return redacted + + def _record_proposed( call_id: str, tool_name: str, @@ -478,6 +500,12 @@ class ToolExecutor: tool_data, action_name, arguments, ) ) + elif tool_data.get("name") == "code_executor": + # The deployment-level ``config.require_approval`` is authoritative + # over the cached action snapshot, so consult the tool directly. + require_approval = self._code_executor_requires_approval( + tool_data, action_name, arguments, + ) or require_approval if require_approval: if self.headless: @@ -552,6 +580,30 @@ class ToolExecutor: ) return True, True + def _code_executor_requires_approval( + self, tool_data: Dict, action_name: str, arguments: Dict, + ) -> bool: + """Live approval decision for a ``code_executor`` invocation. + + Honors the deployment-level ``config.require_approval`` even when the + cached action snapshot is stale. Fails closed (require approval) on any + error so a misconfigured tool never silently runs untrusted code. + """ + try: + from application.agents.tools.code_executor import CodeExecutorTool + + tool = CodeExecutorTool( + tool_config=tool_data.get("config") or {}, + user_id=self.user, + ) + requires_approval, _forced = tool.preview_decision(action_name, arguments) + return requires_approval + except Exception: + logger.exception( + "code_executor preview_decision failed; defaulting to a prompt", + ) + return True + def execute(self, tools_dict: Dict, call, llm_class_name: str): """Execute a tool call. Yields status events, returns (result, call_id).""" parser = ToolActionParser(llm_class_name, name_mapping=self._name_to_tool) @@ -748,11 +800,19 @@ class ToolExecutor: try: if tool_data["name"] == "api_tool": logger.debug( - f"Executing api: {action_name} with query_params: {query_params}, headers: {headers}, body: {body}" + "Executing api: %s with query_params: %s, headers: %s, body: %s", + action_name, + _redact_args_for_log(query_params), + _redact_args_for_log(headers), + _redact_args_for_log(body), ) result = tool.execute_action(action_name, **body) else: - logger.debug(f"Executing tool: {action_name} with args: {call_args}") + logger.debug( + "Executing tool: %s with args: %s", + action_name, + _redact_args_for_log(call_args), + ) result = tool.execute_action(action_name, **parameters) except Exception as exc: if proposed_ok: diff --git a/application/agents/tools/code_executor.py b/application/agents/tools/code_executor.py new file mode 100644 index 00000000..a815f6f2 --- /dev/null +++ b/application/agents/tools/code_executor.py @@ -0,0 +1,476 @@ +"""Code Executor tool: run sandboxed code in a semi-persistent session and capture produced files as artifacts.""" + +from __future__ import annotations + +import hashlib +import io +import logging +import mimetypes +import re +from typing import Any, Dict, List, Optional, Tuple + +from sqlalchemy import text + +from application.agents.tools.base import Tool +from application.core.settings import settings +from application.sandbox.base import ExecResult +from application.sandbox.sandbox_creator import SandboxCreator +from application.storage.db.repositories.artifacts import ArtifactsRepository +from application.storage.db.session import db_readonly, db_session +from application.storage.storage_creator import StorageCreator +from application.utils import safe_filename + +logger = logging.getLogger(__name__) + +# Maximum bytes of stdout/stderr returned to the LLM. The raw stream is never +# forwarded; only this tail keeps binary/runaway output out of the context. +_OUTPUT_TAIL_BYTES = 4000 + +# Session ids become a kernel workspace path component; the gateway only accepts +# [A-Za-z0-9_-]+, so any disallowed character is stripped before binding. +_SESSION_ID_RE = re.compile(r"[^A-Za-z0-9_-]+") + +_DEFAULT_KIND = "file" + +# Coarse mime -> artifact kind mapping for the UI rail; defaults to "file". +_KIND_BY_MIME_PREFIX: Dict[str, str] = { + "image/": "image", + "text/html": "html", + "text/csv": "data", + "application/json": "data", + "application/vnd.openxmlformats-officedocument.presentationml": "presentation", + "application/vnd.openxmlformats-officedocument.spreadsheetml": "spreadsheet", + "application/vnd.ms-excel": "spreadsheet", + "application/vnd.openxmlformats-officedocument.wordprocessingml": "document", + "application/msword": "document", + "application/pdf": "document", +} + + +def _tail(stream: Optional[str]) -> str: + """Return the trailing slice of ``stream`` bounded by ``_OUTPUT_TAIL_BYTES``.""" + if not stream: + return "" + if len(stream) <= _OUTPUT_TAIL_BYTES: + return stream + return stream[-_OUTPUT_TAIL_BYTES:] + + +def _infer_mime(filename: str) -> str: + """Infer a mime type from a filename, falling back to a generic binary type.""" + mime, _ = mimetypes.guess_type(filename) + return mime or "application/octet-stream" + + +def _kind_for_mime(mime: str) -> str: + """Map a mime type to a coarse artifact ``kind`` for the artifact rail.""" + for prefix, kind in _KIND_BY_MIME_PREFIX.items(): + if mime.startswith(prefix): + return kind + return _DEFAULT_KIND + + +class CodeExecutorTool(Tool): + """Code Executor + Run Python (or other) code in a sandboxed, semi-persistent session and capture produced files as artifacts. + """ + + def __init__(self, tool_config: Optional[Dict[str, Any]] = None, user_id: Optional[str] = None) -> None: + """Bind the tool to the invoker and its conversation/run-scoped sandbox session.""" + self.config: Dict[str, Any] = tool_config or {} + self.user_id: Optional[str] = user_id + self.tool_id: Optional[str] = self.config.get("tool_id") + self.conversation_id: Optional[str] = self.config.get("conversation_id") + self.workflow_run_id: Optional[str] = self.config.get("workflow_run_id") + # Static, deployment-level approval gate (mirrors the action metadata flag). + self._require_approval: bool = bool(self.config.get("require_approval", False)) + self._last_artifact_id: Optional[str] = None + + # ------------------------------------------------------------------ + # Tool ABC + # ------------------------------------------------------------------ + def get_actions_metadata(self) -> List[Dict[str, Any]]: + """Return JSON metadata describing the ``run_code`` action for tool schemas.""" + return [ + { + "name": "run_code", + "description": ( + "Execute code in a sandboxed, stateful session bound to this conversation. " + "Files written by the code are captured as downloadable artifacts; only a " + "compact summary (output tail + artifact references) is returned, never raw bytes." + ), + "active": True, + "require_approval": self._require_approval, + "parameters": { + "type": "object", + "properties": { + "code": { + "type": "string", + "description": "Source code to execute in the session.", + }, + "language": { + "type": "string", + "description": "Programming language (default: python).", + }, + "libraries": { + "type": "array", + "items": {"type": "string"}, + "description": "Packages to ensure are installed before running.", + }, + "inputs": { + "type": "array", + "items": {"type": "string"}, + "description": "Artifact ids (from this conversation/run) to materialize into the workspace.", + }, + "timeout": { + "type": "integer", + "description": "Wall-clock timeout in seconds for this execution.", + }, + "ttl": { + "type": "integer", + "description": "Keep-alive lifetime (seconds) for the session; clamped by SANDBOX_MAX_TTL.", + }, + "persist": { + "type": "boolean", + "description": ( + "Keep the session warm after the call (state survives the next run). " + "The session is kept alive when this is true or a positive ttl is given " + "(clamped by SANDBOX_MAX_TTL); otherwise it is closed after the run." + ), + }, + "capture_artifacts": { + "type": "boolean", + "description": "Capture newly written workspace files as artifacts (default: true).", + }, + }, + "required": ["code"], + }, + } + ] + + def get_config_requirements(self) -> Dict[str, Any]: + """Return configuration requirements (approval gate + backend selection).""" + return { + "require_approval": { + "type": "boolean", + "label": "Require approval", + "description": "Pause for human approval before each code execution.", + "required": False, + }, + "sandbox_backend": { + "type": "string", + "label": "Sandbox backend", + "description": "Code-execution backend (defaults to the SANDBOX_BACKEND setting).", + "required": False, + }, + } + + def get_artifact_id(self, action_name: str, **kwargs: Any) -> Optional[str]: + """Return the primary produced artifact id so the UI artifact rail lights up.""" + return self._last_artifact_id + + def preview_decision(self, action_name: str, params: dict) -> Tuple[bool, bool]: + """Return ``(requires_approval, denylist_forced)`` for the approval gate; never denylist-forced here.""" + if action_name != "run_code": + return True, False + return self._require_approval, False + + # ------------------------------------------------------------------ + # Execution + # ------------------------------------------------------------------ + def execute_action(self, action_name: str, **kwargs: Any) -> Dict[str, Any]: + """Dispatch a tool action; only ``run_code`` is supported.""" + if action_name != "run_code": + return {"status": "error", "error": f"unknown action: {action_name}"} + self._last_artifact_id = None + return self._run_code(**kwargs) + + def _run_code(self, **kwargs: Any) -> Dict[str, Any]: + """Bind a session, materialize inputs, execute, and capture produced artifacts.""" + if not self.user_id: + return {"status": "error", "error": "code_executor requires a valid user_id."} + + session_id = self._resolve_session_id() + if session_id is None: + return {"status": "error", "error": "code_executor requires a conversation_id or workflow_run_id."} + + code = kwargs.get("code") + if not isinstance(code, str) or not code.strip(): + return {"status": "error", "error": "code is required."} + + capture_artifacts = kwargs.get("capture_artifacts", True) + ttl = self._coerce_int(kwargs.get("ttl")) + timeout = self._resolve_timeout(kwargs.get("timeout")) + inputs = kwargs.get("inputs") or [] + + manager = SandboxCreator.get_manager() + try: + manager.open(session_id, ttl=ttl) + except Exception as exc: + logger.exception("code_executor: failed to open sandbox session") + return {"status": "error", "error": f"sandbox unavailable: {type(exc).__name__}: {exc}"} + + try: + materialized = self._materialize_inputs(manager, session_id, inputs) + if materialized.get("error"): + return {"status": "error", "error": materialized["error"]} + + pre_signatures: Dict[str, Tuple[int, Optional[str]]] = {} + if capture_artifacts: + pre_signatures = self._snapshot_signatures(manager, session_id) + + try: + result = manager.exec(session_id, code, timeout=timeout) + except Exception as exc: + 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. + artifacts: List[Dict[str, Any]] = [] + if capture_artifacts: + try: + artifacts = self._capture_artifacts(manager, session_id, pre_signatures) + except Exception: + logger.exception("code_executor: artifact capture failed") + + return self._shape_payload(result, artifacts, materialized.get("loaded", [])) + finally: + if not self._keep_alive(kwargs.get("persist"), ttl): + try: + manager.close(session_id) + except Exception: + logger.exception("code_executor: session close failed") + + # ------------------------------------------------------------------ + # Inputs / outputs + # ------------------------------------------------------------------ + def _materialize_inputs(self, manager: Any, session_id: str, inputs: List[Any]) -> Dict[str, Any]: + """Fetch parent-scoped input artifacts and copy their current-version bytes into the workspace.""" + loaded: List[str] = [] + if not inputs: + return {"loaded": loaded} + storage = StorageCreator.get_storage() + for raw_id in inputs: + artifact_id = str(raw_id).strip() + if not artifact_id: + continue + try: + with db_readonly() as conn: + repo = ArtifactsRepository(conn) + artifact = repo.get_artifact_in_parent( + artifact_id, + conversation_id=self.conversation_id, + workflow_run_id=self.workflow_run_id, + ) + if artifact is None: + return {"error": f"input artifact {artifact_id} not found in this conversation/run."} + version = repo.get_version(artifact_id, artifact["current_version"]) + except Exception: + logger.exception("code_executor: failed to load input artifact") + return {"error": f"failed to load input artifact {artifact_id}."} + + if not version or not version.get("storage_path"): + return {"error": f"input artifact {artifact_id} has no stored content."} + + filename = safe_filename(version.get("filename") or artifact_id) + try: + file_obj = storage.get_file(version["storage_path"]) + data = file_obj.read() + except Exception: + logger.exception("code_executor: failed to read input artifact bytes") + return {"error": f"failed to read input artifact {artifact_id}."} + try: + manager.put_file(session_id, f"inputs/{filename}", data) + except Exception: + logger.exception("code_executor: put_file failed for input artifact") + return {"error": f"failed to stage input artifact {artifact_id} into the workspace."} + loaded.append(f"inputs/{filename}") + return {"loaded": loaded} + + # Cap the per-run capture work so a workspace full of pre-existing files + # can't turn one exec into an unbounded read+persist sweep. + _MAX_CAPTURED_FILES = 64 + + def _snapshot_signatures(self, manager: Any, session_id: str) -> Dict[str, Tuple[int, Optional[str]]]: + """Map each non-input workspace file to a (size, sha256) signature for change detection.""" + signatures: Dict[str, Tuple[int, Optional[str]]] = {} + try: + files = manager.list_files(session_id) + except Exception: + logger.exception("code_executor: pre-exec listing failed") + return signatures + for rel_path in files: + if rel_path.startswith("inputs/"): + continue + try: + data = manager.get_file(session_id, rel_path) + except Exception: + logger.exception("code_executor: pre-exec signature read failed") + continue + signatures[rel_path] = (len(data), hashlib.sha256(data).hexdigest()) + return signatures + + def _capture_artifacts( + self, manager: Any, session_id: str, pre_signatures: Dict[str, Tuple[int, Optional[str]]] + ) -> List[Dict[str, Any]]: + """Persist each non-input workspace file that is new or whose content changed.""" + try: + post_files = set(manager.list_files(session_id)) + except Exception: + logger.exception("code_executor: post-exec listing failed") + return [] + + candidates = sorted(f for f in post_files if not f.startswith("inputs/")) + storage = StorageCreator.get_storage() + captured: List[Dict[str, Any]] = [] + for rel_path in candidates: + if len(captured) >= self._MAX_CAPTURED_FILES: + logger.warning("code_executor: capture cap reached; remaining files skipped") + break + try: + data = manager.get_file(session_id, rel_path) + except Exception: + logger.exception("code_executor: get_file failed during capture") + continue + # A pre-existing file is only captured when its content changed; an + # unchanged file is skipped so re-runs don't re-persist stale inputs. + signature = (len(data), hashlib.sha256(data).hexdigest()) + if pre_signatures.get(rel_path) == signature: + continue + ref = self._persist_artifact(storage, rel_path, data, session_id) + if ref is not None: + captured.append(ref) + if captured: + self._last_artifact_id = captured[0]["artifact_id"] + return captured + + def _persist_artifact( + self, storage: Any, rel_path: str, data: bytes, session_id: str + ) -> Optional[Dict[str, Any]]: + """Store ``data`` and create an artifact row; size/sha256 are computed server-side. + + The storage write is the last statement before commit, so a failed + write rolls the row back (bytes are never orphaned). The only remaining + window is a commit that fails after a successful write; that key is + deleted best-effort, but a crash between save and commit can still leak. + """ + # The sandbox filename is display-only; the storage key is derived from + # server-controlled values so a hostile name can't redirect the write. + display_name = rel_path.rsplit("/", 1)[-1] + filename = safe_filename(display_name) + size = len(data) + sha256 = hashlib.sha256(data).hexdigest() + mime_type = _infer_mime(filename) + kind = _kind_for_mime(mime_type) + saved_key: Optional[str] = None + try: + with db_session() as conn: + repo = ArtifactsRepository(conn) + artifact = repo.create_artifact( + self.user_id, + kind, + conversation_id=self.conversation_id, + workflow_run_id=self.workflow_run_id, + title=display_name, + mime_type=mime_type, + filename=filename, + storage_path=None, + size=size, + sha256=sha256, + produced_by={ + "tool": "code_executor", + "action": "run_code", + "session_id": session_id, + }, + ) + artifact_id = str(artifact["id"]) + # ``inputs/{user}/artifacts/...`` is the project storage-namespace + # convention (matches attachments + spec §4); the ``inputs/`` prefix + # is the user's namespace root, not an "input file" marker. + storage_path = f"inputs/{self.user_id}/artifacts/{artifact_id}/v1/{filename}" + # Set the server-derived key on version 1, then write the bytes as + # the LAST statement so a save failure rolls the whole row back. + conn.execute( + text( + "UPDATE artifact_versions SET storage_path = :p " + "WHERE artifact_id = CAST(:aid AS uuid) AND version = 1" + ), + {"p": storage_path, "aid": artifact_id}, + ) + storage.save_file(io.BytesIO(data), storage_path) + saved_key = storage_path + except Exception: + logger.exception("code_executor: failed to persist artifact") + # The bytes landed but the commit failed: drop the now-orphaned key. + if saved_key is not None: + try: + storage.delete_file(saved_key) + except Exception: + logger.exception("code_executor: orphaned-key cleanup failed for %s", saved_key) + return None + return { + "artifact_id": artifact_id, + "version": 1, + "filename": filename, + "mime_type": mime_type, + "size": size, + } + + def _shape_payload( + self, result: ExecResult, artifacts: List[Dict[str, Any]], inputs_loaded: List[str] + ) -> Dict[str, Any]: + """Build the compact LLM-facing payload; raw bytes never appear here.""" + status = "ok" if result.ok else "error" + payload: Dict[str, Any] = { + "status": status, + "stdout_tail": _tail(result.stdout), + "artifacts": artifacts, + } + stderr_tail = _tail(result.stderr) + if stderr_tail: + payload["stderr_tail"] = stderr_tail + if not result.ok: + payload["error"] = ( + f"{result.error_name}: {result.error_value}" + if result.error_name + else (result.error_value or "execution error") + ) + if inputs_loaded: + payload["inputs_loaded"] = inputs_loaded + return payload + + # ------------------------------------------------------------------ + # Helpers + # ------------------------------------------------------------------ + def _resolve_session_id(self) -> Optional[str]: + """Derive a sandbox session id from the bound conversation/run; sanitize to the gateway charset.""" + raw = self.conversation_id or self.workflow_run_id + if not raw: + return None + sanitized = _SESSION_ID_RE.sub("-", str(raw)) + return sanitized or None + + @staticmethod + def _coerce_int(value: Any) -> Optional[int]: + """Coerce a value to a positive int, or None when absent/invalid.""" + if value is None: + return None + try: + parsed = int(value) + except (TypeError, ValueError): + return None + return parsed if parsed > 0 else None + + def _resolve_timeout(self, requested: Any) -> float: + """Return the stricter of the requested timeout and the sandbox's default cap.""" + cap = float(getattr(settings, "SANDBOX_EXEC_TIMEOUT", 60)) + parsed = self._coerce_int(requested) + if parsed is None: + return cap + return float(min(parsed, cap)) + + @staticmethod + def _keep_alive(persist: Any, ttl: Optional[int]) -> bool: + """True when the agent asked to keep the session warm after the call.""" + return bool(persist) or (ttl is not None and ttl > 0) diff --git a/application/agents/tools/tool_manager.py b/application/agents/tools/tool_manager.py index f969e87f..a0dc35e6 100644 --- a/application/agents/tools/tool_manager.py +++ b/application/agents/tools/tool_manager.py @@ -29,7 +29,8 @@ class ToolManager: for member_name, obj in inspect.getmembers(module, inspect.isclass): if issubclass(obj, Tool) and obj is not Tool: if ( - tool_name in {"mcp_tool", "notes", "memory", "todo_list", "scheduler", "remote_device"} + tool_name + in {"mcp_tool", "notes", "memory", "todo_list", "scheduler", "remote_device", "code_executor"} and user_id ): return obj(tool_config, user_id) @@ -40,7 +41,8 @@ class ToolManager: if tool_name not in self.tools: raise ValueError(f"Tool '{tool_name}' not loaded") if ( - tool_name in {"mcp_tool", "memory", "todo_list", "notes", "scheduler", "remote_device"} + tool_name + in {"mcp_tool", "memory", "todo_list", "notes", "scheduler", "remote_device", "code_executor"} and user_id ): tool_config = self.config.get(tool_name, {}) diff --git a/tests/integration/test_code_executor_e2e.py b/tests/integration/test_code_executor_e2e.py new file mode 100644 index 00000000..293b6d85 --- /dev/null +++ b/tests/integration/test_code_executor_e2e.py @@ -0,0 +1,279 @@ +"""End-to-end CodeExecutorTool: live Jupyter gateway + ephemeral Postgres + local storage. + +Launches a real ``jupyter kernelgateway`` (no Docker — the credential helper +hangs on dev machines), wires the tool to the ephemeral pytest-postgresql DB +and a temp-dir ``LocalStorage``, and drives ``run_code`` through the full path: +code writes a file -> artifact row + bytes persisted with server-side size/ +sha256/mime -> compact payload + ``get_artifact_id``. Also covers an input +artifact round-trip and the timeout error path (no hang). + +Skips gracefully when the gateway binary or websocket-client is unavailable. +""" + +from __future__ import annotations + +import hashlib +import io +import shutil +import socket +import subprocess +import time +import uuid + +import pytest +from sqlalchemy import text + +requests = pytest.importorskip("requests") +pytest.importorskip("websocket") # websocket-client + +from application.agents.tools.code_executor import CodeExecutorTool # noqa: E402 +from application.sandbox.jupyter_gateway import JupyterKernelGatewaySandbox # noqa: E402 +from application.sandbox.manager import SandboxManager # noqa: E402 +from application.sandbox.sandbox_creator import SandboxCreator # noqa: E402 +from application.storage.db.repositories.artifacts import ArtifactsRepository # noqa: E402 +from application.storage.local import LocalStorage # noqa: E402 +from application.storage.storage_creator import StorageCreator # noqa: E402 + +_GATEWAY_BIN = shutil.which("jupyter-kernelgateway") or shutil.which("jupyter") + +pytestmark = [ + pytest.mark.integration, + pytest.mark.skipif( + _GATEWAY_BIN is None, + reason="jupyter kernel gateway not installed (pip install jupyter-kernel-gateway)", + ), +] + + +def _free_port() -> int: + with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as sock: + sock.bind(("127.0.0.1", 0)) + return sock.getsockname()[1] + + +def _gateway_cmd(port: int) -> list: + if _GATEWAY_BIN.endswith("jupyter-kernelgateway"): + base = [_GATEWAY_BIN] + else: + base = [_GATEWAY_BIN, "kernelgateway"] + return base + [ + "--KernelGatewayApp.ip=127.0.0.1", + f"--KernelGatewayApp.port={port}", + "--ZMQChannelsWebsocketConnection.limit_rate=False", + ] + + +@pytest.fixture(scope="module") +def gateway_url(): + port = _free_port() + proc = subprocess.Popen(_gateway_cmd(port), stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL) + url = f"http://127.0.0.1:{port}" + deadline = time.time() + 30 + ready = False + try: + while time.time() < deadline: + if proc.poll() is not None: + pytest.skip("jupyter kernelgateway process exited during startup") + try: + if requests.get(f"{url}/api", timeout=1).status_code == 200: + ready = True + break + except requests.RequestException: + time.sleep(0.3) + if not ready: + pytest.skip("jupyter kernelgateway did not become ready in time") + yield url + finally: + proc.terminate() + try: + proc.wait(timeout=10) + except subprocess.TimeoutExpired: + proc.kill() + + +@pytest.fixture() +def wired_tool(gateway_url, pg_engine, tmp_path, monkeypatch): + """A CodeExecutorTool wired to the live gateway, ephemeral PG, and a temp local store.""" + # Sandbox: a fresh manager over the live gateway, installed as the singleton. + backend = JupyterKernelGatewaySandbox(gateway_url=gateway_url, default_timeout=30.0) + SandboxCreator._instance = SandboxManager(backend=backend, max_ttl=1200.0) + + # Storage: temp-dir local storage as the singleton. + storage = LocalStorage(base_dir=str(tmp_path)) + monkeypatch.setattr(StorageCreator, "_instance", storage, raising=False) + + # DB: route db_session()/db_readonly() at the ephemeral PG engine. + monkeypatch.setattr("application.storage.db.session.get_engine", lambda: pg_engine) + + conversation_id = str(uuid.uuid4()) + tool = CodeExecutorTool( + tool_config={"conversation_id": conversation_id, "tool_id": str(uuid.uuid4())}, + user_id="user-e2e", + ) + try: + yield tool, conversation_id, pg_engine, storage + finally: + SandboxCreator.reset() + + +def test_run_code_persists_produced_artifact(wired_tool): + tool, conversation_id, pg_engine, storage = wired_tool + + code = ( + "with open('report.txt', 'w') as f:\n" + " f.write('hello artifact')\n" + "print('wrote report')\n" + ) + payload = tool.execute_action("run_code", code=code, ttl=120) + + assert payload["status"] == "ok", payload + assert "wrote report" in payload["stdout_tail"] + assert len(payload["artifacts"]) == 1 + art = payload["artifacts"][0] + assert art["filename"] == "report.txt" + assert art["mime_type"] == "text/plain" + assert art["size"] == len(b"hello artifact") + assert art["version"] == 1 + + # get_artifact_id points at the produced artifact (UI rail). + assert tool.get_artifact_id("run_code") == art["artifact_id"] + + # DB row exists, parent-scoped, with server-computed size + sha256. + with pg_engine.connect() as conn: + repo = ArtifactsRepository(conn) + artifact = repo.get_artifact_in_parent(art["artifact_id"], conversation_id=conversation_id) + assert artifact is not None + version = repo.get_version(art["artifact_id"], 1) + assert version["size"] == len(b"hello artifact") + assert version["sha256"] == hashlib.sha256(b"hello artifact").hexdigest() + assert version["mime_type"] == "text/plain" + produced = version["produced_by"] + assert produced["tool"] == "code_executor" and produced["action"] == "run_code" + + # Bytes are actually in storage at the server-derived key, and match. + storage_path = version["storage_path"] + assert storage_path.startswith("inputs/user-e2e/artifacts/") + assert art["artifact_id"] in storage_path + assert storage.get_file(storage_path).read() == b"hello artifact" + + +def test_run_code_input_artifact_roundtrip(wired_tool): + tool, conversation_id, pg_engine, storage = wired_tool + + # Seed an input artifact (row + bytes) scoped to this conversation. + seed_bytes = b"seed-value-123" + artifact_id = _seed_artifact(pg_engine, storage, conversation_id, "seed.txt", seed_bytes) + + code = ( + "data = open('inputs/seed.txt', 'rb').read()\n" + "open('echo.txt', 'wb').write(data + b'-processed')\n" + "print('read', len(data), 'bytes')\n" + ) + payload = tool.execute_action("run_code", code=code, inputs=[artifact_id], ttl=120) + + assert payload["status"] == "ok", payload + assert "read 14 bytes" in payload["stdout_tail"] + assert payload["inputs_loaded"] == ["inputs/seed.txt"] + assert len(payload["artifacts"]) == 1 + out = payload["artifacts"][0] + assert out["filename"] == "echo.txt" + assert out["size"] == len(seed_bytes + b"-processed") + + with pg_engine.connect() as conn: + version = ArtifactsRepository(conn).get_version(out["artifact_id"], 1) + assert storage.get_file(version["storage_path"]).read() == seed_bytes + b"-processed" + + +def test_run_code_input_artifact_cross_tenant_blocked(wired_tool): + tool, conversation_id, pg_engine, storage = wired_tool + + # An artifact that belongs to a DIFFERENT conversation must not be reachable. + other_conversation = str(uuid.uuid4()) + foreign_id = _seed_artifact(pg_engine, storage, other_conversation, "secret.txt", b"top-secret") + + payload = tool.execute_action( + "run_code", code="open('x.txt','w').write('x')", inputs=[foreign_id], ttl=60 + ) + assert payload["status"] == "error" + assert "not found in this conversation/run" in payload["error"] + + +def test_run_code_captures_overwritten_file(wired_tool): + tool, conversation_id, pg_engine, storage = wired_tool + + first = tool.execute_action( + "run_code", code="open('out.txt','w').write('first-content')", ttl=120 + ) + assert first["status"] == "ok", first + assert len(first["artifacts"]) == 1 + first_id = first["artifacts"][0]["artifact_id"] + + # The same persisted session overwrites out.txt with new content; the diff + # is content-aware, so the new bytes must be captured as a fresh artifact + # rather than dropped as an "already-seen path". + second = tool.execute_action( + "run_code", code="open('out.txt','w').write('second-content-longer')", ttl=120 + ) + assert second["status"] == "ok", second + assert len(second["artifacts"]) == 1 + second_art = second["artifacts"][0] + assert second_art["artifact_id"] != first_id + assert second_art["filename"] == "out.txt" + assert second_art["size"] == len(b"second-content-longer") + + with pg_engine.connect() as conn: + version = ArtifactsRepository(conn).get_version(second_art["artifact_id"], 1) + assert storage.get_file(version["storage_path"]).read() == b"second-content-longer" + + +def test_run_code_skips_unchanged_file_on_rerun(wired_tool): + tool, _conversation_id, _pg_engine, _storage = wired_tool + + first = tool.execute_action( + "run_code", code="open('keep.txt','w').write('same')", ttl=120 + ) + assert len(first["artifacts"]) == 1 + + # A re-run that leaves keep.txt untouched must not re-persist it. + second = tool.execute_action("run_code", code="print('noop')", ttl=120) + assert second["status"] == "ok", second + assert second["artifacts"] == [] + + +def test_run_code_timeout_returns_clean_error(wired_tool): + tool, _conversation_id, _pg_engine, _storage = wired_tool + + payload = tool.execute_action("run_code", code="import time; time.sleep(10)", timeout=1) + assert payload["status"] == "error" + # A clean structured error, not a hang: the sandbox interrupted the kernel. + assert "TimeoutError" in payload["error"] + assert payload["artifacts"] == [] + + +def _seed_artifact(pg_engine, storage, conversation_id, filename, data) -> str: + """Create an artifact row + version + stored bytes, returning its id.""" + sha256 = hashlib.sha256(data).hexdigest() + with pg_engine.begin() as conn: + repo = ArtifactsRepository(conn) + artifact = repo.create_artifact( + "user-e2e", + "file", + conversation_id=conversation_id, + title=filename, + mime_type="text/plain", + filename=filename, + storage_path=None, + size=len(data), + sha256=sha256, + ) + artifact_id = str(artifact["id"]) + storage_path = f"inputs/user-e2e/artifacts/{artifact_id}/v1/{filename}" + storage.save_file(io.BytesIO(data), storage_path) + conn.execute( + text( + "UPDATE artifact_versions SET storage_path = :p " + "WHERE artifact_id = CAST(:aid AS uuid) AND version = 1" + ), + {"p": storage_path, "aid": artifact_id}, + ) + return artifact_id diff --git a/tests/test_code_executor_tool.py b/tests/test_code_executor_tool.py new file mode 100644 index 00000000..624bce59 --- /dev/null +++ b/tests/test_code_executor_tool.py @@ -0,0 +1,268 @@ +"""Unit tests for CodeExecutorTool: payload shaping, mime/kind inference, and allowlist wiring. + +These tests exercise the pure logic (no sandbox, no DB, no storage) so they +run in the fast unit suite. The end-to-end persistence path is covered by +``tests/integration/test_code_executor_e2e.py`` against a live gateway + PG. +""" + +from __future__ import annotations + +from application.agents.tools.code_executor import ( + CodeExecutorTool, + _infer_mime, + _kind_for_mime, + _tail, + _OUTPUT_TAIL_BYTES, +) +from application.sandbox.base import ExecResult + + +class _FakeManager: + """In-memory sandbox stand-in recording open/close and serving a fixed exec result.""" + + def __init__(self, result: ExecResult) -> None: + self._result = result + self.closed: list = [] + self.opened: list = [] + + def open(self, session_id, ttl=None): + self.opened.append((session_id, ttl)) + return session_id + + def exec(self, session_id, code, timeout=None): + return self._result + + def list_files(self, session_id): + return [] + + def close(self, session_id): + self.closed.append(session_id) + + +def _tool() -> CodeExecutorTool: + return CodeExecutorTool( + tool_config={"conversation_id": "conv-1", "tool_id": "t1"}, + user_id="user-1", + ) + + +# --------------------------------------------------------------------------- +# Output truncation +# --------------------------------------------------------------------------- +def test_tail_returns_short_text_unchanged(): + assert _tail("hello") == "hello" + assert _tail("") == "" + assert _tail(None) == "" + + +def test_tail_truncates_to_trailing_window(): + head = "HEAD_MARKER" + "X" * (_OUTPUT_TAIL_BYTES + 500) + long_text = head + "TAIL_MARKER" + out = _tail(long_text) + assert len(out) == _OUTPUT_TAIL_BYTES + # The tail keeps the END of the stream (where errors/results land), not the head. + assert out.endswith("TAIL_MARKER") + assert "HEAD_MARKER" not in out + + +# --------------------------------------------------------------------------- +# Mime / kind inference +# --------------------------------------------------------------------------- +def test_infer_mime_known_and_unknown(): + assert _infer_mime("deck.pptx").endswith("presentationml.presentation") + assert _infer_mime("report.pdf") == "application/pdf" + assert _infer_mime("out.txt") == "text/plain" + assert _infer_mime("data.csv") == "text/csv" + assert _infer_mime("blob.weirdext") == "application/octet-stream" + + +def test_kind_for_mime_maps_office_and_media(): + assert _kind_for_mime("image/png") == "image" + assert _kind_for_mime("application/pdf") == "document" + assert _kind_for_mime(_infer_mime("deck.pptx")) == "presentation" + assert _kind_for_mime(_infer_mime("sheet.xlsx")) == "spreadsheet" + assert _kind_for_mime("text/html") == "html" + assert _kind_for_mime("application/octet-stream") == "file" + + +# --------------------------------------------------------------------------- +# Payload shaping +# --------------------------------------------------------------------------- +def test_shape_payload_ok_with_artifacts(): + tool = _tool() + artifacts = [{"artifact_id": "a1", "version": 1, "filename": "out.txt", + "mime_type": "text/plain", "size": 9}] + result = ExecResult(status="ok", stdout="done\n", stderr="") + payload = tool._shape_payload(result, artifacts, inputs_loaded=["inputs/seed.txt"]) + assert payload["status"] == "ok" + assert payload["stdout_tail"] == "done\n" + assert payload["artifacts"] == artifacts + assert payload["inputs_loaded"] == ["inputs/seed.txt"] + assert "error" not in payload + # No raw bytes ever leak into the payload. + assert "bytes" not in payload + + +def test_shape_payload_error_carries_clean_message_no_hang(): + tool = _tool() + result = ExecResult( + status="error", error_name="TimeoutError", + error_value="execution exceeded 1.0s", exit_code=-1, + ) + payload = tool._shape_payload(result, artifacts=[], inputs_loaded=[]) + assert payload["status"] == "error" + assert payload["error"] == "TimeoutError: execution exceeded 1.0s" + assert payload["artifacts"] == [] + + +def test_shape_payload_includes_stderr_tail_only_when_present(): + tool = _tool() + with_err = tool._shape_payload( + ExecResult(status="ok", stdout="ok", stderr="warn"), [], [] + ) + assert with_err["stderr_tail"] == "warn" + no_err = tool._shape_payload(ExecResult(status="ok", stdout="ok", stderr=""), [], []) + assert "stderr_tail" not in no_err + + +# --------------------------------------------------------------------------- +# Session id resolution & timeout / ttl coercion +# --------------------------------------------------------------------------- +def test_resolve_session_id_prefers_conversation_then_run(): + conv = CodeExecutorTool({"conversation_id": "conv-1"}, user_id="u") + assert conv._resolve_session_id() == "conv-1" + run = CodeExecutorTool({"workflow_run_id": "run-9"}, user_id="u") + assert run._resolve_session_id() == "run-9" + none = CodeExecutorTool({}, user_id="u") + assert none._resolve_session_id() is None + + +def test_resolve_session_id_sanitizes_disallowed_chars(): + tool = CodeExecutorTool({"conversation_id": "../evil id;rm"}, user_id="u") + sid = tool._resolve_session_id() + # Only [A-Za-z0-9_-] survives; path-traversal / shell chars are collapsed. + import re + + assert re.fullmatch(r"[A-Za-z0-9_-]+", sid) + + +def test_resolve_timeout_picks_the_stricter_cap(monkeypatch): + from application.core import settings as settings_module + + monkeypatch.setattr(settings_module.settings, "SANDBOX_EXEC_TIMEOUT", 60, raising=False) + tool = _tool() + assert tool._resolve_timeout(10) == 10.0 # under the cap + assert tool._resolve_timeout(999) == 60.0 # clamped to cap + assert tool._resolve_timeout(None) == 60.0 # default + assert tool._resolve_timeout("bad") == 60.0 # invalid -> default + assert tool._resolve_timeout(-5) == 60.0 # non-positive -> default + + +def test_coerce_int_and_keep_alive(): + assert CodeExecutorTool._coerce_int("3") == 3 + assert CodeExecutorTool._coerce_int(0) is None + assert CodeExecutorTool._coerce_int(None) is None + assert CodeExecutorTool._keep_alive(True, None) is True + assert CodeExecutorTool._keep_alive(False, 30) is True + assert CodeExecutorTool._keep_alive(False, None) is False + + +# --------------------------------------------------------------------------- +# Action metadata / approval surface +# --------------------------------------------------------------------------- +def test_run_code_metadata_reflects_require_approval(): + gated = CodeExecutorTool({"require_approval": True}, user_id="u") + meta = gated.get_actions_metadata()[0] + assert meta["name"] == "run_code" + assert meta["require_approval"] is True + assert "code" in meta["parameters"]["required"] + assert gated.preview_decision("run_code", {}) == (True, False) + + ungated = CodeExecutorTool({}, user_id="u") + assert ungated.get_actions_metadata()[0]["require_approval"] is False + assert ungated.preview_decision("run_code", {}) == (False, False) + # An unknown action always requires approval (fail closed). + assert ungated.preview_decision("other", {}) == (True, False) + + +def test_config_requirements_expose_approval_and_backend(): + reqs = CodeExecutorTool({}, user_id="u").get_config_requirements() + assert "require_approval" in reqs + assert "sandbox_backend" in reqs + + +def test_execute_action_rejects_unknown_action_and_missing_code(): + tool = _tool() + assert tool.execute_action("nope")["status"] == "error" + missing = tool.execute_action("run_code", code=" ") + assert missing["status"] == "error" + assert "code is required" in missing["error"] + + +def test_execute_action_requires_user_and_parent(): + no_user = CodeExecutorTool({"conversation_id": "c"}, user_id=None) + out = no_user.execute_action("run_code", code="print(1)") + assert out["status"] == "error" and "user_id" in out["error"] + + no_parent = CodeExecutorTool({}, user_id="u") + out2 = no_parent.execute_action("run_code", code="print(1)") + assert out2["status"] == "error" and "conversation_id" in out2["error"] + + +# --------------------------------------------------------------------------- +# Allowlist wiring +# --------------------------------------------------------------------------- +def test_tool_manager_injects_user_and_conversation(): + """code_executor must be in the per-user allowlist so it receives user_id/conversation_id.""" + # Importing the app first resolves the mcp_tool<->api.user import cycle that + # ToolManager's eager tool discovery would otherwise trip in a bare process. + import application.app # noqa: F401 + from application.agents.tools.tool_manager import ToolManager + + tm = ToolManager(config={}) + tool = tm.load_tool( + "code_executor", + {"conversation_id": "conv-xyz", "tool_id": "tool-abc", "require_approval": True}, + user_id="user-42", + ) + assert isinstance(tool, CodeExecutorTool) + assert tool.user_id == "user-42" + assert tool.conversation_id == "conv-xyz" + assert tool.tool_id == "tool-abc" + assert tool._require_approval is True + + +# --------------------------------------------------------------------------- +# Keep-alive vs. close behavior +# --------------------------------------------------------------------------- +def _run_with_fake_manager(monkeypatch, manager, **run_kwargs): + from application.agents.tools import code_executor as ce + + monkeypatch.setattr(ce.SandboxCreator, "get_manager", lambda: manager) + return _tool().execute_action("run_code", **run_kwargs) + + +def test_session_closed_when_not_kept_alive(monkeypatch): + manager = _FakeManager(ExecResult(status="ok", stdout="ok")) + payload = _run_with_fake_manager( + monkeypatch, manager, code="print(1)", capture_artifacts=False + ) + assert payload["status"] == "ok" + # Neither persist nor a positive ttl -> the warm session is torn down. + assert manager.closed == ["conv-1"] + + +def test_session_kept_alive_on_persist(monkeypatch): + manager = _FakeManager(ExecResult(status="ok", stdout="ok")) + _run_with_fake_manager( + monkeypatch, manager, code="print(1)", persist=True, capture_artifacts=False + ) + assert manager.closed == [] + + +def test_session_kept_alive_on_positive_ttl(monkeypatch): + manager = _FakeManager(ExecResult(status="ok", stdout="ok")) + _run_with_fake_manager( + monkeypatch, manager, code="print(1)", ttl=30, capture_artifacts=False + ) + assert manager.closed == []