"""
Ch10 配套代码 2 / 3 —— Cluster 客户端 API 演示

演示：
  1. 用 redis.RedisCluster (redis-py 4.0+) 连接 Redis Cluster
  2. 集群感知客户端的核心特性：
       - 自动 CRC16 计算 + slot 路由
       - 自动处理 MOVED / ASK 重定向
       - 节点拓扑自动感知（CLUSTER SLOTS）
  3. 当无真实 Cluster 可用时，用纯 Python 模拟一个 6 节点（3 主 3 从）集群

使用：
  - 真集群： python 02_cluster_client.py --real --host 127.0.0.1 --port 7000
  - 模拟器： python 02_cluster_client.py            （默认）
"""

import argparse
import random
from typing import Dict, List, Optional, Tuple

# 复用 01 里的 CRC16 + slot 计算
from importlib import import_module
_crc = import_module("01_crc16_slot") if False else None  # Python 不允许 01_ 开头的模块名

# 简单内联：避免动态导入受限于文件名
CRC16TAB = [
    0x0000, 0x1021, 0x2042, 0x3063, 0x4084, 0x50A5, 0x60C6, 0x70E7,
    0x8108, 0x9129, 0xA14A, 0xB16B, 0xC18C, 0xD1AD, 0xE1CE, 0xF1EF,
    0x1231, 0x0210, 0x3273, 0x2252, 0x52B5, 0x4294, 0x72F7, 0x62D6,
    0x9339, 0x8318, 0xB37B, 0xA35A, 0xD3BD, 0xC39C, 0xF3FF, 0xE3DE,
    0x2462, 0x3443, 0x0420, 0x1401, 0x64E6, 0x74C7, 0x44A4, 0x5485,
    0xA56A, 0xB54B, 0x8528, 0x9509, 0xE5EE, 0xF5CF, 0xC5AC, 0xD58D,
    0x3653, 0x2672, 0x1611, 0x0630, 0x76D7, 0x66F6, 0x5695, 0x46B4,
    0xB75B, 0xA77A, 0x9719, 0x8738, 0xF7DF, 0xE7FE, 0xD79D, 0xC7BC,
    0x48C4, 0x58E5, 0x6886, 0x78A7, 0x0840, 0x1861, 0x2802, 0x3823,
    0xC9CC, 0xD9ED, 0xE98E, 0xF9AF, 0x8948, 0x9969, 0xA90A, 0xB92B,
    0x5AF5, 0x4AD4, 0x7AB7, 0x6A96, 0x1A71, 0x0A50, 0x3A33, 0x2A12,
    0xDBFD, 0xCBDC, 0xFBBF, 0xEB9E, 0x9B79, 0x8B58, 0xBB3B, 0xAB1A,
    0x6CA6, 0x7C87, 0x4CE4, 0x5CC5, 0x2C22, 0x3C03, 0x0C60, 0x1C41,
    0xEDAE, 0xFD8F, 0xCDEC, 0xDDCD, 0xAD2A, 0xBD0B, 0x8D68, 0x9D49,
    0x7E97, 0x6EB6, 0x5ED5, 0x4EF4, 0x3E13, 0x2E32, 0x1E51, 0x0E70,
    0xFF9F, 0xEFBE, 0xDFDD, 0xCFFC, 0xBF1B, 0xAF3A, 0x9F59, 0x8F78,
    0x9188, 0x81A9, 0xB1CA, 0xA1EB, 0xD10C, 0xC12D, 0xF14E, 0xE16F,
    0x1080, 0x00A1, 0x30C2, 0x20E3, 0x5004, 0x4025, 0x7046, 0x6067,
    0x83B9, 0x9398, 0xA3FB, 0xB3DA, 0xC33D, 0xD31C, 0xE37F, 0xF35E,
    0x02B1, 0x1290, 0x22F3, 0x32D2, 0x4235, 0x5214, 0x6277, 0x7256,
    0xB5EA, 0xA5CB, 0x95A8, 0x8589, 0xF56E, 0xE54F, 0xD52C, 0xC50D,
    0x34E2, 0x24C3, 0x14A0, 0x0481, 0x7466, 0x6447, 0x5424, 0x4405,
    0xA7DB, 0xB7FA, 0x8799, 0x97B8, 0xE75F, 0xF77E, 0xC71D, 0xD73C,
    0x26D3, 0x36F2, 0x0691, 0x16B0, 0x6657, 0x7676, 0x4615, 0x5634,
    0xD94C, 0xC96D, 0xF90E, 0xE92F, 0x99C8, 0x89E9, 0xB98A, 0xA9AB,
    0x5844, 0x4865, 0x7806, 0x6827, 0x18C0, 0x08E1, 0x3882, 0x28A3,
    0xCB7D, 0xDB5C, 0xEB3F, 0xFB1E, 0x8BF9, 0x9BD8, 0xABBB, 0xBB9A,
    0x4A75, 0x5A54, 0x6A37, 0x7A16, 0x0AF1, 0x1AD0, 0x2AB3, 0x3A92,
    0xFD2E, 0xED0F, 0xDD6C, 0xCD4D, 0xBDAA, 0xAD8B, 0x9DE8, 0x8DC9,
    0x7C26, 0x6C07, 0x5C64, 0x4C45, 0x3CA2, 0x2C83, 0x1CE0, 0x0CC1,
    0xEF1F, 0xFF3E, 0xCF5D, 0xDF7C, 0xAF9B, 0xBFBA, 0x8FD9, 0x9FF8,
    0x6E17, 0x7E36, 0x4E55, 0x5E74, 0x2E93, 0x3EB2, 0x0ED1, 0x1EF0,
]
CLUSTER_SLOTS = 16384


