josie / llm-bench

"""Agentic fix instrument — multi-turn tool loop over a tiny in-memory codebase.

The model gets a failing Go test and a read-only toolbox (list_files / read_file
/ search / submit_fix), must explore the code, reason ACROSS files, and submit
a fix. Scored by compile+run against the hidden test with the real Go toolchain.

Why this shape:
  - Long context: the fix requires tracing Set() in store.go against isExpired()
    in ttl.go; a model that only reads one file cannot find it.
  - Agentic: real tool-call loop (assistant tool_calls -> tool results -> repeat),
    the exact pattern an agent panel drives — not a one-shot prompt.
  - Coherence-under-turns: MAX_TURNS caps runaway loops; we record turns used
    and whether the agent converged, which is where reasoning-loop behavior shows.

Usage (via the bench.py CLI):
  python bench.py agentic-fix validate
  python bench.py agentic-fix run

Env (via .env): ENDPOINT, MODEL, API_KEY, MAX_TOKENS, MAX_TURNS, HTTP_TIMEOUT.
"""
import json
import os
import re
import sys
import time
import shutil
import subprocess

from bench.core import client
from bench.core.reporting import write_json

HERE = os.path.dirname(os.path.abspath(__file__))
WORK = os.path.join(os.path.dirname(os.path.dirname(os.path.abspath(__file__))), "work", "agentic_fix")
MAX_TOKENS = int(os.environ.get("MAX_TOKENS", "4000"))
MAX_TURNS = int(os.environ.get("MAX_TURNS", "24"))
RUN_TIMEOUT = 120

