#!/usr/bin/env python3
"""ShopFlow incremental SELECT + keyed UPSERT simulation, using only SQLite.

Run alone: python3 incremental-selection.py
Test a saved SELECT: python3 incremental-selection.py --candidate-file your-query.sql
Add --json for structured results. No account, network, packages or disk database.
This is NOT a dbt run, adapter MERGE test, warehouse integration or performance test.
Bundled candidate responses are authored teaching examples, not model transcripts.
"""

import argparse
import json
from pathlib import Path
import sqlite3
import time


# ShopFlow adaptation: sale prices and revenue use integer cents. _loaded_at is
# added ingestion metadata on BOTH staging inputs; it is not an original source
# field. Inputs contain one current row per non-null key. Load markers must change
# when values change. All timestamps are non-null fixed-width UTC strings in
# YYYY-MM-DD HH:MM:SS form, so text comparison preserves time order in this fixture.
SOURCE_SELECT = """
    SELECT oi.order_id, oi.product_id, o.customer_id, o.order_ts,
           oi.quantity * oi.unit_price_cents AS line_revenue_cents,
           max(o._loaded_at, oi._loaded_at) AS source_loaded_at,
           o._loaded_at AS header_loaded_at
    FROM stg_order_items oi JOIN stg_orders o ON o.order_id = oi.order_id
"""
OUTPUT_COLUMNS = ("order_id", "product_id", "customer_id", "order_ts",
                  "line_revenue_cents", "source_loaded_at")
PROJECT = ", ".join(OUTPUT_COLUMNS)
FULL_SELECTION = f"WITH source_rows AS ({SOURCE_SELECT}) SELECT {PROJECT} FROM source_rows"
PREDICATES = {
    "overlap": """
        (SELECT max(source_loaded_at) FROM fact_sales) IS NULL
        OR source_loaded_at >= (SELECT datetime(max(source_loaded_at), '-3 days') FROM fact_sales)
    """,
    "event_time": """
        (SELECT max(order_ts) FROM fact_sales) IS NULL
        OR order_ts >= (SELECT datetime(max(order_ts), '-3 days') FROM fact_sales)
    """,
    "strict_new_only": """
        (SELECT max(source_loaded_at) FROM fact_sales) IS NULL
        OR source_loaded_at > (SELECT max(source_loaded_at) FROM fact_sales)
    """,
    "header_only": """
        (SELECT max(source_loaded_at) FROM fact_sales) IS NULL
        OR header_loaded_at >= (SELECT datetime(max(source_loaded_at), '-3 days') FROM fact_sales)
    """,
    "exclusive_cutoff": """
        (SELECT max(source_loaded_at) FROM fact_sales) IS NULL
        OR source_loaded_at > (SELECT datetime(max(source_loaded_at), '-3 days') FROM fact_sales)
    """,
    "missing_empty_guard": """
        source_loaded_at >= (SELECT datetime(max(source_loaded_at), '-3 days') FROM fact_sales)
    """,
}
CANDIDATES = {name: FULL_SELECTION + " WHERE " + predicate
              for name, predicate in PREDICATES.items()}

