"""
第 17 章 - Kafka Streams 与 ksqlDB
windowed_aggregate.py - 用 Faust 演示「1 分钟滚动窗口聚合 GMV」

业务场景：
    实时大屏需要「按分钟聚合的订单总额（GMV）」，按支付状态分组，
    输出到下游 `learn.gmv-per-minute` Topic 给前端拉取。

核心知识点：
    1) Tumbling Window（1 分钟，不重叠）
    2) Aggregate（不仅 count，还累加 amount）
    3) Event Time（用消息里的 ts 而不是处理时间）
    4) Grace Period 处理乱序消息

依赖：
    pip install faust-streaming python-rocksdb

准备：
    kafka-topics.sh ... --create --topic learn.orders          --partitions 3
    kafka-topics.sh ... --create --topic learn.gmv-per-minute  --partitions 1

运行：
    faust -A windowed_aggregate worker -l info --web-port 6069
"""
from __future__ import annotations

import faust
from datetime import timedelta
from typing import NamedTuple

WINDOW_SIZE  = timedelta(minutes=1)    # 1 分钟滚动
GRACE_PERIOD = timedelta(seconds=30)   # 允许 30s 迟到


class Order(faust.Record, serializer="json"):
    order_id: str
    user_id:  str
    amount:   float
    status:   str           # PENDING / PAID / CANCELLED / REFUNDED
    ts:       float         # event time (epoch seconds)


class MinuteStats(faust.Record, serializer="json"):
    bucket_minute: int      # epoch minute
    status:        str
    order_cnt:     int
    gmv:           float
    avg_amount:    float


app = faust.App(
    "windowed-gmv",
    broker="kafka://localhost:9092",
    store="rocksdb://",
)

orders_topic = app.topic("learn.orders", value_type=Order)
gmv_topic    = app.topic("learn.gmv-per-minute", value_type=MinuteStats)


# -------------------------------------------------------------------
# 一个 windowed Table，key = status，value = (count, sum)
# -------------------------------------------------------------------
class Acc(NamedTuple):
    cnt: int
    sum: float


minute_acc = (
    app.Table("minute-acc",
              default=lambda: Acc(0, 0.0),
              partitions=3)
       .tumbling(WINDOW_SIZE, expires=timedelta(hours=1))
       .relative_to_field(Order.ts)   # 按 event time 做窗口
)


# -------------------------------------------------------------------
# Agent：每来一条订单，按 status 累加 cnt + sum
# -------------------------------------------------------------------
@app.agent(orders_topic)
async def consume_orders(stream):
    async for order in stream:
        cur = minute_acc[order.status].value()
        minute_acc[order.status] = Acc(cur.cnt + 1, cur.sum + order.amount)


# -------------------------------------------------------------------
# 周期性扫描当前窗口，输出聚合
# -------------------------------------------------------------------
@app.timer(interval=10.0)
async def flush_to_topic():
    """每 10 秒把当前窗口的快照推到下游 Topic。"""
    for status, w in minute_acc.items():
        try:
            acc: Acc = w.current()
        except Exception:
            continue
        if acc.cnt == 0:
            continue
        # current() 返回当前 wall-clock 落入的窗口聚合
        # 拿当前的 minute bucket
        import time
        bucket = int(time.time() // 60)
        stats = MinuteStats(
            bucket_minute=bucket,
            status=status,
            order_cnt=acc.cnt,
            gmv=round(acc.sum, 2),
            avg_amount=round(acc.sum / acc.cnt, 2),
        )
        await gmv_topic.send(key=f"{bucket}:{status}", value=stats)
        print(f"[GMV] minute={bucket} status={status} "
              f"cnt={stats.order_cnt} gmv={stats.gmv} avg={stats.avg_amount}")


# -------------------------------------------------------------------
# 实时查询接口：当前窗口的实时数字
# -------------------------------------------------------------------
@app.page("/gmv/current")
async def http_current(web, request):
    out = []
    for status, w in minute_acc.items():
        try:
            acc = w.current()
            if acc.cnt > 0:
                out.append({"status": status, "cnt": acc.cnt,
                            "gmv": round(acc.sum, 2)})
        except Exception:
            pass
    return web.json(out)


if __name__ == "__main__":
    app.main()


# ============================================================================
# 等价的 Java Streams DSL 写法（仅参考，不在本文件中运行）
# ----------------------------------------------------------------------------
# StreamsBuilder builder = new StreamsBuilder();
# KStream<String, Order> orders = builder.stream("learn.orders",
#     Consumed.with(Serdes.String(), orderSerde)
#             .withTimestampExtractor(new OrderTsExtractor()));
#
# KTable<Windowed<String>, Acc> agg = orders
#     .groupBy((k, v) -> v.status)
#     .windowedBy(TimeWindows.ofSizeAndGrace(
#         Duration.ofMinutes(1),
#         Duration.ofSeconds(30)))
#     .aggregate(
#         () -> new Acc(0, 0.0),
#         (k, order, acc) -> new Acc(acc.cnt + 1, acc.sum + order.amount),
#         Materialized.<String, Acc, WindowStore<Bytes,byte[]>>as("minute-agg")
#                     .withValueSerde(accSerde));
#
# agg.toStream()
#    .map((wk, acc) -> KeyValue.pair(
#        wk.window().start() / 60_000 + ":" + wk.key(),
#        new MinuteStats(...)))
#    .to("learn.gmv-per-minute", Produced.with(Serdes.String(), statsSerde));
# ============================================================================