# ----------------------------------------------------------------- the codebase
# A small in-memory TTL cache split across files. The bug is cross-file:
# store.go Set() stores expireAt = now instead of now.Add(ttl), so every
# entry is born expired. Finding it requires reading BOTH store.go (writes
# expireAt) and ttl.go (isExpired compares expireAt to now).
PROJECT = {
"go.mod": "module kvstore\n\ngo 1.26\n",
"doc.go": '''// Package kvstore is a small concurrency-safe in-memory key/value cache with
// per-entry TTL expiry and a background sweeper that evicts expired entries.
package kvstore
''',
"store.go": '''package kvstore

import (
	"sync"
	"time"
)

type entry struct {
	value    any
	expireAt time.Time
	created  time.Time
}

type Store struct {
	mu     sync.RWMutex
	items  map[string]*entry
	stats  *Stats
	opts   options
	stopCh chan struct{}
	closed bool
}

func New(optFns ...Option) *Store {
	o := defaultOptions()
	for _, fn := range optFns {
		fn(&o)
	}
	s := &Store{
		items:  make(map[string]*entry),
		stats:  &Stats{},
		opts:   o,
		stopCh: make(chan struct{}),
	}
	go s.sweepLoop()
	return s
}

func (s *Store) Set(key string, value any, ttl time.Duration) {
	s.mu.Lock()
	defer s.mu.Unlock()
	now := time.Now()
	e := &entry{value: value, created: now}
	if ttl > 0 {
		e.expireAt = now // BUG: forgets to add ttl -> entry is born expired.
	}
	s.items[key] = e
}

func (s *Store) Get(key string) (any, bool) {
	s.mu.Lock()
	defer s.mu.Unlock()
	e, ok := s.items[key]
	if !ok {
		s.stats.recordMiss()
		return nil, false
	}
	if isExpired(e, time.Now()) {
		delete(s.items, key)
		s.stats.recordMiss()
		return nil, false
	}
	s.stats.recordHit()
	return e.value, true
}

func (s *Store) Delete(key string) bool {
	s.mu.Lock()
	defer s.mu.Unlock()
	_, ok := s.items[key]
	delete(s.items, key)
	return ok
}

func (s *Store) Len() int {
	s.mu.RLock()
	defer s.mu.RUnlock()
	return len(s.items)
}

func (s *Store) Keys() []string {
	s.mu.RLock()
	defer s.mu.RUnlock()
	ks := make([]string, 0, len(s.items))
	for k := range s.items {
		ks = append(ks, k)
	}
	return ks
}

func (s *Store) Close() {
	s.mu.Lock()
	defer s.mu.Unlock()
	if !s.closed {
		close(s.stopCh)
		s.closed = true
	}
}
''',
"ttl.go": '''package kvstore

import "time"

func isExpired(e *entry, now time.Time) bool {
	if e.expireAt.IsZero() {
		return false
	}
	return now.After(e.expireAt)
}

func remaining(e *entry, now time.Time) time.Duration {
	if e.expireAt.IsZero() {
		return 0
	}
	d := e.expireAt.Sub(now)
	if d < 0 {
		return 0
	}
	return d
}
''',
"sweeper.go": '''package kvstore

import "time"

func (s *Store) sweepLoop() {
	t := time.NewTicker(s.opts.sweepInterval)
	defer t.Stop()
	for {
		select {
		case <-s.stopCh:
			return
		case <-t.C:
			s.sweepOnce()
		}
	}
}

func (s *Store) sweepOnce() {
	now := time.Now()
	s.mu.Lock()
	defer s.mu.Unlock()
	for k, e := range s.items {
		if isExpired(e, now) {
			delete(s.items, k)
			s.stats.recordEviction()
		}
	}
}
''',
"stats.go": '''package kvstore

import "sync/atomic"

type Stats struct {
	hits      atomic.Int64
	misses    atomic.Int64
	evictions atomic.Int64
}

func (s *Stats) recordHit()      { s.hits.Add(1) }
func (s *Stats) recordMiss()     { s.misses.Add(1) }
func (s *Stats) recordEviction() { s.evictions.Add(1) }

func (s *Stats) Snapshot() (hits, misses, evictions int64) {
	return s.hits.Load(), s.misses.Load(), s.evictions.Load()
}

func (s *Store) Stats() *Stats { return s.stats }
''',
"options.go": '''package kvstore

import "time"

type options struct {
	sweepInterval time.Duration
}

type Option func(*options)

func defaultOptions() options {
	return options{sweepInterval: 30 * time.Second}
}

func WithSweepInterval(d time.Duration) Option {
	return func(o *options) {
		if d > 0 {
			o.sweepInterval = d
		}
	}
}
''',
"store_test.go": '''package kvstore

import (
	"testing"
	"time"
)

func TestSetThenGetWithinTTL(t *testing.T) {
	s := New(WithSweepInterval(time.Hour))
	defer s.Close()
	s.Set("k", "v", time.Hour)
	got, ok := s.Get("k")
	if !ok {
		t.Fatalf("Get right after Set: want hit, got miss")
	}
	if got != "v" {
		t.Fatalf("Get value = %v, want v", got)
	}
}

func TestExpiresAfterTTL(t *testing.T) {
	s := New(WithSweepInterval(time.Hour))
	defer s.Close()
	s.Set("k", "v", 10*time.Millisecond)
	time.Sleep(25 * time.Millisecond)
	if _, ok := s.Get("k"); ok {
		t.Fatalf("Get after TTL: want miss, got hit")
	}
}

func TestZeroTTLNeverExpires(t *testing.T) {
	s := New(WithSweepInterval(time.Hour))
	defer s.Close()
	s.Set("k", "v", 0)
	time.Sleep(15 * time.Millisecond)
	if _, ok := s.Get("k"); !ok {
		t.Fatalf("zero-ttl entry: want hit, got miss")
	}
}
''',
}

FIX = {
"store.go": PROJECT["store.go"].replace(
    "\t\te.expireAt = now // BUG: forgets to add ttl -> entry is born expired.",
    "\t\te.expireAt = now.Add(ttl)",
),
}
BUG_FILES = ["store.go"]

