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

import ast
import json
import os
import tempfile
import uuid
from pathlib import Path

import pytest
from alembic import command
from alembic.config import Config
from alembic.script import ScriptDirectory
from alembic.util.exc import AutogenerateDiffsDetected
from sqlalchemy import create_engine, event, inspect, text

from pyrit.memory.alembic.versions import ab8f2c1a9d07_pre_alembic_release_schema
from pyrit.memory.alembic.versions.ab8f2c1a9d07_pre_alembic_release_schema import _CustomUUID
from pyrit.memory.migration import (
    ALEMBIC_OUTPUT_PREFIX,
    _PrefixedTextStream,
    check_schema_migrations,
    generate_schema_migration,
    run_schema_migrations,
)


def test_prefixed_text_stream_prefixes_each_line():
    """_PrefixedTextStream prepends the prefix to the start of each line, across separate writes."""
    import io

    buffer = io.StringIO()
    stream = _PrefixedTextStream(stream=buffer, prefix="[tag] ")

    # Alembic writes the message and the trailing newline as separate calls.
    assert stream.write("first line") == len("first line")
    assert stream.write("\n") == 1
    stream.write("second\nthird\n")

    assert buffer.getvalue() == "[tag] first line\n[tag] second\n[tag] third\n"


def test_prefixed_text_stream_delegates_attributes():
    """_PrefixedTextStream delegates unknown attributes (e.g. encoding) to the wrapped stream."""
    import io

    buffer = io.StringIO()
    stream = _PrefixedTextStream(stream=buffer, prefix="[tag] ")

    assert stream.getvalue() == ""
    stream.flush()  # delegated, should not raise


def test_alembic_env_raises_when_no_connection():
    """Covers env.py line 15: RuntimeError when connection is None."""
    import importlib
    import sys
    from unittest.mock import MagicMock

    # Build a mock alembic context where config.attributes has no "connection"
    mock_config = MagicMock()
    mock_config.attributes = {}  # .get("connection") → None

    mock_context = MagicMock()
    mock_context.config = mock_config

    # Remove cached env module so reload runs module-level code
    env_module_name = "pyrit.memory.alembic.env"
    saved = sys.modules.pop(env_module_name, None)

    # Patch alembic.context to be our mock context
    import alembic

    original_context = getattr(alembic, "context", None)
    alembic.context = mock_context

    try:
        with pytest.raises(RuntimeError, match="No connection found for Alembic migration"):
            importlib.import_module(env_module_name)
    finally:
        # Restore original state
        alembic.context = original_context
        sys.modules.pop(env_module_name, None)
        if saved is not None:
            sys.modules[env_module_name] = saved


def _get_alembic_head_revision(*, config: Config) -> str:
    """Return the current Alembic head revision for the configured script location."""
    head_revision = ScriptDirectory.from_config(config).get_current_head()
    if head_revision is None:
        raise RuntimeError("No Alembic head revision found for memory migrations.")

    return head_revision


def test_custom_uuid_process_bind_param_with_none():
    """Test _CustomUUID.process_bind_param with None value."""
    uuid_type = _CustomUUID()
    result = uuid_type.process_bind_param(None, None)
    assert result is None


def test_custom_uuid_process_bind_param_with_uuid():
    """Test _CustomUUID.process_bind_param with UUID value."""
    uuid_type = _CustomUUID()
    test_uuid = uuid.uuid4()
    result = uuid_type.process_bind_param(test_uuid, None)
    assert result == str(test_uuid)


def test_custom_uuid_process_result_value_with_none():
    """Test _CustomUUID.process_result_value with None value."""
    uuid_type = _CustomUUID()
    result = uuid_type.process_result_value(None, None)
    assert result is None


def test_custom_uuid_process_result_value_with_uuid():
    """Test _CustomUUID.process_result_value with UUID value."""
    uuid_type = _CustomUUID()
    test_uuid = uuid.uuid4()
    result = uuid_type.process_result_value(test_uuid, None)
    assert result == test_uuid


def test_custom_uuid_process_result_value_with_string():
    """Test _CustomUUID.process_result_value with string value."""
    uuid_type = _CustomUUID()
    test_uuid = uuid.uuid4()
    result = uuid_type.process_result_value(str(test_uuid), None)
    assert result == test_uuid


def test_custom_uuid_load_dialect_impl_sqlite():
    """Test _CustomUUID.load_dialect_impl for SQLite dialect."""
    from sqlalchemy.dialects import sqlite

    uuid_type = _CustomUUID()
    dialect = sqlite.dialect()
    result = uuid_type.load_dialect_impl(dialect)
    assert result is not None


def test_custom_uuid_load_dialect_impl_postgresql():
    """Test _CustomUUID.load_dialect_impl for PostgreSQL dialect."""
    from sqlalchemy.dialects import postgresql

    uuid_type = _CustomUUID()
    dialect = postgresql.dialect()
    result = uuid_type.load_dialect_impl(dialect)
    assert result is not None


def test_run_schema_migrations_applies_head_revision():
    """Test that run_schema_migrations applies the current head revision."""
    with tempfile.TemporaryDirectory() as temp_dir:
        db_path = os.path.join(temp_dir, "offline-test.db")
        engine = create_engine(f"sqlite:///{db_path}")
        try:
            pyrit_root = Path(__file__).resolve().parent.parent.parent.parent / "pyrit"
            script_location = pyrit_root / "memory" / "alembic"
            config = Config()
            config.set_main_option("script_location", str(script_location))
            expected_head = _get_alembic_head_revision(config=config)

            run_schema_migrations(engine=engine)

            with engine.connect() as connection:
                version = connection.execute(text("SELECT version_num FROM pyrit_memory_alembic_version")).scalar_one()
            assert version == expected_head
        finally:
            engine.dispose()


def test_scenario_progress_migration_adds_composite_index():
    """The migration head contains the parent/timestamp/id keyset index."""
    with tempfile.TemporaryDirectory() as temp_dir:
        db_path = os.path.join(temp_dir, "scenario-progress-index.db")
        engine = create_engine(f"sqlite:///{db_path}")
        try:
            with engine.begin() as connection:
                config = _config_for(connection)
                command.upgrade(config, "head")
                indexes = {index["name"] for index in inspect(connection).get_indexes("AttackResultEntries")}
            assert "ix_AttackResultEntries_attribution_parent_timestamp_id" in indexes
        finally:
            engine.dispose()


def test_migration_head_removes_additional_initializers_table():
    """The migration head removes the obsolete second initializer configuration source."""
    with tempfile.TemporaryDirectory() as temp_dir:
        db_path = os.path.join(temp_dir, "additional-initializers-removal.db")
        engine = create_engine(f"sqlite:///{db_path}")
        try:
            with engine.begin() as connection:
                config = _config_for(connection)
                command.upgrade(config, "4c9a6e1f2b7d")
                assert "AdditionalInitializers" in set(inspect(connection).get_table_names())

                command.upgrade(config, "head")
                assert "AdditionalInitializers" not in set(inspect(connection).get_table_names())
        finally:
            engine.dispose()


def test_migration_online_mode():
    """
    Test that online migration configuration is valid.
    This tests the run_migrations_online path in env.py.
    """
    with tempfile.TemporaryDirectory() as temp_dir:
        db_path = os.path.join(temp_dir, "online-test.db")
        engine = create_engine(f"sqlite:///{db_path}")
        try:
            with engine.begin() as connection:
                pyrit_root = Path(__file__).resolve().parent.parent.parent.parent / "pyrit"
                script_location = pyrit_root / "memory" / "alembic"
                config = Config()
                config.set_main_option("script_location", str(script_location))
                config.attributes["connection"] = connection
                config.attributes["version_table"] = "pyrit_memory_alembic_version"
                expected_head = _get_alembic_head_revision(config=config)

                command.upgrade(config, "head")

                version = connection.execute(text("SELECT version_num FROM pyrit_memory_alembic_version")).scalar_one()
                assert version == expected_head
        finally:
            engine.dispose()


def test_migration_script_metadata():
    """Test that the initial migration script has correct metadata."""

    assert ab8f2c1a9d07_pre_alembic_release_schema.revision == "ab8f2c1a9d07"
    assert ab8f2c1a9d07_pre_alembic_release_schema.down_revision is None
    assert ab8f2c1a9d07_pre_alembic_release_schema.branch_labels is None
    assert ab8f2c1a9d07_pre_alembic_release_schema.depends_on is None


def test_migration_downgrade_creates_proper_structure():
    """
    Test that downgrade function doesn't corrupt the database.
    This indirectly tests the downgrade path in ab8f2c1a9d07_pre_alembic_release_schema.py.
    """
    with tempfile.TemporaryDirectory() as temp_dir:
        db_path = os.path.join(temp_dir, "downgrade-test.db")
        engine = create_engine(f"sqlite:///{db_path}")
        try:
            with engine.begin() as connection:
                pyrit_root = Path(__file__).resolve().parent.parent.parent.parent / "pyrit"
                script_location = pyrit_root / "memory" / "alembic"
                config = Config()
                config.set_main_option("script_location", str(script_location))
                config.attributes["connection"] = connection
                config.attributes["version_table"] = "pyrit_memory_alembic_version"

                command.upgrade(config, "head")

                tables_before = set(inspect(connection).get_table_names())
                assert len(tables_before) > 0

                command.downgrade(config, "base")

                tables_after = set(inspect(connection).get_table_names())
                assert "pyrit_memory_alembic_version" not in tables_after or len(tables_after) == 1
        finally:
            engine.dispose()


# =============================================================================
# Backfill tests for the attribution_parent_id foreign key migration
# =============================================================================


_SCENARIO_LINKAGE_REV = "9c8b7a6d5e4f"
_PREV_REV = "7a1b2c3d4e5f"


def _seed_pre_migration_scenario(connection, *, scenario_id, manifest_json):
    """Insert a ScenarioResultEntry row at the pre-migration revision."""
    connection.execute(
        text(
            'INSERT INTO "ScenarioResultEntries" '
            "(id, scenario_name, scenario_description, scenario_version, pyrit_version, "
            "objective_target_identifier, scenario_init_data, scenario_run_state, attack_results_json, "
            "number_tries, completion_time, timestamp) "
            "VALUES (:id, :name, '', 1, '0.14.0.dev0', '{}', '{}', 'COMPLETED', :manifest, 0, "
            "'2026-05-18', '2026-05-18')"
        ),
        {"id": scenario_id, "name": "Backfill Test", "manifest": manifest_json},
    )


def _seed_pre_migration_attack_result(connection, *, attack_id, conversation_id):
    """Insert an AttackResultEntry row at the pre-migration revision."""
    connection.execute(
        text(
            'INSERT INTO "AttackResultEntries" '
            "(id, conversation_id, objective, attack_identifier, objective_sha256, executed_turns, "
            "execution_time_ms, outcome, timestamp) "
            "VALUES (:id, :conv, 'obj', '{}', 'sha', 1, 0, 'success', '2026-05-18')"
        ),
        {"id": attack_id, "conv": conversation_id},
    )


def _seed_post_drop_attack_result(connection, *, attack_id, conversation_id):
    """Insert an AttackResultEntry row at the Conversations pre-migration revision.

    By this revision the deprecated ``AttackResultEntries.attack_identifier`` column
    has already been dropped, so it is omitted from the insert.
    """
    connection.execute(
        text(
            'INSERT INTO "AttackResultEntries" '
            "(id, conversation_id, objective, objective_sha256, executed_turns, "
            "execution_time_ms, outcome, timestamp) "
            "VALUES (:id, :conv, 'obj', 'sha', 1, 0, 'success', '2026-05-18')"
        ),
        {"id": attack_id, "conv": conversation_id},
    )


def _config_for(connection):
    pyrit_root = Path(__file__).resolve().parent.parent.parent.parent / "pyrit"
    script_location = pyrit_root / "memory" / "alembic"
    config = Config()
    config.set_main_option("script_location", str(script_location))
    config.attributes["connection"] = connection
    config.attributes["version_table"] = "pyrit_memory_alembic_version"
    return config


