#!/usr/bin/env python3
"""Copy nmdpr_supply SQLite into MySQL. Safe to re-run."""
import os
import sqlite3
import sys
from urllib.parse import quote_plus

sys.path.insert(0, os.path.join(os.path.dirname(__file__), "..", "nmdpr_supply", "backend"))
from dotenv import load_dotenv
load_dotenv(os.path.join(os.path.dirname(__file__), "..", "nmdpr_supply", "backend", ".env"))

from sqlalchemy import create_engine, inspect, text, MetaData, Table

HERE = os.path.join(os.path.dirname(__file__), "..", "nmdpr_supply", "backend")
SQLITE = os.path.join(HERE, "nmdpra.db")


def mysql_url():
    url = os.getenv("DATABASE_URL", "")
    if url.startswith("mysql"):
        return url
    host = os.getenv("MYSQL_HOST")
    if not host:
        raise SystemExit("MySQL is not configured in nmdpr_supply/backend/.env")
    user = os.getenv("MYSQL_USER", "nmdpr")
    password = quote_plus(os.getenv("MYSQL_PASSWORD", ""))
    port = os.getenv("MYSQL_PORT", "3306")
    db = os.getenv("MYSQL_DATABASE", "nmdpr_supply")
    return f"mysql+pymysql://{user}:{password}@{host}:{port}/{db}?charset=utf8mb4"


def main():
    from app.database import engine, Base
    import app.models  # noqa: F401 — register tables on Base.metadata
    Base.metadata.create_all(bind=engine)
    if not os.path.exists(SQLITE):
        print(f"No supply SQLite file at {SQLITE} — schema created, nothing to copy.")
        return

    src = create_engine(f"sqlite:///{SQLITE}")
    dst = create_engine(mysql_url(), pool_pre_ping=True)
    src_meta = MetaData()
    src_meta.reflect(bind=src)
    print("Migrating supply SQLite → MySQL")
    dst_names = set(inspect(dst).get_table_names())
    with dst.begin() as conn:
        conn.execute(text("SET FOREIGN_KEY_CHECKS=0"))
        for table in src_meta.sorted_tables:
            if table.name not in dst_names:
                print(f"  {table.name}: skipped (not in MySQL schema)")
                continue
            rows = list(src.connect().execute(table.select()))
            dst_table = Table(table.name, MetaData(), autoload_with=dst)
            conn.execute(dst_table.delete())
            if rows:
                payload = [dict(r._mapping) for r in rows]
                for i in range(0, len(payload), 200):
                    conn.execute(dst_table.insert(), payload[i:i + 200])
            print(f"  {table.name}: {len(rows)} rows")
        conn.execute(text("SET FOREIGN_KEY_CHECKS=1"))
    print("Supply MySQL is loaded.")


if __name__ == "__main__":
    main()
