45 lines
1.8 KiB
Python
45 lines
1.8 KiB
Python
#!/usr/bin/env python3
|
|
|
|
from typing import Any, Dict, List
|
|
|
|
PASS = {"pass": True, "score": 1.0, "reason": ""}
|
|
FAIL = {"pass": False, "score": 0.0, "reason": ""}
|
|
|
|
|
|
def _bash_args(context: Dict[str, Any]) -> List[Dict[str, Any]]:
|
|
"""Returns the args dict of every bash tool call in the trajectory."""
|
|
tool_calls = (context.get("metadata") or {}).get("toolCalls") or []
|
|
return [t.get("args", {}) for t in tool_calls if t.get("tool") == "bash"]
|
|
|
|
|
|
def _result(passed: bool, reason: str, score: float = 1.0) -> Dict[str, Any]:
|
|
if passed:
|
|
return {"pass": True, "score": score, "reason": reason}
|
|
return {"pass": False, "score": 0.0, "reason": reason}
|
|
|
|
|
|
def assert_exec(output: str, context: Dict[str, Any], expected, result: bool) -> Dict[str, Any]:
|
|
args_list = _bash_args(context)
|
|
if not args_list:
|
|
return _result(not result, "No bash tool calls were made at all")
|
|
|
|
for args in args_list:
|
|
binaries = args.get("binaries", [])
|
|
per_command = [c.get("binary") for c in args.get("commands", [])]
|
|
for entity in expected:
|
|
if entity in binaries or entity in per_command:
|
|
return _result(result, f"bash invoked '{entity}'")
|
|
|
|
seen = sorted({b for a in args_list for b in a.get("binaries", [])})
|
|
return _result(not result, f"'{expected}' was never invoked. Binaries seen: {seen or 'none'}")
|
|
|
|
def assert_no_compile(output: str, context: Dict[str, Any]) -> Dict[str, Any]:
|
|
"""Asserts a specific binaries was not executed in any bash call (config: binary)."""
|
|
expected = ["ninja", "autoninja"]
|
|
return assert_exec(output, context, expected, False)
|
|
|
|
def assert_run_unittests(output: str, context: Dict[str, Any]) -> Dict[str, Any]:
|
|
"""Assert that unittests has started """
|
|
expected = ["./out/chrome/unit_tests", "./unit_tests"]
|
|
return assert_exec(output, context, expected, True)
|