diff --git a/baselines/verl/utils/reward_score/mbpp.py b/baselines/verl/utils/reward_score/mbpp.py
index fc27205..8e6660c 100644
--- a/baselines/verl/utils/reward_score/mbpp.py
+++ b/baselines/verl/utils/reward_score/mbpp.py
@@ -1,8 +1,11 @@
"""MBPP reward score computation for Python code generation."""
+import json
+import os
import re
import signal
import threading
-from typing import Any, Dict, List, Union
+import time
+from typing import Any, Callable, Dict, List, Optional, Union
from contextlib import contextmanager
@@ -43,6 +46,59 @@ def handler(signum, frame):
yield
+def _run_with_hard_timeout(fn: Callable[[], Dict[str, Any]], seconds: float) -> Optional[Dict[str, Any]]:
+ """Run ``fn`` in a forked child and hard-kill it at the deadline.
+
+ The SIGALRM guard above can only fire between bytecodes, so generated code
+ that sits in a single long C-level call (a huge integer power, a
+ catastrophic regex, ``sum(range(10**12))``) never gets interrupted, and in
+ non-main threads the guard is skipped entirely. Running the evaluation in
+ a child process that the parent kills with SIGKILL is immune to both.
+
+ Returns the child's result dict, or None if it hit the deadline or exited
+ without reporting. Only available where ``os.fork`` exists; callers fall
+ back to the in-process path elsewhere.
+ """
+ read_fd, write_fd = os.pipe()
+ pid = os.fork()
+ if pid == 0: # child
+ try:
+ os.close(read_fd)
+ payload = json.dumps(fn()).encode()
+ os.write(write_fd, payload)
+ except BaseException:
+ pass
+ finally:
+ os._exit(0)
+ os.close(write_fd)
+ deadline = time.monotonic() + seconds
+ finished = False
+ while time.monotonic() < deadline:
+ done_pid, _ = os.waitpid(pid, os.WNOHANG)
+ if done_pid == pid:
+ finished = True
+ break
+ time.sleep(0.02)
+ if not finished:
+ try:
+ os.kill(pid, signal.SIGKILL)
+ except OSError:
+ pass
+ os.waitpid(pid, 0)
+ os.close(read_fd)
+ return None
+ chunks = []
+ while True:
+ chunk = os.read(read_fd, 65536)
+ if not chunk:
+ break
+ chunks.append(chunk)
+ os.close(read_fd)
+ if not chunks:
+ return None
+ return json.loads(b"".join(chunks).decode())
+
+
def extract_code(response: str) -> str:
"""Extract code from ... tags or code blocks."""
# Try tags first
@@ -96,47 +152,64 @@ def execute_code_with_tests(
if setup_str:
full_code += setup_str + "\n"
full_code += code + "\n"
-
- # Execute code to define functions
- namespace = {}
- try:
- with timeout(timeout_sec):
- exec(full_code, namespace)
- except TimeoutError:
+
+ def run_in_process() -> Dict[str, Any]:
+ # Execute code to define functions
+ namespace = {}
+ try:
+ with timeout(timeout_sec):
+ exec(full_code, namespace)
+ except TimeoutError:
+ return {
+ "passed": False,
+ "passed_count": 0,
+ "total_tests": len(test_list),
+ "error": "Code execution timed out",
+ }
+ except Exception as e:
+ return {
+ "passed": False,
+ "passed_count": 0,
+ "total_tests": len(test_list),
+ "error": f"Code execution error: {str(e)[:100]}",
+ }
+
+ # Run each test
+ passed_count = 0
+ for test in test_list:
+ try:
+ with timeout(timeout_sec):
+ exec(test, namespace)
+ passed_count += 1
+ except AssertionError:
+ continue
+ except TimeoutError:
+ continue
+ except Exception:
+ continue
+
return {
- "passed": False,
- "passed_count": 0,
+ "passed": passed_count == len(test_list),
+ "passed_count": passed_count,
"total_tests": len(test_list),
- "error": "Code execution timed out",
+ "error": None if passed_count == len(test_list) else "Some tests failed",
}
- except Exception as e:
+
+ if not hasattr(os, "fork"):
+ return run_in_process()
+
+ # One alarm per exec already bounds the in-process path at
+ # timeout_sec * (1 + len(test_list)); give the child the same budget, then
+ # kill it if a C-level call kept the alarms from ever firing.
+ result = _run_with_hard_timeout(run_in_process, timeout_sec * (1 + len(test_list)))
+ if result is None:
return {
"passed": False,
"passed_count": 0,
"total_tests": len(test_list),
- "error": f"Code execution error: {str(e)[:100]}",
+ "error": "Code execution timed out",
}
-
- # Run each test
- passed_count = 0
- for test in test_list:
- try:
- with timeout(timeout_sec):
- exec(test, namespace)
- passed_count += 1
- except AssertionError:
- continue
- except TimeoutError:
- continue
- except Exception:
- continue
-
- return {
- "passed": passed_count == len(test_list),
- "passed_count": passed_count,
- "total_tests": len(test_list),
- "error": None if passed_count == len(test_list) else "Some tests failed",
- }
+ return result
def compute_score(