"""
04_ssi_serialization_failure.py —— SSI 下捕获 serialization_failure 并重试

依赖：
    pip install "psycopg[binary]>=3.1"

场景：
    多个并发线程在 SERIALIZABLE 隔离级别下让两位医生轮流请假，
    必然出现 SSI 检测到的写偏序冲突（SQLSTATE 40001）。
    本脚本演示「指数退避 + 整事务重试」的标准做法。
"""
from __future__ import annotations

import random
import threading
import time
from dataclasses import dataclass

import psycopg
from psycopg import IsolationLevel
from psycopg.errors import SerializationFailure

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

MAX_RETRIES = 8


@dataclass
class Stats:
    success: int = 0
    failed: int = 0
    retries: int = 0


def reset_doctors() -> None:
    with psycopg.connect(DSN, autocommit=True) as conn:
        conn.execute("UPDATE ch7_doctors SET on_duty = TRUE")


def take_leave(doctor_id: int, stats: Stats) -> None:
    """让 doctor_id 请假；若违反业务规则则不操作；遇到序列化失败则重试。"""
    for attempt in range(MAX_RETRIES):
        try:
            with psycopg.connect(DSN) as conn:
                conn.isolation_level = IsolationLevel.SERIALIZABLE
                with conn.transaction():
                    with conn.cursor() as cur:
                        cur.execute(
                            "SELECT count(*) FROM ch7_doctors WHERE on_duty"
                        )
                        n = cur.fetchone()[0]
                        if n <= 1:
                            return
                        cur.execute(
                            "UPDATE ch7_doctors SET on_duty = FALSE "
                            "WHERE id = %s AND on_duty = TRUE",
                            (doctor_id,),
                        )
                stats.success += 1
                return
        except SerializationFailure:
            stats.retries += 1
            backoff = (2 ** attempt) * 0.005 + random.uniform(0, 0.005)
            time.sleep(backoff)
        except Exception as e:
            print(f"  其他异常：{e!r}")
            stats.failed += 1
            return
    stats.failed += 1


def main() -> None:
    rounds = 50
    print(f"=== SSI 重试演示：进行 {rounds} 轮并发请假 ===")
    overall = Stats()

    for r in range(rounds):
        reset_doctors()
        stats = Stats()
        t1 = threading.Thread(target=take_leave, args=(1, stats))
        t2 = threading.Thread(target=take_leave, args=(2, stats))
        t1.start(); t2.start()
        t1.join();  t2.join()

        with psycopg.connect(DSN) as conn:
            n = conn.execute(
                "SELECT count(*) FROM ch7_doctors WHERE on_duty"
            ).fetchone()[0]

        assert n >= 1, "业务规则被破坏！SSI 应该阻止这种情况"
        overall.success += stats.success
        overall.failed += stats.failed
        overall.retries += stats.retries

    print(
        f"\n汇总：成功 = {overall.success}，"
        f"放弃 = {overall.failed}，"
        f"重试次数 = {overall.retries}"
    )
    print("最终所有轮次中，在班医生数始终 ≥ 1，业务规则被 SSI 完美保护 ✓")


if __name__ == "__main__":
    main()
