"""
第 4 章 - 异步发送 + 回调 + flush 边界
=================================

标准的高吞吐姿势：
1) produce() 异步入队；
2) 每次 produce 后 poll(0) 触发已就绪回调；
3) 业务边界处 flush() 阻塞等所有未完成消息出去；
4) BufferError 时退避并继续。

运行：
    bash ../init.sh
    python async_producer.py            # 默认发 10000 条
    python async_producer.py 100000
"""

from __future__ import annotations

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 10000
TOPIC = "learn.04.demo"

success = [0]
failure = [0]
sample = []


def on_delivery(err, msg):
    if err is not None:
        failure[0] += 1
        if failure[0] <= 3:
            print(f"  ❌ {err}")
    else:
        success[0] += 1
        if len(sample) < 5:
            sample.append(
                f"P{msg.partition()}@{msg.offset()} key={msg.key()} val_size={len(msg.value())}"
            )


def main() -> None:
    p = Producer(
        {
            "bootstrap.servers": BOOTSTRAP,
            "client.id": "ch4-async-producer",
            "acks": "all",
            "enable.idempotence": True,
            "compression.type": "zstd",
            "linger.ms": 10,
            "batch.size": 65536,
            # 单连接最多 5 个未 ack 请求（开了幂等的安全上限）
            "max.in.flight.requests.per.connection": 5,
            "queue.buffering.max.messages": 200000,
        }
    )

    print(f"🚀 异步发送 {N} 条到 {TOPIC}（acks=all + idempotence + zstd + linger=10）")
    payload = ("x" * 800).encode()
    start = time.perf_counter()
    for i in range(N):
        while True:
            try:
                p.produce(
                    TOPIC,
                    key=f"k{i % 100}".encode(),
                    value=payload + f"-{i}".encode(),
                    on_delivery=on_delivery,
                )
                break
            except BufferError:
                # 队列满，让回调释放空间
                p.poll(0.5)
        # 每 1000 条触发一次回调，避免回调队列堆积
        if i % 1000 == 0:
            p.poll(0)

    remaining = p.flush(timeout=30)
    elapsed = time.perf_counter() - start
    print("\n📊 结果")
    print(f"  成功 {success[0]} / 失败 {failure[0]} / 未发出 {remaining}")
    print(f"  总耗时 {elapsed:.2f}s，吞吐 {success[0] / elapsed:.0f} msg/s")
    bytes_total = success[0] * (len(payload) + 5)
    print(f"  数据量 {bytes_total / 1024 / 1024:.1f} MB，{bytes_total / elapsed / 1024 / 1024:.1f} MB/s")
    print("\n  样本：")
    for s in sample:
        print(f"    - {s}")


if __name__ == "__main__":
    main()
