TenaOS / training_code /reconstruct_kb_hits.py
beza4588's picture
Add synthetic LoRA training corpus and scripts
fbacbff verified
Raw History Blame Contribute Delete
8.39 kB
#!/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()