#!/usr/bin/env python3
"""
Random test for the planning of set operations.

Makes random queries with nested UNION, INTERSECT and EXCEPT, runs each one
under several planner settings, and compares the output with an answer
worked out in Python from the table contents.  A run fails if EXPLAIN or
the query fails, if the rows are wrong, or if ORDER BY output is not sorted.

usage: setop_check.py N SEED        run N random queries
       setop_check.py N SEED ID     print query ID, expected rows and plans
Runs psql, so set PGHOST, PGPORT, PGDATABASE.  Creates tables d1 and d2.
"""
import random
import subprocess
import sys
from collections import Counter
from multiprocessing import Pool

SETUP = """
DROP TABLE IF EXISTS d1, d2;
CREATE TABLE d1 (a int, b int);
CREATE TABLE d2 (a int, b numeric);
INSERT INTO d1 SELECT nullif(g % 8, 7), nullif(g % 6, 5)
  FROM generate_series(1, 300) g;
INSERT INTO d2 SELECT nullif(g % 7, 6), nullif(g % 5, 4)
  FROM generate_series(1, 200) g;
CREATE INDEX ON d1 (a, b);
CREATE INDEX ON d2 (a);
ANALYZE d1, d2;
"""
NOHASH = "SET enable_hashagg = off; "
PARALLEL = ("SET parallel_setup_cost = 0; SET parallel_tuple_cost = 0; "
            "SET min_parallel_table_scan_size = 0; "
            "SET max_parallel_workers_per_gather = 2; ")
CONFIGS = {             # every query runs once under each of these
    "default": "",
    "nohash": NOHASH,
    "nohash_noseq": NOHASH + "SET enable_seqscan = off;",
    "nosort": "SET enable_sort = off; SET enable_incremental_sort = off;",
    "merge": NOHASH + "SET enable_hashjoin = off; SET enable_nestloop = off;",
    "parallel": PARALLEL,
    "parallel_nohash": PARALLEL + NOHASH,
}
NODES = ["SetOp", "Merge Append", "Parallel Append", "Merge Join", "WindowAgg"]
MARK = "@@"             # separates statement results in the psql output
TABLES = {}             # table name -> list of (a, b) rows

def psql(script):
    """Run a script, return its output lines."""
    p = subprocess.run(["psql", "-X", "-q", "-At", "-F", ","], input=script,
                       capture_output=True, text=True)
    return p.stdout.split("\n")[:-1]

def parse(lines):
    """psql output lines to rows.  A NULL is printed as an empty string."""
    return [tuple(int(f) if f else None for f in line.split(","))
            for line in lines]

# Parts of a leaf SELECT: the SQL text, and the same thing in Python.
COLUMN1 = [("a", lambda a, b: a), ("b", lambda a, b: b),
           ("a + 1", lambda a, b: None if a is None else a + 1),
           ("1", lambda a, b: 1)]
COLUMN2 = [("b", lambda a, b: b), ("a", lambda a, b: a),
           ("2", lambda a, b: 2), ("NULL::int", lambda a, b: None)]
FILTERS = [("true", lambda a, b: True), ("true", lambda a, b: True),
           ("a = 1", lambda a, b: a == 1),
           ("a < 3", lambda a, b: a is not None and a < 3),
           ("a IS NULL", lambda a, b: a is None),
           ("false", lambda a, b: False), ("1 = 2", lambda a, b: False)]

# A result is a Counter mapping each row to how many times it appears.
# Set operations treat NULLs as equal, like Python does with None.
SETOPS = {
    "UNION ALL": lambda l, r: l + r,
    "UNION": lambda l, r: Counter(set(l + r)),
    "INTERSECT ALL": lambda l, r: l & r,
    "INTERSECT": lambda l, r: Counter(set(l & r)),
    "EXCEPT ALL": lambda l, r: l - r,
    "EXCEPT": lambda l, r: Counter(set(l) - set(r)),
}

def leaf(rnd, ncols):
    """SELECT from one table.  Returns (sql, Counter of rows)."""
    table = rnd.choice(["d1", "d2"])
    cols = [rnd.choice(COLUMN1), rnd.choice(COLUMN2)][:ncols]
    where_sql, where = rnd.choice(FILTERS)
    sql = "SELECT %s FROM %s WHERE %s" % (
        ", ".join("%s AS x%d" % (c[0], i + 1) for i, c in enumerate(cols)),
        table, where_sql)
    if rnd.random() < 0.15:
        sql += " ORDER BY 1"                # does not change the rows
    return sql, Counter(tuple(c[1](a, b) for c in cols)
                        for a, b in TABLES[table] if where(a, b))

def tree(rnd, ncols, depth):
    """A set operation over two smaller trees, or a leaf."""
    if depth == 0:
        return leaf(rnd, ncols)
    op = rnd.choice(list(SETOPS))
    lsql, lrows = tree(rnd, ncols, depth - rnd.choice([1, 1, 1, depth]))
    rsql, rrows = tree(rnd, ncols, depth - rnd.choice([1, 1, 1, depth]))
    return "(%s) %s (%s)" % (lsql, op, rsql), SETOPS[op](lrows, rrows)

def sort_key(row):
    """Python sort key that matches ORDER BY, which puts NULLs last."""
    return [(v is None, v or 0) for v in row]

