"""
01_table_lock_compat.py
-----------------------
演示 PostgreSQL 表级锁的兼容性：
  会话 A 持有 X 模式表锁后，会话 B 请求另一个模式时是否被阻塞。

依赖:  pip install psycopg[binary]>=3.1
运行:  python 01_table_lock_compat.py
连接:  host=127.0.0.1 port=5432 dbname=learn_pg user=postgres
"""
import time
import threading
import psycopg

DSN = "host=127.0.0.1 port=5432 dbname=learn_pg user=postgres"

# 待测试的「持有锁 -> 请求锁」组合（一定不会阻塞 / 一定会阻塞 都列出来对比）
CASES = [
    # (说明,                持锁模式,                 请求 SQL,                    预期是否阻塞)
    ("S(SELECT)+W(INSERT)",  "ACCESS SHARE",          "INSERT INTO ch11_seckill(sku_name,stock) VALUES('t',1)", False),
    ("DDL(DROP)+S(SELECT)",  "ACCESS EXCLUSIVE",      "SELECT count(*) FROM ch11_seckill",                      True),
    ("CREATE INDEX + INSERT","SHARE",                 "INSERT INTO ch11_seckill(sku_name,stock) VALUES('t',1)", True),
    ("VACUUM + INSERT",      "SHARE UPDATE EXCLUSIVE","INSERT INTO ch11_seckill(sku_name,stock) VALUES('t',1)", False),
    ("VACUUM + VACUUM",      "SHARE UPDATE EXCLUSIVE","LOCK TABLE ch11_seckill IN SHARE UPDATE EXCLUSIVE MODE NOWAIT", True),
]


def hold_lock(mode: str, hold_seconds: float, ready_evt: threading.Event):
    """会话 A：拿一把表级锁并持有 hold_seconds 秒"""
    with psycopg.connect(DSN, autocommit=False) as conn:
        with conn.cursor() as cur:
            cur.execute(f"LOCK TABLE ch11_seckill IN {mode} MODE")
            ready_evt.set()
            time.sleep(hold_seconds)
        conn.rollback()  # 释放锁


def try_request(sql: str, timeout_seconds: float = 1.5) -> bool:
    """会话 B：尝试在 timeout 内执行 sql；阻塞超过 timeout 视为「被锁住」"""
    blocked = {"v": False}

    def runner():
        try:
            with psycopg.connect(DSN, autocommit=True) as conn:
                with conn.cursor() as cur:
                    cur.execute(sql)
        except Exception as e:
            # NOWAIT 拿不到锁会抛错也算"被阻塞"
            if "could not obtain lock" in str(e) or "lock_not_available" in str(e):
                blocked["v"] = True

    t = threading.Thread(target=runner, daemon=True)
    t.start()
    t.join(timeout_seconds)
    if t.is_alive():
        blocked["v"] = True
    return blocked["v"]


def main():
    print(f"{'场景':30s} | {'持锁模式':28s} | 实际 | 预期 | 结果")
    print("-" * 90)
    for desc, mode, sql, expect_blocked in CASES:
        ready = threading.Event()
        holder = threading.Thread(
            target=hold_lock, args=(mode, 3.0, ready), daemon=True
        )
        holder.start()
        ready.wait(timeout=2.0)

        time.sleep(0.1)  # 让锁稳定持有
        actual_blocked = try_request(sql, timeout_seconds=1.5)
        ok = "✓" if actual_blocked == expect_blocked else "✗"
        print(
            f"{desc:30s} | {mode:28s} | "
            f"{'阻塞' if actual_blocked else '通过'} | "
            f"{'阻塞' if expect_blocked else '通过'} | {ok}"
        )
        holder.join(timeout=5)


if __name__ == "__main__":
    main()
