File size: 2,988 Bytes
71a1c53
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
"""

Grader

------

Sandboxed test runner that executes agent-submitted code against a test suite

and returns structured results.

"""

from __future__ import annotations

import multiprocessing
import traceback
from dataclasses import dataclass, field
from typing import Any, List, Tuple


@dataclass
class GradeResult:
    tests_passed: int = 0
    tests_total: int = 0
    details: str = ""
    stderr: str = ""


_EXEC_TIMEOUT = 5  # seconds per test case


def _run_single_test(

    code: str,

    function_name: str,

    args: tuple,

    expected: Any,

    queue: multiprocessing.Queue,

) -> None:
    """Executed in a child process to isolate side effects."""
    try:
        namespace: dict = {}
        exec(compile(code, "<agent_code>", "exec"), namespace)
        fn = namespace.get(function_name)
        if fn is None:
            queue.put(("error", f"Function '{function_name}' not defined"))
            return
        actual = fn(*args)
        if actual == expected:
            queue.put(("pass", None))
        else:
            queue.put(("fail", f"Expected {expected!r}, got {actual!r}"))
    except Exception:
        queue.put(("error", traceback.format_exc()))


def grade(

    code: str,

    function_name: str,

    tests: List[Tuple[tuple, Any]],

) -> GradeResult:
    """Run *code* against every test case and return a GradeResult."""
    result = GradeResult(tests_total=len(tests))
    lines: list[str] = []
    stderr_parts: list[str] = []

    # Quick syntax check
    try:
        compile(code, "<agent_code>", "exec")
    except SyntaxError as exc:
        result.stderr = f"SyntaxError: {exc}"
        result.details = f"Code failed to compile:\n  {exc}"
        return result

    for idx, (args, expected) in enumerate(tests, 1):
        q: multiprocessing.Queue = multiprocessing.Queue()
        proc = multiprocessing.Process(
            target=_run_single_test,
            args=(code, function_name, args, expected, q),
        )
        proc.start()
        proc.join(timeout=_EXEC_TIMEOUT)

        if proc.is_alive():
            proc.terminate()
            proc.join(timeout=2)
            status, detail = "error", "Timed out"
        elif q.empty():
            status, detail = "error", "No result (process crashed)"
        else:
            status, detail = q.get_nowait()

        if status == "pass":
            result.tests_passed += 1
            lines.append(f"  Test {idx}: PASS  {function_name}{args} == {expected!r}")
        elif status == "fail":
            lines.append(
                f"  Test {idx}: FAIL  {function_name}{args}  {detail}"
            )
        else:
            lines.append(f"  Test {idx}: ERROR {function_name}{args}  {detail}")
            if detail:
                stderr_parts.append(detail)

    result.details = "\n".join(lines)
    result.stderr = "\n---\n".join(stderr_parts)
    return result