"""
04_advisory_lock.py
-------------------
使用 PG 咨询锁实现"分布式定时任务防重入"。
启动多个进程/线程，只有第一个抢到锁的会真正执行，其他直接返回。

依赖: pip install psycopg[binary]>=3.1
运行: python 04_advisory_lock.py            # 单进程多线程演示
     ( 或多终端各开一个: python 04_advisory_lock.py worker1 )
"""
import sys
import threading
import time
import psycopg

DSN = "host=127.0.0.1 port=5432 dbname=learn_pg user=postgres"
LOCK_KEY = 20260417  # 自定义业务 key (bigint)


def cron_task(node_name: str, work_seconds: int = 3):
    """模拟一个分布式 cron 任务，确保同一时刻只有一个节点真正执行。"""
    with psycopg.connect(DSN, autocommit=True) as conn:
        with conn.cursor() as cur:
            # try_advisory_lock = 拿不到立刻返回 false (非阻塞)
            cur.execute("SELECT pg_try_advisory_lock(%s)", (LOCK_KEY,))
            got = cur.fetchone()[0]
            if not got:
                print(f"[{node_name}] 未抢到锁，跳过本次执行")
                return False
            try:
                print(f"[{node_name}] 抢到锁，开始执行任务 ({work_seconds}s)…")
                time.sleep(work_seconds)
                print(f"[{node_name}] 任务完成")
                return True
            finally:
                cur.execute("SELECT pg_advisory_unlock(%s)", (LOCK_KEY,))


def transactional_demo():
    """事务级咨询锁：业务异常自动释放，不会泄漏"""
    print("\n=== 事务级 advisory lock 演示 (异常自动释放) ===")
    try:
        with psycopg.connect(DSN) as conn:
            with conn.transaction():
                with conn.cursor() as cur:
                    cur.execute(
                        "SELECT pg_advisory_xact_lock(%s, %s)", (1, 2)
                    )
                    print("拿到事务级锁")
                    raise RuntimeError("模拟业务异常")
    except RuntimeError as e:
        print(f"业务异常: {e}, 但事务回滚后锁已自动释放")

    # 再次尝试拿同样的锁
    with psycopg.connect(DSN, autocommit=True) as conn:
        with conn.cursor() as cur:
            cur.execute("SELECT pg_try_advisory_xact_lock(%s, %s)", (1, 2))
            print("再次尝试 try_advisory_xact_lock = ", cur.fetchone()[0])


def main():
    if len(sys.argv) > 1:
        cron_task(sys.argv[1], work_seconds=5)
        return

    threads = [
        threading.Thread(target=cron_task, args=(f"node-{i}", 2))
        for i in range(5)
    ]
    for t in threads:
        t.start()
    for t in threads:
        t.join()

    transactional_demo()


if __name__ == "__main__":
    main()
