"""
01_lsn_walk.py
--------------
观察每次 INSERT 之后 LSN 是怎么递增的，以及每行平均产生多少 WAL。

依赖：
    pip install "psycopg[binary]>=3.1"

运行：
    python 01_lsn_walk.py
"""

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 lsn_now(cur: psycopg.Cursor) -> str:
    cur.execute("SELECT pg_current_wal_lsn();")
    return cur.fetchone()[0]


def lsn_diff(cur: psycopg.Cursor, a: str, b: str) -> int:
    cur.execute("SELECT pg_wal_lsn_diff(%s, %s);", (a, b))
    return int(cur.fetchone()[0])


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

        banner("1. LSN 与 WAL Segment 文件名")
        cur.execute(
            "SELECT pg_current_wal_lsn() AS lsn, pg_walfile_name(pg_current_wal_lsn()) AS file;"
        )
        lsn, fname = cur.fetchone()
        print(f"当前 LSN     = {lsn}")
        print(f"对应文件     = {fname}")
        cur.execute(
            "SELECT pg_current_wal_insert_lsn(), pg_current_wal_flush_lsn();"
        )
        ins, flush = cur.fetchone()
        print(f"insert_lsn   = {ins}  (写入 wal_buffers 的位置)")
        print(f"flush_lsn    = {flush} (已 fsync 到磁盘的位置)")

        banner("2. 每次 INSERT 之后 LSN 增量")
        prev = lsn_now(cur)
        print(f"{'op':<30} {'lsn':<16} {'+bytes':>10}")
        print("-" * 60)
        print(f"{'起点':<30} {prev:<16} {'-':>10}")

        for i in range(1, 6):
            cur.execute(
                "INSERT INTO ch10_lsn_demo (payload) VALUES (%s);",
                (f"row-{i:04d}",),
            )
            now = lsn_now(cur)
            delta = lsn_diff(cur, now, prev)
            print(f"{'INSERT row-' + str(i):<30} {now:<16} {delta:>10}")
            prev = now

        banner("3. 一次性插入 1000 行的 WAL 总量")
        before = lsn_now(cur)
        cur.execute(
            "INSERT INTO ch10_lsn_demo (payload) "
            "SELECT 'bulk-' || g FROM generate_series(1, 1000) g;"
        )
        after = lsn_now(cur)
        total = lsn_diff(cur, after, before)
        print(f"插入前 LSN   = {before}")
        print(f"插入后 LSN   = {after}")
        print(f"总写入 WAL   = {total} 字节  ≈ {total/1024:.1f} KB")
        print(f"平均每行     ≈ {total/1000:.1f} 字节")
        print(
            "\n注意：FPI 会让某些行特别贵——尤其是 checkpoint 之后第一次写到的页。"
        )

        banner("4. 不同操作对 WAL 的开销")
        cases = [
            ("INSERT 100 rows",
             "INSERT INTO ch10_lsn_demo (payload) SELECT 'x' FROM generate_series(1,100);"),
            ("UPDATE 100 rows (WAL 含老元组的死指针 + 新元组)",
             "UPDATE ch10_lsn_demo SET payload = 'y' WHERE id <= 100;"),
            ("DELETE 100 rows (只标 t_xmax，不实际删)",
             "DELETE FROM ch10_lsn_demo WHERE id <= 100;"),
            ("CREATE INDEX (大量元组要写)",
             "CREATE INDEX idx_ch10_lsn_demo_payload ON ch10_lsn_demo (payload);"),
            ("DROP INDEX (轻量元数据)",
             "DROP INDEX idx_ch10_lsn_demo_payload;"),
        ]
        for name, sql in cases:
            b = lsn_now(cur)
            cur.execute(sql)
            a = lsn_now(cur)
            d = lsn_diff(cur, a, b)
            print(f"  {name:<55} +{d:>8} B")

        banner("5. WAL Segment 文件 (16MB) 与 LSN 的换算")
        cur.execute(
            """
            SELECT pg_walfile_name('0/00000000') AS f0,
                   pg_walfile_name('0/01000000') AS f1,
                   pg_walfile_name('1/00000000') AS f256,
                   pg_walfile_name('FF/FFFFFFFF') AS f_high;
            """
        )
        f0, f1, f256, fh = cur.fetchone()
        print(f"  LSN 0/00000000 → {f0}")
        print(f"  LSN 0/01000000 → {f1}   (跨过一个 16MB 段)")
        print(f"  LSN 1/00000000 → {f256} (跨过 256 个段)")
        print(f"  LSN FF/FFFFFFFF → {fh}  (基本上是上界)")


if __name__ == "__main__":
    main()
