#!/usr/bin/env python3
"""
static_membership.py — 演示 group.instance.id（Static Membership, KIP-345）

实验流程（建议同时开 4 个终端）：

    终端 1（启用 static membership 的 Consumer A）：
        python static_membership.py --consumer-id static-a --instance-id pod-a
    终端 2（启用 static membership 的 Consumer B）：
        python static_membership.py --consumer-id static-b --instance-id pod-b
    终端 3（启用 static membership 的 Consumer C）：
        python static_membership.py --consumer-id static-c --instance-id pod-c

    然后在终端 4 用 kafka-consumer-groups.sh 看消费组状态：
        kafka-consumer-groups.sh --bootstrap-server localhost:9092 \
            --describe --group static-mem-demo --state

    实验：CTRL+C 关掉终端 1（pod-a），10 秒内重启它：
        * 在 60s session.timeout.ms 内，Coordinator 不会触发 Rebalance；
        * pod-a 重新启动后，原来的分区直接还回来，整组无 Rebalance；
        * 终端 2 / 3 完全不停。

    对比试验：去掉 --instance-id 参数（变成 dynamic membership）再重做：
        * 关掉某个 Consumer 立刻触发 Rebalance；
        * 启动它又触发一次。整组停顿明显。

注意：
    * group.instance.id 必须**全局唯一**；如果你启动两个相同的 instance.id，
      后启动的会顶替先启动的（先启动的会拿到 fenced 错误）；
    * static membership 要求 broker 端 inter.broker.protocol.version >= 2.3。
"""

from __future__ import annotations

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

from confluent_kafka import Consumer


def parse_args():
    p = argparse.ArgumentParser()
    p.add_argument("--bootstrap", default="localhost:9092")
    p.add_argument("--topic", default="learn.11.static")
    p.add_argument("--group", default="static-mem-demo")
    p.add_argument("--consumer-id", default="c1", help="client.id")
    p.add_argument("--instance-id", default=None,
                   help="group.instance.id；不传则退化为 dynamic membership")
    p.add_argument("--session-ms", type=int, default=60000,
                   help="session.timeout.ms（static 模式建议调到 30~60s）")
    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
    iid = args.instance_id

    color = "\033[32m" if iid else "\033[33m"
    tag = f"static:{iid}" if iid else f"dynamic:{cid}"

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

    cfg = {
        "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": args.session_ms,
        "heartbeat.interval.ms": max(3000, args.session_ms // 5),
        "max.poll.interval.ms": 300000,
    }
    if iid:
        # ★关键★ 启用 static membership
        cfg["group.instance.id"] = iid

    def on_assign(consumer, partitions):
        log(f"on_assign  →  +{fmt_tps(partitions)}", "\033[36m")
        consumer.incremental_assign(partitions)

    def on_revoke(consumer, partitions):
        log(f"on_revoke  →  -{fmt_tps(partitions)}", "\033[35m")
        try:
            consumer.commit(asynchronous=False)
        except Exception as e:
            log(f"commit failed: {e}", "\033[31m")
        consumer.incremental_unassign(partitions)

    consumer = Consumer(cfg)
    consumer.subscribe([args.topic], on_assign=on_assign, on_revoke=on_revoke)

    log(f"started.  cfg.session_ms={args.session_ms}  "
        f"static={'YES' if iid else 'NO'}")
    log("提示：CTRL+C 关掉本进程，10s 内 rerun，观察是否触发 Rebalance")

    stop = {"flag": False}
    def _sig(*_):
        stop["flag"] = True
    signal.signal(signal.SIGINT, _sig)
    signal.signal(signal.SIGTERM, _sig)

    msgs = 0
    last_print = time.time()
    try:
        while not stop["flag"]:
            msg = consumer.poll(1.0)
            if msg is None:
                if time.time() - last_print > 5:
                    last_print = time.time()
                    held = consumer.assignment()
                    log(f"alive | holding: {fmt_tps(held)} | msgs={msgs}", "\033[34m")
                continue
            if msg.error():
                log(f"err: {msg.error()}", "\033[31m")
                continue
            msgs += 1
            if msgs <= 3 or msgs % 50 == 0:
                log(f"poll  {msg.topic()}-{msg.partition()}@{msg.offset()}", "\033[36m")
            consumer.commit(message=msg, asynchronous=True)
    finally:
        log("closing…", "\033[35m")
        # 优雅 close 会发 LeaveGroup
        # 但 static membership 的 close() 行为可以通过 leave-group-on-close 控制
        # 此处保持默认（dynamic 会发 LeaveGroup，static 不发）
        consumer.close()
        log(f"closed. msgs={msgs}", "\033[35m")


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