#!/usr/bin/env python3
"""ai-local-lab 読者向け Ollama ベンチマーク。

Python 3.9+ の標準ライブラリだけで動作します。Ollamaを起動し、モデルを取得してから実行してください。

    ollama pull qwen3.5:4b
    python bench_ollama.py qwen3.5:4b --device-label "My PC"

既定ではローカルのOllamaだけに接続し、サイトDBと同じ固定プロンプトでウォーム計測を2回行います。
出力は summary.json / summary.csv / events.jsonl。events.jsonl は応答本文とcontextを
保存しないため、業務文書を誤って持ち出しにくい設計です。本文も保存する場合だけ
--include-text を明示してください。
"""

from __future__ import annotations

import argparse
import csv
import json
import platform
import re
import shutil
import statistics
import subprocess
import sys
import threading
import time
import urllib.error
import urllib.request
from datetime import datetime, timezone
from pathlib import Path
from urllib.parse import urlparse

SCRIPT_VERSION = "1.0.0"
SCHEMA_VERSION = "1.0.0"
DEFAULT_URL = "http://127.0.0.1:11434"
DEFAULT_PROMPT = (
    "あなたはエッジデバイス上で動作するアシスタントです。以下の質問に日本語で詳しく答えてください。\n"
    "質問: ローカル環境でLLMを動かす利点と課題を、プライバシー、レイテンシ、コスト、"
    "運用の4つの観点からそれぞれ具体例を挙げて説明してください。"
)
PRIVATE_EVENT_FIELDS = {"response", "thinking", "context"}


def request_json(base_url: str, path: str, payload: dict | None = None, timeout: int = 120) -> dict:
    url = f"{base_url.rstrip('/')}{path}"
    if payload is None:
        req = urllib.request.Request(url)
    else:
        req = urllib.request.Request(
            url,
            data=json.dumps(payload, ensure_ascii=False).encode("utf-8"),
            headers={"Content-Type": "application/json"},
        )
    with urllib.request.urlopen(req, timeout=timeout) as response:
        return json.load(response)


def read_nvidia(field: str, gpu_index: int | None) -> list[float]:
    """指定GPU（未指定時は全GPU）の値を行ごとに返す。集計は呼び出し側が行う。"""
    if not shutil.which("nvidia-smi"):
        return []
    try:
        cmd = ["nvidia-smi", f"--query-gpu={field}", "--format=csv,noheader,nounits"]
        if gpu_index is not None:
            cmd.extend(["-i", str(gpu_index)])
        result = subprocess.run(
            cmd,
            capture_output=True,
            text=True,
            timeout=5,
            check=False,
        )
        return [float(line) for line in result.stdout.strip().splitlines() if line.strip()]
    except (OSError, ValueError, subprocess.SubprocessError):
        return []


def read_tegra_power() -> float | None:
    if not shutil.which("tegrastats"):
        return None
    process = None
    try:
        process = subprocess.Popen(
            ["tegrastats", "--interval", "250"],
            stdout=subprocess.PIPE,
            stderr=subprocess.DEVNULL,
            text=True,
        )
        line = process.stdout.readline() if process.stdout else ""
        match = re.search(r"VDD_IN\s+(\d+)mW", line)
        return int(match.group(1)) / 1000.0 if match else None
    except (OSError, ValueError, subprocess.SubprocessError):
        return None
    finally:
        if process is not None:
            process.terminate()


def read_pi_power() -> float | None:
    if not shutil.which("vcgencmd"):
        return None
    try:
        result = subprocess.run(
            ["vcgencmd", "pmic_read_adc"],
            capture_output=True,
            text=True,
            timeout=5,
            check=False,
        )
        currents: dict[str, float] = {}
        volts: dict[str, float] = {}
        for line in result.stdout.splitlines():
            match = re.search(r"(\S+?)_(A|V)\s+\w+\(\d+\)=([\d.]+)[AV]", line)
            if match:
                name, kind, value = match.group(1), match.group(2), float(match.group(3))
                (currents if kind == "A" else volts)[name] = value
        total = sum(current * volts[name] for name, current in currents.items() if name in volts)
        return total if total > 0 else None
    except (OSError, ValueError, subprocess.SubprocessError):
        return None


