"""
Ch11 配套代码 2 / 4 —— 完整安全锁实现

特性：
  ① SET key uuid NX PX  原子加锁
  ② Lua 脚本「校验 + 删除」原子释放
  ③ Lua 脚本「校验 + PEXPIRE」原子续期
  ④ 阻塞 / 非阻塞两种获取语义
  ⑤ Python 上下文管理器，with 块自动释放

并发测试：30 个线程对同一个临界区累加 200 次，
正确实现 → 最终值 = 30 × 200 = 6000；若锁错误，会出现「丢失更新」。
"""

import time
import uuid
import threading
import contextlib

try:
    import redis
except ImportError:
    print("请先 pip install redis"); raise SystemExit(1)


RELEASE_LUA = """
if redis.call('GET', KEYS[1]) == ARGV[1] then
    return redis.call('DEL', KEYS[1])
else
    return 0
end
"""

RENEW_LUA = """
if redis.call('GET', KEYS[1]) == ARGV[1] then
    return redis.call('PEXPIRE', KEYS[1], ARGV[2])
else
    return 0
end
"""


class RedisLock:
    """ 安全的 Redis 分布式锁（V3 + Lua 释放）。"""

    def __init__(self, client: "redis.Redis", key: str, ttl_ms: int = 30_000):
        self.client = client
        self.key = key
        self.ttl_ms = ttl_ms
        self.token = uuid.uuid4().hex
        self._release_sha = None
        self._renew_sha = None

    def _load_scripts(self):
        if self._release_sha is None:
            self._release_sha = self.client.script_load(RELEASE_LUA)
        if self._renew_sha is None:
            self._renew_sha = self.client.script_load(RENEW_LUA)

    def acquire(self, blocking: bool = True, timeout: float = 10.0,
                retry_interval: float = 0.05) -> bool:
        """ 获取锁。blocking=True 时最多重试 timeout 秒。"""
        self._load_scripts()
        deadline = time.time() + timeout
        while True:
            ok = self.client.set(self.key, self.token, nx=True, px=self.ttl_ms)
            if ok:
                return True
            if not blocking or time.time() >= deadline:
                return False
            time.sleep(retry_interval)

    def release(self) -> bool:
        """ 释放锁（仅释放自己的）。"""
        self._load_scripts()
        return bool(self.client.evalsha(self._release_sha, 1, self.key, self.token))

    def renew(self, ttl_ms: int = None) -> bool:
        """ 续期，扩展过期时间到 ttl_ms。"""
        self._load_scripts()
        ttl = ttl_ms or self.ttl_ms
        return bool(self.client.evalsha(self._renew_sha, 1, self.key, self.token, ttl))

    @contextlib.contextmanager
    def hold(self, **kwargs):
        if not self.acquire(**kwargs):
            raise TimeoutError(f"无法获取锁 {self.key}")
        try:
            yield self
        finally:
            self.release()


# -------------------------------------------------------------
# 并发自测：30 线程争抢同一锁，每个累加共享变量 200 次
# 期望最终值 = 6000；若锁失效会出现丢失。
# -------------------------------------------------------------
def demo_concurrent_counter() -> None:
    print("\n" + "=" * 60)
    print("并发测试：30 线程 × 200 次累加（共享 dict + 锁保护）")
    print("=" * 60)

    pool = redis.ConnectionPool(host="127.0.0.1", port=6379, decode_responses=True)
    counter_key = "demo:counter"
    lock_key = "demo:lock"

    r = redis.Redis(connection_pool=pool)
    r.set(counter_key, 0)
    r.delete(lock_key)

    THREADS, PER = 30, 200

    def worker():
        rr = redis.Redis(connection_pool=pool)
        lock = RedisLock(rr, lock_key, ttl_ms=2000)
        for _ in range(PER):
            with lock.hold(timeout=10):
                # 临界区：读 → 改 → 写（如果没锁会丢更新）
                cur = int(rr.get(counter_key))
                rr.set(counter_key, cur + 1)

    ts = [threading.Thread(target=worker) for _ in range(THREADS)]
    t0 = time.time()
    for t in ts: t.start()
    for t in ts: t.join()
    elapsed = time.time() - t0

    final = int(r.get(counter_key))
    expected = THREADS * PER
    print(f"  期望 = {expected}")
    print(f"  实际 = {final}")
    print(f"  耗时 = {elapsed:.2f}s")
    print("  ✅ 锁正确" if final == expected else "  ❌ 出现丢失")
    r.delete(counter_key, lock_key)


# -------------------------------------------------------------
# 校验 release 不会误删别人的锁
# -------------------------------------------------------------
def demo_release_safety() -> None:
    print("\n" + "=" * 60)
    print("安全校验：A 持锁过期 → B 拿锁 → A 调 release，不能误删 B 的锁")
    print("=" * 60)
    r = redis.Redis(host="127.0.0.1", port=6379, decode_responses=True)
    key = "demo:safety"
    r.delete(key)

    a = RedisLock(r, key, ttl_ms=500)  # 0.5s 短 TTL
    b = RedisLock(r, key, ttl_ms=5000)

    assert a.acquire(blocking=False)
    print(f"[A] 加锁 token={a.token[:8]}, TTL=500ms")
    time.sleep(0.7)  # 等过期
    assert b.acquire(blocking=False)
    print(f"[B] 锁已过期，B 重新拿到 token={b.token[:8]}")

    deleted_by_a = a.release()
    print(f"[A] 业务跑完，调 release → 删除? {deleted_by_a}")
    print(f"    （应当为 False —— A 没有误删 B 的锁）")
    assert not deleted_by_a, "❌ A 误删了 B 的锁！"

    still_b_holds = (r.get(key) == b.token)
    print(f"    B 仍持有? {still_b_holds}  ✅")
    b.release()


if __name__ == "__main__":
    try:
        redis.Redis(host="127.0.0.1", port=6379).ping()
    except redis.ConnectionError as e:
        print(f"❌ Redis 连接失败：{e}")
        raise SystemExit(1)

    demo_release_safety()
    demo_concurrent_counter()
