"""BM25 + dense vectors fused with Reciprocal Rank Fusion, on the real FastAPI docs (151 Markdown files).
Chunker = heading-aware + heading-path prefix, the winning strategy from the chunking benchmark.
Embeddings = sentence-transformers/all-MiniLM-L6-v2 (384-dim), cosine via normalized dot product.
Two query sets:
  - "heading" queries (773): section headings, same harness as the chunking post. Tests paraphrase-ish recall.
  - "identifier" queries: inline-code identifiers (`response_model`, `BackgroundTasks`, ...) that occur in
    exactly one chunk in the corpus. Tests exact-token recall, BM25's classic strength.
RRF: score(d) = alpha / (k_rrf + rank_bm25(d)) + (1 - alpha) / (k_rrf + rank_vec(d)), k_rrf = 60.
alpha=1.0 is pure BM25, alpha=0.0 is pure vector. Swept over 0.0..1.0 in steps of 0.1.
"""
import re, glob, math, collections, statistics as st, time
import numpy as np
from sentence_transformers import SentenceTransformer

K_RRF = 60
CAND_DEPTH = 50  # how deep each retriever's ranking feeds the fusion

TOK = re.compile(r"[a-z0-9_]+")
def toks(s): return TOK.findall(s.lower())

HEADING = re.compile(r"^(#{1,3}) (.+?)\s*(\{.*\})?$", re.M)
def sections(text):
    ms = list(HEADING.finditer(text)); out = []; stack = []
    for i, m in enumerate(ms):
        lvl = len(m.group(1)); title = m.group(2).strip()
        stack = [s for s in stack if s[0] < lvl] + [(lvl, title)]
        end = ms[i + 1].start() if i + 1 < len(ms) else len(text)
        out.append((m.start(), end, lvl, title, " > ".join(t for _, t in stack)))
    return out

SIZE = 800
def paragraphs(text):
    out, start, cur = [], 0, 0
    for m in re.finditer(r"\n\s*\n", text):
        if m.end() - start > SIZE and cur > start: out.append((start, cur)); start = cur
        cur = m.end()
    out.append((start, len(text)))
    return [c for c in out if text[c[0]:c[1]].strip()]

def heading_aware_chunks(text):
    out = []
    secs = sections(text) or [(0, len(text), 0, '', '')]
    if secs[0][0] > 0: secs = [(0, secs[0][0], 0, '', '')] + secs
    for s, e, _, _, path in secs:
        for a, b in paragraphs(text[s:e]):
            out.append((s + a, s + b, path))
    return out

docs = {p: open(p, encoding='utf-8').read() for p in sorted(glob.glob('../corpus/*.md'))}

chunks = []  # (doc, start, end, path, text, embed_text)
for p, t in docs.items():
    for a, b, path in heading_aware_chunks(t):
        body = t[a:b]
        embed = f"{path}\n{body}" if path else body
        chunks.append((p, a, b, path, body, embed))
print(f"docs={len(docs)} chars={sum(map(len, docs.values())):,} chunks={len(chunks)}")

# --- BM25 index ---
tf = [collections.Counter(toks(c[5])) for c in chunks]
N = len(chunks)
df = collections.Counter(t for c in tf for t in c)
L = [sum(c.values()) for c in tf]
avg = sum(L) / N
idf = {t: math.log(1 + (N - n + .5) / (n + .5)) for t, n in df.items()}
inv = collections.defaultdict(list)
for i, c in enumerate(tf):
    for t, f in c.items(): inv[t].append((i, f))

def bm25_search(q, k=CAND_DEPTH):
    sc = collections.defaultdict(float)
    for t in set(toks(q)):
        for i, f in inv.get(t, ()):
            sc[i] += idf[t] * f * 2.5 / (f + 1.5 * (.25 + .75 * L[i] / avg))
    return [i for i, _ in sorted(sc.items(), key=lambda x: -x[1])[:k]]

# --- vector index ---
t0 = time.time()
model = SentenceTransformer('sentence-transformers/all-MiniLM-L6-v2')
print(f"model loaded in {time.time()-t0:.1f}s")
t0 = time.time()
emb = model.encode([c[5] for c in chunks], batch_size=64, show_progress_bar=False, convert_to_numpy=True)
emb = emb / np.linalg.norm(emb, axis=1, keepdims=True)
print(f"embedded {len(chunks)} chunks in {time.time()-t0:.1f}s")

def vec_search_batch(queries, k=CAND_DEPTH):
    q = model.encode(queries, show_progress_bar=False, convert_to_numpy=True)
    q = q / np.linalg.norm(q, axis=1, keepdims=True)
    sims = q @ emb.T
    return [list(np.argsort(-sims[i])[:k]) for i in range(len(queries))]

