"""05_lateral_join.py —— LATERAL JOIN 实战

场景：
  A) 每个用户最近 3 笔订单
  B) 用窗口函数做同样的事，对比执行计划
  C) LATERAL + 聚合：每个用户近 30 天消费总额
"""
from __future__ import annotations

from _common import connect, print_table, section


def ensure_index(cur) -> None:
    # (user_id, created_at DESC) 是 LATERAL Top-N 的最佳索引
    cur.execute(
        "CREATE INDEX IF NOT EXISTS idx_ch5_orders_user_time "
        "ON ch5_orders (user_id, created_at DESC)"
    )


def lateral_topN(cur) -> None:
    section("A. LATERAL：每个用户最近 3 笔订单")
    cur.execute(
        """
        SELECT u.id, u.name, o.id AS order_id, o.amount, o.created_at
        FROM ch5_users u
        LEFT JOIN LATERAL (
          SELECT id, amount, created_at
          FROM ch5_orders
          WHERE user_id = u.id
          ORDER BY created_at DESC
          LIMIT 3
        ) o ON true
        WHERE u.id <= 5
        ORDER BY u.id, o.created_at DESC NULLS LAST
        """
    )
    print_table(
        cur.fetchall(),
        ["uid", "name", "order_id", "amount", "created_at"],
    )


def window_topN(cur) -> None:
    section("B. 窗口函数版（等价但全表排序）")
    cur.execute(
        """
        WITH ranked AS (
          SELECT u.id AS uid, u.name, o.id AS order_id, o.amount, o.created_at,
                 ROW_NUMBER() OVER (PARTITION BY u.id ORDER BY o.created_at DESC) AS rn
          FROM ch5_users u
          LEFT JOIN ch5_orders o ON o.user_id = u.id
        )
        SELECT uid, name, order_id, amount, created_at
        FROM ranked
        WHERE rn <= 3 AND uid <= 5
        ORDER BY uid, created_at DESC NULLS LAST
        """
    )
    print_table(
        cur.fetchall(),
        ["uid", "name", "order_id", "amount", "created_at"],
    )


def compare_plans(cur) -> None:
    section("C. 两种写法的执行计划对比")
    sqls = {
        "LATERAL": """
            SELECT u.id, o.id FROM ch5_users u
            LEFT JOIN LATERAL (
              SELECT id FROM ch5_orders
              WHERE user_id = u.id
              ORDER BY created_at DESC LIMIT 3
            ) o ON true
        """,
        "Window":  """
            WITH r AS (
              SELECT u.id AS uid, o.id,
                     ROW_NUMBER() OVER (PARTITION BY u.id ORDER BY o.created_at DESC) rn
              FROM ch5_users u LEFT JOIN ch5_orders o ON o.user_id = u.id
            ) SELECT uid, id FROM r WHERE rn <= 3
        """,
    }
    for name, sql in sqls.items():
        cur.execute("EXPLAIN (ANALYZE, BUFFERS, FORMAT JSON) " + sql)
        p = cur.fetchone()[0][0]["Plan"]
        print(f"  {name:8s}  cost={p['Total Cost']:>9.2f}  "
              f"actual={p['Actual Total Time']:>8.3f}ms  "
              f"rows={p['Actual Rows']}")


def lateral_agg(cur) -> None:
    section("D. LATERAL + 聚合：每个用户近 30 天订单数 & 金额")
    cur.execute(
        """
        SELECT u.id, u.name, s.cnt, s.total
        FROM ch5_users u
        LEFT JOIN LATERAL (
          SELECT COUNT(*) AS cnt, COALESCE(SUM(amount), 0) AS total
          FROM ch5_orders
          WHERE user_id = u.id
            AND created_at >= NOW() - INTERVAL '30 days'
        ) s ON true
        WHERE u.id <= 10
        ORDER BY s.total DESC NULLS LAST
        """
    )
    print_table(cur.fetchall(), ["uid", "name", "近30天单数", "近30天金额"])


def main() -> None:
    with connect() as conn:
        with conn.cursor() as cur:
            ensure_index(cur)
            conn.commit()
            lateral_topN(cur)
            window_topN(cur)
            compare_plans(cur)
            lateral_agg(cur)


if __name__ == "__main__":
    main()
