#!/usr/bin/env python3
"""
第 9 章 · JOIN vs Dictionary 对比脚本

  · 自动造数据：100 万订单 + 10 万商品 + 5 万用户
  · 同一个查询分别用 JOIN / dictGet / direct JOIN 三种姿势跑
  · 输出耗时与峰值内存
  · 重复跑 3 次取平均，避免冷启动误差

依赖：pip install clickhouse-connect
准备：先 clickhouse-client < init.sql
"""
from __future__ import annotations

import random
import time
from datetime import datetime, timedelta
from decimal import Decimal

import clickhouse_connect

HOST = "127.0.0.1"
PORT = 8123
ORDERS = 1_000_000
PRODUCTS = 100_000
USERS = 50_000


def seed(client) -> None:
    print("=" * 80)
    print("[seed] truncate target tables ...")
    for t in ("orders", "products", "users"):
        client.command(f"TRUNCATE TABLE learn_ck.{t}")

    print(f"[seed] inserting {USERS:,} users ...")
    cats = ["3C", "服装", "母婴", "美妆", "食品", "家居", "运动"]
    brands = ["Apple", "Nike", "Adidas", "MAC", "蒙牛", "宜家", "李宁"]
    cities = ["北京", "上海", "深圳", "广州", "杭州", "成都"]
    user_rows = [
        (i, f"u{i}", random.choice(cities), random.randint(0, 1), datetime.now())
        for i in range(1, USERS + 1)
    ]
    client.insert("learn_ck.users", user_rows,
                  column_names=["id", "name", "city", "vip", "updated_at"])

    print(f"[seed] inserting {PRODUCTS:,} products ...")
    p_rows = [
        (i, f"商品-{i}", random.choice(cats), random.choice(brands),
         Decimal(f"{random.randint(10, 999)}.{random.randint(0, 99):02d}"),
         datetime.now())
        for i in range(1, PRODUCTS + 1)
    ]
    client.insert("learn_ck.products", p_rows,
                  column_names=["sku_id", "title", "category", "brand", "price", "updated_at"])

    print(f"[seed] inserting {ORDERS:,} orders ...")
    end = datetime.now()
    start = end - timedelta(days=7)
    span = int((end - start).total_seconds())
    rows = []
    cols = ["ts", "order_id", "uid", "sku_id", "qty", "amount"]
    for i in range(1, ORDERS + 1):
        ts = start + timedelta(seconds=random.randint(0, span))
        qty = random.randint(1, 5)
        amt = Decimal(f"{random.randint(10, 9999) * qty}.{random.randint(0, 99):02d}")
        rows.append((ts, i, random.randint(1, USERS),
                     random.randint(1, PRODUCTS), qty, amt))
        if len(rows) >= 100_000:
            client.insert("learn_ck.orders", rows, column_names=cols)
            rows.clear()
    if rows:
        client.insert("learn_ck.orders", rows, column_names=cols)

    print("[seed] reload dictionaries ...")
    client.command("SYSTEM RELOAD DICTIONARIES")
    print("[seed] done.")


def time_query(client, sql: str, runs: int = 3) -> tuple[float, int, int]:
    """返回 (avg_ms, peak_mem_bytes, returned_rows)"""
    times = []
    last_mem = 0
    last_rows = 0
    for _ in range(runs):
        client.command("SYSTEM FLUSH LOGS")
        t0 = time.time()
        res = client.query(sql)
        elapsed = (time.time() - t0) * 1000
        times.append(elapsed)
        last_rows = len(res.result_rows)
    client.command("SYSTEM FLUSH LOGS")
    mem_row = client.query(
        f"""
        SELECT max(memory_usage)
        FROM system.query_log
        WHERE event_time > now() - INTERVAL 1 MINUTE
          AND type = 'QueryFinish'
          AND query LIKE %(snip)s
        """,
        parameters={"snip": "%" + sql.strip().split("\n")[0][:40] + "%"},
    ).result_rows
    if mem_row and mem_row[0][0]:
        last_mem = int(mem_row[0][0])
    return sum(times) / len(times), last_mem, last_rows


def bench(client) -> None:
    print("=" * 80)
    print("[bench] 同一个查询：每个用户每个商品的销售总额 Top 10")
    print("=" * 80)

    Q1 = """
    -- 方式 A：传统 JOIN
    SELECT
        u.name AS user_name,
        p.title AS product,
        sum(o.amount) AS gmv
    FROM learn_ck.orders o
    JOIN learn_ck.users    u ON o.uid    = u.id
    JOIN learn_ck.products p ON o.sku_id = p.sku_id
    GROUP BY user_name, product
    ORDER BY gmv DESC
    LIMIT 10
    SETTINGS join_algorithm = 'hash'
    """

    Q2 = """
    -- 方式 B：dictGet 字典查询
    SELECT
        dictGetString('learn_ck.user_dict',    'name',  uid)    AS user_name,
        dictGetString('learn_ck.product_dict', 'title', sku_id) AS product,
        sum(amount) AS gmv
    FROM learn_ck.orders
    GROUP BY user_name, product
    ORDER BY gmv DESC
    LIMIT 10
    """

    Q3 = """
    -- 方式 C：direct JOIN（左 JOIN 字典 + join_algorithm = 'direct'）
    SELECT
        u.name  AS user_name,
        p.title AS product,
        sum(o.amount) AS gmv
    FROM learn_ck.orders o
    LEFT JOIN learn_ck.user_dict    u ON o.uid    = u.id
    LEFT JOIN learn_ck.product_dict p ON o.sku_id = p.sku_id
    GROUP BY user_name, product
    ORDER BY gmv DESC
    LIMIT 10
    SETTINGS join_algorithm = 'direct'
    """

    for name, sql in [("A. JOIN (hash)", Q1),
                      ("B. dictGet    ", Q2),
                      ("C. direct JOIN", Q3)]:
        try:
            t, mem, rows = time_query(client, sql)
            mem_disp = f"{mem/1024/1024:7.1f} MB" if mem else "  n/a "
            print(f"  {name}  →  avg {t:7.1f} ms   peak mem {mem_disp}   rows {rows}")
        except Exception as e:
            print(f"  {name}  →  FAILED: {e}")

    print()
    print("[bench] 单一字段查询（百万行 + 50K 维表）")
    Q4 = """
    SELECT count() FROM learn_ck.orders o
    JOIN learn_ck.users u ON o.uid = u.id
    WHERE u.vip = 1
    SETTINGS join_algorithm = 'hash'
    """
    Q5 = """
    SELECT count() FROM learn_ck.orders
    WHERE dictGetUInt8('learn_ck.user_dict', 'vip', uid) = 1
    """
    for name, sql in [("A. JOIN ", Q4), ("B. dict ", Q5)]:
        t, mem, _ = time_query(client, sql)
        mem_disp = f"{mem/1024/1024:7.1f} MB" if mem else "  n/a "
        print(f"  {name}  →  avg {t:7.1f} ms   peak mem {mem_disp}")


def main() -> None:
    client = clickhouse_connect.get_client(host=HOST, port=PORT, username="default")
    seed(client)
    bench(client)


if __name__ == "__main__":
    main()
