"""01_pg_stat_statements_topn.py —— 第 17 章配套代码 #1

用途
    通过 pg_stat_statements 找出当前数据库中「最慢 / 最频繁 / 最耗 IO」的 Top N 条 SQL。
    这是生产环境慢 SQL 排查的第一站：与其凭直觉猜哪条 SQL 慢，不如让 PG 自己告诉你。

前置条件
    1. 已安装并启用 pg_stat_statements：
           shared_preload_libraries = 'pg_stat_statements'   # postgresql.conf
           CREATE EXTENSION IF NOT EXISTS pg_stat_statements;
    2. 安装 psycopg：pip install "psycopg[binary]>=3.1"

运行方式
    python 01_pg_stat_statements_topn.py [--top 10] [--reset]

输出说明
    脚本会先打印「按总耗时排序」的 Top N，再打印「按平均耗时排序」的 Top N，
    再打印「按平均缓存命中率排序」最差的 Top N。
"""

from __future__ import annotations

import argparse
import sys

try:
    import psycopg
except ImportError:
    sys.exit('请先安装 psycopg v3：pip install "psycopg[binary]>=3.1"')


CONN_INFO = "host=127.0.0.1 port=5432 dbname=learn_pg user=postgres"


SQL_TOTAL_TIME = """
SELECT
    queryid,
    LEFT(query, 80)                                  AS query,
    calls,
    ROUND(total_exec_time::NUMERIC, 1)               AS total_ms,
    ROUND(mean_exec_time::NUMERIC, 2)                AS mean_ms,
    rows,
    ROUND(
        100.0 * shared_blks_hit
              / NULLIF(shared_blks_hit + shared_blks_read, 0)
        , 2
    )                                                AS hit_pct
FROM pg_stat_statements
ORDER BY total_exec_time DESC
LIMIT %s
"""

SQL_MEAN_TIME = """
SELECT
    queryid,
    LEFT(query, 80) AS query,
    calls,
    ROUND(mean_exec_time::NUMERIC, 2)  AS mean_ms,
    ROUND(stddev_exec_time::NUMERIC,2) AS stddev_ms,
    rows
FROM pg_stat_statements
WHERE calls > 1            -- 调用过 1 次以上才有统计意义
ORDER BY mean_exec_time DESC
LIMIT %s
"""

SQL_LOWEST_HIT = """
SELECT
    queryid,
    LEFT(query, 80) AS query,
    calls,
    shared_blks_read,
    shared_blks_hit,
    ROUND(
        100.0 * shared_blks_hit
              / NULLIF(shared_blks_hit + shared_blks_read, 0)
        , 2
    ) AS hit_pct
FROM pg_stat_statements
WHERE shared_blks_hit + shared_blks_read > 100
ORDER BY hit_pct ASC NULLS LAST
LIMIT %s
"""


def banner(text: str) -> None:
    print()
    print("=" * 80)
    print(f" {text}")
    print("=" * 80)


def show(cur, sql: str, top_n: int, headers: list[str]) -> None:
    cur.execute(sql, (top_n,))
    rows = cur.fetchall()
    if not rows:
        print("  (空) —— 当前 pg_stat_statements 没有满足条件的记录")
        return
    widths = [max(len(str(h)), max((len(str(r[i])) for r in rows), default=0))
              for i, h in enumerate(headers)]
    line = "  " + "  ".join(f"{h:<{widths[i]}}" for i, h in enumerate(headers))
    print(line)
    print("  " + "  ".join("-" * w for w in widths))
    for r in rows:
        print("  " + "  ".join(f"{str(r[i]):<{widths[i]}}" for i in range(len(headers))))


def ensure_extension(cur) -> None:
    cur.execute("CREATE EXTENSION IF NOT EXISTS pg_stat_statements")


def main() -> None:
    parser = argparse.ArgumentParser()
    parser.add_argument("--top", type=int, default=10, help="Top N，默认 10")
    parser.add_argument("--reset", action="store_true", help="先重置 pg_stat_statements 再退出")
    args = parser.parse_args()

    try:
        conn = psycopg.connect(CONN_INFO, autocommit=True)
    except psycopg.OperationalError as exc:
        sys.exit(f"连接失败：{exc}")

    with conn, conn.cursor() as cur:
        try:
            ensure_extension(cur)
        except psycopg.errors.UndefinedFile:
            sys.exit(
                "❌ pg_stat_statements 未编译进 PG，请检查 contrib 模块是否安装：\n"
                "   apt install postgresql-contrib  或  yum install postgresql-contrib"
            )

        if args.reset:
            cur.execute("SELECT pg_stat_statements_reset()")
            print("✅ pg_stat_statements 已重置。")
            return

        # 让脚本本身的 SQL 也产生一些样本，方便没数据时也能看到东西
        for _ in range(3):
            cur.execute("SELECT count(*) FROM ch17_perf_orders WHERE status = 1")
            cur.execute("SELECT * FROM ch17_perf_orders WHERE id = 12345")

        banner(f"Top {args.top} · 总耗时最高（找「占用 DB 时间最多」的 SQL）")
        show(cur, SQL_TOTAL_TIME, args.top,
             ["queryid", "query", "calls", "total_ms", "mean_ms", "rows", "hit_pct"])

        banner(f"Top {args.top} · 平均耗时最高（找「单次最慢」的 SQL）")
        show(cur, SQL_MEAN_TIME, args.top,
             ["queryid", "query", "calls", "mean_ms", "stddev_ms", "rows"])

        banner(f"Top {args.top} · 缓存命中率最低（找「最伤 IO」的 SQL）")
        show(cur, SQL_LOWEST_HIT, args.top,
             ["queryid", "query", "calls", "blks_read", "blks_hit", "hit_pct"])

        print("\n💡 提示：")
        print("   · total_ms = mean_ms × calls，先看 total 找「拖慢系统的元凶」；")
        print("   · 再看 mean_ms 找「单次特别慢」的 SQL（可能是缺索引）；")
        print("   · 命中率 < 90% 说明该 SQL 在频繁访问磁盘，要么缺索引、要么 shared_buffers 太小。")


if __name__ == "__main__":
    main()
