"""Evaluate quality + energy of the agent on the held-out eval set.

    python scripts/evaluate.py --config config/config.yaml     # vLLM server must be running

Quality (per eval example {question, answer, source}):
  - answer_similarity : cosine sim between generated and reference answer (small embedder)
  - has_citation      : did the answer cite a source?
  - source_match      : did it cite the expected source?
Energy:
  - samples GPU power via NVML in a background thread -> integrates to Wh per query
  - reports cache hit-rate (cached answers cost ~0 GPU)

Writes a JSON report to reports/eval_report.json and prints a summary.
"""
from __future__ import annotations

import argparse
import json
import sys
import threading
import time
from pathlib import Path

sys.path.insert(0, str(Path(__file__).resolve().parents[1] / "src"))

import numpy as np  # noqa: E402

from agent.embeddings import Embedder  # noqa: E402
from agent.pipeline import Agent  # noqa: E402
from agent.settings import load_config  # noqa: E402


class PowerSampler:
    """Background NVML power sampler. Integrates watts over time -> Wh. No-op if NVML missing."""

    def __init__(self, interval: float = 0.2):
        self.interval = interval
        self._energy_wh = 0.0
        self._stop = threading.Event()
        self._thread: threading.Thread | None = None
        self._handle = None
        try:
            import pynvml

            pynvml.nvmlInit()
            self._pynvml = pynvml
            self._handle = pynvml.nvmlDeviceGetHandleByIndex(0)
        except Exception as e:  # noqa: BLE001
            self._pynvml = None
            print(f"  (energy: NVML unavailable, skipping power measurement: {e})")

    def _run(self):
        last = time.time()
        while not self._stop.is_set():
            time.sleep(self.interval)
            now = time.time()
            try:
                mw = self._pynvml.nvmlDeviceGetPowerUsage(self._handle)  # milliwatts
                self._energy_wh += (mw / 1000.0) * ((now - last) / 3600.0)
            except Exception:  # noqa: BLE001
                pass
            last = now

    def __enter__(self):
        if self._pynvml is not None:
            self._thread = threading.Thread(target=self._run, daemon=True)
            self._thread.start()
        return self

    def __exit__(self, *exc):
        self._stop.set()
        if self._thread:
            self._thread.join()

    @property
    def energy_wh(self) -> float:
        return self._energy_wh


def cosine(a: np.ndarray, b: np.ndarray) -> float:
    return float(np.dot(a, b))  # embedder returns normalized vectors


def main() -> None:
    ap = argparse.ArgumentParser()
    ap.add_argument("--config", default="config/config.yaml")
    ap.add_argument("--limit", type=int, default=0, help="evaluate only the first N examples")
    args = ap.parse_args()
    cfg = load_config(args.config)

    eval_path = cfg.resolve("data.sft_eval")
    if not eval_path.exists():
        raise SystemExit(f"{eval_path} not found. Run scripts/prepare_data.py first.")
    examples = [json.loads(l) for l in eval_path.read_text(encoding="utf-8").splitlines() if l.strip()]
    if args.limit:
        examples = examples[: args.limit]

    agent = Agent(cfg, lazy_llm=True)
    embedder = Embedder(cfg)

    rows = []
    n_cached = 0
    with PowerSampler() as sampler:
        t0 = time.time()
        for ex in examples:
            resp = agent.answer(ex["question"])
            pred_vec = embedder.embed_query(resp.answer)
            ref_vec = embedder.embed_query(ex["answer"])
            sim = cosine(pred_vec, ref_vec)
            src = (ex.get("source") or "").strip()
            source_match = bool(src) and any(src in c or c in src for c in resp.citations)
            n_cached += int(resp.cached)
            rows.append(
                {
                    "question": ex["question"],
                    "answer_similarity": round(sim, 3),
                    "has_citation": bool(resp.citations),
                    "source_match": source_match,
                    "grounded": resp.grounded,
                    "cached": resp.cached,
                    "route": resp.route,
                    "prompt_tokens": resp.prompt_tokens,
                    "completion_tokens": resp.completion_tokens,
                    "total_tokens": resp.total_tokens,
                    "guardrail_reason": resp.guardrail_reason,
                }
            )
        wall = time.time() - t0

    n = len(rows)
    total_tokens = sum(r["total_tokens"] for r in rows)
    summary = {
        "n": n,
        "avg_answer_similarity": round(float(np.mean([r["answer_similarity"] for r in rows])), 3) if n else 0,
        "citation_rate": round(sum(r["has_citation"] for r in rows) / n, 3) if n else 0,
        "source_match_rate": round(sum(r["source_match"] for r in rows) / n, 3) if n else 0,
        "grounded_rate": round(sum(r["grounded"] for r in rows) / n, 3) if n else 0,
        "cache_hit_rate": round(n_cached / n, 3) if n else 0,
        # --- token efficiency (the cost/energy objective) ---
        "avg_prompt_tokens": round(sum(r["prompt_tokens"] for r in rows) / n, 1) if n else 0,
        "avg_completion_tokens": round(sum(r["completion_tokens"] for r in rows) / n, 1) if n else 0,
        "avg_total_tokens_per_query": round(total_tokens / n, 1) if n else 0,
        "total_tokens": total_tokens,
        "zero_token_queries": sum(1 for r in rows if r["route"] in ("cache", "router")),
        "wall_seconds": round(wall, 2),
        "energy_wh_total": round(sampler.energy_wh, 4),
        "energy_wh_per_query": round(sampler.energy_wh / n, 4) if n else 0,
    }

    report_path = Path(__file__).resolve().parents[1] / "reports" / "eval_report.json"
    report_path.parent.mkdir(parents=True, exist_ok=True)
    report_path.write_text(json.dumps({"summary": summary, "rows": rows}, indent=2), encoding="utf-8")

    print("\n=== Evaluation summary ===")
    for k, v in summary.items():
        print(f"  {k:22} {v}")
    print(f"\nFull report -> {report_path}")


if __name__ == "__main__":
    main()
