"""
Migration script: Drupal MySQL (nirta_letter) → PostgreSQL (jarvis)

Run from backend/ folder:
    source venv/bin/activate
    python migrations/migrate_mysql_to_pg.py

Requirements: pymysql, asyncpg, sqlalchemy[asyncio] (already in requirements.txt)
"""
import asyncio
import os
import sys

import pymysql
import pymysql.cursors
from sqlalchemy.ext.asyncio import create_async_engine, async_sessionmaker
from dotenv import load_dotenv

load_dotenv(os.path.join(os.path.dirname(__file__), "../.env"))

sys.path.insert(0, os.path.join(os.path.dirname(__file__), ".."))

from app.database import Base
from app.models.user import User
from app.models.company import Company
from app.models.contract import Contract
from app.services.auth import hash_password

# ── Config ────────────────────────────────────────────────────────────────────
MYSQL_CONFIG = {
    "host": os.getenv("MYSQL_HOST", "localhost"),
    "port": int(os.getenv("MYSQL_PORT", 3306)),
    "user": os.getenv("MYSQL_USER", "root"),
    "password": os.getenv("MYSQL_PASSWORD", ""),
    "database": os.getenv("MYSQL_DATABASE", "nirta_letter"),
    "charset": "utf8mb4",
    "cursorclass": pymysql.cursors.DictCursor,
}

PG_URL = os.getenv("DATABASE_URL")  # postgresql+asyncpg://...


def fetch_mysql(query: str, params=None) -> list[dict]:
    conn = pymysql.connect(**MYSQL_CONFIG)
    try:
        with conn.cursor() as cur:
            cur.execute(query, params or ())
            return cur.fetchall()
    finally:
        conn.close()


async def run_migration():
    print("=== Jarvis Migration: MySQL → PostgreSQL ===\n")

    # ── 1. Connect PostgreSQL & create tables ─────────────────────────────────
    engine = create_async_engine(PG_URL, echo=False)
    SessionLocal = async_sessionmaker(engine, expire_on_commit=False)

    async with engine.begin() as conn:
        await conn.run_sync(Base.metadata.create_all)
    print("✓ PostgreSQL tables created\n")

    async with SessionLocal() as db:

        # ── 2. Migrate companies ─────────────────────────────────────────────
        print("→ Migrating companies...")
        company_rows = fetch_mysql("""
            SELECT DISTINCT nfv.field_company_target_id AS taxonomy_id,
                   td.name AS company_name
            FROM node__field_company nfv
            JOIN taxonomy_term_field_data td ON td.tid = nfv.field_company_target_id
        """)

        company_map = {}  # drupal taxonomy_id → new pg id
        for row in company_rows:
            existing = await db.execute(
                __import__("sqlalchemy").select(Company).where(Company.name == row["company_name"])
            )
            company = existing.scalar_one_or_none()
            if not company:
                company = Company(name=row["company_name"])
                db.add(company)
                await db.flush()
            company_map[row["taxonomy_id"]] = company.id

        await db.commit()
        print(f"  ✓ {len(company_map)} companies migrated\n")

        # ── 3. Migrate users ─────────────────────────────────────────────────
        print("→ Migrating users...")
        drupal_users = fetch_mysql("""
            SELECT uid, name, mail, status
            FROM users_field_data
            WHERE uid > 0
        """)

        user_map = {}  # drupal uid → new pg id
        for u in drupal_users:
            existing = await db.execute(
                __import__("sqlalchemy").select(User).where(User.username == u["name"])
            )
            user = existing.scalar_one_or_none()
            if not user:
                user = User(
                    username=u["name"],
                    email=u["mail"] or f"{u['name']}@jarvis.local",
                    hashed_password=hash_password("changeme123"),  # temp password
                    role="viewer",
                    is_active=bool(u["status"]),
                )
                db.add(user)
                await db.flush()
            user_map[u["uid"]] = user.id

        await db.commit()
        print(f"  ✓ {len(user_map)} users migrated")
        print("  ⚠ All migrated users have temporary password: changeme123\n")

        # ── 4. Migrate contracts/nodes ───────────────────────────────────────
        print("→ Migrating contracts/nodes...")
        nodes = fetch_mysql("""
            SELECT n.nid, n.type, n.uid, nfd.title, nfd.status,
                   nos.field_no_surat_value      AS no_surat,
                   np.field_perihal_value        AS perihal,
                   nt.field_tanggal_value        AS tanggal,
                   na.field_approve_value        AS is_approved,
                   ns.field_status_value         AS status,
                   nc.field_company_target_id    AS company_tid,
                   nf.field_file_target_id       AS file_fid,
                   n.created
            FROM node n
            JOIN node_field_data nfd ON nfd.nid = n.nid
            LEFT JOIN node__field_no_surat nos ON nos.entity_id = n.nid
            LEFT JOIN node__field_perihal np ON np.entity_id = n.nid
            LEFT JOIN node__field_tanggal nt ON nt.entity_id = n.nid
            LEFT JOIN node__field_approve na ON na.entity_id = n.nid
            LEFT JOIN node__field_status ns ON ns.entity_id = n.nid
            LEFT JOIN node__field_company nc ON nc.entity_id = n.nid
            LEFT JOIN node__field_file nf ON nf.entity_id = n.nid
            WHERE n.type IN ('kontrak', 'page')
            ORDER BY n.nid
        """)

        count = 0
        for row in nodes:
            from datetime import datetime, timezone
            contract = Contract(
                type=row["type"],
                title=row["title"] or "Untitled",
                no_surat=row["no_surat"],
                perihal=row["perihal"],
                tanggal=row["tanggal"],
                company_id=company_map.get(row["company_tid"]) if row["company_tid"] else None,
                is_approved=bool(row["is_approved"]),
                status=int(row["status"] or 0),
                created_by=user_map.get(row["uid"]),
                created_at=datetime.fromtimestamp(row["created"], tz=timezone.utc) if row["created"] else None,
            )
            db.add(contract)
            count += 1

        await db.commit()
        print(f"  ✓ {count} contracts migrated\n")

    await engine.dispose()
    print("=== Migration complete! ===")
    print("\nNext steps:")
    print("  1. Buat admin user: POST /api/v1/auth/register dengan role='admin'")
    print("  2. Reset password user yang dimigrasikan via panel Users")
    print("  3. Verifikasi data di http://localhost:8000/docs")


if __name__ == "__main__":
    asyncio.run(run_migration())
