"""Prepare raw data into (a) SFT instruction JSONL and (b) a RAG chunk file.

    python scripts/prepare_data.py --config config/config.yaml

Reads everything in data.raw_dir:
  - *.jsonl with {question, answer, source?}  -> SFT pairs (train/eval split)
  - *.md / *.txt / *.csv                       -> chunked knowledge base -> data/kb/chunks.jsonl

SFT pairs are written in chat format ({"messages": [...]}) that TRL's SFTTrainer consumes, with
the system prompt from agent.prompt so fine-tuning reinforces the same behavioral contract we serve.
"""
from __future__ import annotations

import argparse
import csv
import json
import re
import sys
from pathlib import Path

# Make the src/ package importable when run as a script.
sys.path.insert(0, str(Path(__file__).resolve().parents[1] / "src"))

from agent.prompt import SYSTEM_PROMPT  # noqa: E402
from agent.settings import load_config  # noqa: E402

_HEADING_RE = re.compile(r"^(#{1,6})\s+(.*)$")


def slugify(text: str) -> str:
    s = re.sub(r"[^a-z0-9]+", "-", text.lower()).strip("-")
    return s or "section"


def chunk_words(text: str, size: int, overlap: int) -> list[str]:
    """Approximate token chunking by words (fast, dependency-free)."""
    words = text.split()
    if not words:
        return []
    chunks, start = [], 0
    step = max(1, size - overlap)
    while start < len(words):
        chunk = " ".join(words[start : start + size]).strip()
        if chunk:
            chunks.append(chunk)
        start += step
    return chunks


def parse_markdown(path: Path, size: int, overlap: int) -> list[dict]:
    """Split markdown into (heading-anchored) chunks for citation."""
    stem = path.stem
    sections: list[tuple[str, list[str]]] = [("overview", [])]
    for line in path.read_text(encoding="utf-8").splitlines():
        m = _HEADING_RE.match(line.strip())
        if m:
            sections.append((slugify(m.group(2)), []))
        else:
            sections[-1][1].append(line)
    out = []
    for anchor, body in sections:
        text = "\n".join(body).strip()
        if not text:
            continue
        for i, chunk in enumerate(chunk_words(text, size, overlap)):
            suffix = f"-{i}" if i else ""
            out.append({"text": chunk, "source": f"{stem}#{anchor}{suffix}"})
    return out


def parse_text(path: Path, size: int, overlap: int) -> list[dict]:
    stem = path.stem
    return [
        {"text": chunk, "source": f"{stem}#part-{i}"}
        for i, chunk in enumerate(chunk_words(path.read_text(encoding="utf-8"), size, overlap))
    ]


def parse_csv(path: Path) -> list[dict]:
    out = []
    with open(path, newline="", encoding="utf-8") as f:
        for i, row in enumerate(csv.DictReader(f)):
            text = (row.get("text") or "").strip()
            if text:
                out.append({"text": text, "source": row.get("source") or f"{path.stem}#row-{i}"})
    return out


def load_qa(path: Path) -> list[dict]:
    pairs = []
    for line in path.read_text(encoding="utf-8").splitlines():
        line = line.strip()
        if not line:
            continue
        obj = json.loads(line)
        q, a = obj.get("question"), obj.get("answer")
        if q and a:
            pairs.append({"question": q.strip(), "answer": a.strip(), "source": obj.get("source", "")})
    return pairs


def to_sft_example(pair: dict) -> dict:
    answer = pair["answer"]
    src = pair.get("source")
    # Teach the citation habit: ensure a bracketed source is present in the target.
    if src and f"[{src}]" not in answer:
        answer = f"{answer} [{src}]"
    return {
        "messages": [
            {"role": "system", "content": SYSTEM_PROMPT},
            {"role": "user", "content": pair["question"]},
            {"role": "assistant", "content": answer},
        ]
    }


def write_jsonl(path: Path, rows: list[dict]) -> None:
    path.parent.mkdir(parents=True, exist_ok=True)
    with open(path, "w", encoding="utf-8") as f:
        for r in rows:
            f.write(json.dumps(r, ensure_ascii=False) + "\n")


def main() -> None:
    ap = argparse.ArgumentParser()
    ap.add_argument("--config", default="config/config.yaml")
    args = ap.parse_args()
    cfg = load_config(args.config)

    raw_dir = cfg.resolve("data.raw_dir")
    size = int(cfg.get("data.chunk_tokens", 400))
    overlap = int(cfg.get("data.chunk_overlap", 60))
    eval_split = float(cfg.get("data.eval_split", 0.12))

    qa_pairs: list[dict] = []
    kb_chunks: list[dict] = []

    for path in sorted(raw_dir.rglob("*")):
        if path.is_dir() or path.name.lower() == "readme.md":
            continue
        ext = path.suffix.lower()
        if ext == ".jsonl":
            qa_pairs.extend(load_qa(path))
        elif ext == ".md":
            kb_chunks.extend(parse_markdown(path, size, overlap))
        elif ext == ".txt":
            kb_chunks.extend(parse_text(path, size, overlap))
        elif ext == ".csv":
            kb_chunks.extend(parse_csv(path))
        else:
            print(f"  (skipping unsupported file: {path.name})")

    # --- SFT split (deterministic: every Nth pair to eval) ---
    n_eval = max(1, int(len(qa_pairs) * eval_split)) if qa_pairs else 0
    stride = max(2, len(qa_pairs) // n_eval) if n_eval else 0
    train, evalset = [], []
    for i, pair in enumerate(qa_pairs):
        (evalset if stride and i % stride == 0 else train).append(pair)

    write_jsonl(cfg.resolve("data.sft_train"), [to_sft_example(p) for p in train])
    write_jsonl(cfg.resolve("data.sft_eval"), evalset)  # eval keeps raw q/a/source for scoring

    kb_path = cfg.resolve("data.kb_dir") / "chunks.jsonl"
    write_jsonl(kb_path, kb_chunks)

    print(f"SFT:  {len(train)} train, {len(evalset)} eval  -> {cfg.resolve('data.sft_train')}")
    print(f"RAG:  {len(kb_chunks)} chunks               -> {kb_path}")
    if not qa_pairs:
        print("  ! No Q&A pairs found — add a .jsonl to data/raw/ for fine-tuning.")
    if not kb_chunks:
        print("  ! No KB documents found — add .md/.txt/.csv to data/raw/ for RAG.")


if __name__ == "__main__":
    main()
