#!/usr/bin/env python3
"""
lag_monitor.py — 周期性打印某个消费组的每分区 Lag

实现思路：
    1) 用 AdminClient.list_consumer_group_offsets() 拿消费组当前提交的 offset；
    2) 用 AdminClient.list_offsets()（confluent-kafka 2.0+）拿每个分区的 LATEST offset；
    3) Lag = LATEST - committed；
    4) 同时打印消费组当前 STATE / 每个 member / 它们绑定的 client.id；
    5) 周期性刷新（每 N 秒一次）。

用法：
    python lag_monitor.py --group pay-service
    python lag_monitor.py --group pay-service --interval 5

优势 vs kafka-consumer-groups.sh：
    * 可以集成到自己的告警 / Prometheus exporter；
    * 输出可以是 JSON 方便后续处理（--json）；
    * 可以同时监控多个组（重复传 --group）。
"""

from __future__ import annotations

import argparse
import json
import sys
import time
from datetime import datetime

from confluent_kafka import (
    OffsetSpec,
    TopicPartition,
)
from confluent_kafka.admin import (
    AdminClient,
    ConsumerGroupState,
)


def parse_args():
    p = argparse.ArgumentParser()
    p.add_argument("--bootstrap", default="localhost:9092")
    p.add_argument("--group", action="append", required=True,
                   help="要监控的消费组（可重复传多个）")
    p.add_argument("--interval", type=float, default=3.0)
    p.add_argument("--json", action="store_true", help="JSON 输出便于 grep")
    p.add_argument("--once", action="store_true")
    return p.parse_args()


def fetch_state(admin: AdminClient, group: str):
    """拉取消费组的状态、members、各分区 committed offset"""
    fut = admin.describe_consumer_groups([group])[group]
    desc = fut.result(timeout=15)
    state = desc.state
    members = []
    for m in desc.members:
        ass_parts = []
        if m.assignment:
            for tp in m.assignment.topic_partitions:
                ass_parts.append(f"{tp.topic}-{tp.partition}")
        members.append({
            "member_id": m.member_id[:30],
            "client_id": m.client_id,
            "host": m.host,
            "instance_id": m.group_instance_id or "-",
            "assignments": ass_parts,
        })

    # Committed offsets
    fut = admin.list_consumer_group_offsets([
        # 不指定 partitions 表示拉全部
        # confluent-kafka API: ListConsumerGroupOffsetsRequest
    ])
    # 这里用 list_consumer_group_offsets 的简化签名
    # 不同版本 API 略有差异，下面用 describe 的方式做兼容
    committed = {}
    try:
        from confluent_kafka.admin import ConsumerGroupTopicPartitions
        req = ConsumerGroupTopicPartitions(group)
        fut = admin.list_consumer_group_offsets([req])[group]
        res = fut.result(timeout=15)
        for tp in res.topic_partitions:
            committed[(tp.topic, tp.partition)] = tp.offset
    except Exception as e:
        print(f"[warn] list_consumer_group_offsets 不可用，跳过 committed: {e}",
              file=sys.stderr)

    return state, members, committed


def fetch_latest_offsets(admin: AdminClient, tps):
    if not tps:
        return {}
    spec = {tp: OffsetSpec.latest() for tp in tps}
    fut = admin.list_offsets(spec)
    out = {}
    for tp, fu in fut.items():
        try:
            res = fu.result(timeout=15)
            out[(tp.topic, tp.partition)] = res.offset
        except Exception as e:
            out[(tp.topic, tp.partition)] = -1
    return out


def render(group, state, members, committed, latest, args):
    """打印一个组的状态"""
    if args.json:
        rows = []
        for (t, p), off in sorted(committed.items()):
            le = latest.get((t, p), -1)
            lag = max(0, le - off) if (off >= 0 and le >= 0) else None
            rows.append({"topic": t, "partition": p, "committed": off,
                         "log_end": le, "lag": lag})
        print(json.dumps({"ts": datetime.now().isoformat(), "group": group,
                          "state": str(state), "members": members,
                          "lags": rows}, ensure_ascii=False))
        return

    print(f"\n=== group={group}  state={state}  members={len(members)}  "
          f"@{datetime.now().strftime('%H:%M:%S')} ===")
    if members:
        print(f"  {'member':<32}{'client':<20}{'host':<20}{'instance':<14}{'assigned'}")
        for m in members:
            print(f"  {m['member_id']:<32}{m['client_id']:<20}{m['host']:<20}"
                  f"{m['instance_id']:<14}{','.join(m['assignments']) or '<empty>'}")
    else:
        print("  (no active members)")
    if not committed:
        print("  (no committed offsets — 该组从未提交过)")
        return
    print(f"\n  {'Topic':<26}{'Part':<6}{'Committed':<14}{'Log-End':<14}"
          f"{'Lag':<10}{'Status'}")
    print("  " + "-" * 90)
    total_lag = 0
    for (t, p), off in sorted(committed.items()):
        le = latest.get((t, p), -1)
        if off < 0 or le < 0:
            lag_str, status = "?", "no-data"
        else:
            lag = max(0, le - off)
            total_lag += lag
            lag_str = str(lag)
            status = "OK" if lag == 0 else (
                "⚠ growing" if lag > 1000 else "small")
        print(f"  {t:<26}{p:<6}{off:<14}{le:<14}{lag_str:<10}{status}")
    print(f"  {'TOTAL LAG':<46} {total_lag}")


def main():
    args = parse_args()
    admin = AdminClient({"bootstrap.servers": args.bootstrap})

    while True:
        for group in args.group:
            try:
                state, members, committed = fetch_state(admin, group)
                tps = [TopicPartition(t, p) for (t, p) in committed.keys()]
                latest = fetch_latest_offsets(admin, tps)
                render(group, state, members, committed, latest, args)
            except Exception as e:
                print(f"[ERR] group={group} {e}", file=sys.stderr)
        if args.once:
            return
        time.sleep(args.interval)


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