"""Same labelled test as eval.py, against a remote embeddings endpoint."""
import json, os, sys, time
import numpy as np, httpx

ROWS = json.load(open("/bench/labelled.json"))[:150]
DOCS = [(r["question_text"] + " " + (r["options"] or ""))[:4000] for r in ROWS]
QUERIES = [r["explanation"][:600] for r in ROWS]
BASE, KEY, MODEL = os.environ["BASE"], os.environ["KEY"], os.environ["MODEL"]

def embed(texts, batch=16):
    out = []
    with httpx.Client(timeout=90) as c:
        for i in range(0, len(texts), batch):
            r = c.post(f"{BASE}/v1/embeddings",
                       headers={"Authorization": f"Bearer {KEY}"},
                       json={"model": MODEL, "input": texts[i:i + batch]})
            r.raise_for_status()
            out.extend(d["embedding"] for d in r.json()["data"])
            print(f"  {min(i+batch,len(texts))}/{len(texts)}", flush=True)
    return np.array(out, dtype=np.float32)

t0 = time.time(); docs = embed(DOCS); doc_s = time.time() - t0
t0 = time.time(); qs = embed(QUERIES); q_s = time.time() - t0
docs /= np.linalg.norm(docs, axis=1, keepdims=True)
qs /= np.linalg.norm(qs, axis=1, keepdims=True)
order = np.argsort(-(qs @ docs.T), axis=1)
ranks = np.array([np.where(order[i] == i)[0][0] + 1 for i in range(len(ROWS))])
print(f"\n{MODEL:28s} dim={docs.shape[1]:5d} R@1={(ranks==1).mean():.3f} "
      f"R@5={(ranks<=5).mean():.3f} MRR={(1/ranks).mean():.3f} n={len(ROWS)} "
      f"{doc_s/len(DOCS)*1000:.0f}ms/doc {q_s/len(QUERIES)*1000:.0f}ms/query")
