darioooooo0o's picture
tiny-agent-112m: base + RL weights, tokenizer, code, model card
4397e12 verified
Raw History Blame Contribute Delete
38.3 kB
"""Workspace tasks with invented facts: the answer exists only in the workspace files.
Each task = a small fake project (configs in yaml/json/toml/env, code, tests, docs, logs, CSVs,
a long job log) built from made-up names and random values, plus one instruction with a
checkable outcome. Kinds:
lookups config_value, code_constant, csv_lookup, doc_fact, long_file, code_search
compute code_eval (run it), log_count (grep -c), aggregate (count across files), csv_sum
multi-step multi_hop (two files), write_fact (look up, then write a file)
change edit (change a config value), fix_test (fix a bug until the test passes)
refusal not_found (answer NOT_FOUND)
`oracle_trajectory` solves a task with real tool calls in the sandbox; its runs are used as
pretraining/decay data so the base model knows the format before RL.
"""
from __future__ import annotations
import json
import random
import re
import tomllib
from dataclasses import dataclass, field
from tiny_agent.chat import SYSTEM_TA_V1
from tiny_agent.tools import Workspace
KINDS = ["config_value", "code_constant", "code_eval", "log_count", "csv_lookup", "doc_fact",
"multi_hop", "not_found", "edit", "fix_test", "long_file", "code_search", "aggregate",
"csv_sum", "write_fact"]
# Never in pretraining trajectories and not trained on by RL by default: eval on these measures
# tool use that transfers to new task shapes, not recall of the training templates.
HELDOUT_KINDS = ("multi_hop", "log_count", "write_fact")
# training-only kinds, kept out of KINDS so the eval set (which cycles through KINDS) never changes.
# They teach primitives the held-out kinds need without copying them: save_value is the only
# demonstration of the write tool (write_fact = config owner -> notes/owner.txt stays held out), and
# service_hop follows a pointer module -> service -> config (multi_hop follows config -> backup host).
EXTRA_TRAIN_KINDS = ["save_value", "service_hop"]
TRAIN_KINDS = [k for k in KINDS if k not in HELDOUT_KINDS] + EXTRA_TRAIN_KINDS
DONE_KINDS = ("edit", "fix_test", "write_fact", "save_value") # success is a workspace state, answer DONE
_ON = ["b", "d", "f", "g", "k", "l", "m", "n", "p", "r", "s", "t", "v", "z", "br", "dr", "kl", "st", "tr", "qu"]
_NU = ["a", "e", "i", "o", "u", "ai", "ei", "ou"]
_CO = ["", "n", "r", "s", "x", "l", "th", "rk", "nd"]
REGIONS = ["north", "south", "east", "west", "central"]
LEVELS = ["INFO", "INFO", "INFO", "WARN", "ERROR"]
FMTS = ["yaml", "json", "toml", "env"]
JOB_STATUS = ["queued", "running", "done", "failed", "cancelled"]
def word(rng, syll=(2, 3)) -> str:
return "".join(rng.choice(_ON) + rng.choice(_NU) + rng.choice(_CO) for _ in range(rng.randint(*syll)))
def person(rng) -> str:
return f"{word(rng, (1, 2)).capitalize()} {word(rng, (2, 3)).capitalize()}"
@dataclass
class Task:
kind: str
question: str
answer: str
files: dict[str, str]
meta: dict = field(default_factory=dict) # facts the oracle and checker need
def messages(self) -> list[dict]:
return [{"role": "system", "content": SYSTEM_TA_V1}, {"role": "user", "content": self.question}]
def _unique_words(rng, n, syll=(2, 3)):
out = set()
while len(out) < n:
out.add(word(rng, syll))
return sorted(out)
# ---------------------------------------------------------------- config formats
def render_config(fmt: str, d: dict) -> str:
if fmt == "yaml":
return "".join(f"{k}: {v}\n" for k, v in d.items())
if fmt == "json":
return json.dumps(d, indent=2) + "\n"
if fmt == "toml":
return "".join(f"{k} = {v}\n" if isinstance(v, int) else f'{k} = "{v}"\n' for k, v in d.items())
if fmt == "env":
return "".join(f'{k.upper()}="{v}"\n' if " " in str(v) else f"{k.upper()}={v}\n" for k, v in d.items())
raise ValueError(fmt)
def config_line(fmt: str, k: str, v) -> str:
"""The exact text holding key k (what an edit replaces)."""
return {"yaml": f"{k}: {v}", "json": f'"{k}": {v}', "toml": f"{k} = {v}", "env": f"{k.upper()}={v}"}[fmt]
def parse_config(fmt: str, text: str) -> dict[str, str]:
if fmt == "json":
return {k: str(v) for k, v in json.loads(text).items()}
if fmt == "toml":
return {k: str(v) for k, v in tomllib.loads(text).items()}
sep = ":" if fmt == "yaml" else "="
out = {}
for line in text.splitlines():
if sep in line:
k, v = line.split(sep, 1)
out[k.strip().lower()] = v.strip().strip('"')
return out
# ---------------------------------------------------------------- buggy-function templates (fix_test)
def _bug_templates(rng, fn):
a, b, c = rng.randint(1, 30), rng.randint(1, 30), rng.randint(1, 30)
k = rng.randint(3, 25)
lim = rng.randint(5, 20)
vals = [lim - 3, lim, lim + 4, lim - 1, lim]
tag, num = word(rng, (1, 2)), rng.randint(1, 99)
return [
dict(doc="Return the sum of values.", sig="values",
good=" total = 0\n for v in values:\n total += v\n return total\n",
bug_old="total += v", bug_new="total -= v",
tests=[(f"{fn}([{a}, {b}, {c}])", a + b + c), (f"{fn}([])", 0)]),
dict(doc="Return the sum of 1..n inclusive.", sig="n",
good=" return sum(range(1, n + 1))\n", bug_old="range(1, n + 1)", bug_new="range(1, n)",
tests=[(f"{fn}({k})", k * (k + 1) // 2), (f"{fn}(1)", 1)]),
dict(doc="Count values that are at least limit.", sig="values, limit",
good=" return sum(1 for v in values if v >= limit)\n", bug_old="v >= limit", bug_new="v > limit",
tests=[(f"{fn}({vals}, {lim})", sum(v >= lim for v in vals))]),
dict(doc="Build an id like name-num.", sig="name, num",
good=' return f"{name}-{num}"\n', bug_old='f"{name}-{num}"', bug_new='f"{name}_{num}"',
tests=[(f"{fn}('{tag}', {num})", f"{tag}-{num}")]),
]
# ---------------------------------------------------------------- project builder
def build_project(rng: random.Random) -> dict:
svcs = _unique_words(rng, rng.randint(3, 8))
services, files = {}, {}
for s in svcs:
services[s] = dict(port=rng.randint(1024, 65000), timeout_seconds=rng.choice([5, 10, 15, 20, 30, 45, 60, 90]),
owner=person(rng), replicas=rng.randint(1, 12), region=rng.choice(REGIONS),
fmt=rng.choice(FMTS))
for s in svcs:
services[s]["backup_host"] = rng.choice([x for x in svcs if x != s])
for s, c in services.items():
d = {"service": s, "port": c["port"], "timeout_seconds": c["timeout_seconds"], "owner": c["owner"],
"replicas": c["replicas"], "region": c["region"], "backup_host": c["backup_host"]}
c["file"] = f"config/{s}.{c['fmt']}"
files[c["file"]] = render_config(c["fmt"], d)
modules = {}
for m in _unique_words(rng, rng.randint(2, 4)):
fn = word(rng, (2, 2))
a, b = rng.randint(2, 19), rng.randint(-50, 99)
consts = {"MAX_RETRIES": rng.randint(1, 15), "BATCH_SIZE": rng.choice([8, 16, 32, 64, 100, 128, 250, 512]),
"CACHE_TTL": rng.randint(30, 7200)}
modules[m] = dict(fn=fn, a=a, b=b, consts=consts)
files[f"src/{m}.py"] = (
f'"""{m}: helpers for the {rng.choice(svcs)} service."""\n\n'
+ "".join(f"{k} = {v}\n" for k, v in consts.items())
+ f'DEFAULT_REGION = "{rng.choice(REGIONS)}"\n\n\n'
f"def {fn}(n):\n \"\"\"Scale n for {m}.\"\"\"\n return n * {a} + ({b})\n\n\n"
f"def describe():\n return \"{m} v{rng.randint(1, 9)}.{rng.randint(0, 20)}\"\n")
files["src/__init__.py"] = ""
# one module with a function and a test (buggy only for fix_test tasks)
tm, tfn = word(rng, (2, 2)), word(rng, (2, 2))
tpl = rng.choice(_bug_templates(rng, tfn))
files[f"src/{tm}.py"] = f'"""{tm} utilities."""\n\n\ndef {tfn}({tpl["sig"]}):\n """{tpl["doc"]}"""\n{tpl["good"]}'
files[f"tests/test_{tm}.py"] = (
f"import sys\nsys.path.insert(0, \".\")\nfrom src.{tm} import {tfn}\n\n\ndef main():\n"
+ "".join(f" assert {call} == {exp!r}, f\"{call} returned {{{call}!r}}, expected {exp!r}\"\n"
for call, exp in tpl["tests"])
+ ' print("OK")\n\n\nmain()\n')
testmod = dict(module=tm, fn=tfn, tpl=tpl)
components = {c: dict(owner=person(rng), codename=word(rng, (2, 2)).capitalize(),
since=rng.randint(2012, 2026)) for c in _unique_words(rng, rng.randint(3, 5))}
paras = []
for c, d in components.items():
paras.append(rng.choice([
f"The {c} component is owned by {d['owner']}. It was introduced in {d['since']} under the codename {d['codename']}.",
f"{d['owner']} maintains {c}, which dates from {d['since']}. Internally it is called {d['codename']}.",
f"Since {d['since']}, {c} (codename {d['codename']}) has been maintained by {d['owner']}.",
f"{c} was added in {d['since']}. Its codename is {d['codename']} and its maintainer is {d['owner']}.",
]))
paras.append(f"Questions about {c} usually concern {rng.choice(svcs)} latency and {word(rng)} retries.")
rng.shuffle(paras)
files["docs/architecture.md"] = "# Architecture\n\n" + "\n\n".join(paras) + "\n"
logs = {}
for s in rng.sample(svcs, k=min(2, len(svcs))):
codes = [f"E{rng.randint(100, 999)}" for _ in range(3)]
lines = []
for i in range(rng.randint(40, 160)):
lvl = rng.choice(LEVELS)
msg = f"{lvl} code={rng.choice(codes)} {word(rng)} took {rng.randint(1, 900)}ms" if lvl != "INFO" \
else f"INFO request {word(rng)} ok"
lines.append(f"2026-09-{rng.randint(1, 30):02d}T{rng.randint(0, 23):02d}:{rng.randint(0, 59):02d}:00 {msg}")
logs[s] = dict(codes=codes, lines=lines)
files[f"logs/{s}.log"] = "\n".join(lines) + "\n"
tables = {}
for t in _unique_words(rng, rng.randint(1, 2), (2, 2)):
rows = [dict(id=i + 1, name=w, quantity=rng.randint(0, 500), price=round(rng.uniform(0.5, 99), 2))
for i, w in enumerate(_unique_words(rng, rng.randint(6, 25), (2, 2)))]
tables[t] = rows
files[f"data/{t}.csv"] = "id,name,quantity,price\n" + "".join(
f"{r['id']},{r['name']},{r['quantity']},{r['price']}\n" for r in rows)
ids = rng.sample(range(10000, 99999), rng.randint(600, 2500))
jobs = {i: dict(status=rng.choice(JOB_STATUS), owner=word(rng, (1, 2)), duration=rng.randint(1, 5000)) for i in ids}
files["data/jobs.txt"] = "".join(f"job-{i} status={j['status']} owner={j['owner']} duration={j['duration']}s\n"
for i, j in jobs.items())
files["README.md"] = (f"# {word(rng).capitalize()} platform\n\nServices: {', '.join(svcs)}.\n"
"Configs live in config/, code in src/, tests in tests/, docs in docs/, logs in logs/, "
"data in data/.\n")
files["scripts/deploy.sh"] = (f"#!/bin/sh\n# deploy all services\nfor s in {' '.join(svcs)}; do\n"
f" echo \"deploying $s\"\ndone\n")
return dict(files=files, services=services, modules=modules, components=components, logs=logs,
tables=tables, testmod=testmod, jobs=jobs)
# ---------------------------------------------------------------- tasks
def _q(rng, *forms):
return rng.choice(forms)
def make_task(rng: random.Random, kind: str | None = None) -> Task:
p = build_project(rng)
kind = kind or rng.choice(KINDS)
S, M = p["services"], p["modules"]
F = p["files"]
if kind == "config_value":
s = rng.choice(sorted(S))
key = rng.choice(["port", "timeout_seconds", "owner", "replicas", "region"])
q = {"port": _q(rng, f"What port does the {s} service listen on?", f"Which port is {s} configured to use?",
f"{s} listens on which port?", f"Look up the port number for {s}."),
"timeout_seconds": _q(rng, f"What is the timeout (in seconds) configured for {s}?",
f"How many seconds is the {s} timeout?", f"What timeout does {s} use?"),
"owner": _q(rng, f"Who is the owner of the {s} service?", f"Who owns {s}?",
f"Which person is listed as owner of {s}?"),
"replicas": _q(rng, f"How many replicas does {s} run?", f"What is the replica count for {s}?",
f"{s} is configured with how many replicas?"),
"region": _q(rng, f"In which region is {s} deployed?", f"What region does {s} run in?",
f"Which region is configured for the {s} service?")}[key]
return Task(kind, q, str(S[s][key]), F, dict(svc=s, key=key, file=S[s]["file"]))
if kind == "code_constant":
m = rng.choice(sorted(M))
k = rng.choice(sorted(M[m]["consts"]))
q = _q(rng, f"What is {k} set to in src/{m}.py?", f"What value does the {m} module use for {k}?",
f"Find the value of {k} in the {m} module.", f"In {m}, what is {k}?")
return Task(kind, q, str(M[m]["consts"][k]), F, dict(module=m, const=k, file=f"src/{m}.py"))
if kind == "code_eval":
m = rng.choice(sorted(M))
d, n = M[m], rng.randint(2, 40)
q = _q(rng, f"What does {d['fn']}({n}) in src/{m}.py return?", f"What is the result of {d['fn']}({n})?",
f"Compute {d['fn']}({n}) using the code in the {m} module.")
return Task(kind, q, str(n * d["a"] + d["b"]), F, dict(module=m, fn=d["fn"], n=n, file=f"src/{m}.py"))
if kind == "log_count":
s = rng.choice(sorted(p["logs"]))
lg = p["logs"][s]
code = rng.choice(lg["codes"])
lvl = rng.choice(["ERROR", "WARN"])
n = sum(1 for l in lg["lines"] if f" {lvl} " in l and f"code={code}" in l)
q = _q(rng, f"How many {lvl} lines with code {code} are in logs/{s}.log?",
f"Count the {lvl} entries with code {code} in the {s} log.")
return Task(kind, q, str(n), F, dict(svc=s, code=code, level=lvl, file=f"logs/{s}.log"))
if kind == "csv_lookup":
t = rng.choice(sorted(p["tables"]))
r = rng.choice(p["tables"][t])
col = rng.choice(["quantity", "price"])
q = _q(rng, f"What is the {col} of {r['name']} in data/{t}.csv?", f"Look up the {col} for {r['name']} in the {t} table.",
f"In data/{t}.csv, what {col} is listed for {r['name']}?")
return Task(kind, q, str(r[col]), F, dict(table=t, name=r["name"], col=col, file=f"data/{t}.csv"))
if kind == "doc_fact":
c = rng.choice(sorted(p["components"]))
key = rng.choice(["owner", "codename", "since"])
q = {"owner": _q(rng, f"According to the docs, who owns the {c} component?", f"Who maintains {c}?"),
"codename": _q(rng, f"What is the internal codename of {c}?", f"What is {c} called internally?"),
"since": _q(rng, f"In which year was {c} introduced?", f"When was {c} added? Give the year.")}[key]
return Task(kind, q, str(p["components"][c][key]), F, dict(component=c, key=key, file="docs/architecture.md"))
if kind == "multi_hop":
s = rng.choice(sorted(S))
b = S[s]["backup_host"]
key = rng.choice(["port", "owner"])
q = _q(rng, f"What is the {key} of the service that {s} uses as its backup host?",
f"{s} has a backup host. What is that host's {key}?")
return Task(kind, q, str(S[b][key]), F, dict(svc=s, backup=b, key=key, file=S[s]["file"], file2=S[b]["file"]))
if kind == "not_found":
ghost = word(rng)
q = rng.choice([f"What port does the {ghost} service listen on?", f"Who owns the {ghost} component?",
f"What is MAX_RETRIES set to in src/{ghost}.py?", f"What is the status of job-{ghost}?"])
return Task(kind, q, "NOT_FOUND", F, dict(ghost=ghost))
if kind == "edit":
s = rng.choice(sorted(S))
old = S[s]["timeout_seconds"]
new = rng.choice([v for v in [5, 10, 15, 20, 25, 30, 40, 45, 60, 90, 120] if v != old])
q = _q(rng, f"Change the timeout_seconds of the {s} service to {new}, then submit DONE.",
f"Set the {s} timeout to {new} seconds in its config and submit DONE.")
return Task(kind, q, "DONE", F, dict(svc=s, old=old, new=new, file=S[s]["file"], fmt=S[s]["fmt"]))
if kind == "fix_test":
tm = p["testmod"]
tpl = tm["tpl"]
path = f"src/{tm['module']}.py"
F = dict(F)
F[path] = F[path].replace(tpl["bug_old"], tpl["bug_new"])
q = _q(rng, f"The test tests/test_{tm['module']}.py fails. Fix the bug in the code (not the test), then submit DONE.",
f"Make `python3 tests/test_{tm['module']}.py` pass by fixing src/{tm['module']}.py. Submit DONE when it passes.")
return Task(kind, q, "DONE", F, dict(module=tm["module"], file=path, test=f"tests/test_{tm['module']}.py",
bug_old=tpl["bug_old"], bug_new=tpl["bug_new"]))
if kind == "long_file":
jid = rng.choice(sorted(p["jobs"]))
j = p["jobs"][jid]
key = rng.choice(["status", "owner", "duration"])
q = {"status": _q(rng, f"What is the status of job-{jid} in data/jobs.txt?", f"Is job-{jid} done? Give its status."),
"owner": _q(rng, f"Who is the owner of job-{jid}?", f"Which owner is listed for job-{jid} in data/jobs.txt?"),
"duration": _q(rng, f"How many seconds did job-{jid} take?", f"What duration (in seconds) is recorded for job-{jid}?")}[key]
return Task(kind, q, str(j[key]), F, dict(job=jid, key=key, file="data/jobs.txt"))
if kind == "code_search":
m = rng.choice(sorted(M))
fn = M[m]["fn"]
q = _q(rng, f"Which file defines the function {fn}?", f"Where is {fn} defined? Give the file path.")
return Task(kind, q, f"src/{m}.py", F, dict(fn=fn, module=m))
if kind == "aggregate":
r = rng.choice(sorted({c["region"] for c in S.values()}))
n = sum(1 for c in S.values() if c["region"] == r)
q = _q(rng, f"How many services are deployed in the {r} region?", f"Count the services configured for region {r}.")
return Task(kind, q, str(n), F, dict(region=r))
if kind == "csv_sum":
t = rng.choice(sorted(p["tables"]))
total = sum(r["quantity"] for r in p["tables"][t])
q = _q(rng, f"What is the total quantity of all items in data/{t}.csv?", f"Sum the quantity column of data/{t}.csv.")
return Task(kind, q, str(total), F, dict(table=t, file=f"data/{t}.csv"))
if kind == "write_fact":
s = rng.choice(sorted(S))
q = _q(rng, f"Find the owner of the {s} service and write just their name to notes/owner.txt, then submit DONE.",
f"Write the name of whoever owns {s} into notes/owner.txt and submit DONE.")
return Task(kind, q, "DONE", F, dict(svc=s, file=S[s]["file"], owner=S[s]["owner"], out="notes/owner.txt"))
if kind == "save_value":
src = rng.choice(["csv_sum", "job", "const", "codename"])
out = rng.choice(["out/result.txt", f"reports/{word(rng, (1, 2))}.txt", "tmp/value.txt", "build/value.txt"])
if src == "csv_sum":
t = rng.choice(sorted(p["tables"]))
val, f1 = sum(r["quantity"] for r in p["tables"][t]), f"data/{t}.csv"
what = _q(rng, f"the total quantity in data/{t}.csv", f"the sum of the quantity column of the {t} table")
elif src == "job":
jid = rng.choice(sorted(p["jobs"]))
key = rng.choice(["status", "owner", "duration"])
val, f1 = p["jobs"][jid][key], "data/jobs.txt"
what = f"the {key} of job-{jid}"
elif src == "const":
m = rng.choice(sorted(M))
key = rng.choice(sorted(M[m]["consts"]))
val, f1 = M[m]["consts"][key], f"src/{m}.py"
what = _q(rng, f"the value of {key} in the {m} module", f"{key} from src/{m}.py")
else:
c = rng.choice(sorted(p["components"]))
val, f1 = p["components"][c]["codename"], "docs/architecture.md"
what = f"the internal codename of the {c} component"
# worded unlike the held-out write_fact ("Find the owner ... and write just their name to notes/owner.txt")
q = _q(rng, f"Store {what} in {out}. Reply DONE once the file exists.",
f"Record {what} as the only content of {out}; then submit DONE.",
f"Look up {what}. Save that value alone to {out}, then answer DONE.")
return Task(kind, q, "DONE", F, dict(src=src, value=str(val), file=f1, out=out, what=what))
if kind == "service_hop":
m = rng.choice(sorted(M))
svc = re.search(r"helpers for the (\w+) service", F[f"src/{m}.py"]).group(1)
key = rng.choice(["port", "replicas", "timeout_seconds", "region"])
label = {"port": "port", "replicas": "replica count", "timeout_seconds": "timeout in seconds", "region": "region"}[key]
q = _q(rng, f"The {m} module contains helpers for one of the services. What is that service's {label}?",
f"Which {label} is configured for the service that src/{m}.py serves?",
f"src/{m}.py says which service it helps. Report that service's {label}.")
return Task(kind, q, str(S[svc][key]), F, dict(module=m, svc=svc, key=key, file=f"src/{m}.py", file2=S[svc]["file"]))
raise ValueError(kind)
_LABEL = {"port": "port", "timeout_seconds": "timeout (seconds)", "owner": "owner", "replicas": "replica count",
"region": "region", "status": "status", "duration": "duration in seconds", "codename": "codename",
"since": "year it was introduced", "quantity": "quantity", "price": "price"}
_PREFIX = ["", "", "Quick question: ", "I'm going through this project. ", "For a status report: ",
"Can you check something for me? ", "Using only the files here: "]
_SUFFIX = ["", "", " Answer with just the value.", " Look it up in the workspace.", " Thanks."]
def vary_question(task: Task, rng: random.Random, p: float = 0.7) -> Task:
"""Training-only rewording (RL and SFT data; the eval keeps make_task's questions): an unseen
core phrasing with probability p, plus optional framing, so the policy reads the request
instead of keying on a fixed template."""
m, k = task.meta, task.kind
L = lambda key: _LABEL.get(key, key)
bank = {
"config_value": lambda: [f"Tell me the configured {L(m['key'])} of {m['svc']}.",
f"I need the {L(m['key'])} setting for the {m['svc']} service.",
f"What does the config say the {L(m['key'])} of {m['svc']} is?",
f"{m['svc']}: what {L(m['key'])} is set?"],
"code_constant": lambda: [f"Open the {m['module']} source and report the number assigned to {m['const']}.",
f"{m['const']} is defined somewhere in {m['module']}. What's its value?",
f"What number is {m['const']} in src/{m['module']}.py?"],
"code_eval": lambda: [f"If you call {m['fn']} with {m['n']}, what comes back?",
f"Evaluate {m['fn']}({m['n']}) from the {m['module']} module.",
f"Run {m['fn']} on {m['n']} and tell me the result."],
"csv_lookup": lambda: [f"I need {m['name']}'s {m['col']} from the {m['table']} spreadsheet.",
f"Report the {m['col']} recorded for the row named {m['name']} in {m['table']}.",
f"In the {m['table']} data, what {m['col']} does {m['name']} have?"],
"doc_fact": lambda: [f"The architecture docs mention {m['component']}. What is its {L(m['key'])}?",
f"Check the documentation: {L(m['key'])} of the {m['component']} component?",
f"What do the docs list as the {L(m['key'])} for {m['component']}?"],
"edit": lambda: [f"Update {m['svc']}'s config so its timeout_seconds is {m['new']}. Submit DONE afterwards.",
f"{m['svc']} needs a timeout of {m['new']} seconds; change the config and answer DONE."],
"fix_test": lambda: [f"{m['test']} is failing. Repair the code under test (leave the test alone) and submit DONE.",
f"Get {m['test']} passing by fixing src/{m['module']}.py, then answer DONE."],
"long_file": lambda: [f"There is a record for job-{m['job']} in the jobs list. Report its {L(m['key'])}.",
f"For job-{m['job']}, what {L(m['key'])} does data/jobs.txt record?",
f"job-{m['job']}: what is its {L(m['key'])}?"],
"code_search": lambda: [f"In which source file is `{m['fn']}` implemented?",
f"Locate the implementation of {m['fn']} and give me the path.",
f"Which .py file has the def for {m['fn']}?"],
"aggregate": lambda: [f"In the {m['region']} region, how many services are there in total?",
f"Give me a count of services whose region is {m['region']}.",
f"Number of services configured with region {m['region']}?"],
"csv_sum": lambda: [f"Add up every quantity in data/{m['table']}.csv. What's the total?",
f"What do the quantities in the {m['table']} table sum to?"],
"service_hop": lambda: [f"src/{m['module']}.py helps some service. What {L(m['key'])} does that service have?",
f"Find the service the {m['module']} module is written for, then give its {L(m['key'])}."],
}.get(k)
q = task.question
if bank and rng.random() < p:
q = rng.choice(bank())
if rng.random() < 0.5:
pre = rng.choice(_PREFIX)
suf = rng.choice(_SUFFIX if task.answer not in ("DONE", "NOT_FOUND") else _SUFFIX[:2] + _SUFFIX[3:])
q = pre + q + suf
return Task(task.kind, q, task.answer, task.files, task.meta)
def normalize(s: str) -> str:
s = str(s).strip().strip(".").strip('"').strip("'").strip("`")
if s.startswith("/work/"):
s = s[len("/work/"):]
if s.startswith("./"):
s = s[2:]
return re.sub(r"\s+", " ", s).lower()
def _read(ws, rel):
try:
return open(f"{ws.root}/{rel}").read()
except OSError:
return None
def check(task: Task, submitted: str | None, ws: Workspace | None = None) -> bool:
if submitted is None:
return False
if task.kind in DONE_KINDS:
if ws is None or normalize(submitted) != "done":
return False
m = task.meta
if task.kind == "edit":
text = _read(ws, m["file"])
if text is None:
return False
try:
got, orig = parse_config(m["fmt"], text), parse_config(m["fmt"], task.files[m["file"]])
except Exception:
return False
orig["timeout_seconds"] = str(m["new"])
return got == orig
if task.kind == "fix_test":
if _read(ws, m["test"]) != task.files[m["test"]]:
return False # editing the test does not count
out = ws.tool_bash(f"python3 {m['test']}")
return out.strip().endswith("OK")
if task.kind == "write_fact":
text = _read(ws, m["out"])
return text is not None and normalize(text) == normalize(m["owner"])
if task.kind == "save_value":
text = _read(ws, m["out"])
return text is not None and normalize(text) == normalize(m["value"])
return normalize(submitted) == normalize(task.answer)
# ---------------------------------------------------------------- oracle (scripted solver)
def _call(name, **args):
return {"name": name, "arguments": args}
def _wrong_path(path, rng):
stem, _, ext = path.rpartition(".")
if path.startswith("config/"):
return f"{stem}.{rng.choice([f for f in FMTS + ['yml'] if f != ext])}"
return {"py": f"{stem}_util.py", "log": f"{stem}.txt", "csv": f"{stem}.tsv"}.get(ext, "notes/" + path.split("/")[-1])
def oracle_plan(task: Task, rng: random.Random) -> list[tuple[str, list[dict]]]:
"""A list of (thought, calls) steps; the final step submits. Results are filled in by running."""
k, m = task.kind, task.meta
steps: list[tuple[str, list[dict]]] = []
if rng.random() < 0.15 and "file" in m and k not in ("fix_test", "long_file"):
# a wrong first guess, so the data also shows recovering from a tool error
steps.append(("Let me try the likely file directly.", [_call("read", path=_wrong_path(m["file"], rng))]))
steps.append(("That path does not exist. List the files to find the right one.",
[_call("bash", command=rng.choice(["find . -type f | sort", "ls -R"]))]))
elif rng.random() < 0.4:
steps.append(("Let me see what is in the workspace first.",
[_call("bash", command=rng.choice(["ls -R", "ls config src docs logs data", "find . -type f | sort",
"ls; cat README.md"]))]))
if k == "config_value":
steps.append(rng.choice([
(f"The setting for {m['svc']} should be in its config file.", [_call("read", path=m["file"])]),
(f"I'll grep the config for {m['key']}.", [_call("bash", command=f"grep -in {m['key']} {m['file']}")]),
(f"I don't know the file name, so search the configs for {m['svc']}.",
[_call("bash", command=f"grep -rl {m['svc']} config")]),
]))
if steps[-1][1][0]["arguments"].get("command", "").startswith("grep -rl"):
steps.append((f"Its config is {m['file']}.", [_call("read", path=m["file"])]))
elif k == "code_constant":
steps.append((f"{m['const']} should be defined near the top of {m['file']}.",
[rng.choice([_call("read", path=m["file"]), _call("bash", command=f"grep -n {m['const']} {m['file']}")])]))
elif k == "code_eval":
steps.append((f"I need the definition of {m['fn']}.", [_call("read", path=m["file"])]))
steps.append(("Running it is more reliable than computing by hand.",
[_call("bash", command=f"python3 -c 'from src.{m['module']} import {m['fn']}; print({m['fn']}({m['n']}))'")]))
elif k == "log_count":
steps.append(("Count matching lines with grep -c.",
[_call("bash", command=f"grep -c '{m['level']} code={m['code']}' {m['file']}")]))
elif k == "csv_lookup":
steps.append((f"Find the row for {m['name']} and the header to know which column is which.",
[_call("bash", command=f"head -1 {m['file']}; grep -n '{m['name']}' {m['file']}")]))
elif k == "doc_fact":
steps.append((f"Search the docs for {m['component']}.", [_call("bash", command=f"grep -n '{m['component']}' {m['file']}")]))
elif k == "multi_hop":
steps.append((f"First find which host {m['svc']} uses as backup.", [_call("read", path=m["file"])]))
steps.append((f"The backup host is {m['backup']}; now read its config.", [_call("read", path=m["file2"])]))
elif k == "not_found":
g = m["ghost"]
steps.append((f"Search everything for {g}.", [_call("bash", command=f"grep -rn '{g}' . | head")]))
elif k == "edit":
steps.append(("Read the config before editing it.", [_call("read", path=m["file"])]))
steps.append(("Replace the timeout value.",
[_call("edit", path=m["file"], old_string=config_line(m["fmt"], "timeout_seconds", m["old"]),
new_string=config_line(m["fmt"], "timeout_seconds", m["new"]))]))
steps.append(("Check the change.", [_call("bash", command=f"grep -in timeout {m['file']}")]))
elif k == "fix_test":
steps.append(("Run the test to see the failure.", [_call("bash", command=f"python3 {m['test']}")]))
steps.append(("Read the code under test.", [_call("read", path=m["file"])]))
steps.append(("The failing assertion points at this line; fix it.",
[_call("edit", path=m["file"], old_string=m["bug_new"], new_string=m["bug_old"])]))
steps.append(("Run the test again.", [_call("bash", command=f"python3 {m['test']}")]))
elif k == "long_file":
if rng.random() < 0.6:
steps.append((f"data/jobs.txt is long, so grep for job-{m['job']}.",
[_call("bash", command=f"grep -n 'job-{m['job']} ' {m['file']}")]))
else:
steps.append(("The file is long. Find the line number first.",
[_call("bash", command=f"grep -n 'job-{m['job']} ' {m['file']} | cut -d: -f1")]))
# the line number is only known after running; the oracle knows it from the file
line = next(i for i, l in enumerate(task.files[m["file"]].split("\n"), 1) if l.startswith(f"job-{m['job']} "))
steps.append(("Read around that line.", [_call("read", path=m["file"], offset=max(1, line - 2), limit=5)]))
elif k == "code_search":
steps.append((f"Search for the definition of {m['fn']}.", [_call("bash", command=f"grep -rn 'def {m['fn']}' .")]))
elif k == "aggregate":
steps.append(("Configs use different formats, so match the region key case-insensitively and count files.",
[_call("bash", command=f"grep -rilE 'region[^a-z]+{m['region']}' config | wc -l")]))
elif k == "csv_sum":
cmd = rng.choice([f"awk -F, 'NR>1 {{s+=$3}} END {{print s}}' {m['file']}",
f"python3 -c \"import csv; print(sum(int(r['quantity']) for r in csv.DictReader(open('{m['file']}'))))\""])
steps.append(("Sum the quantity column with a command instead of by hand.", [_call("bash", command=cmd)]))
elif k == "save_value":
look = {"csv_sum": (f"Sum the quantity column with a command.",
_call("bash", command=f"awk -F, 'NR>1 {{s+=$3}} END {{print s}}' {m['file']}")),
"job": ("data/jobs.txt is long, so grep for the job.",
_call("bash", command=f"grep -n '{m['what'].split()[-1]} ' {m['file']}")),
"const": (f"Read the module to find the constant.", _call("read", path=m["file"])),
"codename": ("Search the docs for the component.",
_call("bash", command=f"grep -n '{m['what'].split()[-2]}' {m['file']}"))}[m["src"]]
steps.append((look[0], [look[1]]))
steps.append((f"The value is {m['value']}. Write it to {m['out']}.", [_call("write", path=m["out"], content=m["value"] + "\n")]))
steps.append(("Check the file.", [_call("bash", command=f"cat {m['out']}")]))
elif k == "service_hop":
steps.append((f"The module docstring names the service it helps.", [_call("bash", command=f"head -3 {m['file']}")]))
steps.append((f"It helps {m['svc']}. Now read that service's config.", [_call("read", path=m["file2"])]))
elif k == "write_fact":
steps.append((f"Look up the owner of {m['svc']}.", [_call("read", path=m["file"])]))
steps.append(("Write the name to the file.", [_call("write", path=m["out"], content=m["owner"] + "\n")]))
final = {"not_found": f"Nothing in the workspace mentions {m.get('ghost')}, so the answer is not available.",
"edit": "The file now has the new timeout.", "fix_test": "The test passes now.",
"write_fact": "The file has been written.", "save_value": "The file has the value."}.get(k, f"The answer is {task.answer}.")
steps.append((final, [_call("submit", answer=task.answer)]))
return steps
def run_trajectory(task: Task, steps, ws: Workspace) -> list[dict]:
msgs = task.messages()
for thought, calls in steps:
msgs.append({"role": "assistant", "think": thought, "content": "", "tool_calls": calls})
msgs.append({"role": "tool", "results": [ws.call(c["name"], c["arguments"]) for c in calls]})
if ws.submitted is not None:
break
return msgs
def oracle_trajectory(task: Task, rng: random.Random) -> tuple[list[dict], bool]:
with Workspace(task.files) as ws:
msgs = run_trajectory(task, oracle_plan(task, rng), ws)
return msgs, check(task, ws.submitted, ws)
_GUTTER = re.compile(r"^\s*\d+\t", re.M) # read's line-number column
def _seen(text: str, messages: list[dict]) -> bool:
ans = normalize(text)
if not ans:
return False
if re.fullmatch(r"-?\d+(\.\d+)?", ans): # numbers may carry units: "4751s", "12ms"
pat = re.compile(r"(?<!\d)(?<!\d\.)" + re.escape(ans) + r"(?!\d|\.\d)")
else:
pat = re.compile(r"(?<![\w.])" + re.escape(ans) + r"(?![\w])")
return any(pat.search(normalize(_GUTTER.sub("", r)))
for m in messages if m["role"] == "tool" for r in m["results"])
def invented(task: Task, submitted, messages: list[dict]) -> bool:
"""A wrong submitted answer that appears in no tool result: recalled or made up instead of read
(DONE/NOT_FOUND exempt). The thing this model must not do."""
if submitted is None or normalize(submitted) in ("done", "not_found"):
return False
return not _seen(str(submitted), messages)
def grounded(task: Task, messages: list[dict]) -> bool:
"""Is the submitted answer visible in some tool result the agent saw? (NOT_FOUND/DONE exempt.)
Line numbers from read are stripped and matching is on token boundaries, so an answer like
"5" does not count as seen just because some file has a line 5."""
if task.kind == "not_found" or task.kind in DONE_KINDS:
return True
return _seen(task.answer, messages)
if __name__ == "__main__":
rng = random.Random(0)
for kind in KINDS:
t = make_task(rng, kind)
msgs, ok = oracle_trajectory(t, rng)
print(f"== {kind}: {t.question} -> {t.answer} oracle_ok={ok} grounded={grounded(t, msgs)}")