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

import io
import logging
import sys
from pathlib import Path
from typing import Any, TextIO, cast

from alembic import command
from alembic.autogenerate.api import compare_metadata
from alembic.config import Config
from alembic.migration import MigrationContext
from alembic.util.exc import AutogenerateDiffsDetected
from sqlalchemy import MetaData, Table, inspect
from sqlalchemy.engine import Connection, Engine

from pyrit.memory.memory_models import Base

logger = logging.getLogger(__name__)

PYRIT_MEMORY_ALEMBIC_VERSION_TABLE = "pyrit_memory_alembic_version"
_HEAD_REVISION = "head"
_INITIAL_REVISION = "ab8f2c1a9d07"
_MEMORY_TABLES = {table.name for table in Base.metadata.sorted_tables}

# Prefix applied to Alembic's own console output (e.g. "No new upgrade operations
# detected.") so it can be disambiguated from PyRIT/REST API output downstream.
ALEMBIC_OUTPUT_PREFIX = "[pyrit:alembic] "

_ERROR_UNVERSIONED_SCHEMA_COMPARISON_FAILED = (
    "Detected an unversioned legacy memory schema (memory tables exist, but "
    "pyrit_memory_alembic_version is missing), "
    "but failed to compare the existing schema against the initial pre-Alembic schema. "
    "This could be transient, or you may need to repair or rebuild the database "
    "before upgrading to this release."
)

_ERROR_UNVERSIONED_SCHEMA_MISMATCH = (
    "Detected an unversioned legacy memory schema (memory tables exist, but "
    "pyrit_memory_alembic_version is missing), "
    "and it does not match the initial pre-Alembic schema. "
    "Repair or rebuild the database before upgrading to this release."
)


def _include_name_for_memory_schema(
    name: str | None,
    type_: str,
    parent_names: dict[str, str],
) -> bool:
    """
    Restrict schema comparisons to PyRIT memory tables and their child objects.

    Args:
        name (str | None): Name of the database object being considered.
        type_ (str): SQLAlchemy object type (e.g., "table", "column", "index").
        parent_names (dict[str, str]): Parent-name context provided by Alembic.

    Returns:
        bool: True when the object should be included in schema comparison.
    """
    if type_ == "table":
        return bool(name and name in _MEMORY_TABLES)

    table_name = parent_names.get("table_name")
    if table_name:
        return table_name in _MEMORY_TABLES

    return True


def _get_initial_metadata() -> MetaData:
    """
    Return the static MetaData representing the initial pre-Alembic baseline schema.

    This is pinned directly in the initial migration script so it never changes
    and avoids the overhead of spinning up a temporary in-memory database.

    Returns:
        MetaData: The pinned initial schema metadata.
    """
    from pyrit.memory.alembic.versions.ab8f2c1a9d07_pre_alembic_release_schema import INITIAL_METADATA

    return INITIAL_METADATA


def _make_config(*, connection: Connection, stdout: TextIO | None = None) -> Config:
    """
    Build an Alembic config for the memory migration scripts.

    Args:
        connection (Connection): Database connection for Alembic commands.
        stdout (TextIO | None): Stream that Alembic command output is written to. When None,
            Alembic uses its default of ``sys.stdout``. Defaults to None.

    Returns:
        Config: Configured Alembic config object.
    """
    script_location = Path(__file__).with_name("alembic")

    config = Config(stdout=stdout) if stdout is not None else Config()
    config.set_main_option("script_location", str(script_location))
    config.attributes["connection"] = connection
    return config


def _validate_and_stamp_unversioned_memory_schema(*, config: Config, connection: Connection) -> None:
    """
    Validate and stamp unversioned legacy memory schemas.

    Args:
        config (Config): Alembic config bound to the current connection.
        connection (Connection): Database connection to inspect.

    Raises:
        RuntimeError: If an unversioned memory schema does not match models.
    """
    # Perform all inspection in one atomic call to avoid race conditions
    inspector = inspect(connection)
    table_names = set(inspector.get_table_names())

    # If version table already exists, migration has been stamped
    if PYRIT_MEMORY_ALEMBIC_VERSION_TABLE in table_names:
        return

    # If no memory tables exist, this is a fresh database
    if not _MEMORY_TABLES.intersection(table_names):
        return

    # Unversioned memory schema detected; validate it matches the initial pre-Alembic schema
    try:
        initial_metadata = _get_initial_metadata()
        migration_context = MigrationContext.configure(
            connection=connection,
            opts={"compare_type": True, "include_name": _include_name_for_memory_schema},
        )
        diffs = compare_metadata(migration_context, initial_metadata)
    except Exception as e:
        raise RuntimeError(_ERROR_UNVERSIONED_SCHEMA_COMPARISON_FAILED) from e

    if diffs:
        logger.warning(
            f"Detected {len(diffs)} schema diff(s) in unversioned legacy memory schema. "
            "Please address these corrections (or remove your PyRIT DB tables and retry):"
        )
        for index, diff in enumerate(diffs, start=1):
            logger.warning(f"Required correction {index}: {diff}")
        raise RuntimeError(_ERROR_UNVERSIONED_SCHEMA_MISMATCH)

    logger.warning(f"Detected matching unversioned memory schema; stamping at initial revision {_INITIAL_REVISION}")
    command.stamp(config, _INITIAL_REVISION)


