"""
python/tools/migrate_timing_per_media.py

更新タイミング設定 3テーブルに media_name カラムを追加し、
既存行（girlsheaven）をバニラ用に複製するワンショット・マイグレーション。

冪等（再実行しても安全）。実行前に C:\\backup\\ へバックアップを書き出す。
バックアップ書き出し失敗時は ALTER を一切実行しない。

Usage（本番VPS Tera Term）:
    C:\\apps\\ranking_system\\venv\\Scripts\\python.exe ^
        C:\\apps\\ranking_system\\python\\tools\\migrate_timing_per_media.py
"""

import os
import sys
from datetime import datetime, timedelta
from pathlib import Path

sys.path.append(str(Path(__file__).resolve().parents[1]))
from config.database import get_connection

DB_NAME    = "lerisa_rankingsystem"
BACKUP_DIR = r"C:\backup"

TABLES = [
    "update_timing_settings",
    "excluded_times",
    "fixed_update_times",
]


# ---------------------------------------------------------------------------
# バックアップ用ユーティリティ
# ---------------------------------------------------------------------------

def _fmt(v) -> str:
    """バックアップSQL用の値フォーマット（型ごとに適切なリテラルを返す）。"""
    if v is None:
        return "NULL"
    if isinstance(v, bool):
        return "1" if v else "0"
    if isinstance(v, int):
        return str(v)
    if isinstance(v, float):
        return repr(v)
    if isinstance(v, timedelta):
        # mysql.connector は TIME 型を timedelta で返す
        total = int(v.total_seconds())
        h, rem = divmod(total, 3600)
        m, s   = divmod(rem, 60)
        return f"'{h:02d}:{m:02d}:{s:02d}'"
    if isinstance(v, datetime):
        return f"'{v.strftime('%Y-%m-%d %H:%M:%S')}'"
    escaped = str(v).replace("\\", "\\\\").replace("'", "\\'")
    return f"'{escaped}'"


def _build_inserts(table: str, rows: list[dict]) -> str:
    """テーブルの全行を復元可能な INSERT 文の文字列として返す。"""
    if not rows:
        return f"-- TABLE: {table} (0 rows)\n"
    cols     = list(rows[0].keys())
    col_list = ", ".join(f"`{c}`" for c in cols)
    lines    = [f"-- TABLE: {table} ({len(rows)} rows)"]
    for row in rows:
        vals = ", ".join(_fmt(row[c]) for c in cols)
        lines.append(f"INSERT INTO `{table}` ({col_list}) VALUES ({vals});")
    return "\n".join(lines) + "\n"


# ---------------------------------------------------------------------------
# スキーマ確認ヘルパー
# ---------------------------------------------------------------------------

def has_column(cur, table: str) -> bool:
    """information_schema を参照して media_name 列の存在を確認する。"""
    cur.execute(
        "SELECT COUNT(*) AS cnt FROM information_schema.COLUMNS "
        "WHERE TABLE_SCHEMA = %s AND TABLE_NAME = %s AND COLUMN_NAME = 'media_name'",
        (DB_NAME, table),
    )
    return cur.fetchone()["cnt"] > 0


def find_client_only_unique(cur, table: str) -> str | None:
    """client_id 列のみで構成される UNIQUE インデックス名を動的に取得する。

    PRIMARY を除き、GROUP_CONCAT で構成カラムが 'client_id' 単独のものを返す。
    """
    cur.execute(
        """
        SELECT INDEX_NAME
        FROM information_schema.STATISTICS
        WHERE TABLE_SCHEMA = %s
          AND TABLE_NAME   = %s
          AND NON_UNIQUE   = 0
          AND INDEX_NAME  != 'PRIMARY'
        GROUP BY INDEX_NAME
        HAVING GROUP_CONCAT(COLUMN_NAME ORDER BY SEQ_IN_INDEX) = 'client_id'
        """,
        (DB_NAME, table),
    )
    rows = cur.fetchall()
    return rows[0]["INDEX_NAME"] if rows else None


# ---------------------------------------------------------------------------
# テーブル別マイグレーション
# ---------------------------------------------------------------------------

