#!/usr/bin/env python3
"""
sop-answers: answer questions from your own SOP binder, with the passages it used.

  python3 sop-answers.py index sample-binder/          # embed every .md/.txt file once
  python3 sop-answers.py ask "what do I do with moldy flower?"
  python3 sop-answers.py ask "..." --dry-run           # show what it found and the prompt; no answer call
  python3 sop-answers.py eval questions.csv            # does search find the right SOP? (recall@4)
  python3 sop-answers.py --self-test                   # offline checks, no key needed

Standard library only. Needs OPENROUTER_API_KEY in the environment for `index`, `ask` and `eval`.
Your documents go to OpenRouter to be embedded, and the passages it picks go to the answering
model. Do not point it at anything you would not paste into that provider.

From Distru's No Bullshit AI Course, module 11. MIT.
"""
import json, math, os, re, sys, urllib.request
from pathlib import Path

EMBED_MODEL = os.environ.get("EMBED_MODEL", "openai/text-embedding-3-small")
ANSWER_MODEL = os.environ.get("ANSWER_MODEL", "anthropic/claude-sonnet-5")
INDEX_FILE = Path(os.environ.get("SOP_INDEX", "sop-index.json"))
TOP_K = 4
MAX_CHARS = 900  # passages longer than this are split at sentence ends

INSTRUCTIONS = """You answer staff questions using ONLY the numbered SOP passages below.
- Cite the passages you used like [1] or [2][3].
- If the passages do not answer the question, say exactly: "The binder does not cover this." and suggest who to ask.
- Never add steps, numbers or rules that are not in the passages.
- Keep it short: the steps, in order."""


def api(path, payload):
    key = os.environ.get("OPENROUTER_API_KEY")
    if not key:
        sys.exit("OPENROUTER_API_KEY is not set.")
    req = urllib.request.Request(
        f"https://openrouter.ai/api/v1/{path}", data=json.dumps(payload).encode(),
        headers={"Authorization": f"Bearer {key}", "Content-Type": "application/json"})
    with urllib.request.urlopen(req, timeout=60) as r:
        return json.load(r)


def embed(texts):
    out = []
    for i in range(0, len(texts), 64):
        data = api("embeddings", {"model": EMBED_MODEL, "input": texts[i:i + 64]})["data"]
        out += [d["embedding"] for d in sorted(data, key=lambda d: d["index"])]
    return [unit(v) for v in out]


def unit(v):
    n = math.sqrt(sum(x * x for x in v)) or 1.0
    return [x / n for x in v]


def passages(text, title):
    """Split a document into paragraph-sized passages, each labelled with its title."""
    out = []
    for para in re.split(r"\n\s*\n", text):
        para = " ".join(para.split())
        if not para or para.startswith("#"):
            continue
        while len(para) > MAX_CHARS:
            cut = para.rfind(". ", 0, MAX_CHARS)
            cut = cut + 1 if cut > 0 else MAX_CHARS
            out.append(para[:cut].strip())
            para = para[cut:].strip()
        out.append(para)
    return [{"title": title, "text": p} for p in out]


def title_of(path, text):
    m = re.search(r"^#\s+(.+)$", text, re.M)
    return m.group(1).strip() if m else path.stem


def top_k(qv, items, k=TOP_K):
    scored = [(sum(a * b for a, b in zip(qv, it["v"])), it) for it in items]
    return sorted(scored, key=lambda s: s[0], reverse=True)[:k]


def cmd_index(folder):
    files = sorted(p for p in Path(folder).rglob("*") if p.suffix in (".md", ".txt"))
    if not files:
        sys.exit(f"No .md or .txt files in {folder}")
    items = []
    for f in files:
        text = f.read_text(encoding="utf-8")
        for p in passages(text, title_of(f, text)):
            items.append({"file": str(f), **p})
    # Embed each passage WITH its title: a passage on its own may not say what it is about.
    vecs = embed([f"{it['title']}: {it['text']}" for it in items])
    for it, v in zip(items, vecs):
        it["v"] = [round(x, 5) for x in v]
    INDEX_FILE.write_text(json.dumps({"model": EMBED_MODEL, "items": items}))
    print(f"indexed {len(items)} passages from {len(files)} files into {INDEX_FILE}")