# ----------------------------------------------------------------- go toolchain
def _materialize(files):
    shutil.rmtree(WORK, ignore_errors=True)
    os.makedirs(WORK, exist_ok=True)
    for path, content in files.items():
        fp = os.path.join(WORK, path)
        os.makedirs(os.path.dirname(fp), exist_ok=True)
        with open(fp, "w") as f:
            f.write(content)
    return WORK

def _go_test(d):
    try:
        p = subprocess.run(["go", "test", "./..."], cwd=d,
                           capture_output=True, text=True, timeout=RUN_TIMEOUT)
    except subprocess.TimeoutExpired:
        return False, "TIMEOUT"
    return p.returncode == 0, (p.stdout + "\n" + p.stderr).strip()

def _apply(base, patch):
    merged = dict(base)
    merged.update(patch)
    return merged

# ----------------------------------------------------------------- tools
TOOLS = [
    {"type": "function", "function": {
        "name": "list_files", "description": "List all file paths in the project.",
        "parameters": {"type": "object", "properties": {}}}},
    {"type": "function", "function": {
        "name": "read_file", "description": "Return the full contents of one file.",
        "parameters": {"type": "object", "properties": {
            "path": {"type": "string", "description": "Project-relative file path"}},
            "required": ["path"]}}},
    {"type": "function", "function": {
        "name": "search", "description": "Regex-search the whole project; returns matching path:line: text.",
        "parameters": {"type": "object", "properties": {"pattern": {"type": "string"}},
            "required": ["pattern"]}},
    },
    {"type": "function", "function": {
        "name": "submit_fix", "description": "Submit the complete corrected contents of ONE file to fix the failing test. Call this exactly once when you have the fix.",
        "parameters": {"type": "object", "properties": {
            "path": {"type": "string"},
            "content": {"type": "string", "description": "Full new file contents"}},
            "required": ["path", "content"]}},
    },
]

def _tool_result(name, args):
    if name == "list_files":
        return "\n".join(sorted(PROJECT.keys()))
    if name == "read_file":
        path = args.get("path", "")
        return PROJECT.get(path, f"ERROR: no such file {path!r}. Use list_files.")
    if name == "search":
        pat = args.get("pattern", "")
        try:
            rx = re.compile(pat)
        except re.error as e:
            return f"ERROR: bad regex: {e}"
        hits = []
        for path, content in sorted(PROJECT.items()):
            for i, line in enumerate(content.splitlines(), 1):
                if rx.search(line):
                    hits.append(f"{path}:{i}: {line.strip()}")
        return "\n".join(hits) if hits else "(no matches)"
    return f"ERROR: unknown tool {name}"

SYSTEM = (
    "You are a coding agent working in a Go project. A test is failing. Use the "
    "tools to explore the codebase, find the root cause (it may span multiple "
    "files), and fix it. When you are confident, call submit_fix with the COMPLETE "
    "corrected contents of the single file that needs changing. Keep every other "
    "file untouched and preserve all names and signatures. Do not submit until you "
    "have read enough to be sure."
)

def validate():
    print("=== VALIDATION: buggy fails, reference fix passes ===\n")
    d = _materialize(PROJECT)
    bug_pass, bug_out = _go_test(d)
    d = _materialize(_apply(PROJECT, FIX))
    fix_pass, fix_out = _go_test(d)
    ok = (not bug_pass) and fix_pass
    print(f"buggy_fails={not bug_pass}  fix_passes={fix_pass}  ->  {'OK' if ok else 'BAD'}")
    if bug_pass:
        print("  !! buggy project unexpectedly PASSED")
    if not fix_pass:
        print("  !! reference fix FAILED:\n" + fix_out[-1200:])
    return ok

