#!/usr/bin/env python3
"""
key_distribution.py
===================

模拟 100 万个 Key 在不同分区数下，经 Murmur2 哈希后的分布均衡度。

输出：
  - 每个分区接收的消息数
  - 最热分区占比、最冷分区占比
  - 变异系数 CV (Coefficient of Variation)
  - 分布柱状图（终端 ASCII）

不依赖 Kafka 库，纯 Python 实现 Murmur2 算法（与 Kafka 客户端默认一致）。

用法：
    python3 key_distribution.py
    python3 key_distribution.py --total 5000000 --partitions 6 12 24 48 --pattern uuid
    python3 key_distribution.py --pattern biased    # 模拟「大客户占 80% 流量」的场景
    python3 key_distribution.py --pattern few       # 基数过小（仅 5 个 Key）
"""

from __future__ import annotations

import argparse
import math
import random
import string
import time
import uuid
from collections import Counter
from typing import Iterable, List


# ---------------------------------------------------------------------------
# Kafka 默认 Partitioner 用的 Murmur2 哈希（与 org.apache.kafka.common.utils.Utils#murmur2 等价）
# ---------------------------------------------------------------------------

_M = 0x5BD1E995
_R = 24
_SEED = 0x9747B28C
_MASK32 = 0xFFFFFFFF


def murmur2(data: bytes) -> int:
    """与 Kafka 客户端 Utils.murmur2 完全一致的实现。返回非负 int（toPositive 后）。"""
    length = len(data)
    h = (_SEED ^ length) & _MASK32
    i = 0
    while length >= 4:
        k = (
            (data[i] & 0xFF)
            | ((data[i + 1] & 0xFF) << 8)
            | ((data[i + 2] & 0xFF) << 16)
            | ((data[i + 3] & 0xFF) << 24)
        )
        k = (k * _M) & _MASK32
        k ^= (k >> _R) & _MASK32
        k = (k * _M) & _MASK32
        h = (h * _M) & _MASK32
        h ^= k
        i += 4
        length -= 4

    if length == 3:
        h ^= (data[i + 2] & 0xFF) << 16
    if length >= 2:
        h ^= (data[i + 1] & 0xFF) << 8
    if length >= 1:
        h ^= data[i] & 0xFF
        h = (h * _M) & _MASK32

    h ^= (h >> 13) & _MASK32
    h = (h * _M) & _MASK32
    h ^= (h >> 15) & _MASK32

    return h & 0x7FFFFFFF  # toPositive


def partition_for(key: str, n: int) -> int:
    return murmur2(key.encode("utf-8")) % n


# ---------------------------------------------------------------------------
# Key 生成器
# ---------------------------------------------------------------------------


def gen_keys(pattern: str, total: int) -> Iterable[str]:
    """根据 pattern 生成 Key 流。yield 节省内存。"""
    if pattern == "seq":
        # user_0 ... user_{total-1}
        for i in range(total):
            yield f"user_{i}"
    elif pattern == "uuid":
        for _ in range(total):
            yield uuid.uuid4().hex
    elif pattern == "tenant":
        # 100 个租户，但访问按齐普夫分布——Top 5 占大头
        tenants = [f"tenant_{i}" for i in range(100)]
        weights = [1.0 / (i + 1) for i in range(100)]   # 1, 1/2, 1/3 ... Zipf
        s = sum(weights)
        weights = [w / s for w in weights]
        for _ in range(total):
            yield random.choices(tenants, weights=weights, k=1)[0]
    elif pattern == "biased":
        # 80% 流量打在 5 个大 Key 上，20% 打在普通 Key
        hot = [f"VIP_{c}" for c in "ABCDE"]
        for i in range(total):
            if random.random() < 0.8:
                yield hot[i % 5]
            else:
                yield "normal_" + "".join(
                    random.choices(string.ascii_lowercase + string.digits, k=10)
                )
    elif pattern == "few":
        for i in range(total):
            yield f"group_{i % 5}"
    else:
        raise ValueError(f"unknown pattern: {pattern}")


