"""
03_checkpoint_observe.py
------------------------
观察 Checkpoint 的触发与效果：
  · 写入大量数据制造脏页
  · 通过 pg_stat_bgwriter 看 buffers_checkpoint / checkpoint_*_time 增量
  · 手动触发 CHECKPOINT 并测耗时
  · 看 pg_control 里的最近 checkpoint LSN
"""

import time

import psycopg

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


def banner(title: str) -> None:
    print("\n" + "=" * 70)
    print(f"  {title}")
    print("=" * 70)


def get_bgwriter_stats(cur: psycopg.Cursor) -> dict:
    cur.execute(
        """
        SELECT checkpoints_timed, checkpoints_req,
               checkpoint_write_time, checkpoint_sync_time,
               buffers_checkpoint, buffers_clean,
               buffers_backend, buffers_alloc
        FROM pg_stat_bgwriter;
        """
    )
    cols = [d.name for d in cur.description]
    return dict(zip(cols, cur.fetchone()))


def diff(after: dict, before: dict) -> dict:
    return {k: after[k] - before[k] for k in after}


def main() -> None:
    with psycopg.connect(CONN_STR, autocommit=True) as conn, conn.cursor() as cur:
        cur.execute("DROP TABLE IF EXISTS ch10_ckpt_demo;")
        cur.execute("CREATE TABLE ch10_ckpt_demo (id BIGSERIAL PRIMARY KEY, body TEXT);")

        # ---------- 1. 控制相关参数 ----------
        banner("1. Checkpoint 相关配置")
        for k in ("checkpoint_timeout", "max_wal_size", "min_wal_size",
                  "checkpoint_completion_target", "wal_compression",
                  "shared_buffers", "bgwriter_delay"):
            cur.execute("SHOW %s;" % k)
            print(f"  {k:<30} = {cur.fetchone()[0]}")

        # ---------- 2. 看最近 checkpoint 信息 ----------
        banner("2. 最近一次 checkpoint 的元数据 (pg_control)")
        cur.execute("SELECT * FROM pg_control_checkpoint();")
        row = cur.fetchone()
        cols = [d.name for d in cur.description]
        for c, v in zip(cols, row):
            print(f"  {c:<30} = {v}")

        # ---------- 3. 写一波脏数据 ----------
        banner("3. 写入 50000 行制造脏页")
        before = get_bgwriter_stats(cur)
        before_lsn = cur.execute("SELECT pg_current_wal_lsn();").fetchone()[0]
        t0 = time.perf_counter()
        cur.execute(
            "INSERT INTO ch10_ckpt_demo (body) "
            "SELECT repeat(md5(g::text), 4) FROM generate_series(1, 50000) g;"
        )
        elapsed = time.perf_counter() - t0
        after_lsn = cur.execute("SELECT pg_current_wal_lsn();").fetchone()[0]
        wal_size = cur.execute(
            "SELECT pg_wal_lsn_diff(%s, %s);", (after_lsn, before_lsn)
        ).fetchone()[0]
        print(f"  写入耗时 = {elapsed:.2f}s")
        print(f"  生成 WAL = {int(wal_size)/1024/1024:.2f} MB")

        # ---------- 4. 手动 CHECKPOINT ----------
        banner("4. 手动触发 CHECKPOINT 并计时")
        t0 = time.perf_counter()
        cur.execute("CHECKPOINT;")
        ckpt_time = time.perf_counter() - t0
        after = get_bgwriter_stats(cur)
        print(f"  CHECKPOINT 耗时 = {ckpt_time:.3f}s")
        d = diff(after, before)
        print("\n  pg_stat_bgwriter 增量：")
        for k, v in d.items():
            print(f"    {k:<25} +{v}")
        print(
            "\n  解读：\n"
            "  · buffers_checkpoint 增量 = 这次 checkpoint 实际刷了多少 buffer\n"
            "  · checkpoint_write_time / sync_time 单位是毫秒\n"
            "  · 如果 buffers_backend 比 buffers_checkpoint 还多，说明 backend\n"
            "    自己刷脏页太频繁，要调大 bgwriter / max_wal_size"
        )

        # ---------- 5. 比较之前/之后的 checkpoint LSN ----------
        banner("5. 最近 checkpoint LSN 已经前移")
        cur.execute("SELECT * FROM pg_control_checkpoint();")
        row2 = cur.fetchone()
        d2 = dict(zip(cols, row2))
        for c in ("checkpoint_lsn", "redo_lsn", "next_xid", "time"):
            print(f"  {c:<25} = {d2[c]}")
        print(
            "\n  下次崩溃恢复将从 redo_lsn 开始重放。\n"
            "  比这个 LSN 老的 WAL 段（且不被复制 / 归档需要的）可被回收。"
        )

        # ---------- 6. checkpoint 触发统计 ----------
        banner("6. 累计 checkpoint 健康度")
        cur.execute(
            """
            SELECT checkpoints_timed AS by_time,
                   checkpoints_req   AS by_size_or_manual,
                   ROUND(checkpoints_req::numeric /
                         NULLIF(checkpoints_timed + checkpoints_req, 0) * 100, 1)
                                       AS req_ratio_pct
            FROM pg_stat_bgwriter;
            """
        )
        bt, br, ratio = cur.fetchone()
        print(f"  checkpoints_timed (按 timeout) = {bt}")
        print(f"  checkpoints_req   (按 size/手动) = {br}")
        print(f"  req 触发占比 = {ratio}%")
        if ratio is not None and float(ratio) > 30:
            print("  ⚠ req 占比 > 30%：建议调大 max_wal_size 让 timeout 主导。")


if __name__ == "__main__":
    main()