# Independently hand-written complete answers: these do not call a candidate or
# use SOURCE_SELECT. In particular, a stale $20 line must become $35, and rows
# exactly on Jan 7 must be included even though the target max is Jan 10.
EXPECTED_ALL = [
    (101, 10, 101, "2026-01-01 00:00:00", 3000, "2026-01-07 00:00:00"),
    (101, 20, 101, "2026-01-01 00:00:00", 7000, "2026-01-07 00:00:00"),
    (102, 10, 102, "2026-01-01 00:00:00", 5000, "2026-01-11 00:00:00"),
    (201, 10, 201, "2026-01-01 00:00:00", 3500, "2026-01-11 00:00:00"),
    (202, 10, 202, "2026-01-01 00:00:00", 1100, "2026-01-06 00:00:00"),
    (900, 10, 900, "2026-01-10 00:00:00", 900, "2026-01-10 00:00:00"),
]
EXPECTED_BOUNDED = [
    (101, 10, 101, "2026-01-01 00:00:00", 3000, "2026-01-07 00:00:00"),
    (101, 20, 101, "2026-01-01 00:00:00", 7000, "2026-01-07 00:00:00"),
    (102, 10, 102, "2026-01-01 00:00:00", 5000, "2026-01-11 00:00:00"),
    (201, 10, 201, "2026-01-01 00:00:00", 3500, "2026-01-11 00:00:00"),
    (900, 10, 900, "2026-01-10 00:00:00", 900, "2026-01-10 00:00:00"),
]
EXPECTED_RETRY_SELECTION = [
    (102, 10, 102, "2026-01-01 00:00:00", 5000, "2026-01-11 00:00:00"),
    (201, 10, 201, "2026-01-01 00:00:00", 3500, "2026-01-11 00:00:00"),
    (900, 10, 900, "2026-01-10 00:00:00", 900, "2026-01-10 00:00:00"),
]
INITIAL_TARGET = [
    (201, 10, 201, "2026-01-01 00:00:00", 2000, "2026-01-02 00:00:00"),
    (900, 10, 900, "2026-01-10 00:00:00", 900, "2026-01-10 00:00:00"),
]
EXPECTED_PASSES = {
    "overlap": [True, True, True],
    "event_time": [True, False, False],
    "strict_new_only": [True, False, False],
    "header_only": [True, False, False],
    "exclusive_cutoff": [True, False, False],
    "missing_empty_guard": [False, True, True],
}


def fixture(empty=False):
    connection = sqlite3.connect(":memory:")
    connection.executescript("""
        CREATE TABLE stg_orders(
            order_id INTEGER PRIMARY KEY, customer_id INTEGER NOT NULL,
            order_ts TEXT NOT NULL, _loaded_at TEXT NOT NULL
        );
        CREATE TABLE stg_order_items(
            order_id INTEGER NOT NULL, product_id INTEGER NOT NULL,
            quantity INTEGER NOT NULL, unit_price_cents INTEGER NOT NULL,
            _loaded_at TEXT NOT NULL, PRIMARY KEY(order_id, product_id)
        );
        CREATE TABLE fact_sales(
            order_id INTEGER NOT NULL, product_id INTEGER NOT NULL,
            customer_id INTEGER NOT NULL, order_ts TEXT NOT NULL,
            line_revenue_cents INTEGER NOT NULL, source_loaded_at TEXT NOT NULL,
            PRIMARY KEY(order_id, product_id)
        );
    """)
    connection.executemany("INSERT INTO stg_orders VALUES(?,?,?,?)", [
        (101, 101, "2026-01-01 00:00:00", "2026-01-06 00:00:00"),
        (102, 102, "2026-01-01 00:00:00", "2026-01-11 00:00:00"),
        (201, 201, "2026-01-01 00:00:00", "2026-01-02 00:00:00"),
        (202, 202, "2026-01-01 00:00:00", "2026-01-06 00:00:00"),
        (900, 900, "2026-01-10 00:00:00", "2026-01-10 00:00:00"),
    ])
    connection.executemany("INSERT INTO stg_order_items VALUES(?,?,?,?,?)", [
        (101, 10, 1, 3000, "2026-01-07 00:00:00"),
        (101, 20, 2, 3500, "2026-01-07 00:00:00"),
        (102, 10, 1, 5000, "2026-01-06 00:00:00"),
        (201, 10, 1, 3500, "2026-01-11 00:00:00"),
        (202, 10, 1, 1100, "2026-01-06 00:00:00"),
        (900, 10, 1, 900, "2026-01-10 00:00:00"),
    ])
    if not empty:
        connection.executemany("INSERT INTO fact_sales VALUES(?,?,?,?,?,?)", INITIAL_TARGET)
    return connection