def test_scorable_content_migration_creates_hash_column():
    with tempfile.TemporaryDirectory() as temp_dir:
        db_path = os.path.join(temp_dir, "scorable-content.db")
        engine = create_engine(f"sqlite:///{db_path}")
        try:
            with engine.begin() as connection:
                command.upgrade(_config_for(connection), "5a1d3c7e9f04")
                columns = {
                    column["name"]: column for column in inspect(connection).get_columns("ScorableContentEntries")
                }

            assert columns["value"]["nullable"] is False
            assert columns["value_sha256"]["nullable"] is False
            assert columns["value_sha256"]["type"].length == 64
        finally:
            engine.dispose()


def test_backfill_links_attack_results_via_conversation_id():
    """Upgrading from the pre-foreign-key revision backfills
    attribution_parent_id + attribution_data on AttackResultEntries by
    matching conversation_id."""
    import json

    with tempfile.TemporaryDirectory() as temp_dir:
        db_path = os.path.join(temp_dir, "backfill-test.db")
        engine = create_engine(f"sqlite:///{db_path}")
        try:
            sid = str(uuid.uuid4())
            ar1_id = str(uuid.uuid4())
            ar2_id = str(uuid.uuid4())
            ar3_id = str(uuid.uuid4())

            with engine.begin() as connection:
                config = _config_for(connection)
                # Step the schema up to JUST before the linkage migration.
                command.upgrade(config, _PREV_REV)

                _seed_pre_migration_attack_result(connection, attack_id=ar1_id, conversation_id="conv-a-0")
                _seed_pre_migration_attack_result(connection, attack_id=ar2_id, conversation_id="conv-a-1")
                _seed_pre_migration_attack_result(connection, attack_id=ar3_id, conversation_id="conv-b-0")
                _seed_pre_migration_scenario(
                    connection,
                    scenario_id=sid,
                    manifest_json=json.dumps({"a": ["conv-a-0", "conv-a-1"], "b": ["conv-b-0"]}),
                )

                command.upgrade(config, _SCENARIO_LINKAGE_REV)

                rows = connection.execute(
                    text(
                        "SELECT conversation_id, attribution_parent_id, attribution_data "
                        'FROM "AttackResultEntries" ORDER BY conversation_id'
                    )
                ).fetchall()

            results_by_conv = {r[0]: (r[1], r[2]) for r in rows}

            # All three rows now point at the scenario via the new foreign key.
            for conv in ("conv-a-0", "conv-a-1", "conv-b-0"):
                assert results_by_conv[conv][0] == sid, f"{conv} should be backfilled"

            # attribution_data carries parent_collection (the atomic attack name).
            sd_a0 = json.loads(results_by_conv["conv-a-0"][1])
            sd_a1 = json.loads(results_by_conv["conv-a-1"][1])
            sd_b0 = json.loads(results_by_conv["conv-b-0"][1])

            assert sd_a0 == {"parent_collection": "a"}
            assert sd_a1 == {"parent_collection": "a"}
            assert sd_b0 == {"parent_collection": "b"}
        finally:
            engine.dispose()


def test_backfill_is_idempotent_and_does_not_clobber_existing_linkage():
    """The backfill is safe to re-run: rows that already carry an
    ``attribution_parent_id`` are not overwritten (the WHERE IS NULL guard). We
    verify by upgrading, manually retargeting a row, then downgrading +
    re-upgrading and asserting the manual retarget survives."""
    import json

    with tempfile.TemporaryDirectory() as temp_dir:
        db_path = os.path.join(temp_dir, "idempotent.db")
        engine = create_engine(f"sqlite:///{db_path}")
        try:
            sid_old = str(uuid.uuid4())
            sid_manual = str(uuid.uuid4())
            ar_id = str(uuid.uuid4())

            with engine.begin() as connection:
                config = _config_for(connection)
                command.upgrade(config, _PREV_REV)
                _seed_pre_migration_attack_result(connection, attack_id=ar_id, conversation_id="conv-shared")
                _seed_pre_migration_scenario(
                    connection, scenario_id=sid_old, manifest_json=json.dumps({"a": ["conv-shared"]})
                )
                command.upgrade(config, _SCENARIO_LINKAGE_REV)

                # Manually retarget the row to a DIFFERENT attribution_parent_id —
                # simulate code that already linked it post-upgrade.
                connection.execute(
                    text('UPDATE "AttackResultEntries" SET attribution_parent_id = :sid WHERE conversation_id = :conv'),
                    {"sid": sid_manual, "conv": "conv-shared"},
                )

                # Downgrade then re-upgrade to re-run the backfill.
                command.downgrade(config, _PREV_REV)

                # After downgrade the foreign key column is gone, but the
                # manifest still references conv-shared. On re-upgrade, the
                # backfill should NOT clobber sid_manual because the column was
                # just re-added as NULL — actually downgrade DROPS the column
                # data, so on re-upgrade the row will start at NULL and get
                # linked again. The test we want is: re-running the backfill
                # while a row already has a non-NULL foreign key does not
                # overwrite it. We exercise that with a fresh second
                # upgrade-then-no-op-re-upgrade.
                command.upgrade(config, _SCENARIO_LINKAGE_REV)

                # First upgrade after downgrade re-links it to sid_old (the
                # manifest source). Now manually retarget again.
                connection.execute(
                    text('UPDATE "AttackResultEntries" SET attribution_parent_id = :sid WHERE conversation_id = :conv'),
                    {"sid": sid_manual, "conv": "conv-shared"},
                )

                # Stamping should NOT happen on re-invocation since the column
                # is already non-NULL. We verify by re-running the backfill
                # logic via downgrade+upgrade is NOT what we want here — we
                # want the IS NULL guard. Simulate by adding another scenario
                # referencing the same conversation_id and re-running the
                # backfill function only.
                connection.execute(
                    text(
                        'INSERT INTO "ScenarioResultEntries" '
                        "(id, scenario_name, scenario_description, scenario_version, pyrit_version, "
                        "objective_target_identifier, scenario_init_data, scenario_run_state, attack_results_json, "
                        "number_tries, completion_time, timestamp) "
                        "VALUES (:id, 'Other', '', 1, '0.14.0.dev0', '{}', '{}', 'COMPLETED', :manifest, 0, "
                        "'2026-05-18', '2026-05-18')"
                    ),
                    {"id": str(uuid.uuid4()), "manifest": json.dumps({"x": ["conv-shared"]})},
                )

                # Manually call the backfill function (loaded via the alembic
                # script directory — modules with leading-digit filenames are
                # not importable through normal Python import).
                from importlib.util import module_from_spec, spec_from_file_location

                migration_path = (
                    Path(__file__).resolve().parent.parent.parent.parent
                    / "pyrit"
                    / "memory"
                    / "alembic"
                    / "versions"
                    / "9c8b7a6d5e4f_add_attribution_to_attack_results.py"
                )
                spec = spec_from_file_location("scenario_linkage_migration", migration_path)
                assert spec is not None and spec.loader is not None
                mig = module_from_spec(spec)
                spec.loader.exec_module(mig)

                from alembic import op as _op_mod

                _original_get_bind = _op_mod.get_bind
                _op_mod.get_bind = lambda: connection
                try:
                    mig._backfill_attribution_linkage()
                finally:
                    _op_mod.get_bind = _original_get_bind

                # The row's manual retarget MUST survive — the IS NULL guard
                # prevents the backfill from overwriting it.
                row = connection.execute(
                    text('SELECT attribution_parent_id FROM "AttackResultEntries" WHERE conversation_id = :conv'),
                    {"conv": "conv-shared"},
                ).scalar_one()
                assert row == sid_manual
        finally:
            engine.dispose()


def test_migration_drops_error_attack_result_ids_json_column():
    """The not-yet-released error_attack_result_ids_json column is removed
    in this migration (no deprecation window needed)."""
    with tempfile.TemporaryDirectory() as temp_dir:
        db_path = os.path.join(temp_dir, "drop-col.db")
        engine = create_engine(f"sqlite:///{db_path}")
        try:
            with engine.begin() as connection:
                config = _config_for(connection)
                command.upgrade(config, _PREV_REV)
                cols_before = {c["name"] for c in inspect(connection).get_columns("ScenarioResultEntries")}
                assert "error_attack_result_ids_json" in cols_before

                command.upgrade(config, _SCENARIO_LINKAGE_REV)
                cols_after = {c["name"] for c in inspect(connection).get_columns("ScenarioResultEntries")}
                assert "error_attack_result_ids_json" not in cols_after
        finally:
            engine.dispose()


def test_migration_downgrade_restores_dropped_column():
    """Downgrading from the linkage revision re-adds error_attack_result_ids_json
    and removes the new AttackResultEntries columns."""
    with tempfile.TemporaryDirectory() as temp_dir:
        db_path = os.path.join(temp_dir, "downgrade-linkage.db")
        engine = create_engine(f"sqlite:///{db_path}")
        try:
            with engine.begin() as connection:
                config = _config_for(connection)
                command.upgrade(config, _SCENARIO_LINKAGE_REV)

                attack_cols_up = {c["name"] for c in inspect(connection).get_columns("AttackResultEntries")}
                assert "attribution_parent_id" in attack_cols_up
                assert "attribution_data" in attack_cols_up

                command.downgrade(config, _PREV_REV)

                attack_cols_down = {c["name"] for c in inspect(connection).get_columns("AttackResultEntries")}
                assert "attribution_parent_id" not in attack_cols_down
                assert "attribution_data" not in attack_cols_down

                scenario_cols = {c["name"] for c in inspect(connection).get_columns("ScenarioResultEntries")}
                assert "error_attack_result_ids_json" in scenario_cols
        finally:
            engine.dispose()


def test_check_schema_migrations_calls_alembic_check():
    with tempfile.TemporaryDirectory() as temp_dir:
        db_path = os.path.join(temp_dir, "check-test.db")
        engine = create_engine(f"sqlite:///{db_path}")
        try:
            # First apply migrations so schema matches models
            run_schema_migrations(engine=engine)
            # Now check should succeed (no diffs)
            check_schema_migrations(engine=engine)
        finally:
            engine.dispose()


def test_generate_schema_migration_force_creates_revision():
    from unittest.mock import patch

    with tempfile.TemporaryDirectory() as temp_dir:
        db_path = os.path.join(temp_dir, "gen-force-test.db")
        engine = create_engine(f"sqlite:///{db_path}")
        try:
            run_schema_migrations(engine=engine)
            with patch("pyrit.memory.migration.command.revision") as mock_revision:
                generate_schema_migration(engine=engine, message="force empty", force=True)
                mock_revision.assert_called_once()
        finally:
            engine.dispose()


def test_generate_schema_migration_no_changes_raises():
    with tempfile.TemporaryDirectory() as temp_dir:
        db_path = os.path.join(temp_dir, "gen-nochange-test.db")
        engine = create_engine(f"sqlite:///{db_path}")
        try:
            run_schema_migrations(engine=engine)
            with pytest.raises(RuntimeError, match="No schema changes detected"):
                generate_schema_migration(engine=engine, message="should fail")
        finally:
            engine.dispose()


def test_generate_schema_migration_with_diffs_creates_revision():
    from unittest.mock import patch

    with tempfile.TemporaryDirectory() as temp_dir:
        db_path = os.path.join(temp_dir, "gen-diffs-test.db")
        engine = create_engine(f"sqlite:///{db_path}")
        try:
            run_schema_migrations(engine=engine)
            exc = AutogenerateDiffsDetected.__new__(AutogenerateDiffsDetected)
            with (
                patch("pyrit.memory.migration.command.check", side_effect=exc),
                patch("pyrit.memory.migration.command.revision") as mock_revision,
            ):
                generate_schema_migration(engine=engine, message="with diffs")
                mock_revision.assert_called_once()
        finally:
            engine.dispose()


def test_check_schema_migrations_silent_suppresses_output(capsys):
    """check_schema_migrations with silent=True must not print the Alembic message."""
    with tempfile.TemporaryDirectory() as temp_dir:
        db_path = os.path.join(temp_dir, "check-silent-test.db")
        engine = create_engine(f"sqlite:///{db_path}")
        try:
            run_schema_migrations(engine=engine, silent=True)
            capsys.readouterr()  # discard any output from setup

            check_schema_migrations(engine=engine, silent=True)

            captured = capsys.readouterr()
            assert captured.out == ""
        finally:
            engine.dispose()


