#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
第 16 章 · 性能调优对照实验脚本
================================

干什么:
    1) 在 events_v1 (反模式) 和 events_v2 (正确做法) 上各跑 N 次同一条业务 SQL.
    2) 从 system.query_log 取出 query_duration_ms / read_rows / read_bytes /
       memory_usage / SelectedParts / SelectedMarks 等指标.
    3) 分别取中位数, 输出 Markdown 格式的对照表.

前置:
    python -m pip install clickhouse-connect
    clickhouse-client < ../init.sql      # 建表 + 灌 1000 万行

运行:
    python perf_tune.py
    python perf_tune.py --rounds 7 --days 30
"""
from __future__ import annotations

import argparse
import statistics
import time
import uuid
from typing import Dict, List

import clickhouse_connect


HOST = "127.0.0.1"
PORT = 8123
DB = "learn_ck"
USER = "default"
PASSWORD = ""


BUSINESS_SQL = """
SELECT country, count() AS pv, uniqExact(user_id) AS uv
FROM {table}
WHERE event_date >= today() - {days}
  AND country IN ('CN','US','JP','DE','BR')
GROUP BY country
ORDER BY pv DESC
"""


def get_client():
    return clickhouse_connect.get_client(
        host=HOST, port=PORT, username=USER,
        password=PASSWORD, database=DB,
    )


def drop_caches(client) -> None:
    """每次实验前清缓存, 避免 page cache 让结果失真."""
    for sysq in ("SYSTEM DROP MARK CACHE",
                 "SYSTEM DROP UNCOMPRESSED CACHE"):
        try:
            client.command(sysq)
        except Exception as e:  # 普通 default 用户可能没权限, 不致命
            print(f"  [warn] {sysq}: {e}")


def run_one(client, table: str, days: int) -> str:
    """跑一次业务 SQL, 返回 query_id 以便事后从 query_log 拿指标."""
    qid = f"perf-{table}-{uuid.uuid4().hex[:10]}"
    sql = BUSINESS_SQL.format(table=table, days=days)
    client.query(sql, settings={"query_id": qid})
    return qid


def fetch_metrics(client, qids: List[str]) -> List[Dict]:
    """从 system.query_log 把这一批 query_id 的指标拉回来."""
    client.command("SYSTEM FLUSH LOGS")
    placeholder = ",".join(f"'{q}'" for q in qids)
    rows = client.query(f"""
        SELECT
            query_id,
            query_duration_ms,
            read_rows,
            read_bytes,
            memory_usage,
            ProfileEvents['SelectedParts'] AS sel_parts,
            ProfileEvents['SelectedMarks'] AS sel_marks
        FROM system.query_log
        WHERE query_id IN ({placeholder}) AND type = 'QueryFinish'
    """).result_rows
    cols = ["query_id", "duration_ms", "read_rows", "read_bytes",
            "memory_usage", "sel_parts", "sel_marks"]
    return [dict(zip(cols, r)) for r in rows]


def median_of(metrics: List[Dict], key: str) -> float:
    return statistics.median([float(m[key]) for m in metrics])


def fmt_bytes(n: float) -> str:
    units = [("B", 1), ("KiB", 1024), ("MiB", 1024 ** 2), ("GiB", 1024 ** 3)]
    for u, base in reversed(units):
        if n >= base or u == "B":
            return f"{n / base:.1f} {u}"
    return f"{n} B"


def fmt_int(n: float) -> str:
    return f"{int(n):,}"


def benchmark(client, table: str, rounds: int, days: int) -> List[Dict]:
    print(f"\n>>> 跑表 {table} 共 {rounds} 轮, days={days}")
    qids: List[str] = []
    for i in range(rounds):
        drop_caches(client)
        t0 = time.perf_counter()
        qid = run_one(client, table, days)
        dt = (time.perf_counter() - t0) * 1000
        print(f"   第 {i + 1}/{rounds} 轮 client_wall={dt:7.0f} ms  qid={qid}")
        qids.append(qid)
    time.sleep(1.5)  # 让 query_log flush
    return fetch_metrics(client, qids)


def render_report(v1: List[Dict], v2: List[Dict]) -> str:
    keys = ["duration_ms", "read_rows", "read_bytes",
            "memory_usage", "sel_parts", "sel_marks"]
    md = []
    md.append("\n=========== 对照报告 (各取中位数) ===========\n")
    md.append("| 指标 | events_v1 (反模式) | events_v2 (正确做法) | 改善倍数 |")
    md.append("|------|--------------------|----------------------|---------|")
    for k in keys:
        m1, m2 = median_of(v1, k), median_of(v2, k)
        imp = (m1 / m2) if m2 > 0 else float("inf")
        if k == "duration_ms":
            disp1, disp2 = f"{m1:.0f} ms", f"{m2:.0f} ms"
        elif k in ("read_bytes", "memory_usage"):
            disp1, disp2 = fmt_bytes(m1), fmt_bytes(m2)
        else:
            disp1, disp2 = fmt_int(m1), fmt_int(m2)
        md.append(f"| `{k}` | {disp1} | {disp2} | **{imp:.1f}×** |")
    md.append("")
    md.append("解读:")
    md.append("  - sel_parts / sel_marks 急剧下降 → 分区裁剪 + 主键裁剪生效.")
    md.append("  - read_bytes 大幅下降 → LowCardinality 字典编码 + 列读减少.")
    md.append("  - memory_usage 下降 → 字典化的 GROUP BY 内存更省.")
    md.append("  - duration_ms 下降 → 上述三件叠加的最终结果.")
    return "\n".join(md)


def main() -> None:
    ap = argparse.ArgumentParser()
    ap.add_argument("--rounds", type=int, default=5,
                    help="每张表跑几轮取中位数, 默认 5")
    ap.add_argument("--days", type=int, default=30,
                    help="WHERE event_date >= today()-days, 默认 30")
    args = ap.parse_args()

    client = get_client()

    print("== 自检表是否就绪 ==")
    rows = client.query(f"""
        SELECT table, sum(rows) AS rows, count() AS parts
        FROM system.parts
        WHERE database = '{DB}'
          AND table IN ('events_v1', 'events_v2')
          AND active
        GROUP BY table
        ORDER BY table
    """).result_rows
    if not rows:
        print("[ERROR] 找不到 events_v1 / events_v2, 请先执行 init.sql.")
        return
    for t, r, p in rows:
        print(f"  {t}: rows={r:,}  active_parts={p}")

    v1 = benchmark(client, "events_v1", args.rounds, args.days)
    v2 = benchmark(client, "events_v2", args.rounds, args.days)

    if not v1 or not v2:
        print("[ERROR] 没拉到 query_log 指标, 请检查 system.query_log 是否启用.")
        return

    print(render_report(v1, v2))


if __name__ == "__main__":
    main()
