"""Sandbox execution for marimo and manim code.""" import ast import json import subprocess import sys import tempfile from dataclasses import dataclass, field from typing import Any from pathlib import Path @dataclass class SandboxResult: """Structured lint/build result returned to the environment.""" fmt: str parses: bool check_passed: bool exec_success: bool message: str errors: list[dict[str, Any]] = field(default_factory=list) @property def error_codes(self) -> list[str]: codes: list[str] = [] for error in self.errors: code = error.get("code") if code and code not in codes: codes.append(str(code)) return codes def render_errors(self) -> str: if not self.errors: return self.message lines = [] for error in self.errors: code = error.get("code", "error") message = error.get("message", "") lines.append(f"{code}: {message}".strip()) return "\n".join(lines) def ast_parses(code: str) -> bool: """Check whether the code is valid Python (AST-parseable).""" try: ast.parse(code) return True except SyntaxError: return False def syntax_error_message(code: str) -> str: """Return a line-level Python syntax error message for repair prompts.""" try: ast.parse(code) except SyntaxError as exc: location = [] if exc.lineno is not None: location.append(f"line {exc.lineno}") if exc.offset is not None: location.append(f"column {exc.offset}") prefix = f" at {', '.join(location)}" if location else "" details = f"{exc.msg}{prefix}" if exc.text: details += f"\n {exc.text.strip()}" if exc.offset: details += f"\n {' ' * max(exc.offset - 1, 0)}^" return details return "" def extract_scene_class(code: str) -> str | None: """Return the first Scene subclass name found in manim code.""" try: tree = ast.parse(code) except SyntaxError: return None for node in ast.walk(tree): if isinstance(node, ast.ClassDef): for base in node.bases: base_name = "" if isinstance(base, ast.Name): base_name = base.id elif isinstance(base, ast.Attribute): base_name = base.attr if "Scene" in base_name: return node.name return None # --------------------------------------------------------------------------- # Marimo validation via `marimo check` CLI # --------------------------------------------------------------------------- def check_marimo(code: str, timeout: int = 8) -> tuple[bool, str, list[str]]: """Run ``marimo check`` on code. Returns (passed, message, rule_codes). Catches breaking rules MB001-MB005: unparsable cells, duplicate definitions, cycle dependencies, setup cell issues, syntax errors. Runs in ~100-200ms — much faster than full ``marimo export html``. """ with tempfile.NamedTemporaryFile(suffix=".py", mode="w", delete=False) as f: f.write(code) f.flush() tmp = f.name try: result = subprocess.run( [sys.executable, "-m", "marimo", "check", "--format", "json", "--select", "MB", tmp], capture_output=True, text=True, timeout=timeout, ) if not result.stdout.strip(): message = _subprocess_error_message(result, "marimo check produced no output") code = "MARIMO_MISSING" if _missing_module(result, "marimo") else "MARIMO_CHECK" return False, message, [code] data = json.loads(result.stdout) issues = data.get("issues", []) if not issues: return True, "marimo check passed", [] codes = [] for issue in issues: code = issue.get("code") if code and code not in codes: codes.append(code) msg = "\n\n".join(_format_marimo_issue(issue) for issue in issues[:3]) return False, msg, codes except FileNotFoundError: return False, "marimo not installed", ["MARIMO_MISSING"] except subprocess.TimeoutExpired: return False, "marimo check timed out", ["MARIMO_TIMEOUT"] except (json.JSONDecodeError, KeyError): message = _subprocess_error_message(result, "marimo check output unparseable") return False, message, ["MARIMO_CHECK_PARSE"] finally: Path(tmp).unlink(missing_ok=True) # --------------------------------------------------------------------------- # Full execution # --------------------------------------------------------------------------- def run_marimo(code: str, timeout: int = 15, *, skip_check: bool = False) -> tuple[bool, str]: """Export a marimo notebook to HTML, optionally running static checks first. Runs ``marimo check`` first (fast static analysis). If that fails, returns immediately without the expensive export step. """ check_timeout = min(8, max(1, timeout // 2)) if not skip_check: passed, msg, _violations = check_marimo(code, timeout=check_timeout) if not passed: return False, msg with tempfile.NamedTemporaryFile(suffix=".py", mode="w", delete=False) as f: f.write(code) f.flush() tmp = f.name try: result = subprocess.run( [sys.executable, "-m", "marimo", "export", "html", tmp], capture_output=True, text=True, timeout=max(1, timeout - check_timeout), ) if result.returncode == 0: return True, "marimo export succeeded" return False, _subprocess_error_message(result, "marimo export failed")[:500] except FileNotFoundError: return False, "marimo not installed" except subprocess.TimeoutExpired: return False, "marimo export timed out" finally: Path(tmp).unlink(missing_ok=True) def run_manim(code: str, timeout: int = 30) -> tuple[bool, str]: """Try rendering a manim scene (low quality). Returns (success, message).""" scene = extract_scene_class(code) if scene is None: return False, "No Scene subclass found in code" with tempfile.TemporaryDirectory() as tmpdir: src = Path(tmpdir) / "scene.py" src.write_text(code) try: result = subprocess.run( [sys.executable, "-m", "manim", "render", "-ql", "--media_dir", tmpdir, str(src), scene], capture_output=True, text=True, timeout=timeout, ) if result.returncode == 0: return True, "manim render succeeded" return False, _subprocess_error_message(result, "manim render failed")[:500] except FileNotFoundError: return False, "manim not installed" except subprocess.TimeoutExpired: return False, "manim render timed out" def validate_code(fmt: str, code: str) -> SandboxResult: """Validate code and return parseable feedback for generation/repair.""" if not ast_parses(code): details = syntax_error_message(code) return SandboxResult( fmt=fmt, parses=False, check_passed=False, exec_success=False, message="Code has syntax errors and cannot be parsed.", errors=[{"code": "PY_SYNTAX", "message": details or "Code has syntax errors."}], ) if fmt == "marimo": check_passed, check_msg, codes = check_marimo(code) if not check_passed: return SandboxResult( fmt=fmt, parses=True, check_passed=False, exec_success=False, message=check_msg, errors=[{"code": code, "message": check_msg} for code in codes], ) exec_success, exec_msg = run_marimo(code, skip_check=True) return SandboxResult( fmt=fmt, parses=True, check_passed=True, exec_success=exec_success, message=exec_msg, errors=[] if exec_success else [{"code": "MARIMO_EXPORT", "message": exec_msg}], ) if fmt == "manim": scene = extract_scene_class(code) if scene is None: return SandboxResult( fmt=fmt, parses=True, check_passed=False, exec_success=False, message="No Scene subclass found in code", errors=[{"code": "MANIM_NO_SCENE", "message": "No Scene subclass found."}], ) exec_success, exec_msg = run_manim(code) return SandboxResult( fmt=fmt, parses=True, check_passed=True, exec_success=exec_success, message=exec_msg, errors=[] if exec_success else [{"code": "MANIM_RENDER", "message": exec_msg}], ) return SandboxResult( fmt=fmt, parses=True, check_passed=False, exec_success=False, message=f"Unknown format: {fmt}", errors=[{"code": "UNKNOWN_FORMAT", "message": f"Unknown format: {fmt}"}], ) def _format_marimo_issue(issue: dict[str, Any]) -> str: code = issue.get("code", "MB") message = issue.get("message", "unknown error") fix_hint = issue.get("fix", "") location = _format_issue_location(issue) rendered = f"{code}{location}: {message}" if fix_hint: rendered += f"\nFix hint: {fix_hint}" return rendered def _missing_module(result: subprocess.CompletedProcess[str], module: str) -> bool: output = f"{result.stderr}\n{result.stdout}" return f"No module named {module}" in output def _subprocess_error_message( result: subprocess.CompletedProcess[str], fallback: str, ) -> str: details = (result.stderr or result.stdout or "").strip() if details: return details return fallback def _format_issue_location(issue: dict[str, Any]) -> str: line = issue.get("line") or issue.get("lineno") or issue.get("start_line") column = issue.get("column") or issue.get("col") or issue.get("start_column") if line and column: return f" at line {line}, column {column}" if line: return f" at line {line}" return ""