"""Check retrieval before you blame the model.

For each question we know which phrase must be in the context for any model
to answer correctly. We retrieve the top-k chunks and ask one question:
is that phrase in there?

  python retrieval_check.py --strategy fixed
  python retrieval_check.py --strategy sections --k 3
  python retrieval_check.py --strategy fixed --show q1     # print the top-5 chunks
  python retrieval_check.py --strategy hybrid              # vector + keyword, merged
"""
import argparse
import re

import psycopg
from pgvector.psycopg import register_vector
from sentence_transformers import SentenceTransformer

from corpus import DOCS, QUESTIONS

MODEL = "sentence-transformers/all-MiniLM-L6-v2"  # 384 dimensions
DSN = "postgresql://postgres:lab@localhost:5544/postgres"


def norm(s: str) -> str:
    return re.sub(r"\s+", " ", s).strip().lower()


def chunk_fixed(doc, size=300):
    """The naive baseline: cut every `size` characters, ignore structure."""
    t = doc["text"]
    return [t[i : i + size] for i in range(0, len(t), size)]


def chunk_sections(doc):
    """Cut at headings; put the doc title and heading in front of each chunk.
    A doc with no headings (the error-code list) stays as one chunk."""
    text = doc["text"].strip()
    if not text.startswith("## "):
        return [f"{doc['title']}\n{text}"]
    out = []
    for part in re.split(r"(?m)^## ", text):
        part = part.strip()
        if part:
            heading, body = part.split("\n", 1)
            out.append(f"{doc['title']} > {heading}\n{body.strip()}")
    return out


CHUNKERS = {"fixed": chunk_fixed, "sections": chunk_sections, "hybrid": chunk_sections}


def build(conn, model, strategy):
    conn.execute("CREATE EXTENSION IF NOT EXISTS vector")
    conn.execute("DROP TABLE IF EXISTS chunks")
    conn.execute(
        """CREATE TABLE chunks (
             id serial PRIMARY KEY, strategy text, doc_id text, ord int,
             content text, embedding vector(384),
             tsv tsvector GENERATED ALWAYS AS (to_tsvector('simple', content)) STORED)"""
    )
    register_vector(conn)
    rows = []
    for d in DOCS:
        for i, c in enumerate(CHUNKERS[strategy](d)):
            rows.append((strategy, d["id"], i, c))
    vecs = model.encode([r[3] for r in rows], normalize_embeddings=True)
    for r, v in zip(rows, vecs):
        conn.execute(
            "INSERT INTO chunks (strategy, doc_id, ord, content, embedding) VALUES (%s,%s,%s,%s,%s)",
            (*r, v),
        )
    conn.commit()
    return len(rows)


def vector_top(conn, model, q, k):
    qv = model.encode([q], normalize_embeddings=True)[0]
    return conn.execute(
        "SELECT id, doc_id, content, embedding <=> %s AS dist FROM chunks ORDER BY dist LIMIT %s",
        (qv, k),
    ).fetchall()


def keyword_top(conn, q, k):
    # match any token of the question; rank by ts_rank. Crude on purpose.
    terms = " | ".join(re.findall(r"[A-Za-z0-9\-]+", q))
    return conn.execute(
        """SELECT id, doc_id, content, 0.0 AS dist
           FROM chunks, to_tsquery('simple', %s) query
           WHERE tsv @@ query ORDER BY ts_rank(tsv, query) DESC LIMIT %s""",
        (terms, k),
    ).fetchall()


def hybrid_top(conn, model, q, k, pool=20):
    """Reciprocal rank fusion of the vector list and the keyword list."""
    score, keep = {}, {}
    for lst in (vector_top(conn, model, q, pool), keyword_top(conn, q, pool)):
        for rank, row in enumerate(lst, 1):
            score[row[0]] = score.get(row[0], 0) + 1 / (60 + rank)
            keep[row[0]] = row
    return [keep[i] for i in sorted(score, key=score.get, reverse=True)[:k]]


def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("--strategy", choices=CHUNKERS, default="fixed")
    ap.add_argument("--k", type=int, default=5)
    ap.add_argument("--show", help="print the top-5 chunks for one question, e.g. q1")
    a = ap.parse_args()

    model = SentenceTransformer(MODEL)
    with psycopg.connect(DSN) as conn:
        n = build(conn, model, a.strategy)
        search = (lambda q, k: hybrid_top(conn, model, q, k)) if a.strategy == "hybrid" else (lambda q, k: vector_top(conn, model, q, k))
        print(f"strategy={a.strategy}  chunks={n}  model={MODEL}\n")

        if a.show:
            qid, q, phrase = next(x for x in QUESTIONS if x[0] == a.show)
            print(f"Q: {q}\nNeeded in context: {phrase!r}\n")
            for rank, (_, doc_id, content, dist) in enumerate(search(q, 5), 1):
                mark = "✓" if phrase in norm(content) else " "
                print(f"[{rank}] {mark} {doc_id}  (distance {dist:.3f})\n    " + content.replace("\n", "\n    ") + "\n")
            return

        hits = {1: 0, 3: 0, 5: 0}
        print(f"{'q':<4}{'first rank with the answer':<28}question")
        for qid, q, phrase in QUESTIONS:
            res = search(q, 10)
            rank = next((i for i, r in enumerate(res, 1) if phrase in norm(r[2])), None)
            for kk in hits:
                if rank and rank <= kk:
                    hits[kk] += 1
            print(f"{qid:<4}{(str(rank) if rank else 'not in top 10'):<28}{q}")
        t = len(QUESTIONS)
        print(f"\nanswer in top-1: {hits[1]}/{t}   top-3: {hits[3]}/{t}   top-5: {hits[5]}/{t}")


if __name__ == "__main__":
    main()
