"""
第 5 章 MergeTree 核心原理 · EXPLAIN + query_log 对比脚本

作用：
    对比 "不合理建表 (bad_mt)" vs "合理建表 (good_mt)" 在同一条查询下的
        - EXPLAIN PIPELINE 输出
        - system.query_log 里的 read_rows / read_bytes / query_duration_ms
    额外演示：
        - use_skip_indexes=0/1 对跳数索引的影响
        - 手动触发 OPTIMIZE 后 Part 的变化

前提：先跑过 seed.py 造数。

运行：
    python explain_demo.py
"""

from __future__ import annotations

import sys
import textwrap
import time
import uuid

import clickhouse_connect


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


def banner(s: str) -> None:
    print("\n" + "=" * 78)
    print("  " + s)
    print("=" * 78)


def run_query(client, label: str, sql: str, settings: dict | None = None
              ) -> str:
    """执行查询并返回 query_id，方便从 query_log 取指标"""
    qid = str(uuid.uuid4())
    opts = {"query_id": qid}
    if settings:
        opts["settings"] = settings
    t0 = time.time()
    res = client.query(sql, **opts)
    elapsed = (time.time() - t0) * 1000
    cnt = len(res.result_rows)
    print(f"[{label}] qid={qid[:8]}  client_elapsed={elapsed:7.1f} ms  "
          f"rows_returned={cnt}")
    return qid


def show_query_log(client, qids: list[tuple[str, str]]) -> None:
    """从 query_log 里拉指标"""
    print("\n--- system.query_log 对照表 ---")
    time.sleep(2)
    client.command("SYSTEM FLUSH LOGS")
    q_ids = "(" + ",".join(f"'{q}'" for _, q in qids) + ")"
    rows = client.query(
        f"""
        SELECT query_id,
               read_rows,
               formatReadableSize(read_bytes) AS read_bytes_h,
               formatReadableSize(memory_usage) AS mem_h,
               query_duration_ms
        FROM   system.query_log
        WHERE  type = 'QueryFinish'
          AND  query_id IN {q_ids}
        """
    ).result_rows
    by_id = {r[0]: r for r in rows}
    print(f"{'Label':<34} {'read_rows':>12} {'read_bytes':>12} "
          f"{'mem':>10} {'dur_ms':>8}")
    print("-" * 78)
    for label, qid in qids:
        r = by_id.get(qid)
        if not r:
            print(f"{label:<34}  (no log, maybe FLUSH LOGS not propagated)")
            continue
        print(f"{label:<34} {r[1]:>12,} {r[2]:>12} {r[3]:>10} {r[4]:>8}")


def main() -> int:
    client = clickhouse_connect.get_client(
        host=HOST, port=PORT, username=USER, password=PASSWORD,
        database=DATABASE,
    )
    print(f"Connected. Database = {DATABASE}")

    # 选一个真实存在的 user_id 作查询条件
    try:
        uid_row = client.query(
            "SELECT user_id FROM learn_ck.good_mt "
            "WHERE event_date BETWEEN '2024-02-01' AND '2024-02-28' "
            "LIMIT 1"
        ).result_rows
        if not uid_row:
            print("!! good_mt 里没有 2024-02 的数据，先跑 seed.py")
            return 1
        uid = uid_row[0][0]
    except Exception as e:
        print(f"!! 查询 good_mt 失败，确认 seed.py 已执行: {e}")
        return 1
    print(f"挑选查询条件：user_id = {uid}")

    # ─────────────────────────────────────────
    banner("① EXPLAIN PIPELINE：好表 vs 差表")
    # ─────────────────────────────────────────
    for tbl in ("bad_mt", "good_mt"):
        print(f"\n-- EXPLAIN PIPELINE for {tbl} --")
        sql = f"""
            EXPLAIN PIPELINE
            SELECT user_id, count(), sum(revenue)
            FROM   learn_ck.{tbl}
            WHERE  user_id = {uid}
              AND  event_date BETWEEN '2024-02-01' AND '2024-02-28'
            GROUP  BY user_id
        """
        rows = client.query(sql).result_rows
        for r in rows:
            print("   " + r[0])

    # ─────────────────────────────────────────
    banner("② 同一查询：bad_mt vs good_mt 指标对照")
    # ─────────────────────────────────────────
    qids = []
    for tbl in ("bad_mt", "good_mt"):
        qids.append((
            f"{tbl:<10}  WHERE user_id + date",
            run_query(
                client,
                tbl,
                f"""
                SELECT user_id, count(), sum(revenue)
                FROM   learn_ck.{tbl}
                WHERE  user_id = {uid}
                  AND  event_date BETWEEN '2024-02-01' AND '2024-02-28'
                GROUP  BY user_id
                """,
            ),
        ))
    show_query_log(client, qids)

    # ─────────────────────────────────────────
    banner("③ 跳数索引开关对比 (events_mt)")
    # ─────────────────────────────────────────
    qids = []
    for flag in (0, 1):
        qids.append((
            f"events_mt  use_skip_indexes={flag}",
            run_query(
                client,
                f"ski={flag}",
                """
                SELECT count(), sum(revenue)
                FROM   learn_ck.events_mt
                WHERE  event_name = 'purchase'
                  AND  event_date BETWEEN '2024-03-01' AND '2024-03-31'
                """,
                settings={"use_skip_indexes": flag},
            ),
        ))
    show_query_log(client, qids)

    # ─────────────────────────────────────────
    banner("④ system.parts / merges 快照")
    # ─────────────────────────────────────────
    for tbl in ("events_mt", "good_mt", "bad_mt"):
        rows = client.query(
            f"""
            SELECT count() AS parts, sum(rows) AS rows,
                   formatReadableSize(sum(bytes_on_disk)) AS size,
                   min(level) AS min_lvl, max(level) AS max_lvl
            FROM   system.parts
            WHERE  database='learn_ck' AND table='{tbl}' AND active
            """
        ).result_rows
        if rows:
            r = rows[0]
            print(f"  {tbl:<10} parts={r[0]:<4} rows={r[1]:>12,} "
                  f"size={r[2]:>10} level=[{r[3]}..{r[4]}]")

    print("\nHint: 手动触发合并（慎用）：")
    print("  OPTIMIZE TABLE learn_ck.good_mt PARTITION '202402' FINAL;")
    return 0


if __name__ == "__main__":
    sys.exit(main())
