explainer-env / rewards /sandbox.py
kgdrathan's picture
Upload folder using huggingface_hub
8fa7af1 verified
Raw History Blame Contribute Delete
10.5 kB
"""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 ""