"""
Ch7 配套代码 1 / 4 —— 事务 + WATCH 实现乐观锁转账

演示：
  1. 简单 MULTI/EXEC 的命令排队 + 一次执行
  2. WATCH 实现「读改写」乐观锁
  3. 故意制造并发冲突，看 EXEC 返回 nil 触发 WatchError 后自动重试
"""

import threading
import time
import redis

POOL = redis.ConnectionPool(host="127.0.0.1", port=6379, decode_responses=True)


def section(title: str) -> None:
    print("\n" + "=" * 62)
    print(title)
    print("=" * 62)


def demo_basic_transaction() -> None:
    section("Demo 1: 基础 MULTI/EXEC —— 命令排队 + 一次执行")
    r = redis.Redis(connection_pool=POOL)
    r.set("balance", 100)

    pipe = r.pipeline(transaction=True)
    pipe.incrby("balance", 50)
    pipe.decrby("balance", 30)
    pipe.get("balance")
    results = pipe.execute()

    print(f"  事务执行返回 = {results}")
    print(f"  balance 最终 = {r.get('balance')}  (期望 120)")
    r.delete("balance")


def demo_exec_runtime_error() -> None:
    section("Demo 2: 执行时错误 —— 其他命令依然执行（无回滚！）")
    r = redis.Redis(connection_pool=POOL)
    r.set("counter", "abc")
    r.delete("k1", "k2")

    pipe = r.pipeline(transaction=True)
    pipe.set("k1", "v1")
    pipe.incr("counter")
    pipe.set("k2", "v2")
    try:
        results = pipe.execute(raise_on_error=False)
    except Exception as e:
        results = [e]

    for i, res in enumerate(results, 1):
        print(f"  cmd{i} -> {res!r}")
    print(f"\n  k1 = {r.get('k1')}  (前面命令照常执行)")
    print(f"  k2 = {r.get('k2')}  (后面命令也照常执行)")
    print("  💡 Redis 事务不支持回滚，中间出错前后命令都生效")
    r.delete("counter", "k1", "k2")


def transfer_with_watch(from_acc: str, to_acc: str, amount: int, label: str) -> bool:
    """用 WATCH 实现乐观锁转账。返回 True = 成功，False = 余额不足。"""
    r = redis.Redis(connection_pool=POOL)
    retries = 0
    while True:
        try:
            with r.pipeline() as pipe:
                pipe.watch(from_acc, to_acc)

                from_bal = int(pipe.get(from_acc) or 0)
                to_bal = int(pipe.get(to_acc) or 0)

                if from_bal < amount:
                    pipe.unwatch()
                    print(f"  [{label}] ❌ 余额不足 ({from_bal} < {amount})")
                    return False

                pipe.multi()
                pipe.set(from_acc, from_bal - amount)
                pipe.set(to_acc, to_bal + amount)
                pipe.execute()
                print(f"  [{label}] ✅ 转账成功 (重试 {retries} 次)")
                return True
        except redis.WatchError:
            retries += 1
            time.sleep(0.001)
            if retries > 50:
                print(f"  [{label}] 🚫 重试过多放弃")
                return False


def demo_watch_optimistic_lock() -> None:
    section("Demo 3: WATCH 乐观锁 —— 单线程顺利转账")
    r = redis.Redis(connection_pool=POOL)
    r.set("acc:alice", 1000)
    r.set("acc:bob", 100)

    transfer_with_watch("acc:alice", "acc:bob", 200, "T1")

    print(f"\n  Alice 余额 = {r.get('acc:alice')}  (期望 800)")
    print(f"  Bob   余额 = {r.get('acc:bob')}    (期望 300)")


def demo_watch_conflict() -> None:
    section("Demo 4: WATCH 并发冲突 —— 50 线程并发转账，最终一致")
    r = redis.Redis(connection_pool=POOL)
    r.set("acc:alice", 10000)
    r.set("acc:bob", 0)

    THREADS = 50
    AMOUNT_PER = 10

    def worker(i: int) -> None:
        transfer_with_watch("acc:alice", "acc:bob", AMOUNT_PER, f"T{i}")

    threads = [threading.Thread(target=worker, args=(i,)) for i in range(THREADS)]
    t0 = time.time()
    for t in threads: t.start()
    for t in threads: t.join()
    elapsed = time.time() - t0

    alice = int(r.get("acc:alice"))
    bob = int(r.get("acc:bob"))
    expected_total = 10000
    print(f"\n  耗时 {elapsed:.2f}s")
    print(f"  Alice + Bob = {alice + bob}  (期望 {expected_total})")
    print(f"  Bob 收到    = {bob}            (期望 {THREADS * AMOUNT_PER})")
    if alice + bob == expected_total and bob == THREADS * AMOUNT_PER:
        print("  ✅ 余额守恒 —— WATCH 乐观锁工作正常")
    else:
        print("  ❌ 出现不一致")

    r.delete("acc:alice", "acc:bob")


if __name__ == "__main__":
    try:
        demo_basic_transaction()
        demo_exec_runtime_error()
        demo_watch_optimistic_lock()
        demo_watch_conflict()
    except redis.ConnectionError as e:
        print(f"❌ Redis 连接失败: {e}")
