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

演示「加分区」操作对 Key 落点的破坏性。

核心结论：分区数从 N1 → N2 后，
    P(同 Key 落点变化) ≈ (N2 - gcd(N1, N2)) / N2

也就是说：分区数加得越多，变化的 Key 就越多。
**这是为什么「加分区」是破坏性操作、必须改用「重建 Topic」的根本原因。**

输出：
  - 一份对比表：N1 个分区 vs N2 个分区下，每个 Key 的落点
  - 总变化率
  - 「按 Key 聚合」的下游业务受影响范围

用法：
    python3 repartition_impact.py
    python3 repartition_impact.py --n1 6 --n2 12 --keys 1000000
    python3 repartition_impact.py --pairs 6:7 6:8 6:12 6:24 12:24 --keys 100000
"""

from __future__ import annotations

import argparse
import math
from collections import Counter, defaultdict

# 复用同目录下的 murmur2
from key_distribution import murmur2, partition_for, gen_keys


def compute_change_rate(n1: int, n2: int, total: int, pattern: str = "seq") -> dict:
    """跑一次实测，返回：
    - changed: 落点变化的 Key 数
    - same:    落点未变的 Key 数
    - rate:    变化率
    - moves:   {(p1, p2): count}  从 P{p1} 流向 P{p2} 的消息数
    """
    moves: dict = defaultdict(int)
    changed = same = 0
    for k in gen_keys(pattern, total):
        p1 = partition_for(k, n1)
        p2 = partition_for(k, n2)
        moves[(p1, p2)] += 1
        if p1 == p2:
            same += 1
        else:
            changed += 1
    return {
        "n1": n1,
        "n2": n2,
        "total": total,
        "changed": changed,
        "same": same,
        "rate": changed / total if total else 0,
        "moves": moves,
    }


def theoretical_rate(n1: int, n2: int) -> float:
    """理论估计：在 Key 哈希均匀的前提下，加分区后的变化比例。

    简化模型：n1 → n2 时落点变化概率 ≈ (n2 - gcd(n1, n2)) / n2
    （非严格，仅做对比直觉。实际还受 Murmur2 输出分布影响。）
    """
    g = math.gcd(n1, n2)
    return (n2 - g) / n2


def print_sample(n1: int, n2: int, sample_keys: int = 20) -> None:
    """随机抽 N 个 Key 打印对比表。"""
    print(f"\n  样例（前 {sample_keys} 个 user_xxxx）：")
    print(f"  {'Key':<14} {'P (N=' + str(n1) + ')':<10} {'P (N=' + str(n2) + ')':<10} {'变化?':<6}")
    print(f"  {'-'*14} {'-'*10} {'-'*10} {'-'*6}")
    for i in range(sample_keys):
        k = f"user_{i}"
        p1 = partition_for(k, n1)
        p2 = partition_for(k, n2)
        flag = "🔴 变" if p1 != p2 else "🟢 同"
        print(f"  {k:<14} {p1:<10} {p2:<10} {flag}")


def print_top_moves(result: dict, top: int = 8) -> None:
    """打印「最常发生的迁移路径 P_i → P_j」"""
    moves = sorted(result["moves"].items(), key=lambda kv: -kv[1])[:top]
    total = result["total"]
    print(f"\n  迁移热力图（最常见 Top-{top}）：")
    print(f"  {'From':<8} {'To':<8} {'Count':>10} {'Pct':>8}")
    print(f"  {'-'*8} {'-'*8} {'-'*10} {'-'*8}")
    for (p1, p2), c in moves:
        flag = "  " if p1 == p2 else "→ "
        print(f"  P{p1:<6} {flag}P{p2:<6} {c:>10,} {c/total*100:>7.2f}%")


def downstream_impact(result: dict) -> None:
    """评估「按 Key 聚合」的下游业务受影响范围。"""
    rate = result["rate"]
    n1, n2 = result["n1"], result["n2"]
    print()
    print(f"  ┌── 下游受影响评估 ──────────────────────────────────────")
    print(f"  │ 实测变化率: {rate*100:.2f}%")
    print(f"  │ 理论估计  : {theoretical_rate(n1, n2)*100:.2f}%   (公式: (n2-gcd)/n2)")
    if rate > 0.5:
        print(f"  │ 等级      : 🔴 重度破坏（>50% 的 Key 迁移）")
        print(f"  │ 影响范围  : Kafka Streams / Flink 状态、Redis 按 Key 缓存、")
        print(f"  │             所有按 Key 聚合的下游统计 —— 全部失效。")
        print(f"  │ 建议      : 不要加分区。建新 Topic learn.06.xx.v2，双写切流。")
    elif rate > 0.2:
        print(f"  │ 等级      : 🟡 部分破坏（20%~50% 的 Key 迁移）")
        print(f"  │ 影响范围  : 同上，但比例较小。仍需通知所有下游做兼容评估。")
        print(f"  │ 建议      : 评估业务能否容忍统计误差，否则同样建议重建 Topic。")
    else:
        print(f"  │ 等级      : 🟢 较小（<20%）")
        print(f"  │ 注意      : 即使比例小，仍存在「同 Key 在新旧分区上短期并存」")
        print(f"  │             的风险，会导致瞬时顺序错乱。")
    print(f"  └─────────────────────────────────────────────────────")


def run_pair(n1: int, n2: int, total: int, pattern: str) -> None:
    print(f"\n{'='*72}")
    print(f"  分区数变更：{n1} → {n2}   (总消息数 {total:,}, pattern={pattern})")
    print(f"{'='*72}")

    result = compute_change_rate(n1, n2, total, pattern)
    print_sample(n1, n2, sample_keys=20)
    print_top_moves(result, top=10)
    downstream_impact(result)


def main() -> None:
    parser = argparse.ArgumentParser(description="Kafka 加分区破坏性演示")
    parser.add_argument("--n1", type=int, default=6, help="分区数（前）")
    parser.add_argument("--n2", type=int, default=8, help="分区数（后）")
    parser.add_argument("--keys", type=int, default=100_000, help="模拟消息数")
    parser.add_argument("--pairs", type=str, nargs="*", default=None,
                        help="多组对比，格式 n1:n2，如 6:7 6:8 6:12 12:24")
    parser.add_argument("--pattern", type=str, default="seq",
                        choices=["seq", "uuid", "tenant", "biased", "few"])
    args = parser.parse_args()

    if args.pairs:
        for p in args.pairs:
            n1, n2 = (int(x) for x in p.split(":"))
            run_pair(n1, n2, args.keys, args.pattern)
    else:
        run_pair(args.n1, args.n2, args.keys, args.pattern)

    print("\n小结：")
    print("  - Kafka 默认 Partitioner = murmur2(key) % N")
    print("  - 分区数变化时，几乎所有 Key 的落点都会变。")
    print("  - 唯一例外是「N1 是 N2 的因子」时，部分 Key 仍会落到 P_i 或 P_(i+N1)，")
    print("    但仍有 (n2-gcd)/n2 的比例发生迁移。")
    print("  - 因此：生产环境想扩容请「重建 Topic」，不要直接 alter --partitions。")


if __name__ == "__main__":
    main()
