#!/usr/bin/env python3
"""
idempotent_consumer.py
======================
业务侧幂等示例：用 sqlite 做「去重表 + 业务表」，模拟「订单已支付」事件的
At Least Once 安全消费。

关键设计：
  1) order_event 表：业务表，order_id 是 PRIMARY KEY，重复插入直接 IntegrityError
  2) consumer_dedup 表：审计表，记录每条消息的 (msg_id, processed_at)
  3) 业务处理 + 去重写入放在同一个 sqlite 事务里，保证「要么都成功，要么都失败」
  4) Kafka commit 仅在事务成功后才执行

无论 Kafka 重投递多少次，业务表里只会有一条记录。
跑两遍可以亲眼看到「dup, skip」日志。

用法：
  pip install confluent-kafka
  # 先生产几条订单
  python idempotent_consumer.py produce 5
  # 启动幂等消费者
  python idempotent_consumer.py consume
  # ⚠️ 把消费者杀掉再重启 / 重置 offset：依然只有一条订单记录
  kafka-consumer-groups.sh --bootstrap-server 127.0.0.1:9092 \
      --group demo-idempotent --reset-offsets --to-earliest --execute --topic learn.12.idem
"""

import os
import sys
import json
import time
import sqlite3
import logging
import hashlib
from datetime import datetime

from confluent_kafka import Producer, Consumer

BOOTSTRAP = os.environ.get("KAFKA_BOOTSTRAP", "127.0.0.1:9092")
TOPIC = "learn.12.idem"
DB = os.environ.get("IDEM_DB", "/tmp/learn_kafka_idempotent.sqlite")

logging.basicConfig(level=logging.INFO, format="%(asctime)s %(message)s", datefmt="%H:%M:%S")
log = logging.getLogger("idem")


def init_db():
    conn = sqlite3.connect(DB)
    conn.execute("""
        CREATE TABLE IF NOT EXISTS order_event (
            order_id    TEXT PRIMARY KEY,
            user_id     INTEGER,
            amount      REAL,
            status      TEXT,
            created_at  TEXT
        )
    """)
    conn.execute("""
        CREATE TABLE IF NOT EXISTS consumer_dedup (
            msg_id        TEXT PRIMARY KEY,
            consumer      TEXT,
            processed_at  TEXT
        )
    """)
    conn.commit()
    return conn


def produce(n):
    p = Producer({"bootstrap.servers": BOOTSTRAP, "linger.ms": 5})
    for i in range(n):
        oid = f"ORD-2026-{i:04d}"
        body = json.dumps({
            "order_id": oid,
            "user_id": 1000 + i,
            "amount": round(99.9 + i, 2),
            "ts": datetime.utcnow().isoformat(),
        })
        p.produce(TOPIC, key=oid, value=body.encode())
    p.flush(10)
    log.info(f"produced {n} order events")


def compute_msg_id(msg, payload):
    # 优先用业务主键；没有的话 fallback 到 (topic, partition, offset)
    if "order_id" in payload:
        raw = f"{payload['order_id']}".encode()
    else:
        raw = f"{msg.topic()}|{msg.partition()}|{msg.offset()}".encode()
    return hashlib.sha1(raw).hexdigest()[:16]


def consume():
    conn = init_db()
    c = Consumer({
        "bootstrap.servers": BOOTSTRAP,
        "group.id": "demo-idempotent",
        "enable.auto.commit": False,
        "auto.offset.reset": "earliest",
    })
    c.subscribe([TOPIC])

    try:
        while True:
            msg = c.poll(1.0)
            if msg is None: continue
            if msg.error():
                log.error(msg.error()); continue

            try:
                payload = json.loads(msg.value())
            except Exception:
                log.exception("bad msg, skip + commit")
                c.commit(message=msg, asynchronous=False)
                continue

            msg_id = compute_msg_id(msg, payload)
            try:
                with conn:    # sqlite context manager = 一个事务
                    conn.execute(
                        "INSERT INTO consumer_dedup(msg_id, consumer, processed_at) VALUES (?,?,?)",
                        (msg_id, "demo-idempotent", datetime.utcnow().isoformat()),
                    )
                    conn.execute(
                        "INSERT INTO order_event(order_id,user_id,amount,status,created_at) "
                        "VALUES (?,?,?,?,?)",
                        (payload["order_id"], payload["user_id"], payload["amount"],
                         "PAID", datetime.utcnow().isoformat()),
                    )
                log.info(f"  ✅ first time, processed order={payload['order_id']} msg_id={msg_id}")
            except sqlite3.IntegrityError:
                # 唯一约束冲突 -> 已处理过，幂等跳过
                log.info(f"  ⏭️  dup, skip order={payload['order_id']} msg_id={msg_id}")

            # 不管首处理还是跳过都要 commit offset
            c.commit(message=msg, asynchronous=False)
    finally:
        c.close()
        # 给个最终统计
        n_orders = conn.execute("SELECT COUNT(*) FROM order_event").fetchone()[0]
        n_dedup  = conn.execute("SELECT COUNT(*) FROM consumer_dedup").fetchone()[0]
        log.info(f"FINAL: {n_orders} unique orders in business table; "
                 f"{n_dedup} dedup records (= 真实处理过的次数)")


def main():
    if len(sys.argv) < 2:
        print(__doc__); return
    cmd = sys.argv[1]
    if cmd == "produce":
        n = int(sys.argv[2]) if len(sys.argv) > 2 else 5
        produce(n)
    elif cmd == "consume":
        consume()
    else:
        print(__doc__)


if __name__ == "__main__":
    main()
