"""Test SQL against a tiny fixture whose right answers were worked out by hand.

Each case holds the expected rows, an AI-style draft and the reviewed query.
Run with: python sql_review_fixture.py (standard library only, uses SQLite).
"""

import sqlite3
from typing import NamedTuple

FIXTURE = """
CREATE TABLE referrals (referral_id INT, patient_id INT, specialty TEXT, referral_date TEXT);
CREATE TABLE appointments (appt_id INT, referral_id INT, appt_start TEXT, status TEXT);
CREATE TABLE diagnoses (referral_id INT, code TEXT);
CREATE TABLE patient_address (patient_id INT, district TEXT, valid_from TEXT, valid_to TEXT);

INSERT INTO referrals VALUES
  (1, 101, 'Cardiology', '2025-02-10'), (2, 102, 'Cardiology', '2025-02-20'),
  (3, 103, 'Diabetes',   '2025-03-01'), (4, 104, 'Diabetes',   '2025-03-05');
INSERT INTO appointments VALUES
  (1, 1,    '2025-03-03 09:00', 'Attended'),
  (2, 1,    '2025-03-17 10:30', 'DNA'),
  (3, 2,    '2025-03-31 14:00', 'Attended'),   -- last day of the month, after midnight
  (4, 3,    '2025-03-12 11:00', 'Cancelled'),
  (5, 3,    '2025-03-26 15:15', NULL),         -- not yet outcomed
  (6, NULL, '2025-04-02 09:00', 'Attended');   -- walk-in, no referral
INSERT INTO diagnoses VALUES (1, 'I10'), (3, 'E11.9'), (3, 'E11.65'), (4, 'E10.9');
INSERT INTO patient_address VALUES
  (101, 'SA1', '2020-01-01', '2025-02-15'), (101, 'SA2', '2025-02-15', '9999-12-31'),
  (102, 'SA1', '2019-06-01', '9999-12-31'), (103, 'SA4', '2021-03-01', '9999-12-31'),
  (104, 'SA6', '2018-01-01', '9999-12-31');
"""

MARCH = "appt_start >= '2025-03-01' AND appt_start < '2025-04-01'"


class Case(NamedTuple):
    name: str
    expected: list[tuple]
    draft: str
    reviewed: str


CASES = [
    Case(
        "march_appointments",
        [(5,)],
        "SELECT COUNT(*) FROM appointments WHERE appt_start BETWEEN '2025-03-01' AND '2025-03-31'",
        f"SELECT COUNT(*) FROM appointments WHERE {MARCH}",
    ),
    Case(
        "type2_diabetes_appointments",
        [(2,)],
        "SELECT COUNT(*) FROM appointments a "
        "JOIN diagnoses d ON d.referral_id = a.referral_id WHERE d.code LIKE 'E11%'",
        "SELECT COUNT(*) FROM appointments a WHERE EXISTS (SELECT 1 FROM diagnoses d "
        "WHERE d.referral_id = a.referral_id AND d.code LIKE 'E11%')",
    ),
    Case(
        "attended_by_specialty",
        [("Cardiology", 2), ("Diabetes", 0)],
        "SELECT r.specialty, COUNT(a.appt_id) FROM referrals r "
        "LEFT JOIN appointments a ON a.referral_id = r.referral_id "
        "WHERE a.status = 'Attended' GROUP BY r.specialty ORDER BY r.specialty",
        "SELECT r.specialty, COUNT(a.appt_id) FROM referrals r "
        "LEFT JOIN appointments a ON a.referral_id = r.referral_id AND a.status = 'Attended' "
        "GROUP BY r.specialty ORDER BY r.specialty",
    ),
    Case(
        "referrals_never_booked",
        [(4,)],
        "SELECT referral_id FROM referrals "
        "WHERE referral_id NOT IN (SELECT referral_id FROM appointments)",
        "SELECT r.referral_id FROM referrals r WHERE NOT EXISTS "
        "(SELECT 1 FROM appointments a WHERE a.referral_id = r.referral_id)",
    ),
    Case(
        "march_dna_rate",
        [(1 / 3,)],
        "SELECT SUM(CASE WHEN status = 'DNA' THEN 1 ELSE 0 END) "
        "/ SUM(CASE WHEN status IN ('Attended', 'DNA') THEN 1 ELSE 0 END) "
        f"FROM appointments WHERE {MARCH}",
        "SELECT 1.0 * SUM(CASE WHEN status = 'DNA' THEN 1 ELSE 0 END) "
        "/ NULLIF(SUM(CASE WHEN status IN ('Attended', 'DNA') THEN 1 ELSE 0 END), 0) "
        f"FROM appointments WHERE {MARCH}",
    ),
    Case(
        "referral_district",
        [(1, "SA1"), (2, "SA1"), (3, "SA4"), (4, "SA6")],
        "SELECT DISTINCT r.referral_id, p.district FROM referrals r "
        "JOIN patient_address p ON p.patient_id = r.patient_id ORDER BY r.referral_id",
        "SELECT r.referral_id, p.district FROM referrals r "
        "JOIN patient_address p ON p.patient_id = r.patient_id "
        "AND p.valid_from <= r.referral_date AND r.referral_date < p.valid_to "
        "ORDER BY r.referral_id",
    ),
]


def verdict(con: sqlite3.Connection, sql: str, expected: list[tuple]) -> str:
    """PASS if the query returns exactly the expected rows, otherwise FAIL and the rows."""
    got = con.execute(sql).fetchall()
    return "PASS" if got == expected else f"FAIL {got}"


def main() -> None:
    con = sqlite3.connect(":memory:")
    con.executescript(FIXTURE)
    for case in CASES:
        reviewed = verdict(con, case.reviewed, case.expected)
        draft = verdict(con, case.draft, case.expected)
        print(f"{case.name:<28} reviewed {reviewed:<5} draft {draft}")


if __name__ == "__main__":
    main()
