"""
第 5 章 - 手动批次提交 Consumer (业界主流写法)
==================================================

特征：
- enable.auto.commit=False
- 处理完一批后同步 commit
- 业务异常时不 commit (等下次重消费)
- finally 里同步 commit 兜底

要求：业务必须幂等（用业务主键去重 / DB 唯一约束 / Redis SETNX）。

运行：
    bash ../init.sh
    python manual_commit_consumer.py
"""

from __future__ import annotations

import json
import os
import signal
import sys
import time

from confluent_kafka import Consumer, KafkaError

BOOTSTRAP = sys.argv[1] if len(sys.argv) > 1 else "127.0.0.1:9092"
TOPIC = "learn.05.orders"
GROUP = "ch5-manual-commit-demo"
BATCH_SIZE = 50

running = True


def stop(*_):
    global running
    print("\n[stop] received signal, exiting after final commit")
    running = False


signal.signal(signal.SIGINT, stop)
signal.signal(signal.SIGTERM, stop)

seen_keys = set()


def process(msg) -> bool:
    try:
        payload = json.loads(msg.value())
        key = f"{msg.topic()}-{payload['order_id']}"
        if key in seen_keys:
            print(f"  [DUP] already processed {key}, skip")
            return True
        seen_keys.add(key)
        time.sleep(0.005)
        return True
    except Exception as e:
        print(f"  [ERR] processing failed: {e} (offset={msg.offset()})")
        return False


def flush(consumer, batch):
    consumer.commit(asynchronous=False)
    last = batch[-1]
    print(
        f"  [commit] {len(batch)} msgs, last P{last.partition()}@{last.offset()} "
        f"(committed offset={last.offset() + 1})"
    )


def main() -> None:
    consumer = Consumer(
        {
            "bootstrap.servers": BOOTSTRAP,
            "group.id": GROUP,
            "client.id": f"manual-commit-{os.getpid()}",
            "auto.offset.reset": "earliest",
            "enable.auto.commit": False,
            "session.timeout.ms": 30000,
            "max.poll.interval.ms": 300000,
            "partition.assignment.strategy": "cooperative-sticky",
        }
    )
    consumer.subscribe([TOPIC])
    print(f"[manual-commit-demo] group={GROUP}, batch_size={BATCH_SIZE}")

    batch = []
    total = 0
    failed = 0

    try:
        while running:
            msg = consumer.poll(1.0)
            if msg is None:
                if batch:
                    flush(consumer, batch)
                    total += len(batch)
                    batch = []
                continue
            if msg.error():
                if msg.error().code() != KafkaError._PARTITION_EOF:
                    print(f"  poll err: {msg.error()}")
                continue

            ok = process(msg)
            if ok:
                batch.append(msg)
            else:
                failed += 1

            if len(batch) >= BATCH_SIZE:
                flush(consumer, batch)
                total += len(batch)
                batch = []
    finally:
        if batch:
            flush(consumer, batch)
            total += len(batch)
        print(f"\n[done] processed={total}, failed={failed}")
        consumer.close()


if __name__ == "__main__":
    main()