def migrate_update_timing_settings(conn, cur) -> dict:
    # a. media_name 列を追加（既存行は DEFAULT 'girlsheaven' になる）
    print("  [a] ALTER TABLE update_timing_settings ADD COLUMN media_name ...")
    cur.execute(
        "ALTER TABLE update_timing_settings "
        "  ADD COLUMN media_name VARCHAR(20) NOT NULL DEFAULT 'girlsheaven' AFTER client_id"
    )
    print("      → 既存行が girlsheaven に設定されました")

    # b. client_id 単独 UNIQUE を動的に取得して削除
    idx = find_client_only_unique(cur, "update_timing_settings")
    if idx:
        print(f"  [b] ALTER TABLE update_timing_settings DROP INDEX `{idx}`")
        cur.execute(f"ALTER TABLE update_timing_settings DROP INDEX `{idx}`")
    else:
        print("  [b] client_id 単独 UNIQUE が見つからない → DROP INDEX スキップ")

    # c. girlsheaven 行をバニラ用に複製
    print("  [c] INSERT vanilla 行（girlsheaven 行の複製）...")
    cur.execute(
        "INSERT INTO update_timing_settings (client_id, media_name, update_mode) "
        "  SELECT client_id, 'vanilla', update_mode "
        "  FROM update_timing_settings WHERE media_name = 'girlsheaven'"
    )
    conn.commit()
    print(f"      → {cur.rowcount} 行追加")

    # d. (client_id, media_name) の複合 UNIQUE を追加
    print("  [d] ALTER TABLE update_timing_settings ADD UNIQUE KEY uq_client_media ...")
    cur.execute(
        "ALTER TABLE update_timing_settings "
        "  ADD UNIQUE KEY uq_client_media (client_id, media_name)"
    )

    cur.execute(
        "SELECT media_name, COUNT(*) AS cnt "
        "FROM update_timing_settings GROUP BY media_name"
    )
    return {r["media_name"]: r["cnt"] for r in cur.fetchall()}


def migrate_excluded_times(conn, cur) -> dict:
    # a. media_name 列を追加
    print("  [a] ALTER TABLE excluded_times ADD COLUMN media_name ...")
    cur.execute(
        "ALTER TABLE excluded_times "
        "  ADD COLUMN media_name VARCHAR(20) NOT NULL DEFAULT 'girlsheaven' AFTER client_id"
    )

    # b. girlsheaven 行をバニラ用に複製
    print("  [b] INSERT vanilla 行（girlsheaven 行の複製）...")
    cur.execute(
        "INSERT INTO excluded_times (client_id, media_name, start_time, end_time) "
        "  SELECT client_id, 'vanilla', start_time, end_time "
        "  FROM excluded_times WHERE media_name = 'girlsheaven'"
    )
    conn.commit()
    print(f"      → {cur.rowcount} 行追加")

    # c. 複合インデックスを追加
    print("  [c] ALTER TABLE excluded_times ADD KEY idx_client_media ...")
    cur.execute(
        "ALTER TABLE excluded_times ADD KEY idx_client_media (client_id, media_name)"
    )

    cur.execute(
        "SELECT media_name, COUNT(*) AS cnt FROM excluded_times GROUP BY media_name"
    )
    return {r["media_name"]: r["cnt"] for r in cur.fetchall()}


def migrate_fixed_update_times(conn, cur) -> dict:
    # a. media_name 列を追加
    print("  [a] ALTER TABLE fixed_update_times ADD COLUMN media_name ...")
    cur.execute(
        "ALTER TABLE fixed_update_times "
        "  ADD COLUMN media_name VARCHAR(20) NOT NULL DEFAULT 'girlsheaven' AFTER client_id"
    )

    # b. girlsheaven 行をバニラ用に複製
    print("  [b] INSERT vanilla 行（girlsheaven 行の複製）...")
    cur.execute(
        "INSERT INTO fixed_update_times (client_id, media_name, update_time) "
        "  SELECT client_id, 'vanilla', update_time "
        "  FROM fixed_update_times WHERE media_name = 'girlsheaven'"
    )
    conn.commit()
    print(f"      → {cur.rowcount} 行追加")

    # c. 複合インデックスを追加
    print("  [c] ALTER TABLE fixed_update_times ADD KEY idx_client_media ...")
    cur.execute(
        "ALTER TABLE fixed_update_times ADD KEY idx_client_media (client_id, media_name)"
    )

    cur.execute(
        "SELECT media_name, COUNT(*) AS cnt FROM fixed_update_times GROUP BY media_name"
    )
    return {r["media_name"]: r["cnt"] for r in cur.fetchall()}


