#!/usr/bin/env python3
"""
02_hash_partition_demo.py
=========================
HASH 分区均匀性验证：
  1. 建 8 个 HASH 分区，按 user_id 分
  2. 灌入 100,000 个不同 user_id 的事件
  3. 统计每个分区行数，验证哈希分布是否均匀
  4. 演示「WHERE user_id = X」能裁剪到 1 个分区
  5. 演示「WHERE user_id BETWEEN A AND B」无法裁剪（哈希后不连续）

依赖：psycopg[binary]>=3.1
"""
import os
import sys
import statistics

import psycopg

DSN = os.environ.get(
    "PG_DSN",
    "host=127.0.0.1 port=5432 dbname=learn_pg user=postgres password=postgres",
)
N_PARTITIONS = 8
N_ROWS = 100_000


def hr(s):
    print("\n" + "=" * 68)
    print(f"  {s}")
    print("=" * 68)


def main():
    with psycopg.connect(DSN, autocommit=True) as conn, conn.cursor() as cur:
        hr("Step 1: 建 HASH 分区表（8 个分区）")
        cur.execute("DROP TABLE IF EXISTS ch16_hash_demo CASCADE;")
        cur.execute("""
            CREATE TABLE ch16_hash_demo (
                id BIGSERIAL,
                user_id BIGINT NOT NULL,
                payload TEXT,
                PRIMARY KEY (id, user_id)
            ) PARTITION BY HASH (user_id);
        """)
        for i in range(N_PARTITIONS):
            cur.execute(f"""
                CREATE TABLE ch16_hash_demo_p{i} PARTITION OF ch16_hash_demo
                    FOR VALUES WITH (MODULUS {N_PARTITIONS}, REMAINDER {i});
            """)
        print(f"  ✓ 已建 {N_PARTITIONS} 个 HASH 分区 ch16_hash_demo_p0 ~ p{N_PARTITIONS-1}")

        hr(f"Step 2: 灌入 {N_ROWS:,} 行（user_id 从 1 到 {N_ROWS}）")
        # 用单条多行 INSERT 提升插入速度
        cur.execute(f"""
            INSERT INTO ch16_hash_demo (user_id, payload)
            SELECT g, 'event-' || g
            FROM generate_series(1, {N_ROWS}) g;
        """)
        cur.execute("ANALYZE ch16_hash_demo;")
        print(f"  ✓ 灌入完成")

        hr("Step 3: 统计各分区行数")
        cur.execute(f"""
            SELECT relname, pg_stat_get_live_tuples(c.oid)::INT AS rows,
                   pg_size_pretty(pg_relation_size(c.oid)) AS size
            FROM pg_class c
            JOIN pg_inherits i ON i.inhrelid = c.oid
            WHERE i.inhparent = 'ch16_hash_demo'::regclass
            ORDER BY relname;
        """)
        rows = cur.fetchall()
        counts = [r[1] for r in rows]
        ideal = N_ROWS / N_PARTITIONS
        print(f"  分区名             行数        大小        与理想值偏差")
        print(f"  --------------------------------------------------------")
        for name, c, sz in rows:
            diff_pct = (c - ideal) / ideal * 100
            bar = "█" * int(c / ideal * 20)
            print(f"  {name:<18} {c:<10,} {sz:<10}  {diff_pct:+.2f}%  {bar}")
        print(f"  --------------------------------------------------------")
        print(f"  理想值: {ideal:,.0f} 行/分区")
        print(f"  实际标准差: {statistics.stdev(counts):.2f}")
        print(f"  最大偏差: {(max(counts)-min(counts))/ideal*100:.2f}%")
        print(f"  → 哈希分布均匀（PG 内置哈希函数质量很好）")

        hr("Step 4: WHERE user_id = X 触发分区裁剪（点查命中 1 个分区）")
        cur.execute("EXPLAIN SELECT * FROM ch16_hash_demo WHERE user_id = 12345;")
        for r in cur.fetchall():
            print(" ", r[0])

        hr("Step 5: WHERE user_id BETWEEN ... 不能裁剪（哈希后不连续）")
        cur.execute("EXPLAIN SELECT * FROM ch16_hash_demo WHERE user_id BETWEEN 100 AND 200;")
        for r in cur.fetchall():
            print(" ", r[0])
        print("  ↑ 注意所有分区都被扫到（HASH 分区的硬伤）")

        hr("Step 6: WHERE user_id IN (a, b, c) 部分裁剪（PG 11+）")
        cur.execute("EXPLAIN SELECT * FROM ch16_hash_demo WHERE user_id IN (1, 100, 1000);")
        for r in cur.fetchall():
            print(" ", r[0])
        print("  ↑ 各 user_id 算 hash 后落到 ≤3 个分区，PG 11+ 能裁剪")

        hr("结论")
        print("  ✓ HASH 分区适合：均匀分散写入热点 / 点查（=）")
        print("  ✗ HASH 分区不适合：范围查询、按时间归档")
        print("  → 时序数据用 RANGE，高基数随机访问用 HASH")

        return 0


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