Spaces:
Sleeping
Sleeping
Download rewards/sandbox.py from kgdrathan/explainer-env: direct link, hf CLI and curl.
- Browser
- Download file 10.5 kB
-
https://huggingface.co/spaces/kgdrathan/explainer-env/resolve/main/rewards/sandbox.py
- Command line
-
hf download hf://spaces/kgdrathan/explainer-env/rewards/sandbox.py
-
curl -L -o sandbox.py https://huggingface.co/spaces/kgdrathan/explainer-env/resolve/main/rewards/sandbox.py
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 | |
| 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) | |
| 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 "" | |