"""Cross-document comparison of hosted (fused order) vs local rerun (selection on)."""
import json, os, re, glob, psycopg, collections
S=os.path.dirname(os.path.abspath(__file__))
conn=psycopg.connect(host=os.environ["POSTGRES_HOST"],user="quber",password=os.environ["POSTGRES_PASSWORD"],dbname="quber_rag")
CELL=re.compile(r"^t(\d+)-(\d+)-(\d+)$")
def holders(ref):
    m=CELL.match(ref)
    return [f"t{m.group(1)}-line-{m.group(2)}", f"#/tables/{m.group(1)}"] if m else [ref]
def position(cited, ids):
    best=None
    for ref in cited:
        for h in holders(ref):
            if h in ids:
                p=ids.index(h); best=p if best is None or p<best else best
    return best
def norm(s): return re.sub(r"[\s$,]|million|thousand","",(s or "").lower()).replace("—","-").replace("–","-").replace("(","-").replace(")","")
def oldval(old): return old["value"].get("value") if isinstance(old.get("value"),dict) else old.get("answer")
def cell_text(doc_id, ref):
    r=conn.execute("select cell_text from ade_playground.groundings where document_id=%s and ref_id=%s",(doc_id,ref)).fetchone()
    return r[0] if r else None
titles={r[0]:r[1] for r in conn.execute("select doc_key, title from ade_playground.documents").fetchall()}
grand=collections.Counter(); grand_pos_old=collections.Counter(); grand_pos_new=collections.Counter(); gp=collections.Counter(); gdup=0; gq=0
print(f"{'document':32s} {'scalar old->new':16s} {'+ans':>4s} {'-ans':>4s} {'same':>4s} {'diff':>4s} | {'src@1 old/new':>14s} {'<=3 old/new':>12s} {'<=5 old/new':>12s} | picks>=10 {'dup':>4s} | picks-only chars")
for f in sorted(glob.glob(f"{S}/rerun_*.json")):
    d=json.load(open(f)); rows=d["rows"]; doc_id=d["doc_id"]; dk=d["doc_key"]
    trans=collections.Counter(); same=diff=0
    pos_old=collections.Counter(); pos_new=collections.Counter(); n_old=n_new=0
    beyond=0; dup=0; chars_full=chars_picks=0; picks=collections.Counter(); verified_new=0; verified_old=0
    for r in rows:
        old=r["old"]; trans[(old["shape"],r["shape"])]+=1
        picks[r["picks"]]+=1
        if any(j>=10 for j in r["pick_fused_ranks"]): beyond+=1
        pk=r["context_ids"][:r["picks"]]
        if any(re.match(r"^t(\d+)-line-",c) and f"#/tables/{re.match(r'^t(\d+)-line-',c).group(1)}" in pk for c in pk): dup+=1
        chars_full+=sum(r["context_chars"]); chars_picks+=sum(r["context_chars"][:max(r["picks"],1)])
        if old["shape"]=="scalar" and r["shape"]=="scalar":
            if norm(oldval(old))==norm(r["payload"].get("value")): same+=1
            else: diff+=1
        if old["shape"]=="scalar":
            p=position([x["ref_id"] for x in (old.get("references") or [])],[x["chunk_id"] for x in old["retrieved"]])
            if p is not None: pos_old[p]+=1; n_old+=1
            refs=old.get("references") or []
            if refs and refs[0].get("text") is not None and norm(refs[0]["text"])==norm(oldval(old)): verified_old+=1
        if r["shape"]=="scalar":
            p=position(r["cited_ids"], r["context_ids"])
            if p is not None: pos_new[p]+=1; n_new+=1
            c=r["cited_ids"][0] if r["cited_ids"] else None
            t=cell_text(doc_id,c) if c else None
            if t is not None and norm(t)==norm(r["payload"].get("value")): verified_new+=1
    def cum(pos,n,k): return f"{100*sum(v for p,v in pos.items() if p<k)/n:.0f}%" if n else "-"
    os_=sum(v for (a,b),v in trans.items() if a=="scalar"); ns=sum(v for (a,b),v in trans.items() if b=="scalar")
    print(f"{titles.get(dk,dk)[:32]:32s} {os_:>3d} -> {ns:<9d} {trans[('unanswerable','scalar')]:>4d} {trans[('scalar','unanswerable')]:>4d} {same:>4d} {diff:>4d} | {cum(pos_old,n_old,1):>6s}/{cum(pos_new,n_new,1):<6s} {cum(pos_old,n_old,3):>5s}/{cum(pos_new,n_new,3):<5s} {cum(pos_old,n_old,5):>5s}/{cum(pos_new,n_new,5):<5s} | {beyond:>5d}/93  {dup:>4d} | {100*chars_picks/chars_full:.0f}%   picks {dict(sorted(picks.items()))}  cell-verified old {verified_old}/{os_} new {verified_new}/{ns}")
    for k,v in trans.items(): grand[k]+=v
    for k,v in pos_old.items(): grand_pos_old[k]+=v
    for k,v in pos_new.items(): grand_pos_new[k]+=v
    for k,v in picks.items(): gp[k]+=v
    gdup+=dup; gq+=len(rows)
no=sum(grand_pos_old.values()); nn=sum(grand_pos_new.values())
print("\nALL DOCS: transitions", {f"{a}->{b}":v for (a,b),v in grand.items()})
print("cumulative source position, hosted fused vs rerun selection:")
co=cn=0
for k in range(10):
    co+=grand_pos_old.get(k,0); cn+=grand_pos_new.get(k,0)
    print(f"  within {k+1:2d}: hosted {100*co/no:3.0f}%   selection {100*cn/nn:3.0f}%")
print("picks per question:", dict(sorted(gp.items())), "| line item + own parent both picked:", gdup, "/", gq)