def read_mac_power() -> float | None:
    if sys.platform != "darwin" or not shutil.which("powermetrics") or not shutil.which("sudo"):
        return None
    try:
        result = subprocess.run(
            ["sudo", "-n", "powermetrics", "--samplers", "cpu_power", "-i", "200", "-n", "1"],
            capture_output=True,
            text=True,
            timeout=10,
            check=False,
        )
        match = re.search(
            r"Combined Power \(CPU \+ GPU \+ ANE\):\s*([\d.]+)\s*mW", result.stdout
        )
        return float(match.group(1)) / 1000.0 if match else None
    except (OSError, ValueError, subprocess.SubprocessError):
        return None


def read_power_w(gpu_index: int | None) -> tuple[float | None, str | None]:
    # 複数GPU搭載機でOllamaの使用GPUが不明な場合、全GPUの合計を上限側の目安として使う。
    # --gpu-indexを指定すればそのGPU1枚だけに絞れる。
    values = read_nvidia("power.draw", gpu_index)
    if values:
        return round(sum(values), 3), "nvidia-smi GPU power.draw"
    value = read_tegra_power()
    if value is not None:
        return value, "tegrastats VDD_IN"
    value = read_pi_power()
    if value is not None:
        return value, "vcgencmd PMIC rails"
    value = read_mac_power()
    if value is not None:
        return value, "powermetrics CPU+GPU+ANE"
    return None, None


def read_temp_c(gpu_index: int | None) -> tuple[float | None, str | None]:
    values = read_nvidia("temperature.gpu", gpu_index)
    if values:
        return max(values), "nvidia-smi GPU temperature"
    if shutil.which("vcgencmd"):
        try:
            result = subprocess.run(
                ["vcgencmd", "measure_temp"], capture_output=True, text=True, timeout=5, check=False
            )
            match = re.search(r"([\d.]+)", result.stdout)
            if match:
                return float(match.group(1)), "vcgencmd SoC temperature"
        except (OSError, ValueError, subprocess.SubprocessError):
            pass
    thermal = Path("/sys/class/thermal")
    values: list[float] = []
    if thermal.exists():
        for item in thermal.glob("thermal_zone*/temp"):
            try:
                values.append(float(item.read_text(encoding="utf-8").strip()) / 1000.0)
            except (OSError, ValueError):
                pass
    return (max(values), "Linux thermal_zone max") if values else (None, None)


class SensorSampler:
    def __init__(self, gpu_index: int | None = None) -> None:
        self.power_samples: list[float] = []
        self.temp_samples: list[float] = []
        self.samples: list[dict] = []
        self.power_source: str | None = None
        self.temp_source: str | None = None
        self.gpu_index = gpu_index
        self._started = time.perf_counter()
        self._stop = threading.Event()
        self._thread = threading.Thread(target=self._loop, daemon=True)

    def _loop(self) -> None:
        while not self._stop.is_set():
            power, power_source = read_power_w(self.gpu_index)
            temp, temp_source = read_temp_c(self.gpu_index)
            if power is not None:
                self.power_samples.append(power)
                self.power_source = power_source
            if temp is not None:
                self.temp_samples.append(temp)
                self.temp_source = temp_source
            if power is not None or temp is not None:
                self.samples.append(
                    {
                        "elapsed_ms": round((time.perf_counter() - self._started) * 1000, 3),
                        "power_w": power,
                        "temperature_c": temp,
                        "power_source": power_source,
                        "temperature_source": temp_source,
                    }
                )
            self._stop.wait(1.0)

    def __enter__(self) -> "SensorSampler":
        self._thread.start()
        return self

    def __exit__(self, *_: object) -> None:
        self._stop.set()
        self._thread.join(timeout=15)


def unload(base_url: str, model: str) -> None:
    request_json(
        base_url,
        "/api/generate",
        {"model": model, "prompt": "", "stream": False, "keep_alive": 0},
        timeout=60,
    )


