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

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

def get_path(doc, path):
    cur = doc
    for part in path.split("."):
        if not isinstance(cur, dict):
            return None
        cur = cur.get(part)
        if cur is None:
            return None
    return cur

def clean(value):
    if value is None:
        return None
    if isinstance(value, (dict, list)):
        return json.dumps(value, ensure_ascii=False, separators=(",", ":"))
    if isinstance(value, bool):
        return 1 if value else 0
    s = str(value).strip()
    return s if s else None

def column_name(path):
    import re
    p = path.replace("[]", "_item")
    p = p.replace(".", "_")
    p = re.sub(r"[^A-Za-z0-9_]+", "_", p)
    return p.lower().strip("_")[:64]

def main():
    ap = argparse.ArgumentParser(description="Stream audited voter JSON into MySQL using generated import map.")
    ap.add_argument("--input", required=True)
    ap.add_argument("--map", required=True, help="mysql_import_map.json generated by audit_voter_json.py")
    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")
    ap.add_argument("--database", default="voter_data")
    ap.add_argument("--table", default="voters_public")
    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("--resume-from", type=int, default=0, help="Skip this many input records before importing")
    args = ap.parse_args()

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

    config = json.loads(map_path.read_text(encoding="utf-8"))
    fields = config.get("include_fields") or []
    if not fields:
        print("ERROR: No include_fields in import map.", file=sys.stderr)
        sys.exit(1)

    cols = [column_name(p) for p in fields]
    quoted_cols = ", ".join(f"`{c}`" for c in cols)
    placeholders = ", ".join(["%s"] * len(cols))
    insert_sql = f"INSERT INTO `{args.table}` ({quoted_cols}) VALUES ({placeholders})"

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

    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()

    if args.truncate:
        cur.execute(f"TRUNCATE TABLE `{args.table}`")
        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 processed <= args.resume_from:
                    continue
                if not isinstance(doc, dict):
                    skipped += 1
                    continue

                row = tuple(clean(get_path(doc, p)) for p in fields)

                # Skip totally empty projected rows only.
                if not any(v is not None for v in row):
                    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:,} Imported={imported:,} "
                        f"Skipped={skipped:,}",
                        flush=True
                    )

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

    except Exception as e:
        conn.rollback()
        print(f"IMPORT FAILED after input record {processed:,}: {e}", file=sys.stderr)
        print(f"You can resume with: --resume-from {processed - len(batch) - 1}", file=sys.stderr)
        raise
    finally:
        cur.close()
        conn.close()

    print("")
    print("IMPORT COMPLETE")
    print(f"Processed: {processed:,}")
    print(f"Imported:  {imported:,}")
    print(f"Skipped:   {skipped:,}")

if __name__ == "__main__":
    main()
