#!/usr/bin/env python3
"""
05_postgres_fdw.py
==================
演示 postgres_fdw（PG 内置扩展）的「跨库联邦查询」能力。

场景：本地库 (learn_pg) 通过 FDW 访问远端库的表，模拟「假分片」。
为了演示方便，我们用「同一个 PG 实例的另一个数据库」当作远端：
   - learn_pg            → 本地库（建外部表）
   - learn_pg_remote     → 远程库（实际数据放这里）

用法：
    python3 05_postgres_fdw.py setup    # 创建远程库 + 远程数据 + FDW 配置
    python3 05_postgres_fdw.py query    # 跨库查询 + 看 EXPLAIN 下推
    python3 05_postgres_fdw.py join     # 跨库 JOIN
    python3 05_postgres_fdw.py teardown # 清理
"""
import os
import sys

import psycopg
from psycopg import sql

LOCAL_DSN = os.environ.get(
    "PG_LOCAL",
    "host=127.0.0.1 port=5432 dbname=learn_pg user=postgres password=postgres",
)
ADMIN_DSN = os.environ.get(
    "PG_ADMIN",
    "host=127.0.0.1 port=5432 dbname=postgres user=postgres password=postgres",
)
REMOTE_DBNAME = "learn_pg_remote"


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


def setup():
    hr("Step 1: 创建「远端」数据库")
    with psycopg.connect(ADMIN_DSN, autocommit=True) as conn, conn.cursor() as cur:
        cur.execute(f"DROP DATABASE IF EXISTS {REMOTE_DBNAME};")
        cur.execute(f"CREATE DATABASE {REMOTE_DBNAME};")
        print(f"  ✓ 已创建数据库 {REMOTE_DBNAME}")

    hr("Step 2: 在远端建表 + 灌数据")
    remote_dsn = LOCAL_DSN.replace("dbname=learn_pg", f"dbname={REMOTE_DBNAME}")
    with psycopg.connect(remote_dsn, autocommit=True) as conn, conn.cursor() as cur:
        cur.execute("""
            CREATE TABLE ch16_archive_orders (
                id BIGSERIAL PRIMARY KEY,
                user_id BIGINT NOT NULL,
                amount NUMERIC(12,2),
                status TEXT,
                created_at TIMESTAMPTZ NOT NULL
            );
        """)
        cur.execute("""
            INSERT INTO ch16_archive_orders (user_id, amount, status, created_at)
            SELECT
                (random()*100+1)::BIGINT,
                (random()*1000)::NUMERIC(12,2),
                (ARRAY['paid','done'])[1+(random()*1)::INT],
                '2024-01-01'::timestamptz + (random() * INTERVAL '364 day')
            FROM generate_series(1, 10000);
        """)
        cur.execute("CREATE INDEX ON ch16_archive_orders (user_id);")
        cur.execute("CREATE INDEX ON ch16_archive_orders (created_at);")
        cur.execute("ANALYZE ch16_archive_orders;")
        print("  ✓ 远端表 ch16_archive_orders 建好，灌入 10,000 行")

    hr("Step 3: 在本地库装 postgres_fdw + 配置远端")
    with psycopg.connect(LOCAL_DSN, autocommit=True) as conn, conn.cursor() as cur:
        cur.execute("CREATE EXTENSION IF NOT EXISTS postgres_fdw;")
        print("  ✓ 装载 postgres_fdw 扩展")

        # 解析 LOCAL_DSN 拿 host/port，假设远端 PG 也在同一个地址
        import re
        host = (re.search(r"host=(\S+)", LOCAL_DSN) or [None, "127.0.0.1"])[1]
        port = (re.search(r"port=(\S+)", LOCAL_DSN) or [None, "5432"])[1]
        user = (re.search(r"user=(\S+)", LOCAL_DSN) or [None, "postgres"])[1]
        pwd  = (re.search(r"password=(\S+)", LOCAL_DSN) or [None, "postgres"])[1]

        cur.execute("DROP SERVER IF EXISTS pg_archive CASCADE;")
        cur.execute(f"""
            CREATE SERVER pg_archive
            FOREIGN DATA WRAPPER postgres_fdw
            OPTIONS (host '{host}', port '{port}', dbname '{REMOTE_DBNAME}',
                     use_remote_estimate 'true', fetch_size '1000');
        """)
        print("  ✓ 远端 server 已注册（开启 use_remote_estimate）")

        cur.execute(f"""
            CREATE USER MAPPING FOR CURRENT_USER
            SERVER pg_archive
            OPTIONS (user '{user}', password '{pwd}');
        """)
        print("  ✓ 用户映射已建")

        cur.execute("DROP FOREIGN TABLE IF EXISTS ch16_archive_orders;")
        cur.execute("""
            CREATE FOREIGN TABLE ch16_archive_orders (
                id BIGINT,
                user_id BIGINT,
                amount NUMERIC(12,2),
                status TEXT,
                created_at TIMESTAMPTZ
            ) SERVER pg_archive
              OPTIONS (schema_name 'public', table_name 'ch16_archive_orders');
        """)
        print("  ✓ 外部表 ch16_archive_orders 已建")

        # 拉远端统计信息（重要！）
        cur.execute("ANALYZE ch16_archive_orders;")
        print("  ✓ ANALYZE 已拉远端统计信息")