def cmd_ask(question, dry_run=False):
    if not INDEX_FILE.exists():
        sys.exit(f"No index at {INDEX_FILE}. Run: python3 sop-answers.py index <folder>")
    idx = json.loads(INDEX_FILE.read_text())
    if idx["model"] != EMBED_MODEL:
        sys.exit(f"Index was built with {idx['model']}, not {EMBED_MODEL}. Rebuild it.")
    found = top_k(embed([question])[0], idx["items"])
    sources = "\n\n".join(f"[{i + 1}] {it['title']}\n{it['text']}" for i, (_, it) in enumerate(found))
    prompt = f"{INSTRUCTIONS}\n\nPASSAGES\n\n{sources}\n\nQUESTION\n{question}"
    print("Passages found:")
    for i, (score, it) in enumerate(found):
        print(f"  [{i + 1}] {score:.2f}  {it['title']}  ({it['file']})")
    if dry_run:
        print("\n--- prompt that would be sent ---\n" + prompt)
        return
    resp = api("chat/completions", {"model": ANSWER_MODEL, "temperature": 0,
                                    "messages": [{"role": "user", "content": prompt}]})
    print("\n" + resp["choices"][0]["message"]["content"].strip())


def cmd_eval(csv_path):
    """questions.csv: question,expected. `expected` is part of the right file's name, e.g. 14-microbial."""
    import csv
    if not INDEX_FILE.exists():
        sys.exit(f"No index at {INDEX_FILE}. Run: python3 sop-answers.py index <folder>")
    idx = json.loads(INDEX_FILE.read_text())
    rows = list(csv.DictReader(open(csv_path, encoding="utf-8")))
    vecs = embed([r["question"] for r in rows])
    misses = []
    for r, qv in zip(rows, vecs):
        files = [it["file"] for _, it in top_k(qv, idx["items"])]
        if not any(r["expected"] in f for f in files):
            misses.append((r["question"], r["expected"], [Path(f).stem for f in files[:2]]))
    hit = len(rows) - len(misses)
    print(f"recall@{TOP_K}: {hit}/{len(rows)} questions found the right SOP in the top {TOP_K}")
    for q, want, got in misses:
        print(f"  miss: {q!r}  wanted {want}, got {', '.join(got)}")
    print("A miss here means the answer cannot be right. Fix search (passages, titles, wording) before the prompt.")


def self_test():
    ps = passages("# SOP 1 · Test\n\nFirst para.\n\nSecond para. " + "x. " * 400, "SOP 1 · Test")
    assert ps[0]["text"] == "First para." and all(len(p["text"]) <= MAX_CHARS for p in ps), "chunking"
    assert title_of(Path("a.md"), "# SOP 9 · Cash\n\nbody") == "SOP 9 · Cash", "title"
    items = [{"v": unit([1, 0]), "t": "a"}, {"v": unit([0, 1]), "t": "b"}, {"v": unit([1, 1]), "t": "c"}]
    assert [it["t"] for _, it in top_k(unit([1, 0.1]), items, 2)] == ["a", "c"], "ranking"
    print("self-test ok · chunking, titles, ranking")


if __name__ == "__main__":
    args = sys.argv[1:]
    if "--self-test" in args:
        self_test()
    elif len(args) >= 2 and args[0] == "index":
        cmd_index(args[1])
    elif len(args) >= 2 and args[0] == "eval":
        cmd_eval(args[1])
    elif len(args) >= 2 and args[0] == "ask":
        cmd_ask(args[1], dry_run="--dry-run" in args)
    else:
        print(__doc__)