def warm_up(base_url: str, model: str, prompt: str) -> None:
    request_json(
        base_url,
        "/api/generate",
        {
            "model": model,
            "prompt": prompt,
            "stream": False,
            "options": {"num_predict": 16, "temperature": 0},
        },
        timeout=600,
    )


def measure_run(
    base_url: str,
    model: str,
    prompt: str,
    num_predict: int,
    run_number: int,
    include_text: bool,
    raw_file,
) -> dict:
    payload = {
        "model": model,
        "prompt": prompt,
        "stream": True,
        "options": {"num_predict": num_predict, "temperature": 0},
    }
    req = urllib.request.Request(
        f"{base_url.rstrip('/')}/api/generate",
        data=json.dumps(payload, ensure_ascii=False).encode("utf-8"),
        headers={"Content-Type": "application/json"},
    )
    started = time.perf_counter()
    first_visible_token_ms: float | None = None
    final_event: dict = {}
    event_count = 0
    with urllib.request.urlopen(req, timeout=1800) as response:
        for raw_line in response:
            if not raw_line.strip():
                continue
            event = json.loads(raw_line)
            if event.get("error"):
                raise ValueError(f"run {run_number}: Ollamaがエラーイベントを返しました: {event['error']}")
            elapsed_ms = (time.perf_counter() - started) * 1000
            event_count += 1
            if first_visible_token_ms is None and event.get("response"):
                first_visible_token_ms = elapsed_ms
            stored_event = event if include_text else {
                key: value for key, value in event.items() if key not in PRIVATE_EVENT_FIELDS
            }
            raw_file.write(
                json.dumps(
                    {
                        "schema_version": SCHEMA_VERSION,
                        "event_type": "ollama",
                        "run": run_number,
                        "elapsed_ms": round(elapsed_ms, 3),
                        "event": stored_event,
                    },
                    ensure_ascii=False,
                )
                + "\n"
            )
            raw_file.flush()
            if event.get("done"):
                final_event = event
    if not final_event.get("done"):
        raise ValueError(
            f"run {run_number}: ストリームがdone:trueを返さずに終了しました（{event_count}イベント受信）"
        )
    wall_ms = (time.perf_counter() - started) * 1000
    eval_count = int(final_event.get("eval_count") or 0)
    eval_duration_ns = int(final_event.get("eval_duration") or 0)
    total_duration_ns = int(final_event.get("total_duration") or 0)
    prompt_count = int(final_event.get("prompt_eval_count") or 0)
    prompt_duration_ns = int(final_event.get("prompt_eval_duration") or 0)
    if (
        eval_count <= 0
        or eval_duration_ns <= 0
        or total_duration_ns <= 0
        or prompt_count <= 0
        or prompt_duration_ns <= 0
    ):
        raise ValueError(
            f"run {run_number}: 計測カウンタが不正です"
            f"（eval_count={eval_count}, eval_duration={eval_duration_ns}, total_duration={total_duration_ns}, "
            f"prompt_eval_count={prompt_count}, prompt_eval_duration={prompt_duration_ns}）"
        )
    return {
        "run": run_number,
        # サイトDB互換値: 既存bench_llm.pyと同じサーバー計測式。
        "ttft_ms": round((total_duration_ns - eval_duration_ns) / 1e6, 3)
        if total_duration_ns and eval_duration_ns
        else None,
        # 読者が実際に見えるresponseの最初の1文字まで。thinkingは数えない。
        "ttft_ms_stream": round(
            first_visible_token_ms if first_visible_token_ms is not None else wall_ms, 3
        ),
        "wall_duration_ms": round(wall_ms, 3),
        "decode_tok_per_sec": round(eval_count / (eval_duration_ns / 1e9), 3)
        if eval_count and eval_duration_ns
        else None,
        "prefill_tok_per_sec": round(prompt_count / (prompt_duration_ns / 1e9), 3)
        if prompt_count and prompt_duration_ns
        else None,
        "input_tokens": prompt_count or None,
        "output_tokens": eval_count or None,
        "server_total_duration_ms": round(total_duration_ns / 1e6, 3),
        "server_load_duration_ms": round(int(final_event.get("load_duration") or 0) / 1e6, 3),
        "done_reason": final_event.get("done_reason"),
        "event_count": event_count,
    }