def run():
    d = _materialize(PROJECT)
    _, fail_out = _go_test(d)
    messages = [
        {"role": "system", "content": SYSTEM},
        {"role": "user", "content":
            "`go test ./...` fails in this project:\n\n```\n" + fail_out[-1500:] +
            "\n```\n\nInvestigate with the tools and submit a fix."},
    ]
    reads, searches, submitted = [], 0, None
    t0 = time.time()
    peak_prompt_tokens = 0
    turns = 0
    for turn in range(1, MAX_TURNS + 1):
        turns = turn
        try:
            resp, dt = client.call_model(messages, tools=TOOLS, max_tokens=MAX_TOKENS)
        except Exception as e:
            print(f"  turn {turn}: REQUEST ERROR {e!r}")
            submitted = ("__error__", "")
            break
        u = client.usage(resp)
        peak_prompt_tokens = max(peak_prompt_tokens, u["prompt_tokens"])
        msg = resp["choices"][0]["message"]
        tcs = msg.get("tool_calls") or []
        if not tcs:
            content = (msg.get("content") or "")[:200]
            print(f"  turn {turn}: no tool call (content={content!r})")
            messages.append({"role": "assistant", "content": msg.get("content") or ""})
            messages.append({"role": "user", "content":
                "You must either call a tool to keep investigating or call "
                "submit_fix. Do not answer in prose."})
            continue
        messages.append({"role": "assistant", "content": msg.get("content") or "",
                         "tool_calls": tcs})
        done = False
        for tc in tcs:
            fn = tc["function"]
            name = fn["name"]
            try:
                args = json.loads(fn["arguments"] or "{}")
            except Exception:
                args = {}
            if name == "submit_fix":
                submitted = (args.get("path", ""), args.get("content", ""))
                messages.append({"role": "tool", "tool_call_id": tc.get("id", ""),
                                 "content": "fix received"})
                done = True
                print(f"  turn {turn}: submit_fix({args.get('path','?')}) "
                      f"[{len(reads)} reads, {searches} searches]")
                break
            if name == "read_file":
                reads.append(args.get("path", ""))
            elif name == "search":
                searches += 1
            result = _tool_result(name, args)
            messages.append({"role": "tool", "tool_call_id": tc.get("id", ""),
                             "content": result})
        if done:
            break
    wall = time.time() - t0
    rec = {"model": client.MODEL, "turns": turns, "reads": reads, "searches": searches,
           "peak_prompt_tokens": peak_prompt_tokens, "wall_s": round(wall, 1)}
    if not submitted or submitted[0] == "__error__":
        rec.update({"passed": False, "reason": "no fix submitted (turn cap or error)"})
    else:
        path, content = submitted
        rec["submitted_path"] = path
        if path not in BUG_FILES:
            rec["note"] = f"submitted {path}, expected one of {BUG_FILES}"
        merged = _apply(PROJECT, {path: content})
        d = _materialize(merged)
        passed, out = _go_test(d)
        rec.update({"passed": bool(passed), "test_output": out[-800:]})
    _report(rec)
    return rec

def _report(rec):
    print("\n===================== AGENTIC FIX RESULT =====================")
    print(f"model            : {rec['model']}")
    print(f"passed           : {'YES' if rec.get('passed') else 'no'}")
    print(f"turns used       : {rec['turns']}/{MAX_TURNS}")
    print(f"files read       : {len(rec['reads'])}  {rec['reads']}")
    print(f"searches         : {rec['searches']}")
    print(f"peak prompt tok  : {rec['peak_prompt_tokens']}")
    print(f"wall clock       : {rec['wall_s']}s")
    if rec.get("note"):
        print(f"note             : {rec['note']}")
    if not rec.get("passed") and rec.get("test_output"):
        print("test output:\n" + rec["test_output"])
    write_json(rec, rec["model"], prefix="agentic-fix")
    print()



def main():
    mode = sys.argv[1] if len(sys.argv) > 1 else 'validate'
    if mode == 'validate':
        sys.exit(0 if validate() else 1)
    elif mode == 'run':
        run()
    else:
        print(f'usage: python bench.py agentic-fix [validate|run]')
        sys.exit(1)


if __name__ == "__main__":
    main()