def test_check_schema_migrations_not_silent_prints_output(capsys):
    """check_schema_migrations without silent prints the Alembic message tagged as Alembic output."""
    with tempfile.TemporaryDirectory() as temp_dir:
        db_path = os.path.join(temp_dir, "check-loud-test.db")
        engine = create_engine(f"sqlite:///{db_path}")
        try:
            run_schema_migrations(engine=engine, silent=True)
            capsys.readouterr()  # discard any output from setup

            check_schema_migrations(engine=engine, silent=False)

            captured = capsys.readouterr()
            assert f"{ALEMBIC_OUTPUT_PREFIX}No new upgrade operations detected." in captured.out
        finally:
            engine.dispose()


def test_memory_interface_check_schema_migration_calls_check():
    """_check_schema_migration on MemoryInterface calls check_schema_migrations without running upgrade."""
    from unittest.mock import MagicMock, patch

    from pyrit.memory.memory_interface import MemoryInterface

    obj = MagicMock(spec=MemoryInterface)
    obj.engine = MagicMock()

    with patch("pyrit.memory.migration.check_schema_migrations") as mock_check:
        MemoryInterface._check_schema_migration(obj, silent=True)
        mock_check.assert_called_once_with(engine=obj.engine, silent=True)


def test_memory_interface_check_schema_migration_raises_on_mismatch():
    """_check_schema_migration raises AutogenerateDiffsDetected when schema mismatches (pure primitive)."""
    from unittest.mock import MagicMock, patch

    from alembic.util.exc import AutogenerateDiffsDetected

    from pyrit.memory.memory_interface import MemoryInterface

    obj = MagicMock(spec=MemoryInterface)
    obj.engine = MagicMock()

    with patch(
        "pyrit.memory.migration.check_schema_migrations",
        side_effect=AutogenerateDiffsDetected(
            "diffs detected",
            revision_context=MagicMock(),
            diffs=[],
        ),
    ):
        with pytest.raises(AutogenerateDiffsDetected):
            MemoryInterface._check_schema_migration(obj, silent=True)


def test_memory_interface_check_schema_migration_raises_without_engine():
    """_check_schema_migration raises RuntimeError when engine is None."""
    from unittest.mock import MagicMock

    from pyrit.memory.memory_interface import MemoryInterface

    obj = MagicMock(spec=MemoryInterface)
    obj.engine = None

    with pytest.raises(RuntimeError, match="Engine must be initialized"):
        MemoryInterface._check_schema_migration(obj, silent=False)


def test_memory_migrations_head_command(capsys):
    """The 'head' subcommand of memory_migrations.py prints the current Alembic head revision."""
    from build_scripts.memory_migrations import _cmd_head

    _cmd_head()
    captured = capsys.readouterr()
    revision = captured.out.strip()
    # Should be a non-empty hex-ish string
    assert len(revision) > 0
    assert all(c in "0123456789abcdef" for c in revision)


# =============================================================================
# v1 compatibility-column removal (24b44ef076b6)
# =============================================================================


_V1_CLEANUP_REV = "24b44ef076b6"
_V1_CLEANUP_PREV_REV = "e5f7a9c1b3d2"


def test_v1_cleanup_migration_drops_compatibility_columns_and_preserves_rows():
    """The v1 upgrade removes only the superseded columns."""
    with tempfile.TemporaryDirectory() as temp_dir:
        engine = create_engine(f"sqlite:///{os.path.join(temp_dir, 'v1-cleanup.db')}")
        try:
            with engine.begin() as connection:
                config = _config_for(connection)
                command.upgrade(config, _V1_CLEANUP_PREV_REV)

                score_id = str(uuid.uuid4())
                scenario_id = str(uuid.uuid4())
                connection.execute(
                    text(
                        'INSERT INTO "ScoreEntries" '
                        "(id, score_value, score_type, score_metadata, scorer_class_identifier, "
                        "timestamp, task, objective) "
                        "VALUES (:id, 'true', 'true_false', '{}', '{}', '2026-07-14', "
                        "'legacy objective', 'canonical objective')"
                    ),
                    {"id": score_id},
                )
                connection.execute(
                    text(
                        'INSERT INTO "ScenarioResultEntries" '
                        "(id, scenario_name, scenario_version, pyrit_version, scenario_identifier, "
                        "objective_target_identifier, scenario_run_state, attack_results_json, "
                        "number_tries, completion_time, timestamp) "
                        "VALUES (:id, 'Cleanup', 1, '1.0.0', '{}', '{}', 'COMPLETED', "
                        "'{}', 0, '2026-07-14', '2026-07-14')"
                    ),
                    {"id": scenario_id},
                )

                command.upgrade(config, _V1_CLEANUP_REV)

                score_columns = {column["name"] for column in inspect(connection).get_columns("ScoreEntries")}
                scenario_columns = {
                    column["name"] for column in inspect(connection).get_columns("ScenarioResultEntries")
                }
                objective = connection.execute(
                    text('SELECT objective FROM "ScoreEntries" WHERE id = :id'), {"id": score_id}
                ).scalar_one()
                scenario_name = connection.execute(
                    text('SELECT scenario_name FROM "ScenarioResultEntries" WHERE id = :id'),
                    {"id": scenario_id},
                ).scalar_one()

            assert "task" not in score_columns
            assert "attack_results_json" not in scenario_columns
            assert objective == "canonical objective"
            assert scenario_name == "Cleanup"
        finally:
            engine.dispose()


def test_v1_cleanup_downgrade_restores_and_backfills_compatibility_columns():
    """Downgrade reconstructs task and attack-results manifests for older code."""
    with tempfile.TemporaryDirectory() as temp_dir:
        engine = create_engine(f"sqlite:///{os.path.join(temp_dir, 'v1-downgrade.db')}")
        try:
            with engine.begin() as connection:
                config = _config_for(connection)
                command.upgrade(config, _V1_CLEANUP_REV)

                score_id = str(uuid.uuid4())
                scenario_id = str(uuid.uuid4())
                empty_scenario_id = str(uuid.uuid4())
                connection.execute(
                    text(
                        'INSERT INTO "ScoreEntries" '
                        "(id, score_value, score_type, score_metadata, scorer_class_identifier, "
                        "timestamp, objective) "
                        "VALUES (:id, 'true', 'true_false', '{}', '{}', '2026-07-14', 'restored objective')"
                    ),
                    {"id": score_id},
                )
                for current_id, scenario_name in (
                    (scenario_id, "With attacks"),
                    (empty_scenario_id, "Without attacks"),
                ):
                    connection.execute(
                        text(
                            'INSERT INTO "ScenarioResultEntries" '
                            "(id, scenario_name, scenario_version, pyrit_version, scenario_identifier, "
                            "objective_target_identifier, scenario_run_state, number_tries, "
                            "completion_time, timestamp) "
                            "VALUES (:id, :name, 1, '1.0.0', '{}', '{}', 'COMPLETED', "
                            "0, '2026-07-14', '2026-07-14')"
                        ),
                        {"id": current_id, "name": scenario_name},
                    )

                attack_rows = (
                    ("conv-alpha-1", "alpha", "2026-07-14 10:00:00"),
                    ("conv-beta", "beta", "2026-07-14 10:01:00"),
                    ("conv-alpha-2", "alpha", "2026-07-14 10:02:00"),
                )
                for conversation_id, parent_collection, timestamp in attack_rows:
                    connection.execute(
                        text(
                            'INSERT INTO "AttackResultEntries" '
                            "(id, conversation_id, objective, executed_turns, execution_time_ms, outcome, "
                            "timestamp, attribution_parent_id, attribution_data) "
                            "VALUES (:id, :conversation_id, 'objective', 1, 0, 'success', "
                            ":timestamp, :scenario_id, :attribution_data)"
                        ),
                        {
                            "id": str(uuid.uuid4()),
                            "conversation_id": conversation_id,
                            "timestamp": timestamp,
                            "scenario_id": scenario_id,
                            "attribution_data": json.dumps({"parent_collection": parent_collection}),
                        },
                    )

                command.downgrade(config, _V1_CLEANUP_PREV_REV)

                restored_task = connection.execute(
                    text('SELECT task FROM "ScoreEntries" WHERE id = :id'), {"id": score_id}
                ).scalar_one()
                manifests = dict(
                    connection.execute(text('SELECT id, attack_results_json FROM "ScenarioResultEntries"')).fetchall()
                )
                score_columns = {column["name"] for column in inspect(connection).get_columns("ScoreEntries")}
                scenario_columns = {
                    column["name"] for column in inspect(connection).get_columns("ScenarioResultEntries")
                }

            assert "task" in score_columns
            assert "attack_results_json" in scenario_columns
            assert restored_task == "restored objective"
            assert json.loads(manifests[scenario_id]) == {
                "alpha": ["conv-alpha-1", "conv-alpha-2"],
                "beta": ["conv-beta"],
            }
            assert json.loads(manifests[empty_scenario_id]) == {}
        finally:
            engine.dispose()


# =============================================================================
# Backfill tests for the Conversations table migration (b2f4c6a8d1e3)
# =============================================================================


_CONVERSATIONS_REV = "b2f4c6a8d1e3"
_CONVERSATIONS_PREV_REV = "f1a2b3c4d5e6"

_TARGET_A = '{"name": "target-a"}'
_TARGET_B = '{"name": "target-b"}'


def _seed_pre_conversations_prompt_piece(connection, *, piece_id, conversation_id, sequence, target_identifier):
    """Insert a PromptMemoryEntry row at the pre-Conversations revision."""
    connection.execute(
        text(
            'INSERT INTO "PromptMemoryEntries" '
            "(id, role, conversation_id, sequence, timestamp, labels, prompt_metadata, "
            "prompt_target_identifier, attack_identifier, original_value_data_type, "
            "original_value, converted_value_data_type, original_prompt_id) "
            "VALUES (:id, 'user', :conv, :seq, '2026-05-20', '{}', '{}', "
            ":target, '{}', 'text', 'hello', 'text', :id)"
        ),
        {"id": piece_id, "conv": conversation_id, "seq": sequence, "target": target_identifier},
    )


def _make_prompt_target_identifier_nullable(connection):
    """Allow legacy null target values in a SQLite migration fixture."""
    import sqlalchemy as sa
    from alembic.migration import MigrationContext
    from alembic.operations import Operations

    operations = Operations(MigrationContext.configure(connection))
    with operations.batch_alter_table("PromptMemoryEntries") as batch_op:
        batch_op.alter_column(
            "prompt_target_identifier",
            existing_type=sa.JSON(),
            existing_nullable=False,
            nullable=True,
        )


def test_conversations_migration_script_metadata():
    """The Conversations migration declares the expected revision chain."""
    from pyrit.memory.alembic.versions import b2f4c6a8d1e3_add_conversations_table as mig

    assert mig.revision == _CONVERSATIONS_REV
    assert mig.down_revision == _CONVERSATIONS_PREV_REV
    assert mig.branch_labels is None
    assert mig.depends_on is None