# --- heading queries (paraphrase-ish) ---
heading_queries = []
for p, t in docs.items():
    for s, e, lvl, title, _ in sections(t):
        if lvl >= 2 and e - s >= 300 and len(toks(title)) >= 2:
            heading_queries.append((p, s, e, title))
print(f"heading queries: {len(heading_queries)}")

def heading_hit(i, p, s, e):
    cp, ca, cb, *_ = chunks[i]
    if cp != p: return False
    return (min(e, cb) - max(s, ca)) >= 0.5 * (cb - ca)

# --- identifier queries (exact-token) ---
CODE = re.compile(r"`([A-Za-z_][A-Za-z0-9_.]{3,})`")
occ = collections.defaultdict(set)  # identifier -> set of chunk indices
for ci, c in enumerate(chunks):
    for m in CODE.finditer(c[4]):
        ident = m.group(1)
        if ('_' in ident or '.' in ident or ident.lower() != ident) :
            occ[ident].add(ci)
identifier_queries = [(ident, next(iter(idxs))) for ident, idxs in occ.items() if len(idxs) == 1]
print(f"identifier queries: {len(identifier_queries)}")

# --- paraphrase queries (hand-written, no verbatim overlap with the target section's heading) ---
# Each maps to one corpus file that is narrowly scoped to that topic. A hit = any chunk from that file.
PARAPHRASE = [
    ("How do I run code after sending the response to the client?", "tutorial_background-tasks.md"),
    ("How do I change the HTTP status code an endpoint returns without raising an exception?", "advanced_response-change-status-code.md"),
    ("How do I add custom headers to an HTTP response?", "advanced_response-headers.md"),
    ("How do I set a cookie on the response?", "advanced_response-cookies.md"),
    ("How do I accept a persistent two-way connection from a browser?", "advanced_websockets.md"),
    ("How do I run my API behind a reverse proxy that adds a path prefix?", "advanced_behind-a-proxy.md"),
    ("How can I load configuration from environment variables?", "advanced_settings.md"),
    ("How do I generate a TypeScript or Python client from my API automatically?", "advanced_generate-clients.md"),
    ("How can I mount one app inside another at a given path?", "advanced_sub-applications.md"),
    ("How do I run startup and shutdown code for my application?", "advanced_events.md"),
    ("How can I add logic that runs for every request, like a timing header?", "tutorial_middleware.md"),
    ("How do I run my app with multiple worker processes?", "deployment_server-workers.md"),
    ("How do I package my app into a container image?", "deployment_docker.md"),
    ("How do I serve my API over an encrypted connection?", "deployment_https.md"),
    ("How do I read extra parameters from the URL's query string?", "tutorial_query-params.md"),
    ("How do I capture a variable part of the URL path?", "tutorial_path-params.md"),
    ("How can a client upload a file to my API?", "tutorial_request-files.md"),
    ("How do I accept both form fields and an uploaded file in one request?", "tutorial_request-forms-and-files.md"),
    ("How do I return a custom error message and status code when something goes wrong?", "tutorial_handling-errors.md"),
    ("How can I restrict which fields are included in the JSON I return?", "tutorial_response-model.md"),
    ("How do I read a value sent in the request headers?", "tutorial_header-params.md"),
    ("How do I connect my API to a relational database?", "tutorial_sql-databases.md"),
    ("How can I organize a large API into multiple files with shared routing?", "tutorial_bigger-applications.md"),
    ("How do I add a summary and description to an endpoint for the interactive docs?", "tutorial_path-operation-configuration.md"),
    ("How can dependencies be shared and reused across many endpoints?", "tutorial_dependencies_index.md"),
    ("How do I let a browser on a different origin call my API?", "tutorial_cors.md"),
    ("How do I validate the shape of a JSON request body against a schema?", "tutorial_body.md"),
    ("How do I return different response types, like a file download?", "advanced_custom-response.md"),
    ("How can I require a logged-in user before an endpoint runs?", "tutorial_security_first-steps.md"),
    ("How can I make a numeric URL parameter require a minimum or maximum value?", "tutorial_path-params-numeric-validations.md"),
]
import os
paraphrase_queries = [(q, f) for q, f in PARAPHRASE]
print(f"paraphrase queries: {len(paraphrase_queries)}")

