"""
第 6 章 MergeTree 家族进阶 · 5 种引擎行为演示

依次展示：
    1. ReplacingMergeTree    同主键去重 + argMax vs FINAL 对比
    2. SummingMergeTree      数值列自动累加（合并前 vs 合并后）
    3. AggregatingMergeTree  聚合状态合并 + uniqMerge 实时 PV/UV
    4. CollapsingMergeTree   +1/-1 折叠
    5. VersionedCollapsing   乱序 + Version 折叠

前置：
    先跑 init.sql 建表，再跑 seed.py 灌数

运行：
    python family_play.py
"""

from __future__ import annotations

import sys
import time
import textwrap

import clickhouse_connect


HOST = "127.0.0.1"
PORT = 8123
USER = "default"
PASSWORD = ""
DB = "learn_ck"


def banner(s: str) -> None:
    print("\n" + "=" * 78)
    print("  " + s)
    print("=" * 78)


def show(client, sql: str, label: str = "") -> None:
    sql = textwrap.dedent(sql).strip()
    if label:
        print(f"\n-- {label} --")
    print(">>> " + sql.replace("\n", "\n    "))
    try:
        res = client.query(sql)
        rows, cols = res.result_rows, res.column_names
        if not rows:
            print("    (0 rows)"); return
        widths = [max(len(str(c)), max((len(str(r[i])) for r in rows), default=0))
                  for i, c in enumerate(cols)]
        line = "+".join("-"*(w+2) for w in widths)
        print("    " + line)
        print("    |" + "|".join(f" {c:<{w}} " for c, w in zip(cols, widths)) + "|")
        print("    " + line)
        for r in rows[:20]:
            print("    |" + "|".join(f" {str(v):<{w}} " for v, w in zip(r, widths)) + "|")
        if len(rows) > 20:
            print(f"    ... ({len(rows)} rows total, showing first 20)")
        print("    " + line)
    except Exception as e:
        print(f"    !! error: {e}")


def demo_replacing(client):
    banner("① ReplacingMergeTree —— 用户最新画像")
    show(client, f"SELECT count() FROM {DB}.replacing_user", "Merge 前的物理行数")
    show(client,
         f"""SELECT user_id, count() AS dup_rows
             FROM {DB}.replacing_user GROUP BY user_id
             ORDER BY dup_rows DESC LIMIT 5""",
         "每个 user_id 的重复行数（说明没去重）")
    show(client,
         f"""SELECT user_id,
                    argMax(nickname, version) AS nickname,
                    argMax(city, version)     AS city,
                    max(version)              AS latest_ver
             FROM {DB}.replacing_user WHERE user_id <= 3
             GROUP BY user_id ORDER BY user_id""",
         "【推荐】用 argMax 取最新 —— 不阻塞、不归并")
    show(client,
         f"""SELECT user_id, nickname, city, version
             FROM {DB}.replacing_user FINAL
             WHERE user_id <= 3 ORDER BY user_id""",
         "【慎用】SELECT … FINAL 查询时即时去重（慢 2-10 倍）")
    print("\n→ 手动触发合并：OPTIMIZE TABLE learn_ck.replacing_user FINAL")
    try:
        client.command(f"OPTIMIZE TABLE {DB}.replacing_user FINAL")
        print("   OK")
    except Exception as e:
        print(f"   !! {e}")
    show(client,
         f"""SELECT count() FROM {DB}.replacing_user""",
         "OPTIMIZE FINAL 后的行数（等于唯一 user_id 数）")


def demo_summing(client):
    banner("② SummingMergeTree —— 数值自动累加")
    show(client,
         f"SELECT count() FROM {DB}.summing_daily",
         "Merge 前的物理行数")
    show(client,
         f"""SELECT event_date, event_name, count() AS raw_rows, sum(pv) AS pv_sum
             FROM {DB}.summing_daily
             WHERE event_name = 'click'
             GROUP BY event_date, event_name
             ORDER BY event_date LIMIT 5""",
         "正确查询仍要 GROUP BY + sum()，因为 Merge 是异步")
    try:
        client.command(f"OPTIMIZE TABLE {DB}.summing_daily FINAL")
    except Exception:
        pass
    show(client,
         f"SELECT count() FROM {DB}.summing_daily",
         "OPTIMIZE FINAL 后的行数（相同 (date,event,user) 被累加成一条）")


