josie / llm-bench

"""File/exec tool implementations for agentic build benches.

These are the tools the agent can call during a multi-turn build loop:
list_files, read_file, write_file, edit_file, search, run_build, run_test,
run_shell. They operate on a work dir set via `set_work_dir()`.

The tool SCHEMA (OpenAI function-calling format) is in TOOLS; the
implementations are in TOOL_IMPLS. Instruments that want a different tool
set (e.g. agentic_bench_v1 uses read-only exploration + submit_fix) define
their own and don't import these.
"""
import os
import re
import subprocess
import shutil

BUILD_TIMEOUT = int(os.environ.get("BUILD_TIMEOUT", "180"))
TEST_TIMEOUT = int(os.environ.get("TEST_TIMEOUT", "180"))
SHELL_TIMEOUT = int(os.environ.get("SHELL_TIMEOUT", "60"))

_WORK = None

def set_work_dir(d):
    global _WORK
    _WORK = d

def get_work_dir():
    return _WORK


def materialize(files, dest=None):
    """Write a dict of {relative_path: content} into a work dir (cleared first)."""
    d = dest or _WORK
    shutil.rmtree(d, ignore_errors=True)
    os.makedirs(d, exist_ok=True)
    for path, content in files.items():
        fp = os.path.join(d, path)
        os.makedirs(os.path.dirname(fp), exist_ok=True)
        with open(fp, "w") as f:
            f.write(content)
    return d


def _run(cmd, timeout):
    try:
        p = subprocess.run(cmd, cwd=_WORK, shell=True, capture_output=True,
                           text=True, timeout=timeout)
    except subprocess.TimeoutExpired:
        return f"TIMEOUT after {timeout}s"
    out = (p.stdout + "\n" + p.stderr).strip()
    return f"exit={p.returncode}\n{out}"


def _list_files(args):
    out = []
    for root, dirs, files in os.walk(_WORK):
        dirs[:] = sorted(d for d in dirs if d not in {".git", "node_modules",
                                                       ".next", "target", "dist"})
        for f in sorted(files):
            full = os.path.join(root, f)
            rel = os.path.relpath(full, _WORK)
            out.append(rel)
    return "\n".join(out)


def _read_file(args):
    p = args.get("path", "")
    fp = os.path.join(_WORK, p)
    if not os.path.isfile(fp):
        return f"ERROR: no such file {p!r}. Use list_files."
    with open(fp) as f:
        return f.read()


def _write_file(args):
    p = args.get("path", "")
    content = args.get("content", "")
    fp = os.path.join(_WORK, p)
    os.makedirs(os.path.dirname(fp), exist_ok=True)
    with open(fp, "w") as f:
        f.write(content)
    return f"wrote {p} ({len(content)} bytes)"


def _edit_file(args):
    p = args.get("path", "")
    old = args.get("old_str", "")
    new = args.get("new_str", "")
    fp = os.path.join(_WORK, p)
    if not os.path.isfile(fp):
        return f"ERROR: no such file {p!r}."
    with open(fp) as f:
        content = f.read()
    count = content.count(old)
    if count == 0:
        return f"ERROR: old_str not found in {p}."
    if count > 1:
        return f"ERROR: old_str appears {count} times in {p}; make it unique."
    new_content = content.replace(old, new, 1)
    with open(fp, "w") as f:
        f.write(new_content)
    return f"edited {p} ({len(old)} -> {len(new)} chars)"


def _search(args):
    pat = args.get("pattern", "")
    try:
        rx = re.compile(pat)
    except re.error as e:
        return f"ERROR: bad regex: {e}"
    hits = []
    for root, dirs, files in os.walk(_WORK):
        dirs[:] = sorted(d for d in dirs if d not in {".git", "node_modules",
                                                       ".next", "target", "dist"})
        for f in sorted(files):
            full = os.path.join(root, f)
            rel = os.path.relpath(full, _WORK)
            with open(full, errors="replace") as fh:
                for i, line in enumerate(fh, 1):
                    if rx.search(line):
                        hits.append(f"{rel}:{i}: {line.strip()}")
    return "\n".join(hits) if hits else "(no matches)"


def _run_build(args):
    return _run("go build ./...", BUILD_TIMEOUT)


def _run_test(args):
    return _run("go test ./...", TEST_TIMEOUT)


def _run_shell(args):
    cmd = args.get("cmd", "")
    if not cmd:
        return "ERROR: empty cmd"
    return _run(cmd, SHELL_TIMEOUT)


TOOL_IMPLS = {
    "list_files": _list_files,
    "read_file": _read_file,
    "write_file": _write_file,
    "edit_file": _edit_file,
    "search": _search,
    "run_build": _run_build,
    "run_test": _run_test,
    "run_shell": _run_shell,
}

TOOLS = [
    {"type": "function", "function": {
        "name": "list_files",
        "description": "List all file paths in the work dir. Returns one path per line.",
        "parameters": {"type": "object", "properties": {}},
    }},
    {"type": "function", "function": {
        "name": "read_file",
        "description": "Return the full contents of one file in the work dir.",
        "parameters": {"type": "object", "properties": {
            "path": {"type": "string", "description": "Project-relative file path"},
        }, "required": ["path"]},
    }},
    {"type": "function", "function": {
        "name": "write_file",
        "description": "Create or overwrite a file in the work dir with the given contents.",
        "parameters": {"type": "object", "properties": {
            "path": {"type": "string"},
            "content": {"type": "string", "description": "Full file contents to write"},
        }, "required": ["path", "content"]},
    }},
    {"type": "function", "function": {
        "name": "edit_file",
        "description": "Apply a patch-style edit: replace the first occurrence of old_str with new_str. Fails if old_str is not found or appears more than once.",
        "parameters": {"type": "object", "properties": {
            "path": {"type": "string"},
            "old_str": {"type": "string", "description": "The exact text to find (must be unique)"},
            "new_str": {"type": "string", "description": "The replacement text"},
        }, "required": ["path", "old_str", "new_str"]},
    }},
    {"type": "function", "function": {
        "name": "search",
        "description": "Regex-search the whole work dir; returns matching path:line: text.",
        "parameters": {"type": "object", "properties": {
            "pattern": {"type": "string"},
        }, "required": ["pattern"]},
    }},
    {"type": "function", "function": {
        "name": "run_build",
        "description": f"Run `go build ./...` in the work dir. Returns exit code + output. Timeout {BUILD_TIMEOUT}s.",
        "parameters": {"type": "object", "properties": {}},
    }},
    {"type": "function", "function": {
        "name": "run_test",
        "description": f"Run `go test ./...` in the work dir. Returns exit code + output. Timeout {TEST_TIMEOUT}s.",
        "parameters": {"type": "object", "properties": {}},
    }},
    {"type": "function", "function": {
        "name": "run_shell",
        "description": f"Run an arbitrary shell command in the work dir (cwd = work dir). Use for `go get`, `go mod tidy`, npm, cargo, etc. Returns exit code + output. Timeout {SHELL_TIMEOUT}s.",
        "parameters": {"type": "object", "properties": {
            "cmd": {"type": "string", "description": "The shell command to run (sh -c)"},
        }, "required": ["cmd"]},
    }},
]