#!/usr/bin/env python3
import argparse
import getpass
import os
import sys
from pathlib import Path

try:
    import ijson
except ImportError:
    print("ERROR: Missing dependency 'ijson'.", file=sys.stderr)
    print("Install with:", file=sys.stderr)
    print("  /usr/bin/python3 -m pip install ijson mysql-connector-python", file=sys.stderr)
    sys.exit(1)

try:
    import mysql.connector
except ImportError:
    print("ERROR: Missing dependency 'mysql-connector-python'.", file=sys.stderr)
    print("Install with:", file=sys.stderr)
    print("  /usr/bin/python3 -m pip install ijson mysql-connector-python", file=sys.stderr)
    sys.exit(1)

CREATE_DATABASE_SQL = """
CREATE DATABASE IF NOT EXISTS public_directory
  CHARACTER SET utf8mb4
  COLLATE utf8mb4_unicode_ci
"""

CREATE_TABLE_SQL = """
CREATE TABLE IF NOT EXISTS people_addresses (
    id BIGINT UNSIGNED NOT NULL AUTO_INCREMENT,

    first_name VARCHAR(40) NULL,
    middle_name VARCHAR(40) NULL,
    last_name VARCHAR(60) NULL,
    full_name VARCHAR(150) NULL,
    title VARCHAR(20) NULL,

    street_number VARCHAR(20) NULL,
    street VARCHAR(100) NULL,
    unit VARCHAR(30) NULL,
    city VARCHAR(100) NULL,
    state VARCHAR(10) NULL,
    zip_code VARCHAR(20) NULL,
    county VARCHAR(100) NULL,
    country VARCHAR(20) NULL,

    PRIMARY KEY (id)
) ENGINE=InnoDB
  DEFAULT CHARSET=utf8mb4
  COLLATE=utf8mb4_unicode_ci
"""

INSERT_SQL = """
INSERT INTO people_addresses (
    first_name,
    middle_name,
    last_name,
    full_name,
    title,
    street_number,
    street,
    unit,
    city,
    state,
    zip_code,
    county,
    country
) VALUES (
    %s, %s, %s, %s, %s,
    %s, %s, %s, %s, %s,
    %s, %s, %s
)
"""

INDEX_SQL = [
    "CREATE INDEX idx_zip ON people_addresses (zip_code)",
    "CREATE INDEX idx_city ON people_addresses (city)",
    "CREATE INDEX idx_state ON people_addresses (state)",
    "CREATE INDEX idx_last_name ON people_addresses (last_name)",
    "CREATE INDEX idx_street ON people_addresses (street)",
    "CREATE INDEX idx_city_state ON people_addresses (city, state)",
    "CREATE INDEX idx_city_zip ON people_addresses (city, zip_code)",
    "CREATE INDEX idx_last_city ON people_addresses (last_name, city)",
    "CREATE INDEX idx_zip_street ON people_addresses (zip_code, street)",
    "CREATE INDEX idx_city_street_number ON people_addresses (city, street, street_number)",
]

FULLTEXT_SQL = """
ALTER TABLE people_addresses
ADD FULLTEXT INDEX ft_people_search (
    first_name,
    middle_name,
    last_name,
    full_name,
    street,
    city,
    county
)
"""

def clean(value):
    if value is None:
        return None
    value = str(value).strip()
    return value if value else None

def make_row(doc):
    address = doc.get("address")
    if not isinstance(address, dict):
        address = {}

    return (
        clean(doc.get("firstName")),
        clean(doc.get("middleName")),
        clean(doc.get("lastName")),
        clean(doc.get("name")),
        clean(doc.get("title")),

        clean(address.get("streetNumber")),
        clean(address.get("street")),
        clean(address.get("unit")),
        clean(address.get("city")),
        clean(address.get("state")),
        clean(address.get("zipCode")),
        clean(address.get("county")),
        clean(address.get("country")),
    )

def index_exists(cur, db_name, table_name, index_name):
    cur.execute(
        """
        SELECT COUNT(*)
        FROM information_schema.statistics
        WHERE table_schema = %s
          AND table_name = %s
          AND index_name = %s
        """,
        (db_name, table_name, index_name),
    )
    return cur.fetchone()[0] > 0

