#!/usr/bin/env python3
"""
inspect_transaction_state.py
============================
直接消费内部 Topic `__transaction_state`，观察事务的状态变迁日志。

每个事务从 Empty → Ongoing → PrepareCommit/PrepareAbort →
CompleteCommit/CompleteAbort 的过程，都会作为一条消息追加到 `__transaction_state`，
Key 是 transactional.id，Value 是当前的事务元数据（PID, Epoch, state, partitions, ...）。

由于 Value 是 Kafka 内部的二进制 schema，比 __consumer_offsets 还复杂（带 partition
list、多个版本），完整解析超出脚本范围；这里我们只解析 Key（transactional.id），
并用启发式打印 Value 大小、是否 tombstone，让你看到「事务状态变化触发的写入流」。

用法：
  pip install confluent-kafka
  python inspect_transaction_state.py

  # 然后另一个终端跑 transactional_producer.py commit/abort，
  # 这里会实时打印每个事务状态变更
"""

import os
import struct
import logging

from confluent_kafka import Consumer

BOOTSTRAP = os.environ.get("KAFKA_BOOTSTRAP", "127.0.0.1:9092")
logging.basicConfig(level=logging.INFO, format="%(asctime)s %(message)s", datefmt="%H:%M:%S")
log = logging.getLogger("txstate")


def parse_key(key):
    """
    TransactionLogKey schema:
      int16  version
      string transactional_id
    """
    if len(key) < 2: return None
    version = struct.unpack(">h", key[:2])[0]
    pos = 2
    if pos + 2 > len(key): return None
    n = struct.unpack(">h", key[pos:pos+2])[0]; pos += 2
    if n < 0: return {"version": version, "tx_id": None}
    tx_id = key[pos:pos+n].decode("utf-8", errors="replace")
    return {"version": version, "tx_id": tx_id}


def heuristic_state(value):
    """
    Value 的 schema（v0 简化）:
      int16  version
      int64  producer_id
      int16  producer_epoch
      int32  txn_timeout_ms
      int8   txn_state    (0=Empty, 1=Ongoing, 2=PrepCommit, 3=PrepAbort,
                           4=CompleteCommit, 5=CompleteAbort, 6=Dead, 7=PrepareEpochFence)
      int32  txn_start_ts (some versions)
      ...    partitions array (复杂)
    """
    if value is None:
        return {"tombstone": True}
    if len(value) < 15: return {"len": len(value), "parse": "too short"}
    try:
        version = struct.unpack(">h", value[0:2])[0]
        pid = struct.unpack(">q", value[2:10])[0]
        epoch = struct.unpack(">h", value[10:12])[0]
        timeout = struct.unpack(">i", value[12:16])[0]
        state = value[16] if len(value) > 16 else None
        STATE_NAMES = {0: "Empty", 1: "Ongoing", 2: "PrepareCommit", 3: "PrepareAbort",
                       4: "CompleteCommit", 5: "CompleteAbort", 6: "Dead",
                       7: "PrepareEpochFence"}
        return {
            "version": version, "PID": pid, "Epoch": epoch,
            "timeout_ms": timeout, "state": STATE_NAMES.get(state, f"unknown({state})"),
            "value_size": len(value),
        }
    except Exception as e:
        return {"len": len(value), "err": str(e)}


def main():
    c = Consumer({
        "bootstrap.servers": BOOTSTRAP,
        "group.id": "inspect-txstate-" + str(os.getpid()),
        "enable.auto.commit": False,
        "auto.offset.reset": "latest",
    })
    c.subscribe(["__transaction_state"])
    log.info("listening to __transaction_state ...")
    log.info("（在另一个终端运行 transactional_producer.py 即可看到事务状态变迁）")
    try:
        while True:
            msg = c.poll(1.0)
            if msg is None: continue
            if msg.error():
                log.error(msg.error()); continue
            k = parse_key(msg.key() or b"")
            v = heuristic_state(msg.value())
            tx_id = k["tx_id"] if k else "?"
            if v.get("tombstone"):
                log.info(f"🗑️  tx_id={tx_id!r} → TOMBSTONE (Coordinator 清理记录)")
            elif "err" in v:
                log.info(f"📝 tx_id={tx_id!r} → value parse error: {v['err']} (size={v['len']})")
            else:
                log.info(
                    f"📝 tx_id={tx_id!r}  PID={v['PID']}  Epoch={v['Epoch']}  "
                    f"state={v['state']}  timeout={v['timeout_ms']}ms  "
                    f"size={v['value_size']}B"
                )
    except KeyboardInterrupt:
        pass
    finally:
        c.close()


if __name__ == "__main__":
    main()