def mean_of(runs: list[dict], key: str) -> float | None:
    values = [float(run[key]) for run in runs if run.get(key) is not None]
    return round(statistics.fmean(values), 3) if values else None


def safe_slug(value: str) -> str:
    slug = re.sub(r"[^a-zA-Z0-9._-]+", "-", value).strip("-.").lower()
    return slug or "model"


def parse_args() -> argparse.Namespace:
    parser = argparse.ArgumentParser(description="Ollamaの速度・TTFT・電力・温度を同一条件で測定")
    parser.add_argument("model", help="Ollamaモデルタグ（例: qwen3.5:4b）")
    parser.add_argument("--device-label", default="", help="公開してよい機材名。ホスト名は自動収集しません")
    parser.add_argument("--ollama-url", default=DEFAULT_URL, help=f"Ollama URL（既定: {DEFAULT_URL}）")
    parser.add_argument("--allow-remote", action="store_true", help="localhost以外への送信を明示許可")
    parser.add_argument("--runs", type=int, default=2, help="本計測回数（既定: 2、サイトDB互換）")
    parser.add_argument("--num-predict", type=int, default=256, help="最大生成トークン数（既定: 256）")
    parser.add_argument("--mode", choices=["warm", "cold"], default="warm", help="warmは事前ロード、coldは各回アンロード")
    parser.add_argument(
        "--gpu-index",
        type=int,
        default=None,
        help="計測対象のnvidia-smi GPU index。複数GPU搭載機でOllamaが使うGPUを特定したい場合に指定します"
        "（未指定時は全GPUの電力合計・温度最大を使用）",
    )
    parser.add_argument("--prompt-file", type=Path, help="UTF-8の独自プロンプト。公開前に内容を確認してください")
    parser.add_argument("--include-text", action="store_true", help="生成文・contextもevents.jsonlへ保存")
    parser.add_argument("--output", type=Path, default=Path("ai-local-lab-results"), help="出力親ディレクトリ")
    parser.add_argument("--version", action="version", version=SCRIPT_VERSION)
    args = parser.parse_args()
    if not 1 <= args.runs <= 30:
        parser.error("--runs は1〜30で指定してください")
    if not 1 <= args.num_predict <= 8192:
        parser.error("--num-predict は1〜8192で指定してください")
    if args.gpu_index is not None and args.gpu_index < 0:
        parser.error("--gpu-index は0以上の整数で指定してください")
    parsed = urlparse(args.ollama_url)
    if parsed.scheme not in {"http", "https"} or not parsed.hostname:
        parser.error("--ollama-url は http(s) URLで指定してください")
    if parsed.hostname not in {"127.0.0.1", "localhost", "::1"} and not args.allow_remote:
        parser.error("localhost以外へ送る場合は --allow-remote を明示してください")
    return args