def demo_aggregating(client):
    banner("③ AggregatingMergeTree —— 实时 PV/UV 大屏")
    show(client,
         f"SELECT count() FROM {DB}.agg_events_raw",
         "明细表行数")
    show(client,
         f"SELECT count() FROM {DB}.agg_events_state",
         "状态表行数（被物化视图写入）")
    show(client,
         f"""SELECT event_date, url,
                    sumMerge(pv_state)  AS pv,
                    uniqMerge(uv_state) AS uv
             FROM {DB}.agg_events_state
             WHERE event_date = '2024-01-15'
             GROUP BY event_date, url
             ORDER BY pv DESC LIMIT 5""",
         "大屏查询：只扫千级状态行，毫秒返回")
    # 对照：直接在原始表上算
    show(client,
         f"""SELECT event_date, url, count() AS pv, uniq(user_id) AS uv
             FROM {DB}.agg_events_raw
             WHERE event_date = '2024-01-15'
             GROUP BY event_date, url
             ORDER BY pv DESC LIMIT 5""",
         "原始明细直接算（对照用，生产中亿级明细会慢很多）")


def demo_collapsing(client):
    banner("④ CollapsingMergeTree —— +1/-1 折叠")
    show(client,
         f"SELECT count() AS total_rows, sum(Sign) AS live_rows FROM {DB}.collapsing_order",
         "合并前：物理行数 vs 折叠后有效行数")
    show(client,
         f"""SELECT order_id,
                    argMax(status, Sign)           AS latest_status,
                    sum(amount * Sign)             AS live_amount
             FROM {DB}.collapsing_order
             GROUP BY order_id HAVING sum(Sign) > 0
             LIMIT 5""",
         "【推荐】sum(col*Sign) + HAVING sum(Sign)>0 —— 通用折叠查询套路")
    try:
        client.command(f"OPTIMIZE TABLE {DB}.collapsing_order FINAL")
    except Exception:
        pass
    show(client,
         f"SELECT count() FROM {DB}.collapsing_order",
         "OPTIMIZE FINAL 后，物理折叠剩余行数")


def demo_versioned(client):
    banner("⑤ VersionedCollapsingMergeTree —— 允许乱序写入")
    show(client,
         f"SELECT count() FROM {DB}.versioned_order",
         "乱序写入后的物理行数")
    show(client,
         f"""SELECT order_id,
                    argMax(status, Version)  AS latest_status,
                    max(Version)             AS ver
             FROM {DB}.versioned_order GROUP BY order_id
             HAVING sum(Sign) > 0 LIMIT 5""",
         "带 Version 的查询（即使乱序，按 Version 取最新）")
    try:
        client.command(f"OPTIMIZE TABLE {DB}.versioned_order FINAL")
    except Exception:
        pass
    show(client,
         f"SELECT count() FROM {DB}.versioned_order",
         "OPTIMIZE FINAL 后，折叠剩余行数（等于独立 order_id 数）")


def main() -> int:
    client = clickhouse_connect.get_client(
        host=HOST, port=PORT, username=USER, password=PASSWORD, database=DB)
    print(f"Connected. Database = {DB}")

    demo_replacing(client)
    time.sleep(1)
    demo_summing(client)
    time.sleep(1)
    demo_aggregating(client)
    time.sleep(1)
    demo_collapsing(client)
    time.sleep(1)
    demo_versioned(client)

    banner("🎉 全部演示完成。要点回顾：")
    print("  · Replacing    → argMax(col, ver) / SELECT ... FINAL（慎用）")
    print("  · Summing      → GROUP BY + sum() （Merge 只是减少扫描行数）")
    print("  · Aggregating  → xxMerge(state) 取最终值，毫秒级大屏")
    print("  · Collapsing   → sum(col*Sign) + HAVING sum(Sign) > 0")
    print("  · VersionedColl→ 允许乱序，按 (pk, Version) 折叠")
    return 0


if __name__ == "__main__":
    sys.exit(main())