def read_only(action, _arg1, function, _database, _trigger):
    safe_functions = {"sum", "total", "avg", "count", "min", "max", "abs", "round",
                      "coalesce", "ifnull", "nullif", "date", "datetime", "strftime"}
    if action == sqlite3.SQLITE_FUNCTION and (function or "").lower() not in safe_functions:
        return sqlite3.SQLITE_DENY
    allowed = (sqlite3.SQLITE_SELECT, sqlite3.SQLITE_READ, sqlite3.SQLITE_FUNCTION,
               sqlite3.SQLITE_RECURSIVE)
    return sqlite3.SQLITE_OK if action in allowed else sqlite3.SQLITE_DENY


def select_rows(connection, sql):
    ticks = 0
    deadline = time.monotonic() + 1

    def stop():
        nonlocal ticks
        ticks += 1
        return ticks >= 1000 or time.monotonic() >= deadline

    connection.set_authorizer(read_only)
    connection.set_progress_handler(stop, 1000)
    try:
        if len(sql) > 100000:
            raise sqlite3.OperationalError("Query exceeds 100000 characters")
        cursor = connection.execute(sql)
        if tuple(column[0] for column in cursor.description or []) != OUTPUT_COLUMNS:
            raise sqlite3.OperationalError("SELECT must return these columns, in order: " + PROJECT)
        rows = cursor.fetchmany(101)
        if len(rows) > 100:
            raise sqlite3.OperationalError("Fixture result exceeds 100 rows")
        if any(any(value is None for value in row) for row in rows):
            raise sqlite3.OperationalError("Null key/value violates the fixture contract")
        if any(any(not isinstance(row[index], int) for index in (0, 1, 2, 4))
               or any(not isinstance(row[index], str) for index in (3, 5)) for row in rows):
            raise sqlite3.OperationalError("Expected integer keys/cents and text UTC timestamps")
        if len({row[:2] for row in rows}) != len(rows):
            raise sqlite3.OperationalError("Duplicate order-line keys; deduplicate staging deterministically")
        return sorted(rows, key=lambda row: row[:2])
    finally:
        # Older Python 3 releases do not accept None here. Only the fixed
        # reference UPSERT/inspection runs after this point, never candidate SQL.
        connection.set_authorizer(lambda *_args: sqlite3.SQLITE_OK)
        connection.set_progress_handler(None, 0)


def apply_rows(connection, rows):
    # This is a deliberately fixed reference application step. The learner's SQL
    # only SELECTs rows. SQLite UPSERT is not dbt/warehouse MERGE validation.
    connection.executemany("""
        INSERT INTO fact_sales VALUES(?,?,?,?,?,?)
        ON CONFLICT(order_id, product_id) DO UPDATE SET
            customer_id=excluded.customer_id, order_ts=excluded.order_ts,
            line_revenue_cents=excluded.line_revenue_cents,
            source_loaded_at=excluded.source_loaded_at
    """, rows)


def target_rows(connection):
    return connection.execute("SELECT * FROM fact_sales ORDER BY order_id, product_id").fetchall()


def keys(rows):
    return [f"{row[0]}/{row[1]}" for row in rows]


def run_case(connection, name, case, sql, expected_selected, expected_target):
    selected = []
    error = None
    try:
        selected = select_rows(connection, sql)
        apply_rows(connection, selected)
    except (sqlite3.Error, sqlite3.Warning) as problem:
        error = str(problem)
    target = target_rows(connection)
    return {"candidate": name, "case": case, "selected_keys": keys(selected),
            "selected": selected, "target": target,
            "expected_selected": expected_selected, "expected_target": expected_target,
            "target_revenue_cents": sum(row[4] for row in target),
            "passed": error is None and selected == expected_selected and target == expected_target,
            "error": error}


