#!/usr/bin/env python
"""
migration_rehearsal.py
======================
Automated migration rehearsal script for the TrueWave platform.

Phases executed (in order):
  1. Pre-flight checks      -- django check, pending migration detection
  2. Backup                 -- snapshot current DB to timestamped file
  3. Migration plan         -- dry-run plan output + timing estimate
  4. Apply migrations       -- timed apply with per-migration logging
  5. Post-migration validation  -- data integrity checks across all affected apps
  6. Rollback drill         -- unapply migrations in reverse order, timed
  7. Re-apply               -- re-apply after rollback confirms reversibility
  8. Final report           -- write MIGRATION_REHEARSAL_REPORT.md

Usage (from backend/):
    python scripts/migration_rehearsal.py [--no-rollback] [--db-path PATH]

Environment:
    Reads DATABASE_URL or defaults to db.sqlite3 in backend directory.
    Set DJANGO_SETTINGS_MODULE if not already set.

Safety:
    - Never runs against a production DATABASE_URL containing 'prod' in the host.
    - Always requires explicit --allow-prod flag for any non-SQLite/non-localhost target.
"""

import argparse
import datetime
import json
import os
import shutil
import subprocess
import sys
import time

# ---------------------------------------------------------------------------
# Bootstrap Django environment
# ---------------------------------------------------------------------------
BACKEND_DIR = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
sys.path.insert(0, BACKEND_DIR)

os.environ.setdefault("DJANGO_SETTINGS_MODULE", "config.settings.dev")

# ---------------------------------------------------------------------------
# CLI
# ---------------------------------------------------------------------------

def parse_args():
    p = argparse.ArgumentParser(description="TrueWave migration rehearsal runner")
    p.add_argument(
        "--no-rollback",
        action="store_true",
        default=False,
        help="Skip rollback drill (apply only)",
    )
    p.add_argument(
        "--db-path",
        default=None,
        help="Explicit path to SQLite DB file (overrides auto-detect)",
    )
    p.add_argument(
        "--report-dir",
        default=os.path.join(BACKEND_DIR, "docs"),
        help="Directory to write MIGRATION_REHEARSAL_REPORT.md into",
    )
    p.add_argument(
        "--allow-prod",
        action="store_true",
        default=False,
        help="Bypass the production-target safety guard (USE WITH EXTREME CAUTION)",
    )
    p.add_argument(
        "--apps",
        nargs="+",
        default=["geo", "marketplace", "orders", "products"],
        help="Apps whose migrations to apply/rollback",
    )
    return p.parse_args()


# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------

def ts():
    return datetime.datetime.now().strftime("%Y-%m-%d %H:%M:%S")


def log(msg, level="INFO"):
    print(f"[{ts()}] [{level}] {msg}")


def run_manage(args_list, capture=False):
    """Run a manage.py command and return (returncode, stdout, stderr, elapsed_ms)."""
    cmd = [sys.executable, os.path.join(BACKEND_DIR, "manage.py")] + args_list
    t0 = time.monotonic()
    if capture:
        result = subprocess.run(cmd, capture_output=True, text=True, cwd=BACKEND_DIR)
        elapsed = int((time.monotonic() - t0) * 1000)
        return result.returncode, result.stdout, result.stderr, elapsed
    else:
        result = subprocess.run(cmd, cwd=BACKEND_DIR)
        elapsed = int((time.monotonic() - t0) * 1000)
        return result.returncode, "", "", elapsed


def detect_db_path(args):
    """Detect SQLite DB path from Django settings or CLI override."""
    if args.db_path:
        return args.db_path
    try:
        import django
        django.setup()
        from django.conf import settings
        db_cfg = settings.DATABASES.get("default", {})
        engine = db_cfg.get("ENGINE", "")
        if "sqlite" in engine:
            return db_cfg.get("NAME", os.path.join(BACKEND_DIR, "db.sqlite3"))
    except Exception:
        pass
    return os.path.join(BACKEND_DIR, "db.sqlite3")


def safety_guard(db_path, allow_prod):
    """Refuse to run against obvious production targets unless explicitly overridden."""
    db_str = str(db_path).lower()
    is_sqlite = "sqlite" in db_str or db_str.endswith(".db") or db_str.endswith(".sqlite3")
    if is_sqlite:
        return True  # SQLite is always local — safe
    if allow_prod:
        log("--allow-prod flag present. Proceeding despite non-local target.", "WARN")
        return True
    # Block any remote-looking database name
    risky_keywords = ["prod", "production", "live", "rds", "amazonaws", "mysql", "postgresql"]
    for kw in risky_keywords:
        if kw in db_str:
            log(
                f"SAFETY BLOCK: Database target '{db_path}' looks like a production environment. "
                "Pass --allow-prod to override.",
                "ERROR",
            )
            return False
    return True