def crc16(data: bytes) -> int:
    crc = 0x0000
    for b in data:
        crc = ((crc << 8) & 0xFFFF) ^ CRC16TAB[((crc >> 8) ^ b) & 0xFF]
    return crc


def extract_hashtag(key: str) -> Optional[str]:
    s = key.find("{")
    if s == -1:
        return None
    e = key.find("}", s + 1)
    if e == -1 or e == s + 1:
        return None
    return key[s + 1:e]


def key_slot(key: str) -> int:
    tag = extract_hashtag(key)
    target = tag if tag is not None else key
    return crc16(target.encode("utf-8")) % CLUSTER_SLOTS


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


# ──────────────────────────────────────────────────────────────
# 模式 A：真实 RedisCluster
# ──────────────────────────────────────────────────────────────
def run_real_cluster(host: str, port: int) -> None:
    try:
        from redis.cluster import RedisCluster, ClusterNode
    except ImportError:
        print("❌ redis-py 未安装或版本过低（需要 4.0+）")
        print("   pip install 'redis>=4.0'")
        return

    try:
        rc = RedisCluster(
            startup_nodes=[ClusterNode(host, port)],
            decode_responses=True,
            require_full_coverage=False,
        )

        section("Demo A1: 集群基本读写（自动路由）")
        for i in range(5):
            k, v = f"user:{1000 + i}", f"name-{i}"
            rc.set(k, v)
            slot = rc.keyslot(k) if hasattr(rc, "keyslot") else key_slot(k)
            print(f"  SET {k:12} = {v:8}  slot={slot:5}  → 自动路由到正确 master")

        section("Demo A2: 集群拓扑")
        nodes = rc.get_nodes()
        for n in nodes:
            print(f"  {n.host}:{n.port}  role={n.server_type}  name={n.name}")

        section("Demo A3: 跨槽 MGET 会失败，HashTag 拯救")
        try:
            rc.mget("user:1000", "user:1001", "user:1002")
            print("  （部分 redis-py 版本会自动拆分 mget，看不到原始 CROSSSLOT 错误）")
        except Exception as e:
            print(f"  ❌ 跨槽 MGET 错：{e}")

        rc.mset({
            "{user:9999}:profile": "p",
            "{user:9999}:cart":    "c",
            "{user:9999}:orders":  "o",
        })
        result = rc.mget(
            "{user:9999}:profile",
            "{user:9999}:cart",
            "{user:9999}:orders",
        )
        print(f"  ✅ 带 HashTag 的 MGET：{result}")

        rc.delete(
            "{user:9999}:profile", "{user:9999}:cart", "{user:9999}:orders",
            *[f"user:{1000 + i}" for i in range(5)],
        )
    except Exception as e:
        print(f"❌ 连接 Redis Cluster 失败: {e}")
        print("   请确认提供的是「集群模式」节点（cluster-enabled yes）")


# ──────────────────────────────────────────────────────────────
# 模式 B：纯 Python 模拟一个 6 节点（3 主 3 从）集群
# ──────────────────────────────────────────────────────────────
class FakeNode:
    """一个 master 或 slave 节点。master 维护 (slot 范围) 内的 KV"""
    def __init__(self, name: str, port: int, role: str):
        self.name = name              # "A", "B", "C", "A1"...
        self.port = port              # 7000~7005
        self.role = role              # "master" | "slave"
        self.alive = True
        self.master_of: Optional["FakeNode"] = None     # slave→master
        self.slaves: List["FakeNode"] = []
        self.slot_range: Optional[Tuple[int, int]] = None
        self.kv: Dict[str, str] = {}

    def __repr__(self) -> str:
        sr = f"{self.slot_range[0]}~{self.slot_range[1]}" if self.slot_range else "-"
        return f"<{self.name}/{self.port}/{self.role}/slot={sr}/alive={self.alive}>"


