"""Rerun one hosted batch run's questions locally with the selection agent on.
Same chunks (hosted RDS, read only), same k, same declared-scalar answer path.
Records picks, final context, and the answer per question."""
import asyncio, json, os, sys, time
import psycopg
from quber.playground import retrieval
from quber.playground.retrieval import WINDOW, PARENT_CAP
from quber.playground.selection import select
from quber.playground.answers import expectation
from quber.playground.tracing import document_key

JOB = sys.argv[1]; OUT = sys.argv[2]; K = 10
conn = psycopg.connect(host=os.environ["POSTGRES_HOST"], user="quber", password=os.environ["POSTGRES_PASSWORD"], dbname="quber_rag")
job = conn.execute("select b.doc_id, d.doc_key, b.questions from ade_playground.batch_runs b join ade_playground.documents d on d.id=b.doc_id where b.job_id=%s", (JOB,)).fetchone()
doc_id, doc_key, questions = job
old = {r[0]: r[1] for r in conn.execute("select question, response from ade_playground.batch_rows where job_id=%s", (JOB,)).fetchall()}
print("doc", doc_key, "questions", len(questions), "stored rows", len(old), flush=True)

sem = asyncio.Semaphore(5)

async def one(i, q):
    async with sem:
        document_key.set(doc_key)
        t0 = time.time()
        window = retrieval._fused_window(doc_key, q, WINDOW, None)
        pool = retrieval._cap_line_records(window, PARENT_CAP)
        err = None
        try:
            picked = await select(q, pool)
        except Exception as exc:
            picked, err = [], repr(exc)
        # compose exactly as retrieve() does
        out, seen = [], set()
        for j in picked:
            c = pool[j]
            if c.chunk_id not in seen:
                seen.add(c.chunk_id); out.append(c)
            if len(out) >= K: break
        n_picks = len(out)
        for c in pool:
            if len(out) >= K: break
            if c.chunk_id not in seen:
                seen.add(c.chunk_id); out.append(c)
        ans = await expectation.answer(q, out, "value")
        rec = {
            "index": i, "question": q, "pool_size": len(pool), "picks": n_picks,
            "pick_fused_ranks": [j for j in picked][:K],
            "context_ids": [c.chunk_id for c in out],
            "context_chars": [len(c.content) for c in out],
            "shape": ans.payload.__class__.__name__.lower(),
            "payload": ans.payload.model_dump(),
            "cited_ids": ans.cited_ids,
            "old": old.get(q),
            "select_error": err, "seconds": round(time.time() - t0, 1),
        }
        print(f"[{i:2d}] picks={n_picks:2d} ranks={rec['pick_fused_ranks'][:5]} {rec['shape']:12s} {str(rec['payload'].get('value') or rec['payload'].get('reason'))[:40]!r} old={old.get(q,{}).get('shape')} {rec['seconds']}s", flush=True)
        return rec

async def main():
    recs = await asyncio.gather(*(one(i, q) for i, q in enumerate(questions)))
    json.dump({"job": JOB, "doc_key": doc_key, "doc_id": doc_id, "rows": recs}, open(OUT, "w"), indent=1, default=str)
    print("wrote", OUT)
asyncio.run(main())