# ---------------------------------------------------------------------------
# 统计 + 可视化
# ---------------------------------------------------------------------------


def analyse(counts: List[int]) -> dict:
    n = len(counts)
    total = sum(counts)
    if total == 0:
        return {"n": n, "total": 0}
    mean = total / n
    variance = sum((c - mean) ** 2 for c in counts) / n
    stddev = math.sqrt(variance)
    cv = stddev / mean if mean else 0
    return {
        "n": n,
        "total": total,
        "mean": mean,
        "min": min(counts),
        "max": max(counts),
        "min_pct": min(counts) / total * 100,
        "max_pct": max(counts) / total * 100,
        "stddev": stddev,
        "cv": cv,
        "empty": sum(1 for c in counts if c == 0),
    }


def print_bar(counts: List[int], width: int = 60) -> None:
    """终端 ASCII 柱状图。横向显示前 N 个分区。"""
    if not counts:
        return
    max_c = max(counts) or 1
    avg = sum(counts) / len(counts)
    print()
    for i, c in enumerate(counts):
        bar_len = int(c / max_c * width)
        marker = "🔥" if c > avg * 1.5 else "  "
        bar = "█" * bar_len
        print(f"  P{i:>3}  {marker}  {bar} {c:>10,}  ({c/sum(counts)*100:5.2f}%)")
    print()


# ---------------------------------------------------------------------------
# 主流程
# ---------------------------------------------------------------------------


def run_one(pattern: str, total: int, partitions: int) -> None:
    counts = [0] * partitions
    t0 = time.time()
    for k in gen_keys(pattern, total):
        counts[partition_for(k, partitions)] += 1
    elapsed = time.time() - t0

    stats = analyse(counts)
    print(
        f"\n========== pattern={pattern}  total={total:,}  partitions={partitions} =========="
    )
    print(
        f"  耗时 {elapsed:.2f}s  | mean={stats['mean']:,.1f}"
        f"  min={stats['min']:,}({stats['min_pct']:.2f}%)"
        f"  max={stats['max']:,}({stats['max_pct']:.2f}%)"
    )
    print(
        f"  stddev={stats['stddev']:,.1f}  CV={stats['cv']:.4f}"
        f"  empty_partitions={stats['empty']}"
    )
    if stats["cv"] < 0.05:
        verdict = "✅ 极佳：分布几乎完全均匀"
    elif stats["cv"] < 0.15:
        verdict = "🟢 良好：可放心上线"
    elif stats["cv"] < 0.3:
        verdict = "🟡 偏斜：可接受，但要监控热点分区"
    else:
        verdict = "🔴 严重失衡：必须重新设计 Key 或拆 Topic"
    print(f"  评估：{verdict}")
    print_bar(counts)


def main() -> None:
    parser = argparse.ArgumentParser(description="Kafka Key 分布均衡度模拟器")
    parser.add_argument("--total", type=int, default=1_000_000,
                        help="模拟消息总数，默认 100 万")
    parser.add_argument("--partitions", type=int, nargs="+",
                        default=[3, 6, 12, 24, 48],
                        help="要测试的分区数列表")
    parser.add_argument("--pattern", type=str,
                        choices=["seq", "uuid", "tenant", "biased", "few"],
                        default="seq",
                        help="Key 生成模式")
    parser.add_argument("--seed", type=int, default=42,
                        help="随机种子，复现实验")
    args = parser.parse_args()

    random.seed(args.seed)

    print("=" * 70)
    print(f" Kafka Key 分布均衡度模拟器")
    print(f"   pattern={args.pattern}  total={args.total:,}")
    print(f"   partitions={args.partitions}")
    print("=" * 70)

    for n in args.partitions:
        run_one(args.pattern, args.total, n)

    print("\n说明：")
    print("  CV  (Coefficient of Variation) = stddev / mean")
    print("  CV ≤ 0.05  极佳")
    print("  CV ≤ 0.15  良好")
    print("  CV ≤ 0.30  偏斜（需监控）")
    print("  CV >  0.30  失衡（建议改 Key）")


if __name__ == "__main__":
    main()