class FakeCluster:
    """6 节点 3 主 3 从模拟集群（一切只在内存里）"""
    def __init__(self) -> None:
        self.nodes: List[FakeNode] = []
        # 3 master
        for i, name in enumerate(["A", "B", "C"]):
            self.nodes.append(FakeNode(name, 7000 + i, "master"))
        # 3 slave，分别挂到 A/B/C
        for i, (name, master_idx) in enumerate(zip(["A1", "B1", "C1"], [0, 1, 2])):
            sl = FakeNode(name, 7003 + i, "slave")
            sl.master_of = self.nodes[master_idx]
            self.nodes[master_idx].slaves.append(sl)
            self.nodes.append(sl)
        # 槽位均分
        masters = [n for n in self.nodes if n.role == "master"]
        ranges = self._distribute_slots(len(masters))
        for m, r in zip(masters, ranges):
            m.slot_range = r

    @staticmethod
    def _distribute_slots(n: int) -> List[Tuple[int, int]]:
        base = CLUSTER_SLOTS // n
        extra = CLUSTER_SLOTS % n
        out, cursor = [], 0
        for i in range(n):
            size = base + (1 if i < extra else 0)
            out.append((cursor, cursor + size - 1))
            cursor += size
        return out

    def find_master(self, slot: int) -> Optional[FakeNode]:
        for n in self.nodes:
            if n.role == "master" and n.alive and n.slot_range:
                if n.slot_range[0] <= slot <= n.slot_range[1]:
                    return n
        return None

    def get(self, key: str) -> Tuple[Optional[str], FakeNode]:
        slot = key_slot(key)
        node = self.find_master(slot)
        if node is None:
            raise RuntimeError(f"slot {slot} 无可用 master！")
        return node.kv.get(key), node

    def set(self, key: str, val: str) -> FakeNode:
        slot = key_slot(key)
        node = self.find_master(slot)
        if node is None:
            raise RuntimeError(f"slot {slot} 无可用 master！")
        node.kv[key] = val
        return node

    def fail_node(self, name: str) -> None:
        """模拟节点宕机 + 故障转移"""
        for n in self.nodes:
            if n.name == name:
                n.alive = False
                if n.role == "master" and n.slaves:
                    # 选第一个还活着的 slave 接管
                    for sl in n.slaves:
                        if sl.alive:
                            print(f"  💥 {name} 宕机，从节点 {sl.name} 当选新 master")
                            sl.role = "master"
                            sl.slot_range = n.slot_range
                            sl.master_of = None
                            n.slot_range = None
                            return
                    print(f"  💥 {name} 宕机，但无可用 slave 接管 → slot 不可用！")

    def topology(self) -> str:
        out = []
        for n in self.nodes:
            sr = f"slot {n.slot_range[0]:5}~{n.slot_range[1]:5}" if n.slot_range else " " * 19
            mark = "✓" if n.alive else "✗"
            mast = f" → master={n.master_of.name}" if n.master_of else ""
            out.append(f"  [{mark}] {n.name:3} :{n.port}  {n.role:6}  {sr}{mast}")
        return "\n".join(out)


def run_simulator() -> None:
    cl = FakeCluster()

    section("Demo B1: 集群拓扑（6 节点 3 主 3 从）")
    print(cl.topology())

    section("Demo B2: 集群读写 —— 客户端自动算 slot 找节点")
    keys = ["user:1001", "order:8888", "product:42",
            "session:abc", "cart:7", "stock:gpu"]
    for k in keys:
        node = cl.set(k, f"val-of-{k}")
        slot = key_slot(k)
        print(f"  SET {k:13} → slot={slot:5} → 路由到 {node.name}({node.port})")

    section("Demo B3: 模拟 master B 宕机 + 故障转移")
    print("  宕机前 B 上的数据：", {k: v for k, v in cl.nodes[1].kv.items()})
    cl.fail_node("B")
    print(cl.topology())
    print("\n  下次写 slot 5461~10922 的 key 自动路由到新 master（原 B1）：")
    n = cl.set("order:8888", "rewritten")
    print(f"  SET order:8888 → 路由到 {n.name}({n.port})")
    print("  ⚠️ 但是宕机前的数据丢了 —— 因为模拟器没实现复制；真实 Cluster 中 slave 是异步复制 master 的")

    section("Demo B4: HashTag 强制同节点 演示")
    keys2 = ["{user:1001}:profile", "{user:1001}:cart", "{user:1001}:orders"]
    nodes_set = set()
    for k in keys2:
        n = cl.set(k, "x")
        nodes_set.add(n.name)
        print(f"  SET {k:30} → 落到 {n.name}({n.port})")
    print(f"\n  ✅ 三个 key 全部落到同一节点：{nodes_set} —— 可以原子地 MGET / MULTI / EVAL")


def main() -> None:
    parser = argparse.ArgumentParser(description="Ch10 Cluster 客户端演示")
    parser.add_argument("--real", action="store_true", help="连接真实 Redis Cluster")
    parser.add_argument("--host", default="127.0.0.1")
    parser.add_argument("--port", type=int, default=7000)
    args = parser.parse_args()

    if args.real:
        print(f">>> 模式：连接真实 Cluster {args.host}:{args.port}")
        run_real_cluster(args.host, args.port)
    else:
        print(">>> 模式：纯 Python 模拟器（无需真实 Redis Cluster）")
        run_simulator()


if __name__ == "__main__":
    main()