def test_conversations_backfill_populates_targets_and_handles_conflicts(caplog):
    """Upgrading to the Conversations revision backfills one row per conversation_id:
    the target comes from PromptMemoryEntries (first non-null wins on conflict),
    attack-only conversations get a null placeholder, and the per-row identifier
    columns are dropped."""
    import logging

    with tempfile.TemporaryDirectory() as temp_dir:
        db_path = os.path.join(temp_dir, "conversations-backfill.db")
        engine = create_engine(f"sqlite:///{db_path}")
        try:
            with engine.begin() as connection:
                config = _config_for(connection)
                command.upgrade(config, _CONVERSATIONS_PREV_REV)
                _make_prompt_target_identifier_nullable(connection)

                # A conversation whose two pieces share one target.
                _seed_pre_conversations_prompt_piece(
                    connection,
                    piece_id=str(uuid.uuid4()),
                    conversation_id="conv-keep",
                    sequence=0,
                    target_identifier=_TARGET_A,
                )
                _seed_pre_conversations_prompt_piece(
                    connection,
                    piece_id=str(uuid.uuid4()),
                    conversation_id="conv-keep",
                    sequence=1,
                    target_identifier=_TARGET_A,
                )
                # A conversation with two distinct non-null targets -> first wins + warning.
                _seed_pre_conversations_prompt_piece(
                    connection,
                    piece_id=str(uuid.uuid4()),
                    conversation_id="conv-conflict",
                    sequence=0,
                    target_identifier=_TARGET_A,
                )
                _seed_pre_conversations_prompt_piece(
                    connection,
                    piece_id=str(uuid.uuid4()),
                    conversation_id="conv-conflict",
                    sequence=1,
                    target_identifier=_TARGET_B,
                )
                # A later non-null target replaces an earlier null target.
                _seed_pre_conversations_prompt_piece(
                    connection,
                    piece_id=str(uuid.uuid4()),
                    conversation_id="conv-null-first",
                    sequence=0,
                    target_identifier=None,
                )
                _seed_pre_conversations_prompt_piece(
                    connection,
                    piece_id=str(uuid.uuid4()),
                    conversation_id="conv-null-first",
                    sequence=1,
                    target_identifier=_TARGET_A,
                )
                # A conversation with only null targets retains a null target.
                _seed_pre_conversations_prompt_piece(
                    connection,
                    piece_id=str(uuid.uuid4()),
                    conversation_id="conv-null-only",
                    sequence=0,
                    target_identifier=None,
                )
                # A conversation referenced only by an AttackResultEntry (no prompt rows).
                _seed_post_drop_attack_result(
                    connection, attack_id=str(uuid.uuid4()), conversation_id="conv-attack-only"
                )

                with caplog.at_level(logging.WARNING):
                    command.upgrade(config, _CONVERSATIONS_REV)

                rows = connection.execute(
                    text(
                        "SELECT conversation_id, target_identifier, pyrit_version "
                        'FROM "Conversations" ORDER BY conversation_id'
                    )
                ).fetchall()
                prompt_cols = {c["name"] for c in inspect(connection).get_columns("PromptMemoryEntries")}

            targets_by_conv = {r[0]: r[1] for r in rows}

            assert set(targets_by_conv) == {
                "conv-keep",
                "conv-conflict",
                "conv-null-first",
                "conv-null-only",
                "conv-attack-only",
            }
            assert targets_by_conv["conv-keep"] == _TARGET_A
            assert targets_by_conv["conv-conflict"] == _TARGET_A  # first non-null wins
            assert targets_by_conv["conv-null-first"] == _TARGET_A
            assert targets_by_conv["conv-null-only"] is None
            assert targets_by_conv["conv-attack-only"] is None  # placeholder for attack-only conversation
            assert all(row[2] is None for row in rows)

            # The conflicting targets produced a warning.
            assert any("multiple distinct" in r.message for r in caplog.records)

            # The per-row identifier columns are gone.
            assert "prompt_target_identifier" not in prompt_cols
            assert "attack_identifier" not in prompt_cols
        finally:
            engine.dispose()


def test_conversations_backfill_batches_rows_and_preserves_existing(caplog):
    """The backfill bounds each statement and does not overwrite existing rows."""
    from unittest.mock import patch

    from pyrit.memory.alembic.versions import b2f4c6a8d1e3_add_conversations_table as mig

    caplog.set_level("INFO")
    with tempfile.TemporaryDirectory() as temp_dir:
        engine = create_engine(f"sqlite:///{os.path.join(temp_dir, 'conversation-batches.db')}")
        insert_metrics = []

        def record_conversation_insert(conn, cursor, statement, parameters, context, executemany):
            if statement.startswith('INSERT INTO "Conversations"'):
                insert_metrics.append((statement.count("), (") + 1, len(parameters)))

        try:
            with engine.begin() as connection:
                config = _config_for(connection)
                command.upgrade(config, _CONVERSATIONS_PREV_REV)
                connection.execute(
                    text(
                        'CREATE TABLE "Conversations" ('
                        "conversation_id VARCHAR(36) NOT NULL PRIMARY KEY, "
                        "target_identifier JSON, pyrit_version VARCHAR)"
                    )
                )
                connection.execute(
                    text(
                        'INSERT INTO "Conversations" '
                        "(conversation_id, target_identifier, pyrit_version) "
                        "VALUES ('existing', :target, 'preserved')"
                    ),
                    {"target": _TARGET_B},
                )

                prompt_parameters = [
                    {
                        "id": str(uuid.uuid4()),
                        "conv": f"conversation-{index:04d}",
                        "seq": 0,
                        "target": _TARGET_A,
                    }
                    for index in range(mig._CONVERSATION_INSERT_BATCH_SIZE * 2 + 1)
                ]
                prompt_parameters.append(
                    {
                        "id": str(uuid.uuid4()),
                        "conv": "existing",
                        "seq": 0,
                        "target": _TARGET_A,
                    }
                )
                connection.execute(
                    text(
                        'INSERT INTO "PromptMemoryEntries" '
                        "(id, role, conversation_id, sequence, timestamp, labels, prompt_metadata, "
                        "prompt_target_identifier, attack_identifier, original_value_data_type, "
                        "original_value, converted_value_data_type, original_prompt_id) "
                        "VALUES (:id, 'user', :conv, :seq, '2026-05-20', '{}', '{}', "
                        ":target, '{}', 'text', 'hello', 'text', :id)"
                    ),
                    prompt_parameters,
                )

                event.listen(engine, "before_cursor_execute", record_conversation_insert)
                with patch.object(mig.op, "get_bind", return_value=connection):
                    mig._backfill_conversations()
                event.remove(engine, "before_cursor_execute", record_conversation_insert)

                existing = connection.execute(
                    text(
                        'SELECT target_identifier, pyrit_version FROM "Conversations" '
                        "WHERE conversation_id = 'existing'"
                    )
                ).one()
                row_count = connection.execute(text('SELECT COUNT(*) FROM "Conversations"')).scalar_one()
                null_versions = connection.execute(
                    text('SELECT COUNT(*) FROM "Conversations" WHERE pyrit_version IS NULL')
                ).scalar_one()

            assert insert_metrics == [(400, 800), (400, 800), (1, 2)]
            assert existing == (_TARGET_B, "preserved")
            assert row_count == mig._CONVERSATION_INSERT_BATCH_SIZE * 2 + 2
            assert null_versions == mig._CONVERSATION_INSERT_BATCH_SIZE * 2 + 1
            assert any("completed insert batch 3/3" in record.message for record in caplog.records)
        finally:
            if event.contains(engine, "before_cursor_execute", record_conversation_insert):
                event.remove(engine, "before_cursor_execute", record_conversation_insert)
            engine.dispose()


def test_conversations_backfill_rolls_back_prior_batches_on_failure():
    """A later batch failure leaves no rows from earlier batches."""
    from unittest.mock import patch

    from pyrit.memory.alembic.versions import b2f4c6a8d1e3_add_conversations_table as mig

    with tempfile.TemporaryDirectory() as temp_dir:
        engine = create_engine(f"sqlite:///{os.path.join(temp_dir, 'conversation-rollback.db')}")
        insert_count = 0

        def fail_second_conversation_insert(conn, cursor, statement, parameters, context, executemany):
            nonlocal insert_count
            if not statement.startswith('INSERT INTO "Conversations"'):
                return
            insert_count += 1
            if insert_count == 2:
                raise RuntimeError("forced second-batch failure")

        try:
            with engine.begin() as connection:
                config = _config_for(connection)
                command.upgrade(config, _CONVERSATIONS_PREV_REV)
                connection.execute(
                    text(
                        'CREATE TABLE "Conversations" ('
                        "conversation_id VARCHAR(36) NOT NULL PRIMARY KEY, "
                        "target_identifier JSON, pyrit_version VARCHAR)"
                    )
                )
                connection.execute(
                    text(
                        'INSERT INTO "PromptMemoryEntries" '
                        "(id, role, conversation_id, sequence, timestamp, labels, prompt_metadata, "
                        "prompt_target_identifier, attack_identifier, original_value_data_type, "
                        "original_value, converted_value_data_type, original_prompt_id) "
                        "VALUES (:id, 'user', :conv, 0, '2026-05-20', '{}', '{}', "
                        ":target, '{}', 'text', 'hello', 'text', :id)"
                    ),
                    [
                        {
                            "id": str(uuid.uuid4()),
                            "conv": f"conversation-{index:04d}",
                            "target": _TARGET_A,
                        }
                        for index in range(mig._CONVERSATION_INSERT_BATCH_SIZE + 1)
                    ],
                )

            event.listen(engine, "before_cursor_execute", fail_second_conversation_insert)
            with pytest.raises(RuntimeError, match="forced second-batch failure"):
                with engine.begin() as connection:
                    with patch.object(mig.op, "get_bind", return_value=connection):
                        mig._backfill_conversations()
            event.remove(engine, "before_cursor_execute", fail_second_conversation_insert)

            with engine.connect() as connection:
                assert connection.execute(text('SELECT COUNT(*) FROM "Conversations"')).scalar_one() == 0
            assert insert_count == 2
        finally:
            if event.contains(engine, "before_cursor_execute", fail_second_conversation_insert):
                event.remove(engine, "before_cursor_execute", fail_second_conversation_insert)
            engine.dispose()


def test_conversations_migration_downgrade_restores_columns():
    """Downgrading drops the Conversations table and re-adds the per-row identifier columns."""
    with tempfile.TemporaryDirectory() as temp_dir:
        db_path = os.path.join(temp_dir, "conversations-downgrade.db")
        engine = create_engine(f"sqlite:///{db_path}")
        try:
            with engine.begin() as connection:
                config = _config_for(connection)
                command.upgrade(config, _CONVERSATIONS_REV)

                assert "Conversations" in set(inspect(connection).get_table_names())
                cols_up = {c["name"] for c in inspect(connection).get_columns("PromptMemoryEntries")}
                assert "prompt_target_identifier" not in cols_up

                command.downgrade(config, _CONVERSATIONS_PREV_REV)

                assert "Conversations" not in set(inspect(connection).get_table_names())
                cols_down = {c["name"] for c in inspect(connection).get_columns("PromptMemoryEntries")}
                assert "prompt_target_identifier" in cols_down
                assert "attack_identifier" in cols_down
        finally:
            engine.dispose()


# =============================================================================
# Backfill tests for identifier persistence (e5f7a9c1b3d2)
# =============================================================================


def test_identifier_migrations_are_nullable_and_best_effort_with_malformed_json():
    """Malformed retained identifiers do not block upgrades and leave nullable links unset."""
    with tempfile.TemporaryDirectory() as temp_dir:
        engine = create_engine(f"sqlite:///{os.path.join(temp_dir, 'identifier-best-effort.db')}")
        try:
            with engine.begin() as connection:
                config = _config_for(connection)
                command.upgrade(config, "d4e6f8a0b2c4")
                connection.execute(
                    text('INSERT INTO "Conversations" (conversation_id, target_identifier) VALUES (:id, :value)'),
                    {"id": "malformed-conversation", "value": "not-json"},
                )
                score_id = str(uuid.uuid4())
                connection.execute(
                    text(
                        'INSERT INTO "ScoreEntries" '
                        "(id, score_value, score_type, score_metadata, scorer_class_identifier, timestamp) "
                        "VALUES (:id, 'True', 'true_false', '{}', 'not-json', '2026-07-13')"
                    ),
                    {"id": score_id},
                )
                result_id = str(uuid.uuid4())
                connection.execute(
                    text(
                        'INSERT INTO "ScenarioResultEntries" '
                        "(id, scenario_name, scenario_version, pyrit_version, scenario_identifier, "
                        "objective_target_identifier, scenario_run_state, attack_results_json, number_tries, "
                        "completion_time, timestamp) VALUES (:id, 'Scenario', 1, '0.10.0', 'not-json', '{}', "
                        "'COMPLETED', '{}', 1, '2026-07-13', '2026-07-13')"
                    ),
                    {"id": result_id},
                )
                prompt_id = str(uuid.uuid4())
                connection.execute(
                    text(
                        'INSERT INTO "PromptMemoryEntries" '
                        "(id, role, conversation_id, sequence, timestamp, labels, prompt_metadata, "
                        "converter_identifiers, original_value_data_type, original_value, converted_value_data_type, "
                        "original_prompt_id) VALUES (:id, 'user', 'conversation', 0, '2026-07-13', '{}', '{}', "
                        "'not-json', 'text', 'prompt', 'text', :id)"
                    ),
                    {"id": prompt_id},
                )
                attack_result_id = str(uuid.uuid4())
                connection.execute(
                    text(
                        'INSERT INTO "AttackResultEntries" '
                        "(id, conversation_id, objective, atomic_attack_identifier, executed_turns, execution_time_ms, "
                        "outcome, timestamp) VALUES (:id, 'conversation', 'objective', 'not-json', 0, 0, "
                        "'undetermined', '2026-07-13')"
                    ),
                    {"id": attack_result_id},
                )

                command.upgrade(config, "e5f7a9c1b3d2")

                assert (
                    connection.execute(
                        text(
                            'SELECT target_identifier_hash FROM "Conversations" '
                            "WHERE conversation_id = 'malformed-conversation'"
                        )
                    ).scalar_one()
                    is None
                )
                assert (
                    connection.execute(
                        text('SELECT scorer_identifier_hash FROM "ScoreEntries" WHERE id = :id'),
                        {"id": score_id},
                    ).scalar_one()
                    is None
                )
                assert (
                    connection.execute(
                        text('SELECT scenario_identifier_hash FROM "ScenarioResultEntries" WHERE id = :id'),
                        {"id": result_id},
                    ).scalar_one()
                    is None
                )
                assert connection.execute(text('SELECT COUNT(*) FROM "PromptConverterIdentifiers"')).scalar_one() == 0
                assert (
                    connection.execute(
                        text('SELECT atomic_attack_identifier_hash FROM "AttackResultEntries" WHERE id = :id'),
                        {"id": attack_result_id},
                    ).scalar_one()
                    is None
                )

                for table_name in (
                    "TargetIdentifiers",
                    "ScorerIdentifiers",
                    "ScenarioIdentifiers",
                    "ConverterIdentifiers",
                    "SeedIdentifiers",
                    "AttackIdentifiers",
                    "AttackTechniqueIdentifiers",
                    "AtomicAttackIdentifiers",
                ):
                    columns = {column["name"]: column for column in inspect(connection).get_columns(table_name)}
                    assert columns["hash"]["nullable"] is False
                    assert columns["class_name"]["nullable"] is True
                    assert columns["class_module"]["nullable"] is True
                    assert columns["identifier_json"]["nullable"] is True
        finally:
            engine.dispose()


