# Copyright (c) Microsoft Corporation.
# Licensed under the MIT license.

import argparse
import sys
import tempfile
from pathlib import Path

from alembic.util.exc import AutogenerateDiffsDetected
from sqlalchemy import create_engine
from sqlalchemy.engine import Engine

from pyrit.memory.migration import check_schema_migrations, generate_schema_migration, run_schema_migrations

# ANSI color codes
_RED = "\033[91m"
_RESET = "\033[0m"


def _print_error(message: str) -> None:
    """Print an error message in red to stderr."""
    print(f"{_RED}{message}{_RESET}", file=sys.stderr)


def _create_temp_engine() -> tuple[Engine, Path]:
    """Create a temp SQLite database upgraded to head and return engine and path."""
    with tempfile.NamedTemporaryFile(suffix=".db", delete=False) as tmp:
        tmp_path = Path(tmp.name)
    engine = create_engine(f"sqlite:///{tmp_path}")
    run_schema_migrations(engine=engine)
    return engine, tmp_path


def _cmd_generate(*, message: str, force: bool = False) -> None:
    """Generate a new Alembic revision from model changes."""
    engine, tmp_path = _create_temp_engine()
    try:
        generate_schema_migration(engine=engine, message=message, force=force)
        print("Migration file generated. Review it carefully before committing.")
    except RuntimeError as e:
        _print_error(str(e))
        raise SystemExit(1) from e
    finally:
        engine.dispose()
        tmp_path.unlink(missing_ok=True)


def _cmd_check() -> None:
    """Verify all migrations apply cleanly and schema matches models."""
    engine, tmp_path = _create_temp_engine()
    try:
        check_schema_migrations(engine=engine)
    except AutogenerateDiffsDetected as e:
        _print_error(f"Migration check failed. Run 'generate' to create a migration. Error: {e}")
        raise SystemExit(1) from e
    finally:
        engine.dispose()
        tmp_path.unlink(missing_ok=True)


def _cmd_head() -> None:
    """Print the current Alembic head revision ID."""
    from alembic.config import Config
    from alembic.script import ScriptDirectory

    script_location = Path(__file__).parent.parent / "pyrit" / "memory" / "alembic"
    config = Config()
    config.set_main_option("script_location", str(script_location))
    head = ScriptDirectory.from_config(config).get_current_head()
    print(head)


def _build_parser() -> argparse.ArgumentParser:
    """Build the CLI argument parser."""
    parser = argparse.ArgumentParser(
        description="PyRIT memory migration tool. Generate and validate migrations based on the current memory models."
    )
    sub = parser.add_subparsers(dest="command", required=True)

    gen = sub.add_parser("generate", help="Generate a new migration from model changes.")
    gen.add_argument("-m", "--message", required=True, help="Migration message.")
    gen.add_argument("--force", action="store_true", help="Generate migration even if no changes detected.")

    sub.add_parser("check", help="Verify all migrations apply cleanly and add up to the current memory models.")

    sub.add_parser("head", help="Print the current Alembic head revision ID.")

    return parser


def main() -> int:
    """Dispatch the selected migration command."""
    args = _build_parser().parse_args()

    if args.command == "generate":
        _cmd_generate(message=args.message, force=args.force)
    elif args.command == "check":
        _cmd_check()
    elif args.command == "head":
        _cmd_head()

    return 0


if __name__ == "__main__":
    raise SystemExit(main())