def create_indexes(cur, conn, db_name):
    index_pairs = [
        ("idx_zip", INDEX_SQL[0]),
        ("idx_city", INDEX_SQL[1]),
        ("idx_state", INDEX_SQL[2]),
        ("idx_last_name", INDEX_SQL[3]),
        ("idx_street", INDEX_SQL[4]),
        ("idx_city_state", INDEX_SQL[5]),
        ("idx_city_zip", INDEX_SQL[6]),
        ("idx_last_city", INDEX_SQL[7]),
        ("idx_zip_street", INDEX_SQL[8]),
        ("idx_city_street_number", INDEX_SQL[9]),
    ]

    for name, sql in index_pairs:
        if not index_exists(cur, db_name, "people_addresses", name):
            print(f"Creating index {name}...", flush=True)
            cur.execute(sql)
            conn.commit()

    if not index_exists(cur, db_name, "people_addresses", "ft_people_search"):
        print("Creating FULLTEXT index ft_people_search...", flush=True)
        cur.execute(FULLTEXT_SQL)
        conn.commit()

def main():
    ap = argparse.ArgumentParser(
        description="Stream MongoDB-style JSON into a generic MySQL people/address directory."
    )
    ap.add_argument("--input", required=True, help="Path to giant JSON array file")
    ap.add_argument("--host", default="127.0.0.1")
    ap.add_argument("--port", type=int, default=3306)
    ap.add_argument("--user", required=True)
    ap.add_argument("--password", help="MySQL password; if omitted, prompt securely")
    ap.add_argument("--database", default="public_directory")
    ap.add_argument("--batch-size", type=int, default=5000)
    ap.add_argument("--progress-every", type=int, default=100000)
    ap.add_argument("--truncate", action="store_true")
    ap.add_argument(
        "--skip-indexes",
        action="store_true",
        help="Do not build indexes automatically after import",
    )
    args = ap.parse_args()

    src = Path(args.input)
    if not src.is_file():
        print(f"ERROR: Input file not found: {src}", file=sys.stderr)
        sys.exit(1)

    password = args.password or os.environ.get("MYSQL_PASSWORD")
    if password is None:
        password = getpass.getpass("MySQL password: ")

    # Connect without selecting a DB first so we can create it if necessary.
    server_conn = mysql.connector.connect(
        host=args.host,
        port=args.port,
        user=args.user,
        password=password,
        autocommit=True,
        charset="utf8mb4",
    )
    server_cur = server_conn.cursor()
    server_cur.execute(CREATE_DATABASE_SQL.replace("public_directory", f"`{args.database}`"))
    server_cur.close()
    server_conn.close()

    conn = mysql.connector.connect(
        host=args.host,
        port=args.port,
        user=args.user,
        password=password,
        database=args.database,
        autocommit=False,
        charset="utf8mb4",
    )
    cur = conn.cursor()
    cur.execute(CREATE_TABLE_SQL)
    conn.commit()

    if args.truncate:
        print("Truncating people_addresses...", flush=True)
        cur.execute("TRUNCATE TABLE people_addresses")
        conn.commit()

    processed = 0
    imported = 0
    skipped = 0
    batch = []

    try:
        with src.open("rb") as f:
            for doc in ijson.items(f, "item"):
                processed += 1

                if not isinstance(doc, dict):
                    skipped += 1
                    continue

                row = make_row(doc)

                # Require at least a usable name OR usable address.
                has_name = any(row[i] for i in (0, 2, 3))
                has_address = any(row[i] for i in (5, 6, 8, 10))

                if not has_name and not has_address:
                    skipped += 1
                    continue

                batch.append(row)

                if len(batch) >= args.batch_size:
                    cur.executemany(INSERT_SQL, batch)
                    conn.commit()
                    imported += len(batch)
                    batch.clear()

                if processed % args.progress_every == 0:
                    print(
                        f"Processed={processed:,} "
                        f"Imported={imported:,} "
                        f"Skipped={skipped:,}",
                        flush=True,
                    )

        if batch:
            cur.executemany(INSERT_SQL, batch)
            conn.commit()
            imported += len(batch)
            batch.clear()

        cur.execute("SELECT COUNT(*) FROM people_addresses")
        mysql_count = cur.fetchone()[0]

        print("")
        print("DATA IMPORT COMPLETE")
        print(f"Processed:       {processed:,}")
        print(f"Imported:        {imported:,}")
        print(f"Skipped:         {skipped:,}")
        print(f"MySQL row count: {mysql_count:,}")

        if not args.skip_indexes:
            print("")
            print("Building search indexes after bulk import...")
            create_indexes(cur, conn, args.database)

            print("Analyzing table...")
            cur.execute("ANALYZE TABLE people_addresses")
            result = cur.fetchall()
            conn.commit()
            for row_result in result:
                print(" | ".join(str(x) for x in row_result))

        print("")
        print("IMPORT VERIFIED")
        print("No voter-status, voter-ID, party, DOB, petition, signature, validation,")
        print("flag, embedding, or MongoDB metadata fields are imported by this script.")

    except Exception:
        conn.rollback()
        raise
    finally:
        cur.close()
        conn.close()

if __name__ == "__main__":
    main()
