"""01_join_subquery.py —— JOIN 与子查询性能对比

演示：
  1) 4 种 JOIN 结果差异
  2) EXISTS vs IN vs JOIN+DISTINCT 的执行计划对比
  3) NOT IN 的 NULL 陷阱

运行前请先执行 init.sql。
"""
from __future__ import annotations

from _common import connect, print_table, section


def demo_join_types(cur) -> None:
    section("1. 4 种 JOIN 行数对比（ch5_orders / ch5_users）")
    for jt in ("INNER", "LEFT", "RIGHT", "FULL"):
        cur.execute(
            f"""
            SELECT COUNT(*) FROM ch5_users u
            {jt} JOIN ch5_orders o ON u.id = o.user_id
            """
        )
        (n,) = cur.fetchone()
        print(f"  {jt:5s} JOIN -> {n:>6d} 行")

    cur.execute("SELECT COUNT(*) FROM ch5_users u CROSS JOIN ch5_orders o")
    (n,) = cur.fetchone()
    print(f"  CROSS JOIN -> {n:>6d} 行 (笛卡儿积)")


def demo_exists_vs_in(cur) -> None:
    section("2. EXISTS / IN / JOIN 执行计划对比")
    sqls = {
        "IN":     "SELECT COUNT(*) FROM ch5_users WHERE id IN (SELECT user_id FROM ch5_orders)",
        "EXISTS": """
            SELECT COUNT(*) FROM ch5_users u
            WHERE EXISTS (SELECT 1 FROM ch5_orders o WHERE o.user_id = u.id)
        """,
        "JOIN":   """
            SELECT COUNT(*) FROM (
              SELECT DISTINCT u.id FROM ch5_users u JOIN ch5_orders o ON o.user_id = u.id
            ) t
        """,
    }
    for name, sql in sqls.items():
        cur.execute("EXPLAIN (ANALYZE, BUFFERS, FORMAT JSON) " + sql)
        plan = cur.fetchone()[0][0]
        top = plan["Plan"]
        print(f"\n[{name}]  Total Cost={top['Total Cost']:.2f}  "
              f"Actual={top['Actual Total Time']:.3f}ms  "
              f"Rows={top['Actual Rows']}")


def demo_not_in_null_trap(cur) -> None:
    section("3. NOT IN 的 NULL 陷阱")
    cur.execute("CREATE TEMP TABLE tmp_u(id INT)")
    cur.execute("INSERT INTO tmp_u VALUES (1),(2),(3)")
    cur.execute("CREATE TEMP TABLE tmp_o(uid INT)")
    cur.execute("INSERT INTO tmp_o VALUES (1), (NULL)")

    cur.execute("SELECT id FROM tmp_u WHERE id NOT IN (SELECT uid FROM tmp_o)")
    rows_notin = cur.fetchall()

    cur.execute("""
        SELECT id FROM tmp_u u
        WHERE NOT EXISTS (SELECT 1 FROM tmp_o o WHERE o.uid = u.id)
    """)
    rows_notexists = cur.fetchall()

    print_table(
        [
            ["NOT IN   (遇到 NULL 全空)",  str(rows_notin)],
            ["NOT EXISTS (正确)",          str(rows_notexists)],
        ],
        ["写法", "结果"],
    )


def main() -> None:
    with connect() as conn:
        with conn.cursor() as cur:
            demo_join_types(cur)
            demo_exists_vs_in(cur)
            demo_not_in_null_trap(cur)


if __name__ == "__main__":
    main()