MIGRATE_FN = {
    "update_timing_settings": migrate_update_timing_settings,
    "excluded_times":         migrate_excluded_times,
    "fixed_update_times":     migrate_fixed_update_times,
}


# ---------------------------------------------------------------------------
# main
# ---------------------------------------------------------------------------

def main() -> None:
    print("=" * 64)
    print("migrate_timing_per_media.py")
    print(f"実行日時: {datetime.now().strftime('%Y-%m-%d %H:%M:%S')}")
    print("=" * 64)

    conn = get_connection()
    cur  = conn.cursor(dictionary=True)

    # ── 1. 実行前行数 ──────────────────────────────────────────────────
    print("\n[実行前 行数]")
    for table in TABLES:
        cur.execute(f"SELECT COUNT(*) AS cnt FROM `{table}`")
        print(f"  {table}: {cur.fetchone()['cnt']} 行")

    # ── 2. バックアップ ────────────────────────────────────────────────
    print("\n[バックアップ作成]")
    ts          = datetime.now().strftime("%Y%m%d_%H%M%S")
    backup_path = os.path.join(BACKUP_DIR, f"kagoya_timing_backup_{ts}.sql")

    backup_parts = [
        f"-- kagoya_timing_backup_{ts}.sql",
        f"-- DB: {DB_NAME}",
        f"-- 作成日時: {datetime.now().strftime('%Y-%m-%d %H:%M:%S')}",
        "",
    ]
    all_rows   = {}
    total_rows = 0
    for table in TABLES:
        cur.execute(f"SELECT * FROM `{table}`")
        rows            = cur.fetchall()
        all_rows[table] = rows
        total_rows     += len(rows)
        backup_parts.append(_build_inserts(table, rows))

    try:
        os.makedirs(BACKUP_DIR, exist_ok=True)
        with open(backup_path, "w", encoding="utf-8") as f:
            f.write("\n".join(backup_parts))
        print(f"  保存先: {backup_path}")
        print(f"  合計 {total_rows} 行をバックアップ")
    except Exception as e:
        print(f"  [ERROR] バックアップ書き込み失敗: {e}")
        print("  → ALTER を実行せず中止します")
        cur.close()
        conn.close()
        sys.exit(1)

    # ── 3. テーブル別マイグレーション ─────────────────────────────────
    migrated: list[tuple[str, dict]] = []
    skipped:  list[str]              = []

    for table in TABLES:
        print(f"\n{'─' * 48}")
        print(f"[{table}]")
        if has_column(cur, table):
            print(f"  already migrated: {table}")
            skipped.append(table)
        else:
            try:
                counts = MIGRATE_FN[table](conn, cur)
                migrated.append((table, counts))
                gh  = counts.get("girlsheaven", 0)
                van = counts.get("vanilla",     0)
                print(f"  → girlsheaven: {gh} 行 / vanilla: {van} 行")
            except Exception as e:
                print(f"  [ERROR] {table} マイグレーション中にエラー: {e}")
                print("  → バックアップから復元してください")
                print(f"  → バックアップ: {backup_path}")
                cur.close()
                conn.close()
                sys.exit(1)

    # ── 4. 完了レポート ────────────────────────────────────────────────
    print(f"\n{'=' * 64}")
    if migrated:
        print("[マイグレーション後 行数]")
        for table, counts in migrated:
            gh  = counts.get("girlsheaven", 0)
            van = counts.get("vanilla",     0)
            print(f"  {table}: girlsheaven={gh} 行 / vanilla={van} 行")
        if skipped:
            print(f"[スキップ（already migrated）]: {', '.join(skipped)}")
        print("\nMIGRATION COMPLETE")
    else:
        print("NOTHING TO DO")
    print("=" * 64)

    cur.close()
    conn.close()


if __name__ == "__main__":
    main()
