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

完整的 Producer → Broker → Consumer 端到端吞吐压测脚本。

核心目的：把第 8 章里讲到的「批量大小 / 压缩算法 / linger.ms / acks」
四个杠杆都暴露成命令行参数，让你在自己的集群上**亲手感受**它们对吞吐 / 延迟的影响。

依赖：confluent-kafka 或 kafka-python。脚本会自动检测，二选一。
  pip install confluent-kafka
  # 或
  pip install kafka-python

用法：
  python3 throughput_bench.py --bootstrap localhost:9092 --topic learn.bench
  python3 throughput_bench.py --records 200000 --size 1024 --batch 65536 --linger 20 --compression lz4
  python3 throughput_bench.py --acks all --compression zstd --no-consume   # 只压 Producer

输出指标：
  - Producer: 发送速率（msg/s, MB/s）、平均延迟、p50/p95/p99 延迟
  - Consumer: 消费速率（msg/s, MB/s）、首条消息延迟（端到端）

提示：
  - 别在生产集群跑大流量压测。
  - 跑之前先用 kafka-topics.sh 创建好 Topic，分区数 ≥ Producer 并发 / Consumer 并发 的最大值。
"""

from __future__ import annotations

import argparse
import os
import statistics
import sys
import threading
import time
import uuid
from dataclasses import dataclass, field
from typing import List, Optional


# ----------------- 客户端封装：先尝试 confluent-kafka，回退 kafka-python -----------------

class KafkaImpl:
    name = "none"

    @classmethod
    def detect(cls) -> "KafkaImpl":
        try:
            import confluent_kafka  # noqa: F401
            return ConfluentImpl()
        except ImportError:
            pass
        try:
            import kafka  # noqa: F401
            return KafkaPythonImpl()
        except ImportError:
            pass
        raise RuntimeError(
            "未找到 confluent-kafka 或 kafka-python。请先：\n"
            "  pip install confluent-kafka   (推荐，性能更好)\n"
            "  或 pip install kafka-python"
        )


class ConfluentImpl(KafkaImpl):
    name = "confluent-kafka"

    def make_producer(self, args):
        from confluent_kafka import Producer
        cfg = {
            "bootstrap.servers": args.bootstrap,
            "acks": args.acks,
            "linger.ms": args.linger,
            "batch.size": args.batch,
            "compression.type": args.compression,
            "enable.idempotence": "true" if args.acks == "all" else "false",
            "queue.buffering.max.messages": 1_000_000,
        }
        return Producer(cfg)

    def produce(self, producer, topic, key, value, on_delivery):
        producer.produce(topic, key=key, value=value, on_delivery=on_delivery)
        producer.poll(0)

    def flush(self, producer):
        producer.flush(60)

    def make_consumer(self, args, group):
        from confluent_kafka import Consumer
        cfg = {
            "bootstrap.servers": args.bootstrap,
            "group.id": group,
            "auto.offset.reset": "earliest",
            "enable.auto.commit": False,
            "fetch.min.bytes": 1,
            "fetch.max.bytes": 50 * 1024 * 1024,
        }
        c = Consumer(cfg)
        c.subscribe([args.topic])
        return c

    def poll(self, consumer, timeout_ms):
        msg = consumer.poll(timeout_ms / 1000.0)
        if msg is None:
            return None
        if msg.error():
            return None
        return msg.value(), msg.timestamp()[1]


class KafkaPythonImpl(KafkaImpl):
    name = "kafka-python"

    def make_producer(self, args):
        from kafka import KafkaProducer
        return KafkaProducer(
            bootstrap_servers=args.bootstrap.split(","),
            acks=args.acks,
            linger_ms=args.linger,
            batch_size=args.batch,
            compression_type=None if args.compression == "none" else args.compression,
        )

    def produce(self, producer, topic, key, value, on_delivery):
        future = producer.send(topic, key=key, value=value)
        future.add_callback(lambda md, cb=on_delivery: cb(None, _Md(md)))
        future.add_errback(lambda exc, cb=on_delivery: cb(exc, None))

    def flush(self, producer):
        producer.flush(60)

    def make_consumer(self, args, group):
        from kafka import KafkaConsumer
        return KafkaConsumer(
            args.topic,
            bootstrap_servers=args.bootstrap.split(","),
            group_id=group,
            auto_offset_reset="earliest",
            enable_auto_commit=False,
            fetch_min_bytes=1,
            fetch_max_bytes=50 * 1024 * 1024,
        )

    def poll(self, consumer, timeout_ms):
        records = consumer.poll(timeout_ms=timeout_ms, max_records=500)
        msgs = []
        for tp, batch in records.items():
            for m in batch:
                msgs.append((m.value, m.timestamp))
        return msgs or None


class _Md:
    def __init__(self, md):
        self._md = md
    def topic(self): return self._md.topic
    def partition(self): return self._md.partition
    def offset(self): return self._md.offset


# ------------------------------- Producer 压测 -------------------------------

@dataclass
class ProducerStats:
    sent: int = 0
    failed: int = 0
    bytes_: int = 0
    latencies_us: List[int] = field(default_factory=list)
    start_ts: float = 0.0
    end_ts: float = 0.0


def run_producer(impl: KafkaImpl, args) -> ProducerStats:
    producer = impl.make_producer(args)
    stats = ProducerStats()
    sent_ts: dict[int, int] = {}
    lock = threading.Lock()

    def on_delivery(err, msg):
        with lock:
            now = time.perf_counter_ns()
            if err is not None:
                stats.failed += 1
                return
            seq_key = id(msg) if msg is not None else None
            t0 = sent_ts.pop(seq_key, None)
            if t0 is not None:
                stats.latencies_us.append((now - t0) // 1000)

    payload = (b"x" * args.size)
    print(f"\n[Producer] 开始发送 {args.records:,} 条消息，每条 {args.size} B …")
    stats.start_ts = time.perf_counter()

    for i in range(args.records):
        key = f"key-{i % max(1, args.keys)}".encode()
        sent_ts[id(payload) + i] = time.perf_counter_ns()
        try:
            impl.produce(producer, args.topic, key, payload,
                         lambda err, msg, k=id(payload) + i: _record(err, k, sent_ts, stats, lock))
            stats.sent += 1
            stats.bytes_ += args.size + len(key)
        except BufferError:
            time.sleep(0.001)
            continue

    impl.flush(producer)
    stats.end_ts = time.perf_counter()
    return stats


def _record(err, k, sent_ts, stats, lock):
    with lock:
        now = time.perf_counter_ns()
        t0 = sent_ts.pop(k, None)
        if err is not None:
            stats.failed += 1
            return
        if t0 is not None:
            stats.latencies_us.append((now - t0) // 1000)


def report_producer(stats: ProducerStats, args) -> None:
    dur = stats.end_ts - stats.start_ts
    if dur <= 0:
        print("⚠️  Producer 时长为 0")
        return
    msg_rate = stats.sent / dur
    mb_rate = stats.bytes_ / dur / 1024 / 1024
    lats = sorted(stats.latencies_us)
    print("=" * 72)
    print(" Producer 报告")
    print("-" * 72)
    print(f"  records sent : {stats.sent:,}  (failed {stats.failed})")
    print(f"  duration     : {dur:.2f} s")
    print(f"  msg rate     : {msg_rate:,.0f} msg/s")
    print(f"  byte rate    : {mb_rate:,.2f} MB/s")
    if lats:
        print(f"  latency avg  : {statistics.mean(lats)/1000:.2f} ms")
        print(f"  latency p50  : {lats[len(lats)//2]/1000:.2f} ms")
        print(f"  latency p95  : {lats[int(len(lats)*0.95)]/1000:.2f} ms")
        print(f"  latency p99  : {lats[int(len(lats)*0.99)]/1000:.2f} ms")
        print(f"  latency max  : {lats[-1]/1000:.2f} ms")
    print("=" * 72)


# ------------------------------- Consumer 压测 -------------------------------

@dataclass
class ConsumerStats:
    received: int = 0
    bytes_: int = 0
    start_ts: float = 0.0
    end_ts: float = 0.0


def run_consumer(impl: KafkaImpl, args, expected: int) -> ConsumerStats:
    group = f"bench-{uuid.uuid4().hex[:8]}"
    consumer = impl.make_consumer(args, group)
    stats = ConsumerStats()
    print(f"\n[Consumer] 开始消费，期望 {expected:,} 条 …")
    stats.start_ts = time.perf_counter()

    last_print = stats.start_ts
    while stats.received < expected:
        result = impl.poll(consumer, 1000)
        if not result:
            if time.perf_counter() - last_print > 5:
                print(f"  …已消费 {stats.received:,} / {expected:,}")
                last_print = time.perf_counter()
            if time.perf_counter() - stats.start_ts > args.consume_timeout:
                print("  ⚠️  消费超时，提前结束")
                break
            continue
        if isinstance(result, list):
            for value, _ts in result:
                stats.received += 1
                stats.bytes_ += len(value) if value else 0
        else:
            value, _ts = result
            stats.received += 1
            stats.bytes_ += len(value) if value else 0

    stats.end_ts = time.perf_counter()
    try:
        consumer.close()
    except Exception:
        pass
    return stats


def report_consumer(stats: ConsumerStats) -> None:
    dur = stats.end_ts - stats.start_ts
    if dur <= 0:
        print("⚠️  Consumer 时长为 0")
        return
    print("=" * 72)
    print(" Consumer 报告")
    print("-" * 72)
    print(f"  records read : {stats.received:,}")
    print(f"  duration     : {dur:.2f} s")
    print(f"  msg rate     : {stats.received/dur:,.0f} msg/s")
    print(f"  byte rate    : {stats.bytes_/dur/1024/1024:,.2f} MB/s")
    print("=" * 72)


# --------------------------------- main ---------------------------------

def main() -> None:
    parser = argparse.ArgumentParser(description="Kafka Producer/Consumer 端到端吞吐压测")
    parser.add_argument("--bootstrap", default=os.getenv("KAFKA_BOOTSTRAP", "localhost:9092"))
    parser.add_argument("--topic", default="learn.bench")
    parser.add_argument("--records", type=int, default=100_000)
    parser.add_argument("--size", type=int, default=1024, help="单条消息字节数")
    parser.add_argument("--keys", type=int, default=1000, help="key 基数（控制分区分布）")
    parser.add_argument("--batch", type=int, default=16384, help="batch.size 字节")
    parser.add_argument("--linger", type=int, default=0, help="linger.ms")
    parser.add_argument("--compression", default="none",
                        choices=["none", "gzip", "snappy", "lz4", "zstd"])
    parser.add_argument("--acks", default="1", choices=["0", "1", "all"])
    parser.add_argument("--no-consume", action="store_true", help="只压 Producer")
    parser.add_argument("--consume-timeout", type=int, default=120)
    args = parser.parse_args()

    impl = KafkaImpl.detect()
    print(f"使用客户端实现：{impl.name}")
    print(f"目标 Topic    ：{args.topic} @ {args.bootstrap}")
    print(f"配置：batch={args.batch} linger={args.linger}ms "
          f"compression={args.compression} acks={args.acks}")

    p_stats = run_producer(impl, args)
    report_producer(p_stats, args)

    if not args.no_consume:
        c_stats = run_consumer(impl, args, expected=p_stats.sent)
        report_consumer(c_stats)

    # 给一些直觉化建议
    print("\n💡 调参建议：")
    if args.linger == 0 and args.batch <= 16384:
        print("  - 你目前是『极致低延迟』模式，吞吐受限。如果是日志采集 / 数仓场景，把 batch 提到 64KB+，linger 提到 20ms+。")
    if args.compression == "none" and args.size > 256:
        print("  - 没开压缩。文本类消息开启 lz4 / zstd 通常吞吐翻倍，CPU 涨 5~15%。")
    if args.acks == "0":
        print("  - acks=0 数据可能丢，仅适用于「丢一点没关系」的指标埋点。")
    if args.acks == "all":
        print("  - acks=all 是金标准。配合 RF=3 + min.insync.replicas=2 才算真的不丢。")


if __name__ == "__main__":
    sys.exit(main())
