#!/usr/bin/env python3
"""UUIDv4 vs UUIDv7 vs bigint as a primary key, measured on PostgreSQL 18.

Each run creates a fresh table, inserts ROWS rows in committed batches,
then records elapsed time, WAL generated, table/index size, leaf density,
and buffer counts for two read queries. Results go to results.json.
"""
import json
import os
import subprocess
import sys
import time

PG = os.path.expanduser("~/products/c13/pg18/install/bin")
PORT = os.environ.get("PGPORT", "5418")
HOST = os.path.expanduser("~/products/c13/pg18/run")
DB = "lab"
ROWS = int(os.environ.get("ROWS", "5000000"))
BATCH = int(os.environ.get("BATCH", "100000"))
RUNS = int(os.environ.get("RUNS", "3"))
HERE = os.path.dirname(os.path.abspath(__file__))

KEYS = {
    "bigint": "bigint GENERATED ALWAYS AS IDENTITY",
    "uuidv4": "uuid DEFAULT gen_random_uuid()",
    "uuidv7": "uuid DEFAULT uuidv7()",
}


def psql(sql):
    out = subprocess.run(
        [f"{PG}/psql", "-X", "-q", "-At", "-v", "ON_ERROR_STOP=1",
         "-h", HOST, "-p", PORT, "-U", "postgres", "-d", DB, "-c", sql],
        check=True, capture_output=True, text=True)
    return out.stdout.strip()


def explain(sql, prefix=""):
    raw = psql(f"{prefix} EXPLAIN (ANALYZE, BUFFERS, FORMAT JSON) {sql}")
    top = json.loads(raw)[0]["Plan"]
    # Top node only: parent nodes already include their children's buffers.
    return {
        "node": top["Node Type"],
        "shared_hit": top["Shared Hit Blocks"],
        "shared_read": top["Shared Read Blocks"],
    }


def one_run(kind):
    psql("DROP TABLE IF EXISTS t")
    psql(f"CREATE TABLE t (id {KEYS[kind]} PRIMARY KEY, payload text NOT NULL)")
    psql("CHECKPOINT")
    lsn0 = psql("SELECT pg_current_wal_lsn()")
    t0 = time.monotonic()
    psql(f"""
DO $$
BEGIN
  FOR i IN 1..{ROWS // BATCH} LOOP
    INSERT INTO t (payload)
    SELECT md5(g::text) FROM generate_series(1, {BATCH}) g;
    COMMIT;
  END LOOP;
END $$""")
    elapsed = round((time.monotonic() - t0) * 1000)
    wal = int(psql(f"SELECT pg_wal_lsn_diff(pg_current_wal_lsn(), '{lsn0}')"))
    psql("VACUUM ANALYZE t")
    sizes = psql("SELECT pg_relation_size('t'), pg_relation_size('t_pkey')").split("|")
    density = float(psql("SELECT avg_leaf_density FROM pgstatindex('t_pkey')"))
    psql("CHECKPOINT")
    # Read 1: 10,000 random existing keys, each fetched through the primary key.
    lookup = explain(
        "SELECT count(*) FROM k JOIN t USING (id)",
        prefix="CREATE TEMP TABLE k AS SELECT id FROM t ORDER BY random() LIMIT 10000; "
               "ANALYZE k; SET enable_hashjoin = off; SET enable_mergejoin = off;")
    # Read 2: the 1,000 newest rows by key order.
    newest = explain("SELECT * FROM t ORDER BY id DESC LIMIT 1000")
    return {
        "elapsed_ms": elapsed,
        "wal_bytes": wal,
        "table_bytes": int(sizes[0]),
        "index_bytes": int(sizes[1]),
        "leaf_density": density,
        "lookup_10k": lookup,
        "newest_1k": newest,
    }


def main():
    psql("CREATE EXTENSION IF NOT EXISTS pgstattuple")
    meta = {
        "version": psql("SELECT version()"),
        "settings": dict(line.split("|") for line in psql(
            "SELECT name, setting FROM pg_settings WHERE name IN "
            "('shared_buffers','wal_level','full_page_writes','max_wal_size',"
            "'checkpoint_timeout','synchronous_commit','work_mem')").splitlines()),
        "rows": ROWS, "batch": BATCH, "runs": RUNS,
    }
    results = {"meta": meta, "runs": {}}
    for kind in sys.argv[1:] or list(KEYS):
        results["runs"][kind] = []
        for n in range(RUNS):
            r = one_run(kind)
            results["runs"][kind].append(r)
            print(kind, n + 1, json.dumps(r), flush=True)
    with open(os.path.join(HERE, "results.json"), "w") as f:
        json.dump(results, f, indent=2)


if __name__ == "__main__":
    main()
