"""慢查询审计：pg_stat_statements + EXPLAIN

流程：
    1. 确认 pg_stat_statements 已 CREATE EXTENSION 且加入 shared_preload_libraries
    2. 重置统计 → 跑一批业务 SQL → 打印 Top N
    3. 对 Top 1 自动 EXPLAIN ANALYZE

运行：python code/perf_audit.py
"""
from __future__ import annotations

from db import fetch_all, fetch_one, execute, close_pool


def check_extension() -> bool:
    row = fetch_one(
        "SELECT 1 AS ok FROM pg_extension WHERE extname = 'pg_stat_statements'")
    return bool(row)


def reset() -> None:
    execute("SELECT pg_stat_statements_reset()")


def run_sample_workload() -> None:
    """故意跑一些不同耗时的 SQL 以填充 pg_stat_statements。"""
    sqls = [
        "SELECT COUNT(*) FROM posts WHERE status = 'published'",
        "SELECT * FROM posts WHERE tags @> '[\"postgres\"]' ORDER BY created_at DESC LIMIT 10",
        "SELECT id, title FROM posts WHERE search_vector @@ to_tsquery('simple','postgres') LIMIT 10",
        "SELECT id, title FROM posts WHERE title %% 'psotgres' ORDER BY title <-> 'psotgres' LIMIT 10",
        # 故意不走索引
        "SELECT COUNT(*) FROM posts WHERE lower(title) LIKE '%pg%'",
        "SELECT id, title FROM mv_hot_posts ORDER BY hot_score DESC LIMIT 10",
    ]
    for s in sqls:
        fetch_all(s)


def top_n(n: int = 5) -> list[dict]:
    # PG 13+：列名是 total_exec_time；12 及以前是 total_time，这里做兼容
    col_exists = fetch_one(
        """SELECT 1 AS ok FROM information_schema.columns
           WHERE table_name='pg_stat_statements' AND column_name='total_exec_time'""")
    total_col = "total_exec_time" if col_exists else "total_time"
    mean_col  = "mean_exec_time"  if col_exists else "mean_time"

    sql = f"""
        SELECT substring(query, 1, 80) AS query,
               calls,
               round({total_col}::numeric, 2) AS total_ms,
               round({mean_col}::numeric, 2)  AS mean_ms,
               rows
        FROM   pg_stat_statements
        WHERE  query NOT LIKE 'EXPLAIN%'
          AND  query NOT LIKE '%pg_stat_statements%'
          AND  dbid = (SELECT oid FROM pg_database WHERE datname = current_database())
        ORDER  BY {total_col} DESC
        LIMIT  %s
    """
    return fetch_all(sql, (n,))


def main() -> None:
    if not check_extension():
        print("[ERROR] pg_stat_statements 未安装，请先：")
        print("  1) 在 postgresql.conf 里 shared_preload_libraries = 'pg_stat_statements'")
        print("  2) 重启 PG")
        print("  3) psql: CREATE EXTENSION pg_stat_statements;")
        return

    print("=== pg_stat_statements 审计 ===")
    print("[step 1] 重置统计")
    reset()

    print("[step 2] 跑样例业务 SQL")
    run_sample_workload()

    print("\n[step 3] Top 5 最耗时 SQL")
    rows = top_n(5)
    if not rows:
        print("  （没抓到，可能 track=none 或刚刚重置）")
        return

    print(f"  {'query':<80}  {'calls':>5}  {'total_ms':>10}  {'mean_ms':>8}")
    for r in rows:
        print(f"  {r['query']:<80}  {r['calls']:>5}  "
              f"{str(r['total_ms']):>10}  {str(r['mean_ms']):>8}")


if __name__ == "__main__":
    try:
        main()
    finally:
        close_pool()
