File size: 10,545 Bytes
8b8e59d
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
"""SOP-driven procedural agent for TinyLiquid.

The model does not freewheel. It is given a procedure (an AGENTS.md-style SOP
from research/sop_library), works the case with RETRIEVE/READ/NOTE actions
against the local library (the "room"), and the loop enforces the plan,
external ledger, max steps, and constrained final decoding. Output is a JSON
report that audits which procedure steps were actually completed.

This is the on-device analog of the Codex loop: durable procedure text in the
prompt (AGENTS.md analog), an explicit step plan ("task bar"), a tool loop,
and guardrails.

Usage:
  .venv/bin/python research/agent.py --case "Verify: the bridge was painted in 2019."
  .venv/bin/python research/agent.py --case "..." --sop claim_verification --ckpt ckpt/distill
  .venv/bin/python research/agent.py --list-sops
"""

import argparse
import json
import re
import sys
from contextlib import nullcontext
from pathlib import Path

import torch

from model.config import TinyLiquidConfig, CONFIGS
from model.utils import latest_ckpt
from model.tiny_liquid import TinyLiquid
from data.tokenizer import load_tokenizer
from research.structured import analyst_report, _decode_phrase
from research import websearch as ws
from research.room import build_index, hit_text, read_doc

ROOT = Path(__file__).resolve().parents[1]
SOP_DIR = ROOT / "research" / "sop_library"

ACTIONS = ["RETRIEVE", "READ", "NOTE", "VERDICT", "WEB"]
MAX_STEPS = 6

SOP_ALIASES = {
    "claim_verification": ["claim", "verify", "check", "fact", "true", "false", "accurate"],
    "cross_source_discrepancy": ["discrepancy", "disagree", "contradict", "two accounts", "conflict", "differ"],
    "pattern_finding": ["pattern", "cluster", "common cause", "recurring", "trend"],
    "timeline_reconstruction": ["timeline", "sequence", "when did", "chronolog", "order of events"],
    "historical_truth": ["history", "past news", "earlier", "later record", "old report", "retraction", "what was hidden"],
    "politics_analysis": ["politics", "politician", "spin", "party", "talking point", "campaign"],
    "dark_web_research": ["dark web", "onion", "deep web", "clearnet", "leak", "forum"],
    "terminal_control": ["terminal", "shell", "command", "directory", "files", "download", "fetch", "search the corpus"],
    "source_triage": ["source", "credibility", "corroborat", "reliable", "weight", "provenance"],
}


def list_sops():
    print("Available procedures (research/sop_library):")
    for p in sorted(SOP_DIR.glob("*.md")):
        if p.stem == "00_common":
            continue
        tag = f"SOP {p.stem}"
        first = next((l for l in p.read_text(encoding="utf-8").splitlines() if l.strip()), "")
        print(f"  {tag:42s} {first}")


def load_sop(name: str | None, task: str) -> str:
    """Pick the procedure text: explicit --sop wins, else keyword match."""
    if name:
        path = SOP_DIR / f"{name}.md"
        if not path.exists():
            raise SystemExit(f"unknown SOP '{name}'; run --list-sops")
        return path.read_text(encoding="utf-8").strip()

    scored = {}
    low = task.lower()
    for stem, kws in SOP_ALIASES.items():
        scored[stem] = sum(1 for kw in kws if kw in low)
    best = max(scored, key=scored.get)
    if scored[best] == 0:
        best = "claim_verification"
    common = (SOP_DIR / "00_common.md").read_text(encoding="utf-8").strip()
    proc = (SOP_DIR / f"{best}.md").read_text(encoding="utf-8").strip()
    return f"{common}\n\n{proc}\n\nTASK: {task}"


def build_prompt(sop_text: str) -> str:
    return (
        "Work the case under the procedure below. Reply with exactly one line: "
        "ACTION: <RETRIEVE|READ|NOTE|VERDICT|WEB> then ARG: <text>. "
        "RETRIEVE <query> searches the library. READ <key> opens a document. "
        "NOTE <text> records a finding. VERDICT ends the case.\n"
        "PROCEDURE:\n" + sop_text
    )


def make_ctx(prompt, ledger, hits):
    ctx = ("You are working a research case. Keep the case file updated.\n"
           f"CASE FILE:\n{'\n'.join(f'[{i+1}] {e}' for i, e in enumerate(ledger[-8:])) or '(empty)'}\n")
    if hits:
        ctx += "LIBRARY HITS:\n" + hits + "\n"
    return ctx + f"TASK: {prompt}"


def _gen_arg(model, tok, ids, max_new=60, lock=None):
    """Generate an action argument; retry once with more heat if degenerate."""
    with lock or nullcontext():
        raw = tok.decode(model.generate(tok, ids, persona_id=1, max_new=max_new,
                                        temperature=0.5, top_k=40,
                                        repetition_penalty=1.5,
                                        no_repeat_ngram_size=4)[len(ids):]).strip()
        words = re.findall(r"[a-z']+", raw.lower())
        if len(words) >= 8 and len(set(words)) / len(words) < 0.25:
            raw = tok.decode(model.generate(tok, ids, persona_id=1, max_new=max_new,
                                            temperature=0.9, top_k=60,
                                            repetition_penalty=1.6,
                                            no_repeat_ngram_size=4)[len(ids):]).strip()
    return raw[:240]


