#!/usr/bin/env python3
"""ShopFlow revenue checks. Python 3 + SQLite only; no network or disk database.

Run this file alone: python3 revenue-grain.py
Test a saved SELECT: python3 revenue-grain.py --candidate-file your-query.sql
Use --json for machine-readable results. Bundled responses are authored teaching
examples, not actual model transcripts or measurements of a model's performance.
"""

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


# This fixture adapts ShopFlow amount/unit_price to integer amount_cents and
# unit_price_cents: 10000 cents = $100. Avoid floating-point money arithmetic.
# Define currency, tax, shipping and refunds before applying this to real data.
CANDIDATES = {
    "joined_order_total": """
        SELECT SUM(o.amount_cents) AS revenue_cents
        FROM orders o JOIN order_items oi ON o.order_id = oi.order_id
        WHERE o.status = 'paid'
    """,
    "distinct_order_total": """
        SELECT SUM(DISTINCT o.amount_cents) AS revenue_cents
        FROM orders o JOIN order_items oi ON o.order_id = oi.order_id
        WHERE o.status = 'paid'
    """,
    "current_product_price": """
        SELECT SUM(oi.quantity * p.unit_price_cents) AS revenue_cents
        FROM orders o JOIN order_items oi ON o.order_id = oi.order_id
        JOIN products p ON p.product_id = oi.product_id
        WHERE o.status = 'paid'
    """,
    "sale_price_lines": """
        SELECT SUM(oi.quantity * oi.unit_price_cents) AS revenue_cents
        FROM orders o JOIN order_items oi ON o.order_id = oi.order_id
        WHERE o.status = 'paid'
    """,
    "order_header_total": """
        SELECT SUM(amount_cents) AS revenue_cents FROM orders WHERE status = 'paid'
    """,
}

# These answers are hand-authored, not calculated by a candidate query.
# Single: one paid order, $30 + $70 = $100. Equal totals: two distinct paid
# orders, each $100, so $200. The cancelled $999 order contributes nothing.
EXPECTED = {"one_paid_order": 10000, "two_equal_paid_orders": 20000}
EXPECTED_PASSES = {
    "joined_order_total": [False, False],
    "distinct_order_total": [True, False],
    "current_product_price": [False, False],
    "sale_price_lines": [True, True],
    "order_header_total": [True, True],
}


def fixture(case):
    connection = sqlite3.connect(":memory:")
    connection.executescript("""
        CREATE TABLE orders(
            order_id INTEGER PRIMARY KEY, customer_id INTEGER NOT NULL,
            order_ts TEXT NOT NULL, status TEXT NOT NULL, amount_cents INTEGER NOT NULL
        );
        CREATE TABLE order_items(
            order_id INTEGER NOT NULL, product_id INTEGER NOT NULL,
            quantity INTEGER NOT NULL, unit_price_cents INTEGER NOT NULL,
            PRIMARY KEY(order_id, product_id)
        );
        CREATE TABLE products(
            product_id INTEGER PRIMARY KEY, name TEXT NOT NULL,
            category TEXT NOT NULL, unit_price_cents INTEGER NOT NULL
        );
    """)
    connection.executemany("INSERT INTO products VALUES(?,?,?,?)", [
        (10, "ShopFlow book", "Books", 3500),
        (20, "ShopFlow headphones", "Electronics", 4500),
    ])
    orders = [
        (1, 101, "2026-01-01 00:00:00", "paid", 10000),
        (2, 102, "2026-01-01 00:00:00", "cancelled", 99900),
    ]
    items = [(1, 10, 1, 3000), (1, 20, 2, 3500), (2, 10, 1, 99900)]
    if case == "two_equal_paid_orders":
        orders.append((3, 103, "2026-01-02 00:00:00", "paid", 10000))
        items.extend([(3, 10, 1, 3000), (3, 20, 2, 3500)])
    connection.executemany("INSERT INTO orders VALUES(?,?,?,?,?)", orders)
    connection.executemany("INSERT INTO order_items VALUES(?,?,?,?)", items)
    return connection


def read_only(action, _arg1, function, _database, _trigger):
    """Accept SELECTs; reject writes, ATTACH, PRAGMA and schema changes."""
    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 bound_query(connection):
    # A million VM instructions or one second, whichever comes first. The time
    # check is supplementary; the instruction budget is reproducible.
    ticks = 0
    deadline = time.monotonic() + 1

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

    connection.set_progress_handler(stop, 1000)


def evaluate(name, sql):
    results = []
    for case, expected in EXPECTED.items():
        connection = fixture(case)
        connection.set_authorizer(read_only)
        bound_query(connection)
        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 []) != ("revenue_cents",):
                raise sqlite3.OperationalError("SELECT must return one column named revenue_cents")
            rows = cursor.fetchmany(2)
            actual = rows[0][0] if len(rows) == 1 and len(rows[0]) == 1 else rows
            passed = rows == [(expected,)]
            result = {"candidate": name, "case": case, "expected_cents": expected,
                      "actual_cents": actual, "passed": passed, "error": None}
        except (sqlite3.Error, sqlite3.Warning) as error:
            result = {"candidate": name, "case": case, "expected_cents": expected,
                      "actual_cents": None, "passed": False, "error": str(error)}
        finally:
            connection.close()
        results.append(result)
    return results


def reference_plan():
    connection = fixture("one_paid_order")
    try:
        return connection.execute("EXPLAIN QUERY PLAN " + CANDIDATES["sale_price_lines"]).fetchall()
    finally:
        connection.close()


def main():
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--candidate-file", type=Path,
                        help="Read one SQLite SELECT from a local UTF-8 file")
    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)]
    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)
    plan = reference_plan()
    report = {"exercise": "revenue-grain", "unit": "USD cents",
              "sqlite_version": sqlite3.sqlite_version, "reference_query_plan": plan,
              "mode": "your_query" if args.candidate_file else "bundled_examples",
              "results": results, "checks_passed": checks_passed}
    if args.json:
        print(json.dumps(report, indent=2))
    else:
        print("SQLite version:", sqlite3.sqlite_version)
        print("ShopFlow revenue: paid orders only; sale-time line prices; one currency.")
        print("Hand-written expectations: one paid order = $100; two equal paid orders = $200.")
        for row in results:
            actual = row["error"] or str(row["actual_cents"])
            state = "PASSED" if row["passed"] else "REJECTED"
            if not args.candidate_file:
                case_index = list(EXPECTED).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']:23} {row['case']:22} "
                  f"actual_cents={actual}, expected_cents={row['expected_cents']}")
        if args.candidate_file:
            print("Your query passed all fixtures." if checks_passed else "Your query failed a fixture; inspect its grain and filters.")
        else:
            print("Teaching checks passed: expected rejections confirmed; corrected queries pass."
                  if checks_passed else "Teaching check mismatch: inspect the exercise.")
        print("Reference EXPLAIN QUERY PLAN (plans, does not run the SELECT; details vary by SQLite version):")
        for step in plan:
            print("  ", step[3])
        print("Passing these fixtures is evidence for this contract, not proof for all business data.")
    return 0 if checks_passed else 1


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