0b37d286065ecbf5f6dfcfafa2d9274d1d21dba3 / bench/tasks/agentic_fix.py · 14552 bytes · raw
"""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()