"""
Ch14 配套代码 4 / 7 —— 完整秒杀方案
================================================================

包含：
  ① 用户维度滑动窗口限流（ZSet）
  ② 库存原子扣减（Lua）
  ③ 用户去重（Set）
  ④ 异步下单消息（Stream）

整套核心逻辑只用 2 个 Lua 脚本：
  - rate_limit.lua  限流
  - seckill.lua     扣库存 + 去重 + 推消息

运行：
  python 01_init_stock.py 1001 1000
  python 04_seckill_full.py
  # 另开终端运行 python 05_consumer.py 看消费
"""

import os
import sys
import time
import uuid
import random
import threading
import redis

ITEM_ID    = 1001
STOCK_KEY  = f"seckill:stock:{ITEM_ID}"
USERS_KEY  = f"seckill:users:{ITEM_ID}"
STREAM_KEY = "seckill:orders"
RATE_LIMIT = 10        # 每用户每秒最多 10 次抢购请求
RATE_WINDOW_MS = 1000

THREADS    = 100
PER_THREAD = 20

LUA_RATE = """
-- KEYS[1]=rate:user:xxx ARGV[1]=now_ms ARGV[2]=window_ms
-- ARGV[3]=limit       ARGV[4]=uuid
local now = tonumber(ARGV[1])
local win = tonumber(ARGV[2])
local lim = tonumber(ARGV[3])
redis.call('ZREMRANGEBYSCORE', KEYS[1], 0, now - win)
local cnt = redis.call('ZCARD', KEYS[1])
if cnt >= lim then return 0 end
redis.call('ZADD', KEYS[1], now, ARGV[4])
redis.call('PEXPIRE', KEYS[1], win)
return 1
"""

LUA_FILE = os.path.join(os.path.dirname(__file__), "lua", "seckill.lua")

stats = {"ok": 0, "no_stock": 0, "duplicate": 0, "rate_limited": 0, "no_item": 0}
stats_lock = threading.Lock()


def do_seckill(r, rate_script, sk_script, user_id):
    now_ms = int(time.time() * 1000)
    rid = uuid.uuid4().hex
    if rate_script(keys=[f"rate:user:{user_id}"],
                   args=[now_ms, RATE_WINDOW_MS, RATE_LIMIT, rid],
                   client=r) == 0:
        return "rate_limited"

    order_id = uuid.uuid4().hex
    ret = sk_script(keys=[STOCK_KEY, USERS_KEY, STREAM_KEY],
                    args=[user_id, ITEM_ID, order_id],
                    client=r)
    return {1: "ok", 0: "no_stock", -1: "duplicate", -2: "no_item"}[ret]


def worker(rate_script, sk_script):
    rr = redis.Redis(host="127.0.0.1", port=6379, decode_responses=True)
    local = {"ok": 0, "no_stock": 0, "duplicate": 0, "rate_limited": 0, "no_item": 0}
    for _ in range(PER_THREAD):
        user_id = f"user_{random.randint(1, THREADS * 5)}"
        local[do_seckill(rr, rate_script, sk_script, user_id)] += 1
    with stats_lock:
        for k, v in local.items():
            stats[k] += v


def main():
    r = redis.Redis(host="127.0.0.1", port=6379, decode_responses=True)
    initial = int(r.get(STOCK_KEY) or 0)
    if initial == 0:
        print("❌ 请先运行: python 01_init_stock.py 1001 1000")
        sys.exit(1)

    rate_script = r.register_script(LUA_RATE)
    sk_script   = r.register_script(open(LUA_FILE).read())

    print("=" * 60)
    print(f"  完整秒杀方案 · {THREADS} 线程 × {PER_THREAD} 次")
    print(f"  初始库存 = {initial}")
    print("=" * 60)

    t0 = time.time()
    ts = [threading.Thread(target=worker, args=(rate_script, sk_script))
          for _ in range(THREADS)]
    for t in ts: t.start()
    for t in ts: t.join()
    cost = time.time() - t0

    final = int(r.get(STOCK_KEY) or 0)
    sold = initial - final
    total = sum(stats.values())
    qps = total / cost if cost > 0 else 0
    queue_len = r.xlen(STREAM_KEY)

    print(f"  总请求    : {total}")
    print(f"  耗时      : {cost:.2f}s   QPS ≈ {qps:.0f}")
    print(f"  抢购成功  : {stats['ok']}")
    print(f"  库存不足  : {stats['no_stock']}")
    print(f"  重复购买  : {stats['duplicate']}")
    print(f"  被限流    : {stats['rate_limited']}")
    print(f"  剩余库存  : {final}  (卖出 {sold})")
    print(f"  Stream 消息 : {queue_len}")
    print(f"  ✅ 一致性   : 卖出={sold} == 成功={stats['ok']} == 消息={queue_len}"
          if sold == stats['ok'] == queue_len else
          "  ❌ 数据不一致，请检查")


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