"""
第 4 章 - 压缩算法对比压测
=================================

对比 none / gzip / snappy / lz4 / zstd 五种压缩在相同数据量下的：
- Producer 端实际吞吐 (msg/s, MB/s)
- 平均 batch 大小（反映压缩率）
- Producer CPU 时间（user + sys）
- 总耗时

数据：每条 ~1KB 的 JSON 风格 payload（字段重复多，可压缩性强）。
默认 N = 100,000 条，可通过参数调整。

运行：
    bash ../init.sh
    python benchmark_compression.py                  # 默认 100k
    python benchmark_compression.py 500000           # 50w
    python benchmark_compression.py 100000 myhost:9092

输出（示例，单 broker 单机 KRaft）：
    | codec  | msgs/s | MB/s  | avg_batch_kb | cpu_s |
    | none   | 152340 | 145.2 | 86.3         | 4.21  |
    | gzip   |  91230 |  87.0 | 12.7         | 9.55  |
    | snappy | 138900 | 132.4 | 28.4         | 5.18  |
    | lz4    | 161200 | 153.7 | 29.1         | 4.45  |
    | zstd   | 156100 | 148.8 | 15.9         | 5.62  |
"""

from __future__ import annotations

import json
import os
import resource
import sys
import time

from confluent_kafka import Producer

BOOTSTRAP = sys.argv[2] if len(sys.argv) > 2 else "127.0.0.1:9092"
N = int(sys.argv[1]) if len(sys.argv) > 1 else 100_000
TOPIC = "learn.04.bench"
CODECS = ["none", "gzip", "snappy", "lz4", "zstd"]


def make_payload(i: int) -> bytes:
    """约 1KB 的 JSON 风格 payload，重复字段多以放大压缩比差异。"""
    return json.dumps(
        {
            "event_id": f"evt-{i:08d}",
            "user_id": i % 1000,
            "action": "purchase",
            "category": "electronics",
            "currency": "CNY",
            "tags": ["promotion", "vip", "newuser", "marketing"] * 5,
            "description": "the quick brown fox jumps over the lazy dog. " * 12,
            "ts": int(time.time() * 1000),
        }
    ).encode()


def get_cpu_time() -> float:
    r = resource.getrusage(resource.RUSAGE_SELF)
    return r.ru_utime + r.ru_stime


def bench(codec: str) -> dict:
    producer = Producer(
        {
            "bootstrap.servers": BOOTSTRAP,
            "client.id": f"ch4-bench-{codec}",
            "acks": "all",
            "enable.idempotence": True,
            "compression.type": codec,
            "linger.ms": 20,
            "batch.size": 1048576,
            "queue.buffering.max.messages": 500_000,
            "queue.buffering.max.kbytes": 1_048_576,
        }
    )

    print(f"\n--- codec = {codec} ---")
    cpu0 = get_cpu_time()
    t0 = time.perf_counter()
    sent = 0
    failed = 0

    def cb(err, msg):
        nonlocal failed
        if err:
            failed += 1

    for i in range(N):
        while True:
            try:
                producer.produce(TOPIC, value=make_payload(i), key=str(i % 50).encode(), on_delivery=cb)
                break
            except BufferError:
                producer.poll(0.5)
        sent += 1
        if i % 5000 == 0:
            producer.poll(0)

    remain = producer.flush(120)
    elapsed = time.perf_counter() - t0
    cpu_used = get_cpu_time() - cpu0

    msg_size_avg = sum(len(make_payload(i)) for i in range(min(100, N))) / min(100, N)
    bytes_total = sent * msg_size_avg

    return {
        "codec": codec,
        "sent": sent - remain,
        "failed": failed,
        "elapsed": elapsed,
        "msg_per_s": (sent - remain) / elapsed,
        "mb_per_s": bytes_total / elapsed / 1024 / 1024,
        "cpu_s": cpu_used,
    }


def main() -> None:
    print(f"== Compression benchmark ==")
    print(f"N = {N}, msg ~1KB each, acks=all + idempotence + linger=20")
    results = []
    for c in CODECS:
        r = bench(c)
        results.append(r)
        print(f"  -> sent={r['sent']} failed={r['failed']} {r['elapsed']:.2f}s")

    print("\n=== summary ===")
    print(f"{'codec':<8} {'msgs/s':>10} {'MB/s':>8} {'cpu_s':>8} {'failed':>8}")
    for r in results:
        print(f"{r['codec']:<8} {r['msg_per_s']:>10.0f} {r['mb_per_s']:>8.1f} {r['cpu_s']:>8.2f} {r['failed']:>8d}")

    print("\nTip: 实际数据中重复模式越多，gzip / zstd 的压缩比越大；"
          "随机 / 加密数据用 none 或 lz4 即可，硬压缩纯浪费 CPU。")


if __name__ == "__main__":
    main()
