#!/usr/bin/env python3
"""
cooperative_consumer.py — 演示 cooperative-sticky 分配策略

观察重点：
    1) 启动一个 Consumer 时打印 on_assign 的分区列表；
    2) 启动第二个 Consumer（同 group.id）时，第一个 Consumer 只会 revoke 一部分分区，
       不会经历「全部撤销 → 重新分配」的 Stop-the-World；
    3) 关闭其中一个 Consumer，剩下的仅增量接管，不全停。

用法：
    # 终端 1
    python cooperative_consumer.py --consumer-id c1
    # 终端 2（5 秒后启动）
    python cooperative_consumer.py --consumer-id c2
    # 终端 3
    python cooperative_consumer.py --consumer-id c3

前置：
    Topic 至少 3 分区
    kafka-topics.sh --bootstrap-server localhost:9092 \
        --create --topic learn.11.coop --partitions 6 --replication-factor 3 \
        --if-not-exists

    生产端可以用 kafka-console-producer.sh 持续写入：
    kafka-console-producer.sh --bootstrap-server localhost:9092 --topic learn.11.coop
"""

from __future__ import annotations

import argparse
import signal
import sys
import time
from datetime import datetime

from confluent_kafka import Consumer, TopicPartition


def parse_args():
    p = argparse.ArgumentParser()
    p.add_argument("--bootstrap", default="localhost:9092")
    p.add_argument("--topic", default="learn.11.coop")
    p.add_argument("--group", default="coop-demo-group")
    p.add_argument("--consumer-id", default="c1")
    return p.parse_args()


def fmt_tps(tps):
    return ", ".join(f"{tp.topic}-{tp.partition}" for tp in tps) or "<empty>"


def main():
    args = parse_args()
    cid = args.consumer_id

    def log(msg, color="\033[36m"):
        print(f"{color}[{datetime.now().strftime('%H:%M:%S')}][{cid}] {msg}\033[0m",
              flush=True)

    def on_assign(consumer, partitions):
        # cooperative 模式下：partitions 是「新增」的分区，不是「全集」
        log(f"on_assign  →  +{fmt_tps(partitions)}", "\033[32m")
        consumer.incremental_assign(partitions)

    def on_revoke(consumer, partitions):
        # cooperative 模式下：partitions 是「失去」的分区，必须用 incremental_unassign
        log(f"on_revoke  →  -{fmt_tps(partitions)}", "\033[33m")
        try:
            # 同步 commit 当前 offset，防止重复消费
            consumer.commit(asynchronous=False)
        except Exception as e:
            log(f"commit failed: {e}", "\033[31m")
        consumer.incremental_unassign(partitions)

    def on_lost(consumer, partitions):
        # 仅 cooperative 才会调用：所有权已经被 Coordinator 强制收回，不能 commit
        log(f"on_lost    →  -{fmt_tps(partitions)} (forced)", "\033[31m")
        consumer.incremental_unassign(partitions)

    consumer = Consumer({
        "bootstrap.servers": args.bootstrap,
        "group.id": args.group,
        "client.id": cid,
        "partition.assignment.strategy": "cooperative-sticky",  # ★关键★
        "enable.auto.commit": False,
        "auto.offset.reset": "latest",
        "session.timeout.ms": 45000,
        "heartbeat.interval.ms": 3000,
        "max.poll.interval.ms": 300000,
    })
    consumer.subscribe([args.topic],
                       on_assign=on_assign,
                       on_revoke=on_revoke,
                       on_lost=on_lost)

    log(f"started, subscribing to {args.topic} (cooperative-sticky)")

    # 优雅退出
    stop = {"flag": False}
    def _sig(*_):
        log("got signal, draining…", "\033[35m")
        stop["flag"] = True
    signal.signal(signal.SIGINT, _sig)
    signal.signal(signal.SIGTERM, _sig)

    last_print = time.time()
    msg_count = 0
    try:
        while not stop["flag"]:
            msg = consumer.poll(timeout=1.0)
            if msg is None:
                if time.time() - last_print > 5:
                    last_print = time.time()
                    held = consumer.assignment()
                    log(f"heartbeat OK | holding: {fmt_tps(held)} | "
                        f"msgs since start: {msg_count}", "\033[34m")
                continue
            if msg.error():
                log(f"err: {msg.error()}", "\033[31m")
                continue
            msg_count += 1
            if msg_count % 20 == 0 or msg_count <= 5:
                log(f"poll  {msg.topic()}-{msg.partition()}@{msg.offset()}  "
                    f"key={msg.key()!r}  val={msg.value()!r}", "\033[36m")
            # 模拟业务处理时间（很短）
            # 业务处理后异步提交（实际生产推荐定期 commit + 关停时同步 commit）
            consumer.commit(message=msg, asynchronous=True)
    finally:
        log("closing…")
        consumer.close()
        log(f"closed. total msgs = {msg_count}")


if __name__ == "__main__":
    try:
        main()
    except KeyboardInterrupt:
        sys.exit(0)