def query():
    hr("跨库 SELECT + EXPLAIN VERBOSE 看下推")
    with psycopg.connect(LOCAL_DSN, autocommit=True) as conn, conn.cursor() as cur:
        sql_q = "SELECT * FROM ch16_archive_orders WHERE user_id = 50 LIMIT 5"
        print(f"\nSQL: {sql_q}")
        cur.execute(sql_q)
        for r in cur.fetchall():
            print(" ", r)

        print("\n=== EXPLAIN VERBOSE ===")
        cur.execute(f"EXPLAIN (VERBOSE) {sql_q}")
        for r in cur.fetchall():
            print(" ", r[0])
        print("\n  ↑ 注意 'Remote SQL'：WHERE / LIMIT 都被下推到远程库执行")

        print("\n=== EXPLAIN 聚合下推（PG 10+）===")
        cur.execute("EXPLAIN (VERBOSE) SELECT count(*) FROM ch16_archive_orders WHERE status = 'paid'")
        for r in cur.fetchall():
            print(" ", r[0])


def join_demo():
    hr("跨库 JOIN：本地 users JOIN 远端 ch16_archive_orders")
    with psycopg.connect(LOCAL_DSN, autocommit=True) as conn, conn.cursor() as cur:
        # 建一张本地小表
        cur.execute("DROP TABLE IF EXISTS ch16_local_users;")
        cur.execute("""
            CREATE TABLE ch16_local_users (
                id BIGINT PRIMARY KEY,
                name TEXT
            );
        """)
        cur.execute("""
            INSERT INTO ch16_local_users (id, name)
            SELECT g, 'user-' || g FROM generate_series(1, 100) g;
        """)
        cur.execute("ANALYZE ch16_local_users;")

        sql_q = """
            SELECT u.name, count(o.id) AS order_cnt, sum(o.amount) AS total
            FROM ch16_local_users u
            JOIN ch16_archive_orders o ON o.user_id = u.id
            WHERE u.id <= 5
            GROUP BY u.name
            ORDER BY u.name;
        """
        print(f"\nSQL:\n{sql_q}")
        cur.execute(sql_q)
        for r in cur.fetchall():
            print(" ", r)

        print("\n=== EXPLAIN VERBOSE ===")
        cur.execute(f"EXPLAIN (VERBOSE, ANALYZE) {sql_q}")
        for r in cur.fetchall():
            print(" ", r[0])
        print("\n  ↑ 跨 server 的 JOIN 不能整体下推，PG 会从远端拉数据回来本地 JOIN")
        print("    所以远端 WHERE 越精确越好（这里 u.id <= 5 限制只拉 5 个用户的订单）")


def teardown():
    hr("清理")
    with psycopg.connect(LOCAL_DSN, autocommit=True) as conn, conn.cursor() as cur:
        cur.execute("DROP FOREIGN TABLE IF EXISTS ch16_archive_orders;")
        cur.execute("DROP SERVER IF EXISTS pg_archive CASCADE;")
        cur.execute("DROP TABLE IF EXISTS ch16_local_users;")
        cur.execute("DROP EXTENSION IF EXISTS postgres_fdw;")
    with psycopg.connect(ADMIN_DSN, autocommit=True) as conn, conn.cursor() as cur:
        cur.execute(f"DROP DATABASE IF EXISTS {REMOTE_DBNAME};")
    print("  ✓ 已清理")


def main():
    cmd = sys.argv[1] if len(sys.argv) > 1 else ""
    fns = {"setup": setup, "query": query, "join": join_demo, "teardown": teardown}
    if cmd not in fns:
        print(f"Usage: {sys.argv[0]} {{{'|'.join(fns.keys())}}}", file=sys.stderr)
        return 1
    fns[cmd]()
    return 0


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