def main() -> int:
    args = parse_args()
    prompt = (
        args.prompt_file.read_text(encoding="utf-8") if args.prompt_file else DEFAULT_PROMPT
    )
    run_started = datetime.now(timezone.utc)
    stamp = run_started.strftime("%Y%m%dT%H%M%SZ")
    out_dir = args.output / f"{stamp}-{safe_slug(args.model)}"
    out_dir.mkdir(parents=True, exist_ok=False)
    try:
        version = str(request_json(args.ollama_url, "/api/version").get("version", "unknown"))
        show = request_json(args.ollama_url, "/api/show", {"model": args.model})
        details = show.get("details") or {}
        if args.mode == "warm":
            print("[1/3] モデルをウォームアップしています...", file=sys.stderr)
            warm_up(args.ollama_url, args.model, prompt)
        print(f"[2/3] {args.runs}回を計測しています...", file=sys.stderr)
        runs: list[dict] = []
        raw_path = out_dir / "events.jsonl"
        with raw_path.open(
            "w", encoding="utf-8", newline="\n"
        ) as raw_file, SensorSampler(args.gpu_index) as sensors:
            for index in range(1, args.runs + 1):
                if args.mode == "cold":
                    unload(args.ollama_url, args.model)
                runs.append(
                    measure_run(
                        args.ollama_url,
                        args.model,
                        prompt,
                        args.num_predict,
                        index,
                        args.include_text,
                        raw_file,
                    )
                )
                print(f"  run {index}/{args.runs}: {runs[-1]['decode_tok_per_sec']} tok/s", file=sys.stderr)
        # APIイベントと同じJSONLへ生センサーサンプルを残し、平均・最大を再計算可能にする。
        with raw_path.open("a", encoding="utf-8", newline="\n") as raw_file:
            for sample in sensors.samples:
                raw_file.write(
                    json.dumps(
                        {
                            "schema_version": SCHEMA_VERSION,
                            "event_type": "sensor",
                            **sample,
                        },
                        ensure_ascii=False,
                    )
                    + "\n"
                )
        summary = {
            "schema_version": SCHEMA_VERSION,
            "script": "ai-local-lab bench_ollama.py",
            "script_version": SCRIPT_VERSION,
            "measured_at": run_started.isoformat().replace("+00:00", "Z"),
            "device_label": args.device_label or None,
            "platform": {
                "system": platform.system(),
                "release": platform.release(),
                "machine": platform.machine(),
                "python": platform.python_version(),
            },
            "ollama": {"url": args.ollama_url, "version": version},
            "model": {
                "tag": args.model,
                "family": details.get("family"),
                "parameter_size": details.get("parameter_size"),
                "quantization_level": details.get("quantization_level"),
            },
            "protocol": {
                "mode": args.mode,
                "runs": args.runs,
                "num_predict": args.num_predict,
                "temperature": 0,
                "stream": True,
                "prompt": prompt if args.include_text else None,
                "text_in_raw_log": bool(args.include_text),
            },
            "average": {
                "ttft_ms": mean_of(runs, "ttft_ms"),
                "ttft_ms_stream": mean_of(runs, "ttft_ms_stream"),
                "decode_tok_per_sec": mean_of(runs, "decode_tok_per_sec"),
                "prefill_tok_per_sec": mean_of(runs, "prefill_tok_per_sec"),
                "wall_duration_ms": mean_of(runs, "wall_duration_ms"),
                "power_w_avg": round(statistics.fmean(sensors.power_samples), 3)
                if sensors.power_samples
                else None,
                "temp_max_c": round(max(sensors.temp_samples), 3) if sensors.temp_samples else None,
            },
            "measurement_boundary": {
                "power": sensors.power_source,
                "temperature": sensors.temp_source,
                "gpu_index": (
                    (args.gpu_index if args.gpu_index is not None else "all")
                    if "nvidia-smi" in (sensors.power_source or sensors.temp_source or "")
                    else None
                ),
                "note": "電力の測定境界は機材ごとに異なるため、異方式間の比較は目安です。",
            },
            "runs": runs,
        }
        summary_path = out_dir / "summary.json"
        summary_path.write_text(json.dumps(summary, ensure_ascii=False, indent=2) + "\n", encoding="utf-8")
        csv_path = out_dir / "summary.csv"
        csv_row = {
            "schema_version": SCHEMA_VERSION,
            "measured_at": summary["measured_at"],
            "device_label": summary["device_label"],
            "model": args.model,
            "quant": details.get("quantization_level"),
            "runtime": f"ollama {version}",
            "mode": args.mode,
            **summary["average"],
            "power_measurement_boundary": sensors.power_source,
        }
        with csv_path.open("w", encoding="utf-8-sig", newline="") as csv_file:
            writer = csv.DictWriter(csv_file, fieldnames=list(csv_row))
            writer.writeheader()
            writer.writerow(csv_row)
        print("[3/3] 完了", file=sys.stderr)
        print(summary_path)
        print(csv_path)
        print(raw_path)
        return 0
    except (OSError, ValueError, KeyError, urllib.error.URLError, json.JSONDecodeError) as exc:
        print(f"計測に失敗しました: {exc}", file=sys.stderr)
        print("Ollamaが起動済みで、指定モデルを ollama pull 済みか確認してください。", file=sys.stderr)
        return 1


if __name__ == "__main__":
    raise SystemExit(main())