def test_target_identifier_backfill_batches_conversation_links():
    """Target links use bounded statements even when every identifier is unique."""
    from io import StringIO

    with tempfile.TemporaryDirectory() as temp_dir:
        engine = create_engine(f"sqlite:///{os.path.join(temp_dir, 'identifier-link-batches.db')}")
        update_metrics = []
        target_insert_count = 0
        progress_stream = StringIO()

        def record_target_backfill(conn, cursor, statement, parameters, context, executemany):
            nonlocal target_insert_count
            if statement.startswith('INSERT INTO "TargetIdentifiers"'):
                target_insert_count += 1
            if statement.startswith('INSERT INTO "_PyritConversationTargetLinks"'):
                update_metrics.append((statement.count("), (") + 1, len(parameters)))

        try:
            with engine.begin() as connection:
                config = _config_for(connection)
                config.stdout = progress_stream
                command.upgrade(config, "d4e6f8a0b2c4")
                conversation_parameters = []
                for index in range(602):
                    identifier_index = 0 if index == 601 else index
                    identifier = {
                        "class_module": "tests.unit.memory.test_migration",
                        "class_name": f"Target{identifier_index}",
                        "hash": f"{identifier_index:064x}",
                    }
                    conversation_parameters.append(
                        {
                            "conversation_id": f"conversation-{index:04d}",
                            "target_identifier": json.dumps(identifier, sort_keys=True),
                        }
                    )
                connection.execute(
                    text(
                        'INSERT INTO "Conversations" (conversation_id, target_identifier) '
                        "VALUES (:conversation_id, :target_identifier)"
                    ),
                    conversation_parameters,
                )

                event.listen(engine, "before_cursor_execute", record_target_backfill)
                command.upgrade(config, "e5f7a9c1b3d2")
                event.remove(engine, "before_cursor_execute", record_target_backfill)

                linked_count = connection.execute(
                    text('SELECT COUNT(*) FROM "Conversations" WHERE target_identifier_hash IS NOT NULL')
                ).scalar_one()
                identifier_count = connection.execute(text('SELECT COUNT(*) FROM "TargetIdentifiers"')).scalar_one()
                temp_tables = set(inspect(connection).get_temp_table_names())

            progress_output = progress_stream.getvalue()
            assert update_metrics == [(400, 800), (202, 404)]
            assert target_insert_count == 601
            assert linked_count == 602
            assert identifier_count == 601
            assert "_PyritConversationTargetLinks" not in temp_tables
            assert "processed 601 unique identifier graph(s)" in progress_output
            assert "completed staging batch 2/2" in progress_output
            assert "staged target-link update completed" in progress_output
        finally:
            if event.contains(engine, "before_cursor_execute", record_target_backfill):
                event.remove(engine, "before_cursor_execute", record_target_backfill)
            engine.dispose()


def test_target_identifier_link_batch_failure_retries_individually(caplog):
    """A failed target-link batch retries per row without losing healthy links."""
    from unittest.mock import patch

    from pyrit.memory.alembic.versions import e5f7a9c1b3d2_add_identifiers_tables as mig

    with tempfile.TemporaryDirectory() as temp_dir:
        engine = create_engine(f"sqlite:///{os.path.join(temp_dir, 'identifier-link-fallback.db')}")
        batch_attempts = 0

        def fail_first_target_link_batch(conn, cursor, statement, parameters, context, executemany):
            nonlocal batch_attempts
            if not statement.startswith('INSERT INTO "_PyritConversationTargetLinks"'):
                return
            batch_attempts += 1
            if batch_attempts == 1:
                raise RuntimeError("forced target-link batch failure")

        try:
            with engine.begin() as connection:
                config = _config_for(connection)
                command.upgrade(config, "e5f7a9c1b3d2")
                connection.execute(
                    text(
                        'INSERT INTO "TargetIdentifiers" (hash, class_name, class_module, identifier_json) '
                        "VALUES (:hash, 'Target', 'tests.unit.memory.test_migration', '{}')"
                    ),
                    [{"hash": f"{index:064x}"} for index in range(3)],
                )
                connection.execute(
                    text(
                        'INSERT INTO "Conversations" (conversation_id, target_identifier) '
                        "VALUES (:conversation_id, NULL)"
                    ),
                    [{"conversation_id": f"conversation-{index}"} for index in range(3)],
                )

                event.listen(engine, "before_cursor_execute", fail_first_target_link_batch)
                with patch.object(mig, "_TARGET_LINK_INSERT_BATCH_SIZE", 2):
                    linked, skipped = mig._update_conversation_target_links(
                        bind=connection,
                        links=[(f"conversation-{index}", f"{index:064x}") for index in range(3)],
                    )
                event.remove(engine, "before_cursor_execute", fail_first_target_link_batch)

                stored_links = connection.execute(
                    text('SELECT conversation_id, target_identifier_hash FROM "Conversations" ORDER BY conversation_id')
                ).fetchall()

            assert linked == 3
            assert skipped == 0
            assert batch_attempts == 4
            assert stored_links == [
                ("conversation-0", f"{0:064x}"),
                ("conversation-1", f"{1:064x}"),
                ("conversation-2", f"{2:064x}"),
            ]
            assert any("staging insert failed" in record.message for record in caplog.records)
        finally:
            if event.contains(engine, "before_cursor_execute", fail_first_target_link_batch):
                event.remove(engine, "before_cursor_execute", fail_first_target_link_batch)
            engine.dispose()


def test_target_identifier_bulk_update_failure_retries_individually(caplog):
    """A failed staged update retries per row and removes its staging table."""
    from unittest.mock import patch

    from pyrit.memory.alembic.versions import e5f7a9c1b3d2_add_identifiers_tables as mig

    with tempfile.TemporaryDirectory() as temp_dir:
        engine = create_engine(f"sqlite:///{os.path.join(temp_dir, 'identifier-update-fallback.db')}")

        def fail_one_individual_update(conn, cursor, statement, parameters, context, executemany):
            if not statement.startswith('UPDATE "Conversations" SET target_identifier_hash = ?'):
                return
            if parameters[1] == "conversation-1":
                raise RuntimeError("forced individual target-link failure")

        try:
            with engine.begin() as connection:
                config = _config_for(connection)
                command.upgrade(config, "e5f7a9c1b3d2")
                connection.execute(
                    text(
                        'INSERT INTO "TargetIdentifiers" (hash, class_name, class_module, identifier_json) '
                        "VALUES (:hash, 'Target', 'tests.unit.memory.test_migration', '{}')"
                    ),
                    [{"hash": f"{index:064x}"} for index in range(3)],
                )
                connection.execute(
                    text(
                        'INSERT INTO "Conversations" (conversation_id, target_identifier) '
                        "VALUES (:conversation_id, NULL)"
                    ),
                    [{"conversation_id": f"conversation-{index}"} for index in range(3)],
                )

                event.listen(engine, "before_cursor_execute", fail_one_individual_update)
                with patch.object(
                    mig,
                    "_update_conversation_target_links_from_staging",
                    side_effect=RuntimeError("forced staged target-link update failure"),
                ):
                    linked, skipped = mig._update_conversation_target_links(
                        bind=connection,
                        links=[(f"conversation-{index}", f"{index:064x}") for index in range(3)],
                    )
                event.remove(engine, "before_cursor_execute", fail_one_individual_update)

                stored_links = connection.execute(
                    text('SELECT conversation_id, target_identifier_hash FROM "Conversations" ORDER BY conversation_id')
                ).fetchall()
                temp_tables = set(inspect(connection).get_temp_table_names())

            assert linked == 2
            assert skipped == 1
            assert stored_links == [
                ("conversation-0", f"{0:064x}"),
                ("conversation-1", None),
                ("conversation-2", f"{2:064x}"),
            ]
            assert "_PyritConversationTargetLinks" not in temp_tables
            assert any("staged link update failed" in record.message for record in caplog.records)
            assert any("skipped conversation 'conversation-1'" in record.message for record in caplog.records)
        finally:
            if event.contains(engine, "before_cursor_execute", fail_one_individual_update):
                event.remove(engine, "before_cursor_execute", fail_one_individual_update)
            engine.dispose()


