"""Re-answer a rerun file's questions with a different answer model, same context."""
import asyncio, json, os, sys, time
import psycopg
import quber.playground.agent as agent_mod
from quber.playground.answers import declared
MODEL = sys.argv[1]; IN = sys.argv[2]; OUT = sys.argv[3]
agent_mod.DEFAULT_MODEL = MODEL; declared.DEFAULT_MODEL = MODEL
from quber.playground.answers import expectation
from quber.playground.retrieval import RetrievedChunk
from quber.playground.tracing import document_key
d = json.load(open(IN)); doc_id = d["doc_id"]; doc_key = d["doc_key"]
conn = psycopg.connect(host=os.environ["POSTGRES_HOST"], user="quber", password=os.environ["POSTGRES_PASSWORD"], dbname="quber_rag")
cache = {}
def chunk(cid):
    if cid not in cache:
        r = conn.execute("select chunk_id, chunk_type, page, content, parent_chunk_id from ade_playground.chunks where document_id=%s and chunk_id=%s", (doc_id, cid)).fetchone()
        cache[cid] = RetrievedChunk(chunk_id=r[0], chunk_type=r[1], page=r[2], content=r[3], score=0.0, parent_chunk_id=r[4])
    return cache[cid]
sem = asyncio.Semaphore(5)
async def one(r):
    async with sem:
        document_key.set(doc_key); t0 = time.time()
        chunks = [chunk(c) for c in r["context_ids"]]
        try:
            ans = await expectation.answer(r["question"], chunks, "value")
            out = {"shape": ans.payload.__class__.__name__.lower(), "payload": ans.payload.model_dump(), "cited_ids": ans.cited_ids, "error": None}
        except Exception as exc:
            out = {"shape": "error", "payload": {}, "cited_ids": [], "error": repr(exc)}
        out.update({"index": r["index"], "question": r["question"], "context_ids": r["context_ids"], "picks": r["picks"],
                    "opus": {"shape": r["shape"], "payload": r["payload"], "cited_ids": r["cited_ids"]}, "old": r["old"], "seconds": round(time.time() - t0, 1)})
        print(f"[{r['index']:2d}] {out['shape']:12s} {str(out['payload'].get('value') or out['payload'].get('reason') or out['error'])[:40]!r} opus={r['shape']}:{str(r['payload'].get('value'))[:20]!r}", flush=True)
        return out
async def main():
    recs = await asyncio.gather(*(one(r) for r in d["rows"]))
    json.dump({"model": MODEL, "doc_key": doc_key, "doc_id": doc_id, "rows": recs}, open(OUT, "w"), indent=1, default=str)
    print("wrote", OUT)
asyncio.run(main())
