"""Acceptance filter for trajectory records (stage 2). python3 train/accept.py [--traj runs/traj] [--min-score 80] [--out runs/traj/accepted.jsonl] A record is accepted when: it exists and the run had no harness error and no setup failure; end reason is "report" (the model ended itself); the score reaches the threshold; no raw tool result shows a harness problem (RFC lock, session or HTTP error). Loop parts are trimmed: when a call repeats with the identical arguments and the identical result right after itself, only the first pair stays (the world did not change, so the later context is the same). Repair trajectories (a failed write followed by a good write) stay and are marked. Output: one JSON record per line (messages trimmed, metadata, reasons). """ import argparse import glob import json import os import re HARNESS_ERRORS = re.compile(r"Concurrent call detected|\[LOCK\]|HTTP (?:4|5)\d\d|Connection refused|timed out|" r"session (?:dropped|expired)|Traceback|RFC_(?:COMMUNICATION|SYSTEM)_FAILURE", re.I) ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) def tool_pairs(messages): """[(assistant index, [tool message indices])] for assistant messages with tool calls.""" out = [] for i, m in enumerate(messages): if m["role"] == "assistant" and m.get("tool_calls"): j = i + 1 while j < len(messages) and messages[j]["role"] == "tool": j += 1 out.append((i, list(range(i + 1, j)))) return out def call_key(messages, i, tools): calls = messages[i]["tool_calls"] sig = [(c["function"]["name"], json.dumps(_args(c), sort_keys=True)) for c in calls] return sig, [messages[t]["content"] for t in tools] def _args(c): a = c["function"].get("arguments") or "{}" try: return json.loads(a) if isinstance(a, str) else a except ValueError: return {"_raw": a} def trim_loops(messages): """Remove an assistant tool call + its results when it repeats the previous call and result exactly.""" drop, last_key, last_end, removed = set(), None, -1, 0 for i, tools in tool_pairs(messages): key = call_key(messages, i, tools) end = (tools[-1] if tools else i) + 1 if key == last_key and i == last_end: # directly after the same call with the same result drop.update([i] + tools) removed += 1 else: last_key = key last_end = end kept = [m for k, m in enumerate(messages) if k not in drop] return kept, removed def is_failed_write(content): c = content.replace(" ", "") return c.startswith("ERROR:") or '"success":false' in c def has_repair(messages): """A write tool result that failed, followed later by a write result that succeeded.""" failed = False names = {} for m in messages: if m["role"] == "assistant": for c in m.get("tool_calls") or []: names[c.get("id")] = c["function"]["name"] elif m["role"] == "tool" and re.match(r"(sap_push|sap_activate|write_source|activate)", names.get(m.get("tool_call_id"), "")): if is_failed_write(m["content"]): failed = True elif failed and '"success":true' in m["content"].replace(" ", ""): return True return False def judge(rec, summary, min_score): """Return (accepted, reasons).""" md = rec["metadata"] why = [] if summary.get("harness_error") or md.get("harness_error"): why.append("harness_error") if md.get("setup_failed"): why.append("setup_failed") if md.get("end_reason") != "report": why.append(f"end_reason={md.get('end_reason')}") if md.get("score") is None or md["score"] < min_score: why.append(f"score={md.get('score')}") bad = [r for r in rec["tool_results_raw"] if HARNESS_ERRORS.search(r.get("result") or "") and not r.get("tool", "").startswith("sap_run_unit")] if bad: why.append("harness_text_in_result") if not (rec["messages"][-1].get("content") or "").strip(): why.append("no_final_report") return not why, why def main(): ap = argparse.ArgumentParser() ap.add_argument("--traj", default=os.path.join(ROOT, "runs", "traj")) ap.add_argument("--min-score", type=float, default=80) ap.add_argument("--out") a = ap.parse_args() summary = {} sp = os.path.join(a.traj, "summary.jsonl") if os.path.exists(sp): for r in map(json.loads, open(sp)): summary[r.get("run_dir")] = r out = a.out or os.path.join(a.traj, "accepted.jsonl") n_all = n_ok = 0 rejected = {} with open(out, "w") as f: for p in sorted(glob.glob(os.path.join(a.traj, "*", "record.json"))): rec = json.load(open(p)) n_all += 1 ok, why = judge(rec, summary.get(os.path.basename(os.path.dirname(p)), {}), a.min_score) if not ok: for w in why: rejected[w.split("=")[0]] = rejected.get(w.split("=")[0], 0) + 1 continue msgs, removed = trim_loops(rec["messages"]) n_ok += 1 f.write(json.dumps({"id": f"{rec['task']['id']}_r{rec['run']}", "run_dir": os.path.dirname(p), "task": rec["task"], "teacher": rec["teacher"], "metadata": rec["metadata"], "messages": msgs, "tools": rec["tools"], "loop_pairs_removed": removed, "repair": has_repair(msgs)}) + "\n") print(f"{n_ok} accepted of {n_all} records -> {out}; rejected by reason: {rejected}") if __name__ == "__main__": main()