#!/usr/bin/env python3
"""
04_partman_simulation.py
========================
用纯 Python + SQL 模拟 pg_partman 的核心功能：
  - 自动预建未来 N 个月分区
  - 自动 detach + drop 超过保留期的老分区
  - 维护一张配置表记录每张分区表的策略

适合无法装 pg_partman 扩展（如云托管 PG）的场景，思路完全一致。

用法：
    python3 04_partman_simulation.py setup       # 初始化配置 + 创建 demo 分区表
    python3 04_partman_simulation.py maintain    # 跑一次维护（可加到 cron）
    python3 04_partman_simulation.py status      # 查看当前分区
    python3 04_partman_simulation.py teardown    # 清理
"""
import os
import sys
from datetime import date, timedelta

import psycopg

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

CONFIG_TABLE = "ch16_my_partman_config"
PARENT_TABLE = "ch16_ts_logs"
PREMAKE = 4              # 预建未来 4 个月
RETENTION_MONTHS = 12    # 保留 12 个月
INTERVAL = "month"


def add_months(d: date, n: int) -> date:
    """日期加 n 个月（n 可为负数），返回月初"""
    y = d.year
    m = d.month + n
    while m > 12:
        m -= 12; y += 1
    while m < 1:
        m += 12; y -= 1
    return date(y, m, 1)


def setup(cur):
    """初始化：建配置表 + 父表 + 当月分区"""
    print("[setup] 创建配置表...")
    cur.execute(f"""
        CREATE TABLE IF NOT EXISTS {CONFIG_TABLE} (
            parent_table TEXT PRIMARY KEY,
            partition_col TEXT NOT NULL,
            interval_kind TEXT NOT NULL,
            premake INT NOT NULL DEFAULT 4,
            retention_months INT NOT NULL DEFAULT 12,
            last_maintenance TIMESTAMPTZ
        );
    """)

    print(f"[setup] 创建父表 {PARENT_TABLE}（按月 RANGE 分区）...")
    cur.execute(f"DROP TABLE IF EXISTS {PARENT_TABLE} CASCADE;")
    cur.execute(f"""
        CREATE TABLE {PARENT_TABLE} (
            id BIGSERIAL,
            ts TIMESTAMPTZ NOT NULL,
            level TEXT,
            msg TEXT,
            PRIMARY KEY (id, ts)
        ) PARTITION BY RANGE (ts);
    """)

    # 建当月分区
    today = date.today()
    cur_start = today.replace(day=1)
    cur_end = add_months(cur_start, 1)
    pname = f"{PARENT_TABLE}_{cur_start.strftime('%Y_%m')}"
    cur.execute(f"""
        CREATE TABLE {pname} PARTITION OF {PARENT_TABLE}
            FOR VALUES FROM ('{cur_start}') TO ('{cur_end}');
    """)
    print(f"  ✓ 已建当月分区 {pname}")

    # 注册到配置表
    cur.execute(f"""
        INSERT INTO {CONFIG_TABLE}
            (parent_table, partition_col, interval_kind, premake, retention_months)
        VALUES ('{PARENT_TABLE}', 'ts', '{INTERVAL}', {PREMAKE}, {RETENTION_MONTHS})
        ON CONFLICT (parent_table) DO UPDATE
            SET premake = EXCLUDED.premake, retention_months = EXCLUDED.retention_months;
    """)
    print("  ✓ 配置已注册")


def list_partitions(cur, parent: str):
    cur.execute(f"""
        SELECT c.relname,
               pg_get_expr(c.relpartbound, c.oid) AS bound
        FROM pg_class c
        JOIN pg_inherits i ON i.inhrelid = c.oid
        WHERE i.inhparent = %s::regclass
        ORDER BY c.relname;
    """, (parent,))
    return cur.fetchall()


def maintain(cur):
    """逐一处理配置表中所有分区表"""
    cur.execute(f"SELECT parent_table, premake, retention_months FROM {CONFIG_TABLE};")
    for parent, premake, retention in cur.fetchall():
        print(f"\n[maintain] 处理 {parent}（premake={premake}, retention={retention}个月）")
        existing = {row[0] for row in list_partitions(cur, parent)}
        print(f"  当前分区数: {len(existing)}")

        # ========= 1. 预建未来分区 =========
        today = date.today().replace(day=1)
        for i in range(premake + 1):     # 包含当月，所以 +1
            start = add_months(today, i)
            end = add_months(start, 1)
            name = f"{parent}_{start.strftime('%Y_%m')}"
            if name in existing:
                continue
            print(f"  + 创建未来分区 {name} [{start} ~ {end})")
            cur.execute(f"""
                CREATE TABLE {name} PARTITION OF {parent}
                    FOR VALUES FROM ('{start}') TO ('{end}');
            """)

        # ========= 2. 清理过期分区 =========
        cutoff = add_months(today, -retention)
        print(f"  保留截止日期: {cutoff}（更早的将被 detach + drop）")
        for name, bound in list_partitions(cur, parent):
            # 解析 bound 字符串：FOR VALUES FROM ('2024-01-01 ...') TO ('2024-02-01 ...')
            import re
            m = re.search(r"FROM \('([\d-]+)", bound or "")
            if not m:
                continue
            try:
                p_start = date.fromisoformat(m.group(1))
            except ValueError:
                continue
            if p_start < cutoff:
                print(f"  - 删除过期分区 {name}（start={p_start}）")
                # PG 14+ 可以 CONCURRENTLY 不锁表
                try:
                    cur.execute(f"ALTER TABLE {parent} DETACH PARTITION {name} CONCURRENTLY;")
                except psycopg.Error:
                    cur.execute(f"ALTER TABLE {parent} DETACH PARTITION {name};")
                cur.execute(f"DROP TABLE {name};")

        cur.execute(f"UPDATE {CONFIG_TABLE} SET last_maintenance = now() WHERE parent_table = '{parent}';")
        print(f"  ✓ {parent} 维护完成")


def status(cur):
    cur.execute(f"""
        SELECT parent_table, premake, retention_months, last_maintenance
        FROM {CONFIG_TABLE};
    """)
    print("\n=== 配置 ===")
    for r in cur.fetchall():
        print(f"  {r}")
    print(f"\n=== {PARENT_TABLE} 分区列表 ===")
    for name, bound in list_partitions(cur, PARENT_TABLE):
        print(f"  {name}  {bound}")


def teardown(cur):
    cur.execute(f"DROP TABLE IF EXISTS {PARENT_TABLE} CASCADE;")
    cur.execute(f"DROP TABLE IF EXISTS {CONFIG_TABLE};")
    print("✓ 已清理")


def main():
    cmd = sys.argv[1] if len(sys.argv) > 1 else "status"
    fns = {"setup": setup, "maintain": maintain, "status": status, "teardown": teardown}
    if cmd not in fns:
        print(f"Usage: {sys.argv[0]} {{{'|'.join(fns.keys())}}}", file=sys.stderr)
        return 1
    with psycopg.connect(DSN, autocommit=True) as conn, conn.cursor() as cur:
        fns[cmd](cur)
    return 0


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