def run_eval(query_texts, hit_fn, bm25_ranks_list, vec_ranks_list, alpha):
    h1 = h5 = rr = 0
    n = len(query_texts)
    for qi in range(n):
        bm = bm25_ranks_list[qi]; ve = vec_ranks_list[qi]
        rank_bm = {i: r for r, i in enumerate(bm)}
        rank_ve = {i: r for r, i in enumerate(ve)}
        cands = set(bm) | set(ve)
        scored = []
        for i in cands:
            rb = rank_bm.get(i, CAND_DEPTH * 4)
            rv = rank_ve.get(i, CAND_DEPTH * 4)
            score = alpha / (K_RRF + rb + 1) + (1 - alpha) / (K_RRF + rv + 1)
            scored.append((score, i))
        top = [i for _, i in sorted(scored, key=lambda x: -x[0])[:5]]
        hits = [r for r, i in enumerate(top) if hit_fn(i, qi)]
        if hits:
            h5 += 1; rr += 1 / (hits[0] + 1); h1 += hits[0] == 0
    return h1 / n, h5 / n, rr / n

print("\nPre-computing BM25 and vector rankings for all query sets...")
h_bm25_ranks = [bm25_search(title) for _, _, _, title in heading_queries]
h_vec_ranks = vec_search_batch([title for _, _, _, title in heading_queries])
i_bm25_ranks = [bm25_search(ident) for ident, _ in identifier_queries]
i_vec_ranks = vec_search_batch([ident for ident, _ in identifier_queries])
p_bm25_ranks = [bm25_search(q) for q, _ in paraphrase_queries]
p_vec_ranks = vec_search_batch([q for q, _ in paraphrase_queries])

def h_hit_fn(i, qi):
    p, s, e, _ = heading_queries[qi]
    return heading_hit(i, p, s, e)

def i_hit_fn(i, qi):
    return i == identifier_queries[qi][1]

def p_hit_fn(i, qi):
    return os.path.basename(chunks[i][0]) == paraphrase_queries[qi][1]

print(f"\n{'alpha':>6} | {'heading MRR':>11} | {'ident MRR':>9} | {'paraphrase hit@1':>16} {'hit@5':>6} {'MRR':>6}")
rows = []
for alpha in [round(x * 0.1, 1) for x in range(11)]:
    hh1, hh5, hrr = run_eval([q[3] for q in heading_queries], h_hit_fn, h_bm25_ranks, h_vec_ranks, alpha)
    ih1, ih5, irr = run_eval([q[0] for q in identifier_queries], i_hit_fn, i_bm25_ranks, i_vec_ranks, alpha)
    ph1, ph5, prr = run_eval([q for q, _ in paraphrase_queries], p_hit_fn, p_bm25_ranks, p_vec_ranks, alpha)
    rows.append((alpha, hh1, hh5, hrr, ih1, ih5, irr, ph1, ph5, prr))
    print(f"{alpha:6.1f} | {hrr:11.3f} | {irr:9.3f} | {ph1:16.3f} {ph5:6.3f} {prr:6.3f}")

print(f"\n{'alpha':>6} | {'heading hit@1':>13} {'hit@5':>6} {'MRR':>6} | {'ident hit@1':>11} {'hit@5':>6} {'MRR':>6} | {'para hit@1':>10} {'hit@5':>6} {'MRR':>6}")
for r in rows:
    alpha, hh1, hh5, hrr, ih1, ih5, irr, ph1, ph5, prr = r
    print(f"{alpha:6.1f} | {hh1:13.3f} {hh5:6.3f} {hrr:6.3f} | {ih1:11.3f} {ih5:6.3f} {irr:6.3f} | {ph1:10.3f} {ph5:6.3f} {prr:6.3f}")

best_h = max(rows, key=lambda r: r[3])
best_i = max(rows, key=lambda r: r[6])
best_p = max(rows, key=lambda r: r[9])
best_sum = max(rows, key=lambda r: r[3] + r[6] + r[9])
print(f"\nbest alpha for heading MRR: {best_h[0]} ({best_h[3]:.3f})")
print(f"best alpha for identifier MRR: {best_i[0]} ({best_i[6]:.3f})")
print(f"best alpha for paraphrase MRR: {best_p[0]} ({best_p[9]:.3f})")
print(f"best alpha for sum of all three MRR: {best_sum[0]} (heading {best_sum[3]:.3f} + identifier {best_sum[6]:.3f} + paraphrase {best_sum[9]:.3f})")

# --- per-query latency (median of 50 single-query lookups, this corpus size, this machine) ---
sample_qs = [q for q, _ in paraphrase_queries][:1] * 50
t = []
for q in sample_qs:
    s = time.perf_counter(); bm25_search(q); t.append(time.perf_counter() - s)
bm25_ms = st.median(t) * 1000
t = []
for q in sample_qs:
    s = time.perf_counter(); vec_search_batch([q]); t.append(time.perf_counter() - s)
vec_ms = st.median(t) * 1000
print(f"\nmedian single-query latency over {len(chunks)} chunks: BM25={bm25_ms:.2f}ms, vector(encode+search)={vec_ms:.2f}ms")