def backup_db(db_path):
    """Copy SQLite file to timestamped backup. Returns backup path or None."""
    if not os.path.isfile(db_path):
        log(f"No DB file found at {db_path} — skipping backup.", "WARN")
        return None
    stamp = datetime.datetime.now().strftime("%Y%m%d_%H%M%S")
    backup_path = f"{db_path}.backup_{stamp}"
    shutil.copy2(db_path, backup_path)
    size_mb = os.path.getsize(backup_path) / 1_048_576
    log(f"DB backed up to: {backup_path} ({size_mb:.2f} MB)")
    return backup_path


def restore_db(db_path, backup_path):
    """Restore from backup (used in rollback drill)."""
    if not backup_path or not os.path.isfile(backup_path):
        log("No backup to restore from.", "WARN")
        return False
    shutil.copy2(backup_path, db_path)
    log(f"DB restored from: {backup_path}")
    return True


# ---------------------------------------------------------------------------
# Migration states
# ---------------------------------------------------------------------------

# Apps and their rollback targets (the migration applied just BEFORE our new ones)
ROLLBACK_TARGETS = {
    "geo": "0011",
    "marketplace": "zero",
    "orders": "0008",
    "products": "0011",
}


# ---------------------------------------------------------------------------
# Phase functions
# ---------------------------------------------------------------------------

def phase_preflight():
    log("=== PHASE 1: PRE-FLIGHT CHECKS ===")
    results = {}

    # System check
    rc, out, err, ms = run_manage(["check", "--deploy", "--fail-level", "WARNING"], capture=True)
    if rc != 0:
        # deploy check may warn about DEBUG=True etc. on local — escalate only on ERROR
        rc2, out2, err2, ms2 = run_manage(["check"], capture=True)
        if rc2 != 0:
            log(f"System check FAILED:\n{err2}", "ERROR")
            results["system_check"] = {"status": "FAIL", "output": err2, "elapsed_ms": ms2}
            return results, False
        log("System check passed (deploy-level warnings present but no errors).")
        results["system_check"] = {"status": "PASS_WITH_WARNINGS", "output": out2, "elapsed_ms": ms2}
    else:
        log("System check passed.")
        results["system_check"] = {"status": "PASS", "elapsed_ms": ms}

    # Pending migrations check
    rc, out, err, ms = run_manage(["makemigrations", "--check", "--dry-run"], capture=True)
    if rc == 0:
        log("No pending model changes detected.")
        results["pending_check"] = {"status": "PASS", "elapsed_ms": ms}
    else:
        log("WARNING: pending model changes not yet captured in a migration file.", "WARN")
        log(f"Output: {out or err}", "WARN")
        results["pending_check"] = {"status": "WARN", "output": out or err, "elapsed_ms": ms}

    return results, True


def phase_backup(db_path):
    log("=== PHASE 2: DATABASE BACKUP ===")
    backup_path = backup_db(db_path)
    return {"backup_path": backup_path}


def phase_migration_plan(apps):
    log("=== PHASE 3: MIGRATION PLAN ===")
    rc, out, err, ms = run_manage(["migrate", "--plan"] + apps, capture=True)
    if rc != 0:
        rc, out, err, ms = run_manage(["migrate", "--plan"], capture=True)
    log(f"Plan output:\n{out or err}")
    rc2, show_out, _, ms2 = run_manage(["showmigrations", "--plan"], capture=True)
    pending_lines = [l for l in show_out.splitlines() if l.strip().startswith("[ ]")]
    log(f"{len(pending_lines)} pending migration(s) detected.")
    for line in pending_lines:
        log(f"  {line.strip()}")
    return {
        "plan_output": out,
        "pending_count": len(pending_lines),
        "pending_migrations": pending_lines,
        "elapsed_ms": ms,
    }


def phase_apply(apps):
    log("=== PHASE 4: APPLY MIGRATIONS ===")
    results = {}
    total_start = time.monotonic()

    for app in apps:
        log(f"  Applying: {app}")
        rc, out, err, ms = run_manage(["migrate", app, "--verbosity", "2"], capture=True)
        status = "PASS" if rc == 0 else "FAIL"
        log(f"  {app}: {status} ({ms}ms)")
        if rc != 0:
            log(f"  Error:\n{err}", "ERROR")
        results[app] = {"status": status, "output": out, "error": err, "elapsed_ms": ms}

    total_ms = int((time.monotonic() - total_start) * 1000)
    results["__total_elapsed_ms"] = total_ms
    log(f"All app migrations applied in {total_ms}ms total.")
    return results