def evaluate(name, sql):
    results = []
    connection = fixture(empty=True)
    try:
        results.append(run_case(connection, name, "empty_target", sql, EXPECTED_ALL, EXPECTED_ALL))
    finally:
        connection.close()
    connection = fixture()
    try:
        results.append(run_case(connection, name, "bounded_run", sql,
                                EXPECTED_BOUNDED, EXPECTED_BOUNDED))
        # Re-evaluate the query against the updated target, so the cutoff moves.
        results.append(run_case(connection, name, "repeat_run", sql,
                                EXPECTED_RETRY_SELECTION, EXPECTED_BOUNDED))
    finally:
        connection.close()
    return results


def demonstrate_backfill():
    connection = fixture()
    try:
        apply_rows(connection, select_rows(connection, CANDIDATES["overlap"]))
        # Explicitly removing the filter is a separate full-selection backfill
        # here, not an automatic property of incremental selection or UPSERT.
        return run_case(connection, "full_selection", "explicit_backfill",
                        FULL_SELECTION, EXPECTED_ALL, EXPECTED_ALL)
    finally:
        connection.close()


def main():
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--candidate-file", type=Path,
                        help="One local UTF-8 SQLite SELECT returning the six documented columns")
    parser.add_argument("--json", action="store_true")
    args = parser.parse_args()
    candidates = CANDIDATES
    if args.candidate_file:
        try:
            candidates = {"your_query": args.candidate_file.read_text(encoding="utf-8")}
        except (OSError, UnicodeError) as error:
            parser.error(str(error))
    results = [row for name, sql in candidates.items() for row in evaluate(name, sql)]
    backfill = demonstrate_backfill()
    if args.candidate_file:
        checks_passed = all(row["passed"] for row in results)
    else:
        checks_passed = all(
            [row["passed"] for row in results if row["candidate"] == name] == expected
            for name, expected in EXPECTED_PASSES.items()
        ) and all(row["error"] is None for row in results)
    checks_passed = checks_passed and backfill["passed"]
    report = {"exercise": "incremental-selection", "unit": "USD cents",
              "sqlite_version": sqlite3.sqlite_version,
              "simulation": "SQLite SELECT + keyed UPSERT; NOT dbt adapter integration",
              "columns": OUTPUT_COLUMNS,
              "mode": "your_query" if args.candidate_file else "bundled_examples",
              "results": results, "backfill": backfill, "checks_passed": checks_passed}
    if args.json:
        print(json.dumps(report, indent=2))
    else:
        print("SQLite version:", sqlite3.sqlite_version)
        print("SQLite SELECT + keyed UPSERT simulation — NOT a dbt/adapter integration test.")
        print("UTC load metadata on BOTH inputs; integer cents; unique non-null order-line keys.")
        print("Initial watermark Jan 10; inclusive three-day cutoff Jan 7. Retry cutoff Jan 8.")
        for row in results:
            state = "PASSED" if row["passed"] else "REJECTED"
            if not args.candidate_file:
                case_index = ["empty_target", "bounded_run", "repeat_run"].index(row["case"])
                expected_pass = EXPECTED_PASSES[row["candidate"]][case_index]
                prefix = "Expected" if row["passed"] == expected_pass else "UNEXPECTED"
                state = f"{prefix} {state.lower()}"
            print(f"{state:19} {row['candidate']:19} {row['case']:12} "
                  f"selected={','.join(row['selected_keys']) or '(none)'} "
                  f"target_rows={len(row['target'])}, target_cents={row['target_revenue_cents']}")
            if row["error"]:
                print("  Query error:", row["error"])
        print("Separate explicit backfill: 6 rows, 20500 cents; recovers 202/10 outside the window."
              if backfill["passed"] else "Backfill demonstration failed its expected answer.")
        print("Retry selection can shrink while target contents stay identical. Hard deletes are out of scope.")
        print(("Your query passed all fixtures." if checks_passed else "Your query failed a fixture.")
              if args.candidate_file else
              ("Teaching checks passed: expected rejections confirmed; overlap candidate passes."
               if checks_passed else "Teaching check mismatch: inspect the exercise."))
    return 0 if checks_passed else 1


if __name__ == "__main__":
    raise SystemExit(main())