def test_remaining_identifier_backfills_batch_row_links():
    """Scorer, scenario, and attack links use bounded staging statements."""
    from pyrit.memory.alembic.versions import e5f7a9c1b3d2_add_identifiers_tables as mig

    row_count = mig._IDENTIFIER_LINK_INSERT_BATCH_SIZE + 1
    identifiers = {
        "scorer": {
            "hash": "d" * 64,
            "class_name": "BatchScorer",
            "class_module": "tests.unit.memory.test_migration",
        },
        "scenario": {
            "hash": "e" * 64,
            "class_name": "BatchScenario",
            "class_module": "tests.unit.memory.test_migration",
        },
        "attack": {
            "hash": "f" * 64,
            "class_name": "BatchAttack",
            "class_module": "tests.unit.memory.test_migration",
        },
    }

    with tempfile.TemporaryDirectory() as temp_dir:
        engine = create_engine(f"sqlite:///{os.path.join(temp_dir, 'remaining-identifier-batches.db')}")
        staging_metrics = []
        update_count = 0

        def record_identifier_links(conn, cursor, statement, parameters, context, executemany):
            nonlocal update_count
            if statement.startswith('INSERT INTO "_PyritIdentifierLinks"'):
                staging_metrics.append((statement.count("), (") + 1, len(parameters)))
            if statement.startswith(
                (
                    'UPDATE "ScoreEntries" SET "scorer_identifier_hash"',
                    'UPDATE "ScenarioResultEntries" SET "scenario_identifier_hash"',
                    'UPDATE "AttackResultEntries" SET "atomic_attack_identifier_hash"',
                )
            ):
                update_count += 1

        try:
            with engine.begin() as connection:
                config = _config_for(connection)
                command.upgrade(config, "d4e6f8a0b2c4")
                connection.execute(
                    text(
                        'INSERT INTO "ScoreEntries" '
                        "(id, score_value, score_type, score_metadata, scorer_class_identifier, timestamp) "
                        "VALUES (:id, 'True', 'true_false', '{}', :identifier, '2026-07-23')"
                    ),
                    [
                        {
                            "id": str(uuid.UUID(int=100_000 + index)),
                            "identifier": json.dumps(identifiers["scorer"]),
                        }
                        for index in range(row_count)
                    ],
                )
                connection.execute(
                    text(
                        'INSERT INTO "ScenarioResultEntries" '
                        "(id, scenario_name, scenario_version, pyrit_version, scenario_identifier, "
                        "objective_target_identifier, objective_scorer_identifier, scenario_run_state, "
                        "attack_results_json, number_tries, completion_time, timestamp) "
                        "VALUES (:id, 'BatchScenario', 1, '1.0.0', :identifier, '{}', '{}', "
                        "'COMPLETED', '{}', 1, '2026-07-23', '2026-07-23')"
                    ),
                    [
                        {
                            "id": str(uuid.UUID(int=200_000 + index)),
                            "identifier": json.dumps(identifiers["scenario"]),
                        }
                        for index in range(row_count)
                    ],
                )
                connection.execute(
                    text(
                        'INSERT INTO "AttackResultEntries" '
                        "(id, conversation_id, objective, atomic_attack_identifier, objective_sha256, "
                        "executed_turns, execution_time_ms, outcome, timestamp, pyrit_version) "
                        "VALUES (:id, :conversation_id, 'objective', :identifier, 'sha', 1, 0, "
                        "'success', '2026-07-23', '1.0.0')"
                    ),
                    [
                        {
                            "id": str(uuid.UUID(int=300_000 + index)),
                            "conversation_id": f"attack-{index}",
                            "identifier": json.dumps(identifiers["attack"]),
                        }
                        for index in range(row_count)
                    ],
                )

                event.listen(engine, "before_cursor_execute", record_identifier_links)
                command.upgrade(config, "e5f7a9c1b3d2")
                event.remove(engine, "before_cursor_execute", record_identifier_links)

                linked_counts = (
                    connection.execute(
                        text('SELECT COUNT(*) FROM "ScoreEntries" WHERE scorer_identifier_hash IS NOT NULL')
                    ).scalar_one(),
                    connection.execute(
                        text('SELECT COUNT(*) FROM "ScenarioResultEntries" WHERE scenario_identifier_hash IS NOT NULL')
                    ).scalar_one(),
                    connection.execute(
                        text(
                            'SELECT COUNT(*) FROM "AttackResultEntries" WHERE atomic_attack_identifier_hash IS NOT NULL'
                        )
                    ).scalar_one(),
                )
                temp_tables = set(inspect(connection).get_temp_table_names())

            assert staging_metrics == [(400, 800), (1, 2)] * 3
            assert update_count == 3
            assert linked_counts == (row_count, row_count, row_count)
            assert "_PyritIdentifierLinks" not in temp_tables
        finally:
            if event.contains(engine, "before_cursor_execute", record_identifier_links):
                event.remove(engine, "before_cursor_execute", record_identifier_links)
            engine.dispose()


def test_converter_identifier_backfill_batches_links_and_skips_empty_lists():
    """Prompt associations use bounded statements without savepoints for empty lists."""
    from pyrit.memory.alembic.versions import e5f7a9c1b3d2_add_identifiers_tables as mig

    linked_prompt_count = mig._CONVERTER_LINK_INSERT_BATCH_SIZE * 2 + 1
    identifier = {
        "hash": "c" * 64,
        "class_name": "BatchConverter",
        "class_module": "tests.unit.memory.test_migration",
    }

    with tempfile.TemporaryDirectory() as temp_dir:
        engine = create_engine(f"sqlite:///{os.path.join(temp_dir, 'converter-link-batches.db')}")
        insert_metrics = []
        savepoint_count = 0

        def record_converter_links(conn, cursor, statement, parameters, context, executemany):
            nonlocal savepoint_count
            if statement.startswith('INSERT INTO "PromptConverterIdentifiers"'):
                insert_metrics.append((statement.count("), (") + 1, len(parameters)))
            if statement.startswith("SAVEPOINT"):
                savepoint_count += 1

        try:
            with engine.begin() as connection:
                config = _config_for(connection)
                command.upgrade(config, "d4e6f8a0b2c4")
                prompt_statement = text(
                    'INSERT INTO "PromptMemoryEntries" '
                    "(id, role, conversation_id, sequence, timestamp, labels, prompt_metadata, "
                    "converter_identifiers, original_value_data_type, original_value, "
                    "converted_value_data_type, original_prompt_id, pyrit_version) "
                    "VALUES (:id, 'user', :conversation_id, 0, '2026-07-23', '{}', '{}', :identifiers, "
                    "'text', 'prompt', 'text', :id, '1.0.0')"
                )
                connection.execute(
                    prompt_statement,
                    [
                        {
                            "id": str(uuid.UUID(int=400_000 + index)),
                            "conversation_id": f"linked-{index}",
                            "identifiers": json.dumps([identifier]),
                        }
                        for index in range(linked_prompt_count)
                    ],
                )
                connection.execute(
                    prompt_statement,
                    [
                        {
                            "id": str(uuid.UUID(int=500_000 + index)),
                            "conversation_id": f"empty-{index}",
                            "identifiers": "[]",
                        }
                        for index in range(1_000)
                    ],
                )

                event.listen(engine, "before_cursor_execute", record_converter_links)
                command.upgrade(config, "e5f7a9c1b3d2")
                event.remove(engine, "before_cursor_execute", record_converter_links)

                association_count = connection.execute(
                    text('SELECT COUNT(*) FROM "PromptConverterIdentifiers"')
                ).scalar_one()
                converter_count = connection.execute(text('SELECT COUNT(*) FROM "ConverterIdentifiers"')).scalar_one()

            assert insert_metrics == [(300, 900), (300, 900), (1, 3)]
            assert association_count == linked_prompt_count
            assert converter_count == 1
            assert savepoint_count < 20
        finally:
            if event.contains(engine, "before_cursor_execute", record_converter_links):
                event.remove(engine, "before_cursor_execute", record_converter_links)
            engine.dispose()


def test_identifier_row_link_batch_failure_retries_individually(caplog):
    """A failed generic staging batch retries each healthy row."""
    from unittest.mock import patch

    from pyrit.memory.alembic.versions import e5f7a9c1b3d2_add_identifiers_tables as mig

    with tempfile.TemporaryDirectory() as temp_dir:
        engine = create_engine(f"sqlite:///{os.path.join(temp_dir, 'identifier-row-link-fallback.db')}")
        batch_attempts = 0

        def fail_first_identifier_link_batch(conn, cursor, statement, parameters, context, executemany):
            nonlocal batch_attempts
            if not statement.startswith('INSERT INTO "_PyritIdentifierLinks"'):
                return
            batch_attempts += 1
            if batch_attempts == 1:
                raise RuntimeError("forced identifier-link batch failure")

        try:
            with engine.begin() as connection:
                config = _config_for(connection)
                command.upgrade(config, "e5f7a9c1b3d2")
                connection.execute(
                    text(
                        'INSERT INTO "ScorerIdentifiers" (hash, class_name, class_module, identifier_json) '
                        "VALUES (:hash, 'Scorer', 'tests.unit.memory.test_migration', '{}')"
                    ),
                    [{"hash": f"{index:064x}"} for index in range(3)],
                )
                score_ids = [str(uuid.UUID(int=600_000 + index)) for index in range(3)]
                connection.execute(
                    text(
                        'INSERT INTO "ScoreEntries" '
                        "(id, score_value, score_type, score_metadata, scorer_class_identifier, timestamp) "
                        "VALUES (:id, 'True', 'true_false', '{}', '{}', '2026-07-23')"
                    ),
                    [{"id": score_id} for score_id in score_ids],
                )

                event.listen(engine, "before_cursor_execute", fail_first_identifier_link_batch)
                with patch.object(mig, "_IDENTIFIER_LINK_INSERT_BATCH_SIZE", 2):
                    linked, skipped = mig._update_identifier_row_links(
                        bind=connection,
                        links=[(score_ids[index], f"{index:064x}") for index in range(3)],
                        name="ScorerIdentifiers",
                        target_table="ScoreEntries",
                        target_id_column="id",
                        target_hash_column="scorer_identifier_hash",
                    )
                event.remove(engine, "before_cursor_execute", fail_first_identifier_link_batch)

                linked_count = connection.execute(
                    text('SELECT COUNT(*) FROM "ScoreEntries" WHERE scorer_identifier_hash IS NOT NULL')
                ).scalar_one()
                temp_tables = set(inspect(connection).get_temp_table_names())

            assert linked == 3
            assert skipped == 0
            assert batch_attempts == 4
            assert linked_count == 3
            assert "_PyritIdentifierLinks" not in temp_tables
            assert any("staging insert failed" in record.message for record in caplog.records)
        finally:
            if event.contains(engine, "before_cursor_execute", fail_first_identifier_link_batch):
                event.remove(engine, "before_cursor_execute", fail_first_identifier_link_batch)
            engine.dispose()


def test_converter_link_batch_failure_retries_at_prompt_boundary(caplog):
    """A failed association batch preserves healthy prompts and skips one conflicting prompt."""
    from pyrit.memory.alembic.versions import e5f7a9c1b3d2_add_identifiers_tables as mig

    desired_hash = "a" * 64
    existing_hash = "b" * 64
    with tempfile.TemporaryDirectory() as temp_dir:
        engine = create_engine(f"sqlite:///{os.path.join(temp_dir, 'converter-prompt-fallback.db')}")
        try:
            with engine.begin() as connection:
                config = _config_for(connection)
                command.upgrade(config, "e5f7a9c1b3d2")
                connection.execute(
                    text(
                        'INSERT INTO "ConverterIdentifiers" (hash, class_name, class_module, identifier_json) '
                        "VALUES (:hash, 'Converter', 'tests.unit.memory.test_migration', '{}')"
                    ),
                    [{"hash": desired_hash}, {"hash": existing_hash}],
                )
                prompt_ids = [str(uuid.UUID(int=700_000 + index)) for index in range(3)]
                connection.execute(
                    text(
                        'INSERT INTO "PromptMemoryEntries" '
                        "(id, role, conversation_id, sequence, timestamp, labels, prompt_metadata, "
                        "original_value_data_type, original_value, converted_value_data_type, original_prompt_id) "
                        "VALUES (:id, 'user', :conversation_id, 0, '2026-07-23', '{}', '{}', "
                        "'text', 'prompt', 'text', :id)"
                    ),
                    [
                        {"id": prompt_id, "conversation_id": f"conversation-{index}"}
                        for index, prompt_id in enumerate(prompt_ids)
                    ],
                )
                connection.execute(
                    text(
                        'INSERT INTO "PromptConverterIdentifiers" '
                        "(prompt_memory_entry_id, position, converter_identifier_hash) "
                        "VALUES (:prompt_id, 0, :hash)"
                    ),
                    {"prompt_id": prompt_ids[1], "hash": existing_hash},
                )

                inserted, skipped = mig._insert_prompt_converter_links(
                    bind=connection,
                    prompt_links=[(prompt_id, [(prompt_id, 0, desired_hash)]) for prompt_id in prompt_ids],
                )

                stored_links = connection.execute(
                    text(
                        "SELECT prompt_memory_entry_id, converter_identifier_hash "
                        'FROM "PromptConverterIdentifiers" ORDER BY prompt_memory_entry_id'
                    )
                ).fetchall()

            assert inserted == 2
            assert skipped == 1
            assert stored_links == [
                (prompt_ids[0], desired_hash),
                (prompt_ids[1], existing_hash),
                (prompt_ids[2], desired_hash),
            ]
            assert any("association batch failed" in record.message for record in caplog.records)
            assert any(f"skipped prompt {prompt_ids[1]}" in record.message for record in caplog.records)
        finally:
            engine.dispose()


