"""Precision@k before and after reranking, on labels neither ranker produced.

Questions: query = a disease tag's name, relevant = the questions carrying it.
Sections:  query = an article's title,   relevant = that article's sections.
"""
import random, time, sys
from sqlalchemy import text
from app.database import SessionLocal
from app.services import search_service as ss

db = SessionLocal()
random.seed(11)
MODEL = sys.argv[1] if len(sys.argv) > 1 else None
if MODEL:
    from app.services import rerank_service
    rerank_service.rerank_model = lambda: MODEL
    print("model:", MODEL)
K = 3

def prec(ids, relevant, k=K):
    top = ids[:k]
    return sum(1 for i in top if i in relevant) / max(1, len(top))

def run(kind, cases):
    before_p = after_p = before_mrr = after_mrr = 0.0
    moved = improved = worsened = 0
    times = []
    for query, relevant in cases:
        ranked, _ = ss.hybrid_ids(db, query, kind, limit=200)
        if not ranked:
            continue
        t = time.perf_counter()
        new = ss.rerank_ids(db, query, kind, ranked)
        times.append((time.perf_counter() - t) * 1000)
        b, a = prec(ranked, relevant), prec(new, relevant)
        before_p += b; after_p += a
        if a > b: improved += 1
        elif a < b: worsened += 1
        if ranked[0] != new[0]: moved += 1
        for name, lst in (("b", ranked), ("a", new)):
            rank = next((i + 1 for i, rid in enumerate(lst[:10]) if rid in relevant), 0)
            score = 1 / rank if rank else 0.0
            if name == "b": before_mrr += score
            else: after_mrr += score
    n = len(cases)
    print(f"\n{kind}: {n} queries")
    print(f"  precision@{K}   {before_p/n:.3f} -> {after_p/n:.3f}")
    print(f"  MRR@10         {before_mrr/n:.3f} -> {after_mrr/n:.3f}")
    print(f"  top-1 moved    {moved}/{n};  better {improved}, worse {worsened}")
    if times:
        times.sort()
        print(f"  rerank ms      median {times[len(times)//2]:.0f}, p90 {times[int(len(times)*0.9)]:.0f}")

# --- questions, labelled by disease tag -------------------------------------
rows = db.execute(text("""
    SELECT t.name, array_agg(l.question_id) AS ids
    FROM question_tags t JOIN question_tag_links l ON l.tag_id = t.id
    WHERE t.type = 'disease'
    GROUP BY t.id, t.name HAVING count(*) BETWEEN 3 AND 40
""")).fetchall()
cases = [(r.name, set(r.ids)) for r in rows]
random.shuffle(cases)
run("question", cases[:60])

# --- sections, labelled by their article ------------------------------------
rows = db.execute(text("""
    SELECT a.title, array_agg(s.id) AS ids
    FROM articles a JOIN article_section_index s ON s.article_id = a.id
    GROUP BY a.id, a.title HAVING count(*) >= 3
""")).fetchall()
cases = [(r.title, set(r.ids)) for r in rows]
random.shuffle(cases)
run("article_section", cases[:60])
