#!/usr/bin/env python3
"""
热点 Key 模拟器：复现「单分区热点」踩坑（案例 2）。

3 种模式：
  good     : 用 user_id (高基数)，分布均匀
  bad_date : 用 today() 做 key，所有消息进同一分区
  bad_city : 5 个城市做 key，5 个分区被打爆，其余空闲

跑完会打印每个分区的消息数，让你直观看到倾斜。

依赖：pip install confluent-kafka
前置：先建一个 12 分区的 Topic
  kafka-topics.sh --create --bootstrap-server localhost:9092 \
    --topic learn.20.heavy --partitions 12 --replication-factor 1

运行：
  MODE=bad_date python heavy_key_simulator.py
  MODE=good     python heavy_key_simulator.py
  MODE=bad_city python heavy_key_simulator.py
"""

import os
import random
import time
from collections import Counter
from datetime import date

from confluent_kafka import Producer, Consumer, TopicPartition

BOOTSTRAP = os.getenv("BOOTSTRAP", "localhost:9092")
TOPIC = os.getenv("TOPIC", "learn.20.heavy")
MODE = os.getenv("MODE", "bad_date")  # good / bad_date / bad_city
COUNT = int(os.getenv("COUNT", "10000"))

CITIES = ["beijing", "shanghai", "shenzhen", "hangzhou", "chengdu"]


def gen_key(i):
    if MODE == "good":
        return f"user-{random.randint(1, 100000)}"
    if MODE == "bad_date":
        return str(date.today())
    if MODE == "bad_city":
        return random.choice(CITIES)
    raise ValueError(MODE)


def produce():
    p = Producer({"bootstrap.servers": BOOTSTRAP, "linger.ms": 10})
    print(f"模式 = {MODE}, 发送 {COUNT} 条消息到 {TOPIC} ...")
    t0 = time.time()
    for i in range(COUNT):
        k = gen_key(i)
        p.produce(TOPIC, key=k.encode(), value=f"msg-{i}".encode())
        if (i + 1) % 1000 == 0:
            p.poll(0)
    p.flush(20)
    print(f"发送完成，耗时 {time.time()-t0:.1f}s")


def show_distribution():
    """从 Topic 元数据读分区数，遍历所有分区算 hi-lo 算消息数。"""
    c = Consumer({
        "bootstrap.servers": BOOTSTRAP,
        "group.id": "heavy-key-stat",
        "enable.auto.commit": False,
    })
    md = c.list_topics(TOPIC, timeout=10)
    parts = list(md.topics[TOPIC].partitions.keys())
    print(f"\n分区数 = {len(parts)}")
    counter = Counter()
    for pid in parts:
        lo, hi = c.get_watermark_offsets(TopicPartition(TOPIC, pid), timeout=5)
        counter[pid] = hi - lo
    c.close()

    total = sum(counter.values())
    print(f"\n=== 分区消息数分布 (共 {total} 条) ===")
    print(f"{'分区':>4} | {'消息数':>10} | {'占比':>8} | 直方图")
    print("-" * 60)
    for pid in sorted(parts):
        n = counter[pid]
        pct = n / total * 100 if total else 0
        bar = "█" * int(pct / 2)
        print(f"{pid:>4} | {n:>10} | {pct:>7.2f}% | {bar}")

    # 倾斜度
    if total:
        mx = max(counter.values())
        avg = total / len(parts)
        print(f"\n最大/平均 = {mx/avg:.2f}x" + (" 严重倾斜！" if mx/avg > 3 else ""))


if __name__ == "__main__":
    produce()
    show_distribution()
