import datetime as dt
import decimal
import json
import os
from pathlib import Path

import pymysql
import psycopg
from dotenv import load_dotenv
from psycopg import sql

load_dotenv(Path(__file__).resolve().parent / ".env")

MYSQL_CONFIG = {
    "host": os.getenv("MYSQL_SOURCE_HOST", os.getenv("DB_HOST", "127.0.0.1")),
    "port": int(os.getenv("MYSQL_SOURCE_PORT", os.getenv("DB_PORT", "3306"))),
    "user": os.getenv("MYSQL_SOURCE_USER", os.getenv("DB_USER", "root")),
    "password": os.getenv("MYSQL_SOURCE_PASSWORD", os.getenv("DB_PASSWORD", "")),
    "database": os.getenv("MYSQL_SOURCE_NAME", "carpay-dev"),
    "charset": "utf8mb4",
    "cursorclass": pymysql.cursors.SSCursor,
}

POSTGRES_CONFIG = {
    "host": os.getenv("DB_HOST", "127.0.0.1"),
    "port": int(os.getenv("DB_PORT", "5432")),
    "user": os.getenv("DB_USER", "postgres"),
    "password": os.getenv("DB_PASSWORD", ""),
}

TARGET_DB = "paython_db"


def map_mysql_to_postgres(data_type, column_type, char_len, num_precision, num_scale):
    t = data_type.lower()
    ct = (column_type or "").lower()

    if t == "tinyint" and ct == "tinyint(1)":
        return "boolean"
    if t in {"tinyint", "smallint"}:
        return "smallint"
    if t in {"mediumint", "int", "integer"}:
        return "integer"
    if t == "bigint":
        return "bigint"
    if t in {"decimal", "numeric"}:
        if num_precision is not None and num_scale is not None:
            return f"numeric({num_precision},{num_scale})"
        return "numeric"
    if t in {"float", "double"}:
        return "double precision"
    if t in {"varchar", "char"} and char_len:
        return f"varchar({char_len})"
    if t in {"text", "mediumtext", "longtext", "tinytext"}:
        return "text"
    if t in {"json"}:
        return "jsonb"
    if t in {"datetime", "timestamp"}:
        return "timestamp"
    if t == "date":
        return "date"
    if t == "time":
        return "time"
    if t == "year":
        return "integer"
    if t in {"blob", "mediumblob", "longblob", "tinyblob", "binary", "varbinary"}:
        return "bytea"
    if t in {"enum", "set"}:
        return "text"
    return "text"


def ensure_target_db():
    conn = psycopg.connect(dbname="postgres", autocommit=True, **POSTGRES_CONFIG)
    try:
        with conn.cursor() as cur:
            cur.execute("SELECT 1 FROM pg_database WHERE datname = %s", (TARGET_DB,))
            exists = cur.fetchone() is not None
            if not exists:
                cur.execute(
                    sql.SQL("CREATE DATABASE {}").format(sql.Identifier(TARGET_DB))
                )
    finally:
        conn.close()


def get_mysql_tables(mysql_conn):
    with mysql_conn.cursor() as cur:
        cur.execute(
            """
            SELECT table_name
            FROM information_schema.tables
            WHERE table_schema = %s AND table_type = 'BASE TABLE'
            ORDER BY table_name
            """,
            (MYSQL_CONFIG["database"],),
        )
        return [row[0] for row in cur.fetchall()]


def get_mysql_columns(mysql_conn, table_name):
    with mysql_conn.cursor() as cur:
        cur.execute(
            """
            SELECT
                column_name,
                data_type,
                column_type,
                is_nullable,
                character_maximum_length,
                numeric_precision,
                numeric_scale
            FROM information_schema.columns
            WHERE table_schema = %s AND table_name = %s
            ORDER BY ordinal_position
            """,
            (MYSQL_CONFIG["database"], table_name),
        )
        return cur.fetchall()


def normalize_value(value, pg_type):
    if value is None:
        return None
    if pg_type == "boolean":
        return bool(value)
    if isinstance(value, (dict, list)):
        return json.dumps(value, ensure_ascii=False)
    if isinstance(value, decimal.Decimal):
        return value
    if isinstance(value, (dt.date, dt.datetime, dt.time)):
        return value
    return value


def migrate_table(mysql_conn, pg_conn, table_name):
    columns = get_mysql_columns(mysql_conn, table_name)

    drop_stmt = sql.SQL("DROP TABLE IF EXISTS {} CASCADE").format(sql.Identifier(table_name))
    with pg_conn.cursor() as cur:
        cur.execute(drop_stmt)

    col_defs = []
    col_names = []
    col_pg_types = []
    for col in columns:
        col_name, data_type, column_type, is_nullable, char_len, num_precision, num_scale = col
        pg_type = map_mysql_to_postgres(
            data_type, column_type, char_len, num_precision, num_scale
        )
        nullable = "" if is_nullable == "YES" else " NOT NULL"
        col_defs.append(
            sql.SQL("{} {}{}").format(
                sql.Identifier(col_name),
                sql.SQL(pg_type),
                sql.SQL(nullable),
            )
        )
        col_names.append(col_name)
        col_pg_types.append(pg_type)

    create_stmt = sql.SQL("CREATE TABLE {} ({})").format(
        sql.Identifier(table_name), sql.SQL(", ").join(col_defs)
    )
    with pg_conn.cursor() as cur:
        cur.execute(create_stmt)

    with mysql_conn.cursor() as src_cur, pg_conn.cursor() as dst_cur:
        src_cur.execute(f"SELECT * FROM `{table_name}`")
        placeholders = sql.SQL(", ").join(sql.Placeholder() for _ in col_names)
        insert_stmt = sql.SQL("INSERT INTO {} ({}) VALUES ({})").format(
            sql.Identifier(table_name),
            sql.SQL(", ").join(sql.Identifier(c) for c in col_names),
            placeholders,
        )

        batch = []
        while True:
            rows = src_cur.fetchmany(1000)
            if not rows:
                break
            for row in rows:
                batch.append(
                    tuple(
                        normalize_value(value, col_pg_types[idx])
                        for idx, value in enumerate(row)
                    )
                )
            dst_cur.executemany(insert_stmt, batch)
            batch.clear()

    pg_conn.commit()


def main():
    print("Ensuring target PostgreSQL database exists...")
    ensure_target_db()

    print("Connecting to MySQL and PostgreSQL...")
    mysql_conn = pymysql.connect(**MYSQL_CONFIG)
    pg_conn = psycopg.connect(dbname=TARGET_DB, **POSTGRES_CONFIG)

    try:
        tables = get_mysql_tables(mysql_conn)
        print(f"Found {len(tables)} tables in MySQL.")
        for i, table in enumerate(tables, start=1):
            print(f"[{i}/{len(tables)}] Migrating table: {table}")
            migrate_table(mysql_conn, pg_conn, table)
        print("Migration finished successfully.")
    finally:
        mysql_conn.close()
        pg_conn.close()


if __name__ == "__main__":
    main()