def test_identifier_migration_reuses_edges_and_rolls_back_conflicting_row():
    """Repeated graphs reuse edges while a conflicting row is isolated in its savepoint."""
    child_hash = "a" * 64
    conflicting_child_hash = "b" * 64
    parent_hash = "c" * 64
    healthy_hash = "d" * 64
    child = {"hash": child_hash, "class_name": "Child", "class_module": "test"}
    parent = {
        "hash": parent_hash,
        "class_name": "Parent",
        "class_module": "test",
        "children": {"targets": [child]},
    }
    conflicting_parent = {
        **parent,
        "children": {"targets": [{"hash": conflicting_child_hash, "class_name": "OtherChild", "class_module": "test"}]},
    }
    healthy = {"hash": healthy_hash, "class_name": "Healthy", "class_module": "test"}

    with tempfile.TemporaryDirectory() as temp_dir:
        engine = create_engine(f"sqlite:///{os.path.join(temp_dir, 'identifier-row-savepoints.db')}")
        try:
            with engine.begin() as connection:
                config = _config_for(connection)
                command.upgrade(config, "d4e6f8a0b2c4")
                insert_statement = text(
                    'INSERT INTO "Conversations" (conversation_id, target_identifier) VALUES (:id, :identifier)'
                )
                for conversation_id, identifier in (
                    ("01-original", parent),
                    ("02-duplicate", parent),
                    ("03-conflict", conflicting_parent),
                    ("04-healthy", healthy),
                ):
                    connection.execute(
                        insert_statement,
                        {"id": conversation_id, "identifier": json.dumps(identifier)},
                    )

                command.upgrade(config, "e5f7a9c1b3d2")

                links = connection.execute(
                    text('SELECT conversation_id, target_identifier_hash FROM "Conversations" ORDER BY conversation_id')
                ).fetchall()
                target_hashes = set(connection.execute(text('SELECT hash FROM "TargetIdentifiers"')).scalars())
                edges = connection.execute(
                    text('SELECT parent_hash, position, child_hash FROM "TargetIdentifierChildren"')
                ).fetchall()

            assert links == [
                ("01-original", parent_hash),
                ("02-duplicate", parent_hash),
                ("03-conflict", None),
                ("04-healthy", healthy_hash),
            ]
            assert target_hashes == {child_hash, parent_hash, healthy_hash}
            assert edges == [(parent_hash, 0, child_hash)]
        finally:
            engine.dispose()


def test_identifier_migration_downgrade_drops_link_constraints_and_columns():
    """Downgrade removes identifier foreign keys before removing their columns."""
    with tempfile.TemporaryDirectory() as temp_dir:
        engine = create_engine(f"sqlite:///{os.path.join(temp_dir, 'identifier-downgrade.db')}")
        try:
            with engine.begin() as connection:
                config = _config_for(connection)
                command.upgrade(config, "e5f7a9c1b3d2")
                command.downgrade(config, "d4e6f8a0b2c4")

                assert "target_identifier_hash" not in {
                    column["name"] for column in inspect(connection).get_columns("Conversations")
                }
                assert "TargetIdentifiers" not in set(inspect(connection).get_table_names())
        finally:
            engine.dispose()


def test_identifier_migrations_do_not_import_domain_models():
    """Frozen identifier migrations operate on retained JSON rather than current domain models."""
    versions_dir = Path(__file__).resolve().parents[3] / "pyrit" / "memory" / "alembic" / "versions"
    revision_names = ("e5f7a9c1b3d2_add_identifiers_tables.py",)

    for revision_name in revision_names:
        source = (versions_dir / revision_name).read_text(encoding="utf-8")
        assert "from pyrit.models" not in source
        assert "import pyrit.models" not in source


def test_scorer_identifier_migration_backfills_graph_and_score_link():
    """Existing scorer JSON is materialized and linked without changing its content identity."""
    from pyrit.models import ComponentIdentifier

    prompt_target = ComponentIdentifier(
        class_name="OpenAIChatTarget",
        class_module="pyrit.prompt_target",
        params={"endpoint": "https://example.test"},
    )
    sub_scorer = ComponentIdentifier(
        class_name="SelfAskScorer",
        class_module="pyrit.score",
        children={"prompt_target": prompt_target},
    )
    composite = ComponentIdentifier(
        class_name="CompositeScorer",
        class_module="pyrit.score",
        params={"scorer_type": "true_false", "score_aggregator": "AND_"},
        children={"sub_scorers": [sub_scorer]},
    )
    score_id = str(uuid.uuid4())

    with tempfile.TemporaryDirectory() as temp_dir:
        engine = create_engine(f"sqlite:///{os.path.join(temp_dir, 'scorer-backfill.db')}")
        try:
            with engine.begin() as connection:
                config = _config_for(connection)
                command.upgrade(config, "d4e6f8a0b2c4")
                connection.execute(
                    text(
                        'INSERT INTO "ScoreEntries" '
                        "(id, score_value, score_type, score_metadata, scorer_class_identifier, timestamp) "
                        "VALUES (:id, 'True', 'true_false', '{}', :identifier, '2026-07-13')"
                    ),
                    {"id": score_id, "identifier": json.dumps(composite.model_dump())},
                )

                command.upgrade(config, "e5f7a9c1b3d2")

                score_hash = connection.execute(
                    text('SELECT scorer_identifier_hash FROM "ScoreEntries" WHERE id = :id'),
                    {"id": score_id},
                ).scalar_one()
                scorer_rows = connection.execute(
                    text('SELECT hash, scorer_type, score_aggregator, prompt_target_hash FROM "ScorerIdentifiers"')
                ).fetchall()
                target_hashes = connection.execute(text('SELECT hash FROM "TargetIdentifiers"')).scalars().all()
                scorer_child = connection.execute(
                    text('SELECT parent_hash, position, child_hash FROM "ScorerIdentifierChildren"')
                ).one()

            assert score_hash == composite.hash
            assert {row[0] for row in scorer_rows} == {composite.hash, sub_scorer.hash}
            root_row = next(row for row in scorer_rows if row[0] == composite.hash)
            sub_scorer_row = next(row for row in scorer_rows if row[0] == sub_scorer.hash)
            assert root_row[1:] == ("true_false", "AND_", None)
            assert sub_scorer_row[3] == prompt_target.hash
            assert target_hashes == [prompt_target.hash]
            assert scorer_child == (composite.hash, 0, sub_scorer.hash)
        finally:
            engine.dispose()


# =============================================================================
# Backfill tests for converter identifier persistence
# =============================================================================


def test_converter_identifier_migration_backfills_graph_and_prompt_links():
    """Existing converter JSON is normalized with dependencies and ordered prompt links."""
    from pyrit.models import ConverterIdentifier, TargetIdentifier

    target = TargetIdentifier(
        class_name="ConverterTarget",
        class_module="pyrit.prompt_target",
        model_name="converter-model",
    )
    nested = ConverterIdentifier(
        class_name="NestedConverter",
        class_module="pyrit.prompt_converter",
        supported_input_types=["text"],
        supported_output_types=["text"],
        converter_target=target,
    )
    converter = ConverterIdentifier(
        class_name="CompositeConverter",
        class_module="pyrit.prompt_converter",
        supported_input_types=["text"],
        supported_output_types=["text"],
        sub_converter=nested,
    )
    prompt_id = uuid.uuid4()

    with tempfile.TemporaryDirectory() as temp_dir:
        engine = create_engine(f"sqlite:///{os.path.join(temp_dir, 'converter-backfill.db')}")
        try:
            with engine.begin() as connection:
                config = _config_for(connection)
                command.upgrade(config, "d4e6f8a0b2c4")
                connection.execute(
                    text(
                        'INSERT INTO "PromptMemoryEntries" '
                        "(id, role, conversation_id, sequence, timestamp, labels, prompt_metadata, "
                        "converter_identifiers, original_value_data_type, original_value, "
                        "converted_value_data_type, original_prompt_id, pyrit_version) "
                        "VALUES (:id, 'user', 'conversation', 0, '2026-07-13', '{}', '{}', :identifiers, "
                        "'text', 'prompt', 'text', :id, '0.10.0')"
                    ),
                    {
                        "id": str(prompt_id),
                        "identifiers": json.dumps([converter.model_dump(), nested.model_dump()]),
                    },
                )

                command.upgrade(config, "e5f7a9c1b3d2")

                converter_rows = connection.execute(
                    text('SELECT hash, converter_target_hash, sub_converter_hash FROM "ConverterIdentifiers"')
                ).fetchall()
                target_hashes = connection.execute(text('SELECT hash FROM "TargetIdentifiers"')).scalars().all()
                links = connection.execute(
                    text(
                        'SELECT position, converter_identifier_hash FROM "PromptConverterIdentifiers" ORDER BY position'
                    )
                ).fetchall()

            assert {row[0] for row in converter_rows} == {converter.hash, nested.hash}
            root_row = next(row for row in converter_rows if row[0] == converter.hash)
            nested_row = next(row for row in converter_rows if row[0] == nested.hash)
            assert root_row[2] == nested.hash
            assert nested_row[1] == target.hash
            assert target_hashes == [target.hash]
            assert links == [(0, converter.hash), (1, nested.hash)]
        finally:
            engine.dispose()


# =============================================================================
# Backfill tests for scenario identifier persistence
# =============================================================================


def test_scenario_identifier_migration_backfills_dependencies_and_result_link():
    """Scenario-only target and scorer graphs are materialized before the scenario row."""
    from pyrit.models import ScenarioIdentifier, ScorerIdentifier, TargetIdentifier

    target = TargetIdentifier(
        class_name="ObjectiveTarget",
        class_module="pyrit.prompt_target",
        model_name="objective-model",
    )
    scorer_target = TargetIdentifier(
        class_name="ScorerTarget",
        class_module="pyrit.prompt_target",
        model_name="scorer-model",
    )
    scorer = ScorerIdentifier(
        class_name="SelfAskScorer",
        class_module="pyrit.score",
        scorer_type="true_false",
        prompt_target=scorer_target,
    )
    scenario = ScenarioIdentifier(
        class_name="TestScenario",
        class_module="pyrit.scenario",
        version=3,
        techniques=["TechniqueA"],
        datasets=["DatasetA"],
        objective_target=target,
        objective_scorer=scorer,
    )
    result_id = str(uuid.uuid4())

    with tempfile.TemporaryDirectory() as temp_dir:
        engine = create_engine(f"sqlite:///{os.path.join(temp_dir, 'scenario-backfill.db')}")
        try:
            with engine.begin() as connection:
                config = _config_for(connection)
                command.upgrade(config, "d4e6f8a0b2c4")
                connection.execute(
                    text(
                        'INSERT INTO "ScenarioResultEntries" '
                        "(id, scenario_name, scenario_version, pyrit_version, scenario_identifier, "
                        "objective_target_identifier, objective_scorer_identifier, scenario_run_state, "
                        "attack_results_json, number_tries, completion_time, timestamp) "
                        "VALUES (:id, 'TestScenario', 3, '0.10.0', :scenario, :target, :scorer, "
                        "'COMPLETED', '{}', 1, '2026-07-13', '2026-07-13')"
                    ),
                    {
                        "id": result_id,
                        "scenario": json.dumps(scenario.model_dump()),
                        "target": json.dumps(target.model_dump()),
                        "scorer": json.dumps(scorer.model_dump()),
                    },
                )

                command.upgrade(config, "e5f7a9c1b3d2")

                result_hash = connection.execute(
                    text('SELECT scenario_identifier_hash FROM "ScenarioResultEntries" WHERE id = :id'),
                    {"id": result_id},
                ).scalar_one()
                scenario_row = connection.execute(
                    text(
                        "SELECT hash, version, techniques, datasets, objective_target_hash, objective_scorer_hash "
                        'FROM "ScenarioIdentifiers"'
                    )
                ).one()
                target_hashes = set(connection.execute(text('SELECT hash FROM "TargetIdentifiers"')).scalars())
                scorer_row = connection.execute(text('SELECT hash, prompt_target_hash FROM "ScorerIdentifiers"')).one()

            assert result_hash == scenario.hash
            assert scenario_row == (
                scenario.hash,
                3,
                json.dumps(["TechniqueA"]),
                json.dumps(["DatasetA"]),
                target.hash,
                scorer.hash,
            )
            assert target_hashes == {target.hash, scorer_target.hash}
            assert scorer_row == (scorer.hash, scorer_target.hash)
        finally:
            engine.dispose()


# =============================================================================
# Backfill tests for attack identifier persistence
# =============================================================================


