#!/usr/bin/env python3 """Reconstruct full wire-format search_guidelines results for cds/patient_education. The captured production traces only journal a UI-summary of each KB search (`middleware_result` event payload: {"query", "hits_returned", "top_hit"}), never the literal {"hits": [...], "total_accumulated": N} payload the model actually received over the wire (tool_loop.py::_handle_search / material_loop.py's equivalent). See the "Training vs Production ReAct Loop" audit (2026-07-02) and prepare_lora_corpus.py's _kb_agent_loop_turns. This script re-runs each trace's recorded queries (in original order, same `k` the model requested) against the live who_msf_guidelines KB and replays the exact same accumulate-and-dedupe-by-frame_id logic tool_loop.py uses, so the reconstructed payload is faithful to how such a query would be answered today. It is NOT a byte-identical replay of what the model saw when the trace was originally captured -- the KB index may have changed since, and `k` for calls where the model omitted it defaults to 7 (the same default tool_loop.py/material_loop.py use, so this matches the *logic*, just not necessarily bit-for-bit against a KB snapshot that no longer exists). The KB daemon (who_msf_guidelines, port 4276) is only bound to 127.0.0.1 inside the TenaOS_v1 container (see kb_guidelines/daemon.py), not published to the host. Rather than reconfigure the running production container (risk to demo.tenaos.com, which must stay up), this script pipes a JSON manifest into a `docker exec -i` python3 one-liner that queries the daemon from inside the container's own network namespace, and reads the JSON result back over stdout. Usage: python3 reconstruct_kb_hits.py \ --normalized-root artifacts_v2/normalized \ --out artifacts_v2/kb_hits_reconstructed.json \ --container TenaOS_v1 """ from __future__ import annotations import argparse import json import subprocess import sys from pathlib import Path from typing import Any _WORKER_SCRIPT = r""" import json, sys, urllib.request from concurrent.futures import ThreadPoolExecutor def kb_search(query, k): k = min(max(int(k or 7), 1), 10) body = json.dumps({"query": query, "k": k, "search_mode": "rrf", "snippet_chars": 1200}).encode("utf-8") req = urllib.request.Request( "http://127.0.0.1:4276/search", data=body, headers={"Content-Type": "application/json"}, method="POST", ) try: with urllib.request.urlopen(req, timeout=20) as resp: payload = json.loads(resp.read().decode("utf-8")) return {"ok": True, "hits": payload.get("hits") or []} except Exception as exc: return {"ok": False, "error": str(exc)} def compact_hit(h): content = h.get("content") or h.get("snippet") or "" return { "title": (h.get("title") or "")[:150], "source": h.get("source", "WHO Guidelines"), "content_type": h.get("content_type", ""), "recommendation_strength": h.get("recommendation_strength"), "evidence_certainty": h.get("evidence_certainty"), "score": round(float(h.get("score") or 0.0), 4), "content": content[:1000], } def dedup_key(h): return h.get("frame_id") or h.get("uri") or json.dumps(h.get("title")) def process_trace(item): trace_id, calls = item all_hits = [] seen = set() results = [] for call in calls: resp = kb_search(call["query"], call.get("k")) if not resp["ok"]: results.append({"error": resp["error"]}) continue hits = resp["hits"] for h in hits: fid = dedup_key(h) if fid not in seen: seen.add(fid) all_hits.append(h) results.append({"hits": [compact_hit(h) for h in hits], "total_accumulated": len(all_hits)}) return trace_id, results manifest = json.load(sys.stdin) out = {} with ThreadPoolExecutor(max_workers=12) as pool: for trace_id, results in pool.map(process_trace, manifest.items()): out[trace_id] = results json.dump(out, sys.stdout) """ def collect_manifest(normalized_root: Path, kinds: list[str]) -> dict[str, list[dict[str, Any]]]: """Extract, per trace_id, the ordered list of genuine (non-empty, non-duplicate) search_guidelines calls -- mirroring the exact counting/dedup logic in prepare_lora_corpus.py::_kb_agent_loop_turns / tool_loop.py::KbAgentLoop.run, so ordinal N here lines up with ordinal N there.""" manifest: dict[str, list[dict[str, Any]]] = {} for kind in kinds: path = normalized_root / kind / f"{kind}_accepted.jsonl" if not path.exists(): print(f"WARNING: {path} missing, skipping {kind}", file=sys.stderr) continue with path.open() as handle: for line in handle: line = line.strip() if not line: continue record = json.loads(line) trace_id = str(record.get("trace_id") or "") trace = record.get("trace") if isinstance(record.get("trace"), dict) else {} events = trace.get("events") if isinstance(trace.get("events"), list) else [] calls: list[dict[str, Any]] = [] searched: set[str] = set() for event in events: if str(event.get("type") or "") != "model_tool_call": continue if str(event.get("title") or "") != "search_guidelines": continue payload = event.get("payload") if isinstance(event.get("payload"), dict) else {} arguments = payload.get("arguments") if isinstance(payload.get("arguments"), dict) else {} query = str(arguments.get("query") or "").strip() norm_query = " ".join(query.lower().split()) if not norm_query or norm_query in searched: continue searched.add(norm_query) calls.append({"query": query, "k": arguments.get("k")}) if calls: manifest[trace_id] = calls return manifest def run_worker(container: str, manifest: dict[str, list[dict[str, Any]]]) -> dict[str, list[dict[str, Any]]]: if not manifest: return {} proc = subprocess.run( ["docker", "exec", "-i", container, "python3", "-c", _WORKER_SCRIPT], input=json.dumps(manifest).encode("utf-8"), capture_output=True, timeout=3600, ) if proc.returncode != 0: raise RuntimeError(f"worker failed (exit {proc.returncode}): {proc.stderr.decode('utf-8', 'replace')[:2000]}") return json.loads(proc.stdout.decode("utf-8")) def main() -> None: parser = argparse.ArgumentParser(description=__doc__) parser.add_argument("--normalized-root", type=Path, required=True) parser.add_argument("--out", type=Path, required=True) parser.add_argument("--container", default="TenaOS_v1") parser.add_argument("--kinds", nargs="+", default=["cds", "patient_education"]) parser.add_argument( "--batch-size", type=int, default=250, help="Traces per docker-exec invocation, to keep any single call's blast radius small.", ) args = parser.parse_args() manifest = collect_manifest(args.normalized_root, args.kinds) total_calls = sum(len(v) for v in manifest.values()) print(f"Collected {len(manifest)} traces, {total_calls} genuine search_guidelines calls to reconstruct.") items = list(manifest.items()) out: dict[str, list[dict[str, Any]]] = {} error_traces = 0 for start in range(0, len(items), args.batch_size): batch = dict(items[start : start + args.batch_size]) result = run_worker(args.container, batch) out.update(result) done = min(start + args.batch_size, len(items)) print(f" {done}/{len(items)} traces reconstructed...") for trace_id, results in out.items(): if any("error" in r for r in results): error_traces += 1 args.out.parent.mkdir(parents=True, exist_ok=True) args.out.write_text(json.dumps(out), encoding="utf-8") print(f"Wrote {args.out} ({len(out)} traces, {error_traces} with at least one query error).") if __name__ == "__main__": main()