def run_case(model, tok, task, idx, sop_text, max_steps=MAX_STEPS, lock=None, angle=None):
    """Run one SOP agent loop. `angle` narrows the worker's lens; `lock`
    serializes model inference so several workers can share one brain while
    network/dark-web retrieval still runs in parallel."""
    if angle:
        sop_text = f"{sop_text}\n\nAGENT ANGLE: {angle}"
    ledger, plan = [], []
    prompt = build_prompt(sop_text)
    web_hits = ""
    for step in range(max_steps):
        hits = ""
        if web_hits:
            hits = web_hits
        if ledger:
            last = ledger[-1]
            if last.startswith("RETRIEVE:"):
                hits = hit_text(idx, last.split(":", 1)[1].strip())
        ctx = make_ctx(prompt, ledger, hits)
        ids = tok.encode("<|analyst|><|user|>" + ctx + "<|assistant|>ACTION:").ids
        pre = len(ids)
        with lock or nullcontext():
            ids = _decode_phrase(model, tok, ids, 1, ACTIONS)
        action = tok.decode(ids[pre:]).strip().upper()
        if action not in ACTIONS:
            action = "NOTE"
        ids = ids + tok.encode(" ARG:").ids
        arg = _gen_arg(model, tok, ids, lock=lock)
        if action == "WEB" and arg:
            try:
                res = ws.pull(arg, n=1, library_dir="data/library")
                if res["saved"]:
                    idx = build_index("data/library")
                    web_hits = hit_text(idx, arg, k=2)
                else:
                    web_hits = "WEB_ERROR: " + (res["errors"][0]["err"] if res["errors"] else "no results")
            except Exception as e:
                web_hits = "WEB_ERROR: " + str(e)[:160]
        ledger.append(f"{action}: {arg}")
        plan.append({"step": step + 1, "action": action, "arg": arg})
        print(f"  [{step+1}] {action}: {arg}", flush=True)
        if action == "VERDICT":
            break
    return plan, ledger


def audit_sop(sop_text: str, ledger, report) -> list:
    """Report which numbered SOP steps have evidence in the work product."""
    steps = []
    for line in sop_text.splitlines():
        m = re.match(r"^(\d+)\.\s*([A-Z][A-Z ]{2,})", line.strip())
        if not m:
            continue
        num, label = m.group(1), m.group(2).strip()
        key = label.split(" ")[0].lower()
        blob = " ".join(ledger + [report.get("scratchpad", ""), report.get("reasoning", "")]).lower()
        covered = key in blob
        steps.append({"step": num, "label": label, "covered": covered})
    return steps


def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("--case", default=None)
    ap.add_argument("--sop", default=None, help="procedure stem, e.g. claim_verification")
    ap.add_argument("--list-sops", action="store_true")
    ap.add_argument("--ckpt", default="ckpt/distill")
    ap.add_argument("--tok", default="data/tokenizer.json")
    ap.add_argument("--library", default="data/library")
    ap.add_argument("--max-new", type=int, default=200)
    ap.add_argument("--threads", type=int, default=8)
    args = ap.parse_args()

    if args.list_sops:
        list_sops()
        return

    task = args.case or sys.stdin.read().strip()
    assert task, "no case provided (--case or stdin)"

    torch.set_num_threads(args.threads)
    tok = load_tokenizer(args.tok)
    ckpt = latest_ckpt(args.ckpt)
    assert ckpt, f"no checkpoints in {args.ckpt}"
    sd = torch.load(ckpt, map_location="cpu")
    cfg = TinyLiquidConfig(vocab_size=tok.get_vocab_size(),
                           **{k: v for k, v in sd["config"].items() if k != "vocab_size"})
    model = TinyLiquid(cfg)
    model.load_state_dict(sd["model"])
    model.eval()
    print(f"loaded {ckpt} (step {sd.get('step', '?')})", flush=True)

    sop_text = load_sop(args.sop, task)
    used_sop = "unknown (matched: claim_verification)" if args.sop is None else args.sop
    if args.sop is None:
        low = task.lower()
        used_sop = max(SOP_ALIASES, key=lambda k: sum(1 for w in SOP_ALIASES[k] if w in low))

    idx = build_index(args.library)
    print(f"SOP in effect: {used_sop} | library docs: {len(idx.docs)}", flush=True)

    plan, ledger = run_case(model, tok, task, idx, sop_text, max_steps=MAX_STEPS)
    report = analyst_report(model, tok, task, persona_id=1,
                            max_scratch=args.max_new // 2, max_reason=args.max_new // 4)
    audit = audit_sop(sop_text, ledger, report)

    # skeptic pass over the analyst's final report
    skeptic_prompt = (
        "Act as the skeptic. The analyst reached this conclusion; attack it: "
        f"Claim: {task}\nConclusion: {report.get('verdict', '')} "
        f"{report.get('reasoning', '')}"
    )
    p_token = "<|skeptic|>"
    ids = tok.encode(p_token + "<|user|>" + skeptic_prompt + "<|assistant|>").ids
    skeptic = tok.decode(model.generate(tok, ids, persona_id=2, max_new=args.max_new // 2,
                                        temperature=0.6, top_k=40,
                                        repetition_penalty=1.4,
                                        no_repeat_ngram_size=4)[len(ids):]).strip()

    out = {
        "task": task,
        "sop": used_sop,
        "plan": plan,
        "steps_total": len(plan),
        "analyst": report,
        "skeptic": skeptic,
        "sop_audit": audit,
    }
    print("\n=== REPORT ===")
    print(json.dumps(out, indent=2, ensure_ascii=False))


if __name__ == "__main__":
    main()