def test_attack_identifier_migration_backfills_graph_and_result_link():
    """Atomic attack JSON is normalized with dependencies and ordered seed links."""
    from pyrit.models import (
        AtomicAttackIdentifier,
        AttackIdentifier,
        AttackTechniqueIdentifier,
        ConverterIdentifier,
        ScorerIdentifier,
        SeedIdentifier,
        TargetIdentifier,
    )

    target = TargetIdentifier(class_name="Target", class_module="pyrit.prompt_target", model_name="model")
    scorer = ScorerIdentifier(class_name="Scorer", class_module="pyrit.score", scorer_type="true_false")
    converter = ConverterIdentifier(
        class_name="Converter",
        class_module="pyrit.prompt_converter",
        supported_input_types=["text"],
        supported_output_types=["text"],
    )
    technique_seed = SeedIdentifier(
        class_name="Seed",
        class_module="pyrit.models",
        value="technique seed",
        data_type="text",
    )
    dataset_seed = SeedIdentifier(
        class_name="Seed",
        class_module="pyrit.models",
        value="dataset seed",
        data_type="text",
    )
    attack = AttackIdentifier(
        class_name="Attack",
        class_module="pyrit.executor.attack",
        objective_target=target,
        objective_scorer=scorer,
        request_converters=[converter],
    )
    technique = AttackTechniqueIdentifier(
        class_name="AttackTechnique",
        class_module="pyrit.scenario.core.attack_technique",
        attack=attack,
        technique_seeds=[technique_seed],
    )
    atomic = AtomicAttackIdentifier(
        class_name="AtomicAttack",
        class_module="pyrit.scenario.core.atomic_attack",
        attack_technique=technique,
        seed_identifiers=[technique_seed, dataset_seed],
    )
    result_id = str(uuid.uuid4())

    with tempfile.TemporaryDirectory() as temp_dir:
        engine = create_engine(f"sqlite:///{os.path.join(temp_dir, 'attack-backfill.db')}")
        try:
            with engine.begin() as connection:
                config = _config_for(connection)
                command.upgrade(config, "d4e6f8a0b2c4")
                connection.execute(
                    text(
                        'INSERT INTO "AttackResultEntries" '
                        "(id, conversation_id, objective, atomic_attack_identifier, objective_sha256, "
                        "executed_turns, execution_time_ms, outcome, timestamp, pyrit_version) "
                        "VALUES (:id, 'conversation', 'objective', :identifier, 'sha', 1, 0, "
                        "'success', '2026-07-13', '0.10.0')"
                    ),
                    {"id": result_id, "identifier": json.dumps(atomic.model_dump())},
                )

                command.upgrade(config, "e5f7a9c1b3d2")

                result_hash = connection.execute(
                    text('SELECT atomic_attack_identifier_hash FROM "AttackResultEntries" WHERE id = :id'),
                    {"id": result_id},
                ).scalar_one()
                atomic_row = connection.execute(
                    text('SELECT hash, attack_technique_identifier_hash FROM "AtomicAttackIdentifiers"')
                ).one()
                technique_row = connection.execute(
                    text('SELECT hash, attack_identifier_hash FROM "AttackTechniqueIdentifiers"')
                ).one()
                attack_row = connection.execute(
                    text('SELECT hash, objective_target_hash, objective_scorer_hash FROM "AttackIdentifiers"')
                ).one()
                seed_hashes = set(connection.execute(text('SELECT hash FROM "SeedIdentifiers"')).scalars())
                atomic_seed_hashes = (
                    connection.execute(
                        text('SELECT seed_identifier_hash FROM "AtomicAttackSeedIdentifiers" ORDER BY position')
                    )
                    .scalars()
                    .all()
                )
                request_converter_hash = connection.execute(
                    text('SELECT converter_identifier_hash FROM "AttackRequestConverterIdentifiers"')
                ).scalar_one()

            assert result_hash == atomic.hash
            assert atomic_row == (atomic.hash, technique.hash)
            assert technique_row == (technique.hash, attack.hash)
            assert attack_row == (attack.hash, target.hash, scorer.hash)
            assert seed_hashes == {technique_seed.hash, dataset_seed.hash}
            assert atomic_seed_hashes == [technique_seed.hash, dataset_seed.hash]
            assert request_converter_hash == converter.hash
        finally:
            engine.dispose()


# =============================================================================
# Attack recency index + timestamp backfill migration (d7e9f1a3b5c6)
# =============================================================================


_ATTACK_RECENCY_REV = "d7e9f1a3b5c6"
_ATTACK_RECENCY_PREV_REV = "3f6e8a0c2d4b"


def _seed_attack_result_with_metadata(connection, *, attack_id, conversation_id, timestamp, attack_metadata):
    """Insert an AttackResultEntry row carrying legacy JSON recency metadata."""
    connection.execute(
        text(
            'INSERT INTO "AttackResultEntries" '
            "(id, conversation_id, objective, executed_turns, execution_time_ms, outcome, "
            "timestamp, attack_metadata) "
            "VALUES (:id, :conv, 'obj', 1, 0, 'success', :timestamp, :metadata)"
        ),
        {
            "id": attack_id,
            "conv": conversation_id,
            "timestamp": timestamp,
            "metadata": json.dumps(attack_metadata) if attack_metadata is not None else None,
        },
    )


def test_attack_recency_migration_script_metadata():
    """The attack-recency migration declares the expected revision chain."""
    from pyrit.memory.alembic.versions import d7e9f1a3b5c6_index_attack_result_recency as mig

    assert mig.revision == _ATTACK_RECENCY_REV
    assert mig.down_revision == _ATTACK_RECENCY_PREV_REV
    assert mig.branch_labels is None
    assert mig.depends_on is None


def test_attack_recency_upgrade_creates_indexes_and_backfills_timestamp():
    """Upgrading adds both indexes, backfills timestamp from the legacy JSON keys, and
    strips the redundant ``updated_at`` key while preserving ``created_at``."""
    edited_id = str(uuid.uuid4())
    prog_id = str(uuid.uuid4())

    with tempfile.TemporaryDirectory() as temp_dir:
        engine = create_engine(f"sqlite:///{os.path.join(temp_dir, 'attack-recency.db')}")
        try:
            with engine.begin() as connection:
                config = _config_for(connection)
                command.upgrade(config, _ATTACK_RECENCY_PREV_REV)

                # A manually-edited GUI conversation: timestamp column is stale (creation time),
                # recency lives in metadata.updated_at.
                _seed_attack_result_with_metadata(
                    connection,
                    attack_id=edited_id,
                    conversation_id="conv-edited",
                    timestamp="2020-01-01 00:00:00",
                    attack_metadata={
                        "created_at": "2020-01-01T00:00:00+00:00",
                        "updated_at": "2026-06-01T12:00:00+00:00",
                    },
                )
                # A programmatic attack: no JSON recency keys, timestamp column is authoritative.
                _seed_attack_result_with_metadata(
                    connection,
                    attack_id=prog_id,
                    conversation_id="conv-prog",
                    timestamp="2023-03-03 08:00:00",
                    attack_metadata=None,
                )

                command.upgrade(config, _ATTACK_RECENCY_REV)

                index_names = {ix["name"] for ix in inspect(connection).get_indexes("AttackResultEntries")}
                conversation_id_column = next(
                    column
                    for column in inspect(connection).get_columns("AttackResultEntries")
                    if column["name"] == "conversation_id"
                )
                rows = dict(
                    connection.execute(text('SELECT id, attack_metadata FROM "AttackResultEntries"')).fetchall()
                )
                timestamps = dict(
                    connection.execute(text('SELECT id, timestamp FROM "AttackResultEntries"')).fetchall()
                )

            assert "ix_AttackResultEntries_conversation_id" in index_names
            assert "ix_AttackResultEntries_timestamp_id" in index_names
            assert conversation_id_column["type"].length == 36

            edited_metadata = json.loads(rows[edited_id])
            assert "updated_at" not in edited_metadata
            assert edited_metadata["created_at"] == "2020-01-01T00:00:00+00:00"
            # Timestamp advanced to the former updated_at so the edited row keeps its recency.
            assert str(timestamps[edited_id]).startswith("2026-06-01 12:00:00")

            # Programmatic row is untouched: no metadata, original timestamp preserved.
            assert rows[prog_id] is None
            assert str(timestamps[prog_id]).startswith("2023-03-03 08:00:00")
        finally:
            engine.dispose()


def test_attack_recency_downgrade_restores_updated_at_and_drops_indexes():
    """Downgrading drops both indexes and rewrites ``metadata.updated_at`` from the
    timestamp column so the legacy JSON recency sort works again."""
    edited_id = str(uuid.uuid4())

    with tempfile.TemporaryDirectory() as temp_dir:
        engine = create_engine(f"sqlite:///{os.path.join(temp_dir, 'attack-recency-down.db')}")
        try:
            with engine.begin() as connection:
                config = _config_for(connection)
                command.upgrade(config, _ATTACK_RECENCY_REV)

                _seed_attack_result_with_metadata(
                    connection,
                    attack_id=edited_id,
                    conversation_id="conv-edited",
                    timestamp="2026-06-01 12:00:00",
                    attack_metadata={"created_at": "2026-06-01T00:00:00+00:00"},
                )

                index_names_up = {ix["name"] for ix in inspect(connection).get_indexes("AttackResultEntries")}
                assert "ix_AttackResultEntries_timestamp_id" in index_names_up

                command.downgrade(config, _ATTACK_RECENCY_PREV_REV)

                index_names_down = {ix["name"] for ix in inspect(connection).get_indexes("AttackResultEntries")}
                conversation_id_column_down = next(
                    column
                    for column in inspect(connection).get_columns("AttackResultEntries")
                    if column["name"] == "conversation_id"
                )
                restored = json.loads(
                    connection.execute(
                        text('SELECT attack_metadata FROM "AttackResultEntries" WHERE id = :id'),
                        {"id": edited_id},
                    ).scalar_one()
                )

            assert "ix_AttackResultEntries_conversation_id" not in index_names_down
            assert "ix_AttackResultEntries_timestamp_id" not in index_names_down
            assert conversation_id_column_down["type"].length is None
            assert restored["created_at"] == "2026-06-01T00:00:00+00:00"
            assert restored["updated_at"].startswith("2026-06-01T12:00:00")
        finally:
            engine.dispose()


_STRING_TYPES_REQUIRING_LENGTH = {"String", "VARCHAR", "NVARCHAR", "Unicode"}


def _is_truthy_primary_key(*, call: ast.Call) -> bool:
    for keyword in call.keywords:
        if keyword.arg == "primary_key" and isinstance(keyword.value, ast.Constant) and keyword.value.value is True:
            return True
    return False


def _is_unbounded_string_type(*, node: ast.AST) -> bool:
    if not isinstance(node, ast.Call):
        return False

    func_name = None
    if isinstance(node.func, ast.Attribute):
        func_name = node.func.attr
    elif isinstance(node.func, ast.Name):
        func_name = node.func.id

    if func_name not in _STRING_TYPES_REQUIRING_LENGTH:
        return False

    if node.args:
        return False

    for keyword in node.keywords:
        if keyword.arg == "length":
            return isinstance(keyword.value, ast.Constant) and keyword.value.value is None

    return True


def _find_unbounded_string_pk_columns(*, migration_path: Path) -> list[str]:
    tree = ast.parse(migration_path.read_text(encoding="utf-8"), filename=str(migration_path))
    violations: list[str] = []

    for node in ast.walk(tree):
        if not isinstance(node, ast.Call):
            continue

        func_name = None
        if isinstance(node.func, ast.Attribute):
            func_name = node.func.attr
        elif isinstance(node.func, ast.Name):
            func_name = node.func.id

        if func_name != "Column" or not _is_truthy_primary_key(call=node):
            continue

        column_name = "<unknown>"
        if node.args and isinstance(node.args[0], ast.Constant) and isinstance(node.args[0].value, str):
            column_name = node.args[0].value

        type_node = node.args[1] if len(node.args) >= 2 else None
        if type_node is None:
            for keyword in node.keywords:
                if keyword.arg in {"type_", "type"}:
                    type_node = keyword.value
                    break

        if type_node is not None and _is_unbounded_string_type(node=type_node):
            violations.append(f"{migration_path.name}:{node.lineno} column={column_name}")

    return violations


def test_migrations_do_not_use_unbounded_string_primary_keys() -> None:
    """
    Guard against MSSQL-incompatible primary keys.

    ``sa.String()`` without length can map to ``VARCHAR(MAX)/NVARCHAR(MAX)``,
    which SQL Server rejects for key/index columns.
    """
    versions_dir = Path(__file__).resolve().parent.parent.parent.parent / "pyrit" / "memory" / "alembic" / "versions"

    violations: list[str] = []
    for migration_path in sorted(versions_dir.glob("*.py")):
        if migration_path.name == "__init__.py":
            continue
        violations.extend(_find_unbounded_string_pk_columns(migration_path=migration_path))

    assert not violations, "Found unbounded string primary keys in migrations (SQL Server incompatible):\n" + "\n".join(
        violations
    )