class _PrefixedTextStream:
    """
    Wrap a text stream and prefix each written line with a fixed tag.

    Alembic writes its console messages directly to the configured stream (e.g. via
    ``Config.print_stdout``). Wrapping the underlying stream lets PyRIT tag that output
    so it is identifiable as Alembic-originated rather than PyRIT or REST API output.
    Attribute access (e.g. ``encoding``, ``flush``) delegates to the wrapped stream.
    """

    def __init__(self, *, stream: TextIO, prefix: str) -> None:
        """
        Initialize the prefixing stream wrapper.

        Args:
            stream (TextIO): The underlying stream that prefixed text is written to.
            prefix (str): The prefix prepended to the start of each line.
        """
        self._stream = stream
        self._prefix = prefix
        self._at_line_start = True

    def write(self, text: str) -> int:
        """
        Write text to the underlying stream, prefixing the start of each line.

        Args:
            text (str): The text to write.

        Returns:
            int: The number of characters of the original text written.
        """
        if not text:
            return 0

        chunks: list[str] = []
        for char in text:
            if self._at_line_start and char != "\n":
                chunks.append(self._prefix)
                self._at_line_start = False
            chunks.append(char)
            if char == "\n":
                self._at_line_start = True

        self._stream.write("".join(chunks))
        return len(text)

    def __getattr__(self, name: str) -> Any:
        """
        Delegate any other attribute access to the wrapped stream.

        Args:
            name (str): The attribute name being accessed.

        Returns:
            Any: The corresponding attribute from the wrapped stream.
        """
        return getattr(self._stream, name)


def _migration_stdout(*, silent: bool) -> TextIO:
    """
    Resolve the stream Alembic command output is written to.

    When silent, output is swallowed by an in-memory buffer. Otherwise the current
    ``sys.stdout`` is used (resolved at call time so redirections are respected),
    wrapped so each Alembic line is tagged with ``ALEMBIC_OUTPUT_PREFIX``.

    Args:
        silent (bool): If True, returns an in-memory buffer that discards output.

    Returns:
        TextIO: The stream to pass to the Alembic config.
    """
    if silent:
        return io.StringIO()
    return cast("TextIO", _PrefixedTextStream(stream=sys.stdout, prefix=ALEMBIC_OUTPUT_PREFIX))


def run_schema_migrations(*, engine: Engine, silent: bool = False) -> None:
    """
    Upgrade the database schema to the latest Alembic revision.

    Args:
        engine (Engine): SQLAlchemy engine bound to the target database.
        silent (bool): If True, suppresses Alembic console output. Defaults to False.

    Raises:
        Exception: If Alembic fails to apply migrations.
    """
    with engine.begin() as connection:
        config = _make_config(connection=connection, stdout=_migration_stdout(silent=silent))
        _validate_and_stamp_unversioned_memory_schema(config=config, connection=connection)
        command.upgrade(config, _HEAD_REVISION)


def check_schema_migrations(*, engine: Engine, silent: bool = False) -> None:
    """
    Verify that the current schema matches the models.

    Creates a fresh connection, builds an Alembic config, and runs
    ``alembic check`` to detect any unapplied model changes.

    Args:
        engine (Engine): SQLAlchemy engine bound to the target database.
        silent (bool): If True, suppresses Alembic console output (e.g. the
            "No new upgrade operations detected." message). Defaults to False.

    Raises:
        AutogenerateDiffsDetected: If schema does not match models.
    """
    with engine.begin() as connection:
        config = _make_config(connection=connection, stdout=_migration_stdout(silent=silent))
        command.check(config)


def generate_schema_migration(*, engine: Engine, message: str, force: bool = False) -> None:
    """
    Generate a new Alembic revision from model changes.

    Args:
        engine (Engine): SQLAlchemy engine upgraded to head.
        message (str): Human-readable migration message.
        force (bool): If True, generate even if no changes detected. Defaults to False.

    Raises:
        RuntimeError: If no changes detected and force is False.
    """
    with engine.begin() as connection:
        config = _make_config(connection=connection)

        if force:
            command.revision(config, autogenerate=True, message=message)
            return

        try:
            command.check(config)
        except AutogenerateDiffsDetected:
            command.revision(config, autogenerate=True, message=message)
            return

        raise RuntimeError("No schema changes detected. Use force=True to generate an empty migration.")


def reset_database(*, engine: Engine) -> None:
    """
    Drop all tables and recreate the database schema at the latest Alembic revision.

    This destroys all existing data.

    Args:
        engine (Engine): SQLAlchemy engine bound to the target database.
    """
    logger.debug("Resetting database using Alembic migrations")
    with engine.begin() as connection:
        # Drop version table first (not part of Base.metadata)
        inspector = inspect(connection)
        if PYRIT_MEMORY_ALEMBIC_VERSION_TABLE in inspector.get_table_names():
            version_table = Table(PYRIT_MEMORY_ALEMBIC_VERSION_TABLE, MetaData(), autoload_with=connection)
            version_table.drop(connection)

        # Drop all application tables defined in models
        Base.metadata.drop_all(connection)
    # Rebuild schema from migrations
    run_schema_migrations(engine=engine)
