"""
Ch12 配套代码 2 / 4 —— 互斥锁防缓存击穿

场景：
  100 个并发请求同时查 sku:1001，缓存 miss。
  - 无防护：100 个请求全打到 DB
  - 加互斥锁：只有 1 个查 DB 并回写，其他 99 个等回写后命中缓存

依赖：
  pip install redis
"""

import time
import threading
import uuid
import redis


r = redis.Redis(host="127.0.0.1", port=6379, decode_responses=True)


db_query_count = 0
db_lock = threading.Lock()


def fake_db_query(sku_id: int) -> str:
    """模拟一次较慢的 DB 查询。"""
    global db_query_count
    with db_lock:
        db_query_count += 1
    time.sleep(0.2)
    return f"sku-{sku_id}-data"


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


def reset() -> None:
    global db_query_count
    db_query_count = 0
    r.delete("sku:1001", "lock:sku:1001")


def get_sku_no_protection(sku_id: int) -> str:
    """无任何保护的 Cache Aside —— 击穿场景的反例。"""
    key = f"sku:{sku_id}"
    v = r.get(key)
    if v is not None:
        return v
    v = fake_db_query(sku_id)
    r.set(key, v, ex=600)
    return v


def get_sku_with_mutex(sku_id: int, retry_interval: float = 0.05) -> str:
    """互斥锁版本：拿到锁的查 DB，没拿到的等一会重试。"""
    key = f"sku:{sku_id}"
    v = r.get(key)
    if v is not None:
        return v

    lock_key = f"lock:sku:{sku_id}"
    lock_val = uuid.uuid4().hex

    while True:
        if r.set(lock_key, lock_val, nx=True, ex=10):
            try:
                v = r.get(key)
                if v is not None:
                    return v
                v = fake_db_query(sku_id)
                r.set(key, v, ex=600)
                return v
            finally:
                cur = r.get(lock_key)
                if cur == lock_val:
                    r.delete(lock_key)
        else:
            time.sleep(retry_interval)
            v = r.get(key)
            if v is not None:
                return v


def run_concurrent(getter, n_threads: int) -> float:
    threads = []
    start = time.time()
    for _ in range(n_threads):
        t = threading.Thread(target=getter, args=(1001,))
        threads.append(t)
        t.start()
    for t in threads:
        t.join()
    return time.time() - start


def demo() -> None:
    n_threads = 100

    section(f"Demo: {n_threads} 并发线程同时查询冷缓存的 sku:1001")

    reset()
    elapsed_no = run_concurrent(get_sku_no_protection, n_threads)
    db_no = db_query_count
    print(f"  ❌ 无防护      : DB 查询次数 = {db_no:3}  耗时 = {elapsed_no:.2f}s")
    print(f"     → {n_threads} 个请求全打到 DB（击穿现象）")

    reset()
    elapsed_lock = run_concurrent(get_sku_with_mutex, n_threads)
    db_lock_n = db_query_count
    print(f"  ✅ 互斥锁防护  : DB 查询次数 = {db_lock_n:3}  耗时 = {elapsed_lock:.2f}s")
    print(f"     → 只有 {db_lock_n} 个请求真正查 DB，其他都被锁挡住后命中缓存")

    print("\n  说明：")
    print("    - 无防护场景下 DB 被打 100 次（实际生产可能是 10 万次 → DB 雪崩）")
    print("    - 互斥锁后 DB 仅被打 1 次，性能提升数十倍")
    print("    - 代价：拿不到锁的线程多了 retry_interval × N 的延迟")

    r.delete("sku:1001", "lock:sku:1001")


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