def make_query(rnd, qid):
    """Put something above a set operation that depends on its output."""
    ncols = rnd.choice([1, 2])
    sql, counter = tree(rnd, ncols, rnd.choice([1, 2, 3]))
    rows = list(counter.elements())
    first = Counter(row[0] for row in rows)     # first column -> count
    all_columns = ", ".join(str(i + 1) for i in range(ncols))
    order = None                                # "asc", "desc" or None
    shape = rnd.choice(["plain", "order", "order desc", "limit", "group",
                        "window", "in", "join table", "join setop"])
    if shape == "order":
        sql, order = sql + " ORDER BY " + all_columns, "asc"
    elif shape == "order desc":
        sql, order = sql + " ORDER BY %s DESC" % all_columns.replace(
            ", ", " DESC, "), "desc"
    elif shape == "limit":
        sql, order = sql + " ORDER BY %s LIMIT 7" % all_columns, "asc"
        rows = sorted(rows, key=sort_key)[:7]
    elif shape == "group":
        sql = "SELECT x1, count(*) FROM (%s) s GROUP BY x1 ORDER BY 1, 2" % sql
        rows, order = list(first.items()), "asc"
    elif shape == "window":
        sql = "SELECT x1, count(*) OVER (PARTITION BY x1) FROM (%s) s" % sql
        rows = [(row[0], first[row[0]]) for row in rows]
    elif shape == "in":
        sql = "SELECT a, b FROM d1 WHERE a IN (SELECT x1 FROM (%s) s)" % sql
        rows = [(a, b) for a, b in TABLES["d1"]
                if a is not None and a in first]
    elif shape == "join table":
        sql = "SELECT s.x1, d1.b FROM (%s) s JOIN d1 ON s.x1 = d1.a" % sql
        rows = [(row[0], b) for row in rows for a, b in TABLES["d1"]
                if a is not None and a == row[0]]
    elif shape == "join setop":
        sql2, counter2 = tree(rnd, 1, rnd.choice([1, 2]))
        sql = ("SELECT s1.x1 FROM (%s) s1, (%s) s2 WHERE s1.x1 = s2.x1"
               % (sql, sql2))
        rows = [(row[0],) for row in rows if row[0] is not None
                for _ in range(counter2[(row[0],)])]
    return {"id": qid, "sql": sql, "expect": rows, "order": order}

def check(test, got):
    if Counter(got) != Counter(test["expect"]):
        return "WRONG ROWS"
    keys = [sort_key(row) for row in got]
    if test["order"] == "desc":
        keys.reverse()
    if test["order"] and keys != sorted(keys):
        return "WRONG ORDER"
    return "ok"

def run(test):
    """Returns (query id, [(config, verdict, plan lines)])."""
    script = ""
    for settings in CONFIGS.values():
        script += "RESET ALL; %s\n" % settings
        for stmt in "EXPLAIN (VERBOSE, COSTS OFF) " + test["sql"], test["sql"]:
            # after each statement, print its SQLSTATE and error message
            script += "%s;\n\\echo %s:SQLSTATE :LAST_ERROR_MESSAGE\n" % (
                stmt, MARK)
    outputs, lines = [], []         # one (lines, status) per statement
    for line in psql(script):
        if line.startswith(MARK):
            outputs.append((lines, line[len(MARK):]))
            lines = []
        else:
            lines.append(line)
    outputs += [([], "connection lost")] * (2 * len(CONFIGS) - len(outputs))
    results = []
    for i, config in enumerate(CONFIGS):
        (plan, explain_status), (rows, status) = outputs[2 * i:2 * i + 2]
        if not status.startswith("00000"):
            verdict = "ERROR: " + status
        else:
            verdict = check(test, parse(rows))
        if verdict == "ok" and not explain_status.startswith("00000"):
            verdict = "EXPLAIN ERROR: " + explain_status
        results.append((config, verdict, plan))
    return test["id"], results

def main(n, seed, show=None):
    # Create the tables, and read them back so Python sees the same data.
    psql(SETUP)
    for table in "d1", "d2":
        TABLES[table] = parse(psql("SELECT a, b FROM %s" % table))
    rnd = random.Random(seed)
    tests = [make_query(rnd, qid) for qid in range(n)]
    if show is not None:
        test = tests[show]
        print(test["sql"])
        print("\nexpected rows:", sorted(test["expect"], key=sort_key))
        for config, verdict, plan in run(test)[1]:
            print("\n-- %s: %s\n%s" % (config, verdict, "\n".join(plan)))
        return
    stats, failed = Counter(), []
    with Pool(8) as pool:
        for qid, results in pool.imap_unordered(run, tests, chunksize=8):
            stats["queries"] += 1
            stats["queries, expected result not empty"] += bool(
                tests[qid]["expect"])
            if any(verdict != "ok" for _, verdict, _ in results):
                failed.append(qid)
            for _, verdict, plan in results:
                stats["runs: " + verdict] += 1
                stats.update("runs with %s in the plan" % node
                             for node in NODES if any(
                                 line.lstrip(" ->").startswith(node)
                                 for line in plan))
    for key in sorted(stats):
        print("%7d  %s" % (stats[key], key))
    print("failed query ids:", *sorted(failed)[:50])

if __name__ == "__main__":
    main(*map(int, sys.argv[1:]))