def phase_rollback(apps, rollback_targets):
    log("=== PHASE 6: ROLLBACK DRILL ===")
    results = {}

    # Rollback in reverse dependency order
    rollback_order = list(reversed(apps))
    for app in rollback_order:
        target = rollback_targets.get(app, "zero")
        log(f"  Rolling back: {app} -> {target}")
        rc, out, err, ms = run_manage(["migrate", app, target, "--verbosity", "2"], capture=True)
        status = "PASS" if rc == 0 else "FAIL"
        log(f"  {app} rollback: {status} ({ms}ms)")
        if rc != 0:
            log(f"  Error:\n{err}", "ERROR")
        results[app] = {
            "target": target,
            "status": status,
            "output": out,
            "error": err,
            "elapsed_ms": ms,
        }

    return results


def phase_reapply(apps):
    log("=== PHASE 7: RE-APPLY (POST-ROLLBACK CONFIRMATION) ===")
    return phase_apply(apps)


# ---------------------------------------------------------------------------
# Entrypoint
# ---------------------------------------------------------------------------

def main():
    args = parse_args()
    db_path = detect_db_path(args)

    log("=" * 60)
    log("TrueWave MIGRATION REHEARSAL")
    log(f"Date:      {ts()}")
    log(f"DB target: {db_path}")
    log(f"Apps:      {', '.join(args.apps)}")
    log("=" * 60)

    if not safety_guard(db_path, args.allow_prod):
        sys.exit(1)

    rehearsal = {
        "started_at": ts(),
        "db_path": str(db_path),
        "apps": args.apps,
        "phases": {},
    }

    # Phase 1
    preflight_results, ok = phase_preflight()
    rehearsal["phases"]["preflight"] = preflight_results
    if not ok:
        log("Pre-flight failed. Aborting rehearsal.", "ERROR")
        sys.exit(2)

    # Phase 2
    backup_result = phase_backup(db_path)
    rehearsal["phases"]["backup"] = backup_result
    backup_path = backup_result.get("backup_path")

    # Phase 3
    plan_result = phase_migration_plan(args.apps)
    rehearsal["phases"]["plan"] = plan_result

    if plan_result["pending_count"] == 0:
        log("No pending migrations to apply. Rehearsal complete (nothing to apply).")
        rehearsal["result"] = "SKIPPED_NO_PENDING"
        _write_report(rehearsal, args.report_dir)
        return

    # Phase 4
    apply_result = phase_apply(args.apps)
    rehearsal["phases"]["apply"] = apply_result
    apply_ok = all(v.get("status") == "PASS" for k, v in apply_result.items() if not k.startswith("__"))

    if not apply_ok:
        log("Migration apply phase had failures. See report for details.", "WARN")

    # Phase 5 — post-migration validation (calls validation module)
    log("=== PHASE 5: POST-MIGRATION VALIDATION ===")
    try:
        from scripts.post_migration_validation import run_all_validations
        validation_results = run_all_validations()
        rehearsal["phases"]["validation"] = validation_results
        log(f"Validation complete. Passed: {validation_results.get('passed', 0)}, Failed: {validation_results.get('failed', 0)}")
    except ImportError:
        log("post_migration_validation module not found — skipping validation phase.", "WARN")
        rehearsal["phases"]["validation"] = {"status": "SKIPPED"}

    # Phase 6 — rollback drill
    if not args.no_rollback:
        rollback_result = phase_rollback(args.apps, ROLLBACK_TARGETS)
        rehearsal["phases"]["rollback"] = rollback_result
        rollback_ok = all(v.get("status") == "PASS" for v in rollback_result.values())

        if rollback_ok:
            log("Rollback drill succeeded. Re-applying migrations.")
            reapply_result = phase_reapply(args.apps)
            rehearsal["phases"]["reapply"] = reapply_result
        else:
            log("Rollback drill had failures. Review output before proceeding.", "WARN")
            rehearsal["phases"]["reapply"] = {"status": "SKIPPED_ROLLBACK_FAILED"}
    else:
        log("Rollback drill skipped (--no-rollback).")
        rehearsal["phases"]["rollback"] = {"status": "SKIPPED"}
        rehearsal["phases"]["reapply"] = {"status": "SKIPPED"}

    rehearsal["completed_at"] = ts()
    rehearsal["result"] = "COMPLETE"

    _write_report(rehearsal, args.report_dir)
    log("=" * 60)
    log("REHEARSAL COMPLETE. Report written.")
    log("=" * 60)


def _write_report(rehearsal, report_dir):
    """Write a machine-readable JSON summary alongside the markdown report path."""
    os.makedirs(report_dir, exist_ok=True)
    json_path = os.path.join(report_dir, "migration_rehearsal_result.json")
    with open(json_path, "w") as f:
        json.dump(rehearsal, f, indent=2, default=str)
    log(f"JSON result written to: {json_path}")


if __name__ == "__main__":
    main()
