# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
"""Tests for attention backend selectors."""

from types import SimpleNamespace
from unittest.mock import MagicMock, patch

import pytest
import torch

from vllm.platforms import current_platform
from vllm.v1.attention.backends.registry import AttentionBackendEnum
from vllm.v1.attention.selector import AttentionSelectorConfig

# ROCm-specific attention backend selection tests
pytestmark = pytest.mark.skipif(
    not current_platform.is_rocm(), reason="ROCm-specific tests"
)


@pytest.fixture
def mock_vllm_config():
    """Create a mock VllmConfig for testing."""
    config = MagicMock()
    config.model_config.dtype = torch.float16
    config.model_config.hf_config.architectures = ["LlamaForCausalLM"]
    config.cache_config.block_size = 16
    return config


@pytest.fixture
def mock_get_cdna_version():
    """Mock cdna version arch detection to return True."""
    with patch("vllm.platforms.rocm.get_cdna_version", return_value=3):
        yield


@pytest.fixture
def cleared_attention_selector_cache():
    from vllm.v1.attention.selector import _cached_get_attn_backend

    _cached_get_attn_backend.cache_clear()
    yield
    _cached_get_attn_backend.cache_clear()


def test_aiter_unified_attention_uses_dedicated_metadata_builder():
    from vllm.v1.attention.backends.rocm_aiter_unified_attn import (
        RocmAiterUnifiedAttentionBackend,
        RocmAiterUnifiedAttentionMetadataBuilder,
    )
    from vllm.v1.attention.backends.rocm_attn import (
        RocmAttentionBackend,
        RocmAttentionMetadataBuilder,
    )

    assert RocmAttentionBackend.get_builder_cls() is RocmAttentionMetadataBuilder
    assert (
        RocmAiterUnifiedAttentionBackend.get_builder_cls()
        is RocmAiterUnifiedAttentionMetadataBuilder
    )


def test_aiter_unified_attention_capture_preserves_query_start_locations():
    from vllm.v1.attention.backends.rocm_aiter_unified_attn import (
        RocmAiterUnifiedAttentionMetadataBuilder,
    )

    builder = object.__new__(RocmAiterUnifiedAttentionMetadataBuilder)
    metadata = MagicMock()
    metadata.seq_lens = torch.tensor([1048576, 524288], dtype=torch.int32)
    builder.build = MagicMock(return_value=metadata)
    common = MagicMock()
    expected_query_start_loc = torch.tensor([0, 2, 5], dtype=torch.int32)
    common.query_start_loc = expected_query_start_loc.clone()
    metadata.query_start_loc = common.query_start_loc

    actual = builder.build_for_cudagraph_capture(common)

    builder.build.assert_called_once_with(0, common)
    assert actual is metadata
    assert torch.equal(actual.seq_lens, torch.ones_like(actual.seq_lens))
    assert actual.query_start_loc is common.query_start_loc
    assert torch.equal(actual.query_start_loc, expected_query_start_loc)


@pytest.mark.parametrize("use_dcp", [False, True])
@pytest.mark.parametrize(
    "env_vars, selected_backend, expected_backend_path",
    [
        # Test Case: Explicit FLEX_ATTENTION backend
        (
            {},
            "FLEX_ATTENTION",
            AttentionBackendEnum.FLEX_ATTENTION.get_path(),
        ),
        # Test Case 1: Default (no env vars, no explicit backend)
        (
            {},
            None,
            AttentionBackendEnum.ROCM_ATTN.get_path(),
        ),
        # Test Case 2: Explicit TRITON_ATTN backend
        (
            {},
            "TRITON_ATTN",
            AttentionBackendEnum.TRITON_ATTN.get_path(),
        ),
        # Test Case 3: Explicit ROCM_ATTN backend
        (
            {},
            "ROCM_ATTN",
            AttentionBackendEnum.ROCM_ATTN.get_path(),
        ),
        # Test Case 4: Explicit ROCM_AITER_FA backend
        (
            {},
            "ROCM_AITER_FA",
            AttentionBackendEnum.ROCM_AITER_FA.get_path(),
        ),
        # Test Case 5: Explicit ROCM_AITER_UNIFIED_ATTN backend
        (
            {},
            "ROCM_AITER_UNIFIED_ATTN",
            AttentionBackendEnum.ROCM_AITER_UNIFIED_ATTN.get_path(),
        ),
        # Test Case 6: VLLM_ROCM_USE_AITER=1
        (
            {"VLLM_ROCM_USE_AITER": "1"},
            None,
            AttentionBackendEnum.ROCM_ATTN.get_path(),
        ),
        # Test Case 7: VLLM_ROCM_USE_AITER=1 + explicit TRITON_ATTN
        (
            {"VLLM_ROCM_USE_AITER": "1"},
            "TRITON_ATTN",
            AttentionBackendEnum.TRITON_ATTN.get_path(),
        ),
        # Test Case 8: VLLM_ROCM_USE_AITER=1 + VLLM_ROCM_USE_AITER_MHA=0
        (
            {"VLLM_ROCM_USE_AITER": "1", "VLLM_ROCM_USE_AITER_MHA": "0"},
            None,
            AttentionBackendEnum.ROCM_ATTN.get_path(),
        ),
        # Test Case 9: VLLM_ROCM_USE_AITER=1 + explicit ROCM_ATTN
        (
            {"VLLM_ROCM_USE_AITER": "1"},
            "ROCM_ATTN",
            AttentionBackendEnum.ROCM_ATTN.get_path(),
        ),
    ],
)
def test_standard_attention_backend_selection(
    env_vars,
    selected_backend,
    expected_backend_path,
    use_dcp,
    mock_vllm_config,
    mock_get_cdna_version,
    monkeypatch,
):
    """Standard ROCm backends remain selectable without DCP and reject DCP."""
    # Set environment variables
    for key, value in env_vars.items():
        monkeypatch.setenv(key, value)

    # Import after setting env vars to ensure they're picked up
    # Reload envs to pick up new environment variables
    import importlib

    import vllm.envs as envs

    importlib.reload(envs)

    # Convert string backend to enum if provided
    backend_enum = None
    if selected_backend:
        backend_enum = getattr(AttentionBackendEnum, selected_backend)

    # Get the backend class path
    from vllm.platforms.rocm import RocmPlatform

    attn_selector_config = AttentionSelectorConfig(
        head_size=128,
        dtype=torch.float16,
        kv_cache_dtype="auto",
        block_size=16,
        use_mla=False,
        has_sink=False,
        use_sparse=False,
        use_dcp=use_dcp,
    )

    if use_dcp:
        with pytest.raises(ValueError, match="DCP not supported"):
            RocmPlatform.get_attn_backend_cls(backend_enum, attn_selector_config)
        return

    backend_path = RocmPlatform.get_attn_backend_cls(
        selected_backend=backend_enum, attn_selector_config=attn_selector_config
    )

    assert backend_path == expected_backend_path


@pytest.mark.parametrize("use_dcp", [False, True])
@pytest.mark.parametrize(
    "env_vars, selected_backend, block_size, expected_backend_path, should_raise",
    [
        # Test Case 1: TRITON_MLA with block_size != 1
        (
            {},
            "TRITON_MLA",
            16,
            AttentionBackendEnum.TRITON_MLA.get_path(),
            False,
        ),
        # Test Case 2: TRITON_MLA with block_size == 1 (should raise)
        (
            {},
            "TRITON_MLA",
            1,
            None,
            True,
        ),
        # Test Case 3: ROCM_AITER_MLA with block_size == 1
        (
            {},
            "ROCM_AITER_MLA",
            1,
            AttentionBackendEnum.ROCM_AITER_MLA.get_path(),
            False,
        ),
        # Test Case 4: ROCM_AITER_MLA with block_size != 1 (should raise)
        (
            {},
            "ROCM_AITER_MLA",
            16,
            AttentionBackendEnum.ROCM_AITER_MLA.get_path(),
            False,
        ),
        # Test Case 5: VLLM_ROCM_USE_AITER=1 with block_size == 1
        (
            {"VLLM_ROCM_USE_AITER": "1"},
            None,
            1,
            AttentionBackendEnum.ROCM_AITER_MLA.get_path(),
            False,
        ),
        # Test Case 6: VLLM_ROCM_USE_AITER=1 with block_size == 16
        # (should use ROCM_AITER_MLA now, as it supports block_size 16)
        (
            {"VLLM_ROCM_USE_AITER": "1"},
            None,
            16,
            AttentionBackendEnum.ROCM_AITER_MLA.get_path(),
            False,
        ),
        # Test Case 7: VLLM_ROCM_USE_AITER=1 + explicit TRITON_MLA
        (
            {"VLLM_ROCM_USE_AITER": "1"},
            "TRITON_MLA",
            16,
            AttentionBackendEnum.TRITON_MLA.get_path(),
            False,
        ),
        # Test Case 8: Explicit ROCM_AITER_TRITON_MLA
        (
            {},
            "ROCM_AITER_TRITON_MLA",
            16,
            AttentionBackendEnum.ROCM_AITER_TRITON_MLA.get_path(),
            False,
        ),
    ],
)
def test_mla_backend_selection(
    env_vars,
    selected_backend,
    block_size,
    expected_backend_path,
    should_raise,
    use_dcp,
    mock_vllm_config,
    monkeypatch,
):
    """Dense MLA remains selectable with DCP for valid block sizes."""
    # Set environment variables
    for key, value in env_vars.items():
        monkeypatch.setenv(key, value)

    # Import after setting env vars
    # Reload envs
    import importlib

    import vllm.envs as envs

    importlib.reload(envs)

    # Mock is_aiter_mla_enabled based on env vars and block_size
    aiter_enabled = env_vars.get("VLLM_ROCM_USE_AITER") == "1"

    mock_rocm_ops = MagicMock()
    mock_rocm_ops.is_mla_enabled.return_value = aiter_enabled
    mock_aiter_module = MagicMock()
    mock_aiter_module.rocm_aiter_ops = mock_rocm_ops

    with patch.dict("sys.modules", {"vllm._aiter_ops": mock_aiter_module}):
        # Convert string backend to enum if provided
        backend_enum = None
        if selected_backend:
            backend_enum = getattr(AttentionBackendEnum, selected_backend)

        from vllm.platforms.rocm import RocmPlatform

        if should_raise:
            with pytest.raises(ValueError):
                attn_selector_config = AttentionSelectorConfig(
                    head_size=128,
                    dtype=torch.float16,
                    kv_cache_dtype="auto",
                    block_size=block_size,
                    use_mla=True,
                    has_sink=False,
                    use_sparse=False,
                    use_dcp=use_dcp,
                )
                attn_selector_config = AttentionSelectorConfig(
                    head_size=128,
                    dtype=torch.float16,
                    kv_cache_dtype="auto",
                    block_size=block_size,
                    use_mla=True,
                    has_sink=False,
                    use_sparse=False,
                    use_dcp=use_dcp,
                )
                backend_path = RocmPlatform.get_attn_backend_cls(
                    selected_backend=backend_enum,
                    attn_selector_config=attn_selector_config,
                )

        else:
            attn_selector_config = AttentionSelectorConfig(
                head_size=128,
                dtype=torch.float16,
                kv_cache_dtype="auto",
                block_size=block_size,
                use_mla=True,
                has_sink=False,
                use_sparse=False,
                use_dcp=use_dcp,
            )

            backend_path = RocmPlatform.get_attn_backend_cls(
                selected_backend=backend_enum, attn_selector_config=attn_selector_config
            )

            assert backend_path == expected_backend_path


@pytest.mark.parametrize("use_dcp", [False, True])
@pytest.mark.parametrize(
    "selected_backend", [None, AttentionBackendEnum.ROCM_AITER_MLA_SPARSE]
)
def test_sparse_mla_backend_rejects_dcp(selected_backend, use_dcp):
    """Sparse MLA remains selectable without DCP and fails early with DCP."""
    from vllm.platforms.rocm import RocmPlatform

    selector_config = AttentionSelectorConfig(
        head_size=576,
        dtype=torch.bfloat16,
        kv_cache_dtype="auto",
        block_size=16,
        use_mla=True,
        use_sparse=True,
        use_dcp=use_dcp,
    )
    if use_dcp:
        with pytest.raises(ValueError, match="DCP not supported"):
            RocmPlatform.get_attn_backend_cls(selected_backend, selector_config)
    else:
        assert RocmPlatform.get_attn_backend_cls(selected_backend, selector_config) == (
            AttentionBackendEnum.ROCM_AITER_MLA_SPARSE.get_path()
        )


@pytest.mark.parametrize("use_dcp", [False, True])
@pytest.mark.parametrize(
    "selected_backend, head_size, kv_cache_dtype",
    [
        (AttentionBackendEnum.TRITON_ATTN_DIFFKV, 192, "bfloat16"),
        # A 128-dimensional head uses 118 packed bytes, or 59 fp16 elements.
        (AttentionBackendEnum.TURBOQUANT, 59, "turboquant_k3v4_nc"),
    ],
)
def test_specialized_attention_backends_reject_dcp(
    selected_backend, head_size, kv_cache_dtype, use_dcp
):
    """Valid DiffKV and compressed-cache configurations must reject DCP."""
    from vllm.platforms.rocm import RocmPlatform

    selector_config = AttentionSelectorConfig(
        head_size=head_size,
        dtype=torch.bfloat16,
        kv_cache_dtype=kv_cache_dtype,
        block_size=16,
        use_dcp=use_dcp,
    )
    if use_dcp:
        with pytest.raises(ValueError, match="DCP not supported"):
            RocmPlatform.get_attn_backend_cls(selected_backend, selector_config)
    else:
        assert RocmPlatform.get_attn_backend_cls(selected_backend, selector_config) == (
            selected_backend.get_path()
        )


def test_aiter_fa_requires_mi3xx(mock_vllm_config):
    """Test that ROCM_AITER_FA requires CDNA3+ architecture."""
    from vllm.platforms.rocm import RocmPlatform

    # Mock cdna version to return 1 (used by supports_compute_capability)
    with (
        patch("vllm.platforms.rocm.get_cdna_version", return_value=1),
        pytest.raises(
            ValueError,
            match="compute capability not supported",
        ),
    ):
        attn_selector_config = AttentionSelectorConfig(
            head_size=128,
            dtype=torch.float16,
            kv_cache_dtype="auto",
            block_size=16,
            use_mla=False,
            has_sink=False,
            use_sparse=False,
        )

        RocmPlatform.get_attn_backend_cls(
            selected_backend=AttentionBackendEnum.ROCM_AITER_FA,
            attn_selector_config=attn_selector_config,
        )


@pytest.fixture
def turboquant_run_config():
    """Current config of a run whose KV cache dtype is a turboquant_* preset."""
    config = SimpleNamespace(
        cache_config=SimpleNamespace(cache_dtype="turboquant_k8v4")
    )
    with patch(
        "vllm.config.get_current_vllm_config_or_none",
        return_value=config,
    ):
        yield config


def test_turboquant_boundary_selection_is_not_cached_from_ordinary_run(
    cleared_attention_selector_cache,
):
    from vllm.config import CacheConfig, VllmConfig, set_current_vllm_config
    from vllm.v1.attention.backends.utils import get_supported_kv_cache_layouts
    from vllm.v1.attention.selector import get_attn_backend
    from vllm.v1.kv_cache_layout import KVCacheLayout

    ordinary_config = VllmConfig(cache_config=CacheConfig(cache_dtype="auto"))
    turboquant_config = VllmConfig(
        cache_config=CacheConfig(cache_dtype="turboquant_k8v4")
    )

    backend_priorities = [
        AttentionBackendEnum.ROCM_ATTN,
        AttentionBackendEnum.TRITON_ATTN,
        AttentionBackendEnum.TURBOQUANT,
    ]
    with patch(
        "vllm.platforms.rocm._get_backend_priorities",
        return_value=backend_priorities,
    ):
        with set_current_vllm_config(ordinary_config):
            ordinary_backend = get_attn_backend(128, torch.float16, "auto")

        with set_current_vllm_config(turboquant_config):
            boundary_backend = get_attn_backend(128, torch.float16, "auto")
            quantized_backend = get_attn_backend(128, torch.float16, "turboquant_k8v4")

    assert ordinary_backend is AttentionBackendEnum.ROCM_ATTN.get_class()
    assert quantized_backend is AttentionBackendEnum.TURBOQUANT.get_class()
    assert get_supported_kv_cache_layouts([boundary_backend, quantized_backend]) == [
        KVCacheLayout.LBNHC
    ]


def test_turboquant_run_does_not_mask_unsupported_selected_backend(
    turboquant_run_config,
):
    from vllm.platforms.rocm import RocmPlatform

    attn_selector_config = AttentionSelectorConfig(
        head_size=128,
        dtype=torch.float16,
        kv_cache_dtype="auto",
        block_size=16,
    )

    with (
        patch("vllm.platforms.rocm.get_cdna_version", return_value=1),
        pytest.raises(ValueError, match="compute capability not supported"),
    ):
        RocmPlatform.get_attn_backend_cls(
            selected_backend=AttentionBackendEnum.ROCM_AITER_FA,
            attn_selector_config=attn_selector_config,
        )


def test_turboquant_layout_check_respects_backend_override():
    from vllm.platforms.rocm import _shares_layout_with_turboquant
    from vllm.v1.kv_cache_layout import KVCacheLayout

    boundary_backend = MagicMock()
    boundary_backend.supported_kv_cache_layouts.return_value = (KVCacheLayout.LBHNC,)
    turboquant_override = MagicMock()
    turboquant_override.supported_kv_cache_layouts.return_value = (KVCacheLayout.LBHNC,)

    with patch.object(
        AttentionBackendEnum.TURBOQUANT,
        "get_class",
        return_value=turboquant_override,
    ) as get_turboquant_class:
        assert _shares_layout_with_turboquant(boundary_backend)
    get_turboquant_class.assert_called_once_with()


@pytest.mark.parametrize(
    "selected_backend",
    [None, "ROCM_ATTN", "ROCM_AITER_FA", "TURBOQUANT"],
)
def test_turboquant_boundary_layers_share_a_layout(
    selected_backend,
    turboquant_run_config,
    mock_get_cdna_version,
):
    """A turboquant_* run must resolve to backends with a common KV layout.

    The boundary layers keep the native dtype and pick their own backend while
    every other layer picks TURBOQUANT. Engine startup hard-errors when the two
    share no layout, so the boundary layers must not land on a native ROCm or
    AITER backend, whichever backend the run asked for.
    """
    from vllm.platforms.rocm import RocmPlatform
    from vllm.utils.import_utils import resolve_obj_by_qualname
    from vllm.v1.attention.backends.utils import get_supported_kv_cache_layouts

    backend_enum = (
        getattr(AttentionBackendEnum, selected_backend) if selected_backend else None
    )

    def resolve(kv_cache_dtype):
        path = RocmPlatform.get_attn_backend_cls(
            selected_backend=backend_enum,
            attn_selector_config=AttentionSelectorConfig(
                head_size=128,
                dtype=torch.float16,
                kv_cache_dtype=kv_cache_dtype,
                block_size=16,
            ),
        )
        return resolve_obj_by_qualname(path)

    boundary_backend = resolve("auto")
    quantized_backend = resolve("turboquant_k8v4")

    assert quantized_backend is AttentionBackendEnum.TURBOQUANT.get_class()
    # Raises when the intersection is empty, which is the startup failure.
    assert get_supported_kv_cache_layouts([boundary_backend, quantized_backend])


def test_sparse_not_supported(mock_vllm_config):
    """Test that sparse MLA without use_mla flag raises an error."""
    from vllm.platforms.rocm import RocmPlatform

    with pytest.raises(
        ValueError,
        match="No valid attention backend found",
    ):
        attn_selector_config = AttentionSelectorConfig(
            head_size=128,
            dtype=torch.float16,
            kv_cache_dtype="auto",
            block_size=16,
            use_mla=False,
            has_sink=False,
            use_sparse=True,
        )

        RocmPlatform.get_attn_backend_cls(
            selected_backend=None, attn_selector_config=attn_selector_config
        )


def _kv_connector_selector_config() -> AttentionSelectorConfig:
    return AttentionSelectorConfig(
        head_size=128,
        dtype=torch.float16,
        kv_cache_dtype="auto",
        block_size=16,
        use_mla=False,
        has_sink=False,
        use_sparse=False,
        use_kv_connector=True,
    )


def test_unified_attn_declares_kv_connector_support():
    """ROCM_AITER_UNIFIED_ATTN opts into KV connectors and ROCM_ATTN does not."""
    from vllm.v1.attention.backends.rocm_aiter_unified_attn import (
        RocmAiterUnifiedAttentionBackend,
    )
    from vllm.v1.attention.backends.rocm_attn import RocmAttentionBackend

    assert RocmAiterUnifiedAttentionBackend.supports_kv_connector() is True
    assert RocmAttentionBackend.supports_kv_connector() is False


def test_unified_attn_supports_kv_connector(mock_vllm_config, mock_get_cdna_version):
    """ROCM_AITER_UNIFIED_ATTN can be selected with KV connectors."""
    from vllm.platforms.rocm import RocmPlatform

    backend_path = RocmPlatform.get_attn_backend_cls(
        selected_backend=AttentionBackendEnum.ROCM_AITER_UNIFIED_ATTN,
        attn_selector_config=_kv_connector_selector_config(),
    )

    assert backend_path == AttentionBackendEnum.ROCM_AITER_UNIFIED_ATTN.get_path()


def test_rocm_attn_rejects_kv_connector(mock_vllm_config, mock_get_cdna_version):
    """Selecting ROCM_ATTN with a KV connector is illegal."""
    from vllm.platforms.rocm import RocmPlatform

    attn_selector_config = _kv_connector_selector_config()

    with pytest.raises(ValueError, match="KV connector not supported"):
        RocmPlatform.get_attn_backend_cls(
            selected_backend=AttentionBackendEnum.ROCM_ATTN,
            attn_selector_config=attn_selector_config,
        )


@pytest.mark.parametrize(
    "aiter_found, expected_backend",
    [
        (True, AttentionBackendEnum.ROCM_AITER_UNIFIED_ATTN),
        (False, AttentionBackendEnum.TRITON_ATTN),
    ],
)
def test_auto_selection_for_kv_connector(
    aiter_found, expected_backend, mock_vllm_config, mock_get_cdna_version
):
    """Auto-selection with a KV connector and AITER enabled resolves to unified attn,
    and to triton attn if AITER not enabled."""
    from vllm.platforms.rocm import RocmPlatform

    with patch(
        "vllm._aiter_ops.is_aiter_found_and_supported", return_value=aiter_found
    ):
        backend_path = RocmPlatform.get_attn_backend_cls(
            selected_backend=None,
            attn_selector_config=_kv_connector_selector_config(),
        )

    assert backend_path == expected_backend.get_path()


def test_unified_attn_prefers_block_contiguous_layout():
    """Unified attn prefers a block-first KV layout, hence ok with kv connectors."""
    from vllm.v1.attention.backends.rocm_aiter_unified_attn import (
        RocmAiterUnifiedAttentionBackend,
    )
    from vllm.v1.attention.backends.rocm_attn import RocmAttentionBackend

    unified_preferred = RocmAiterUnifiedAttentionBackend.supported_kv_cache_layouts()[0]
    rocm_attn_preferred = RocmAttentionBackend.supported_kv_cache_layouts()[0]

    assert unified_preferred.is_block_contiguous is True
    assert rocm_attn_preferred.is_block_contiguous is False


def test_unified_attn_drops_lhbnc_with_kv_connector():
    """Connectors move a block as one contiguous byte range, which LHBNC breaks."""
    from vllm.config import KVTransferConfig, VllmConfig, set_current_vllm_config
    from vllm.v1.attention.backends.rocm_aiter_unified_attn import (
        RocmAiterUnifiedAttentionBackend,
    )
    from vllm.v1.kv_cache_interface import KVCacheLayout

    assert KVCacheLayout.LHBNC in (
        RocmAiterUnifiedAttentionBackend.supported_kv_cache_layouts()
    )

    config = VllmConfig(
        kv_transfer_config=KVTransferConfig(
            kv_connector="ExampleConnector", kv_role="kv_both"
        )
    )
    with set_current_vllm_config(config):
        layouts = RocmAiterUnifiedAttentionBackend.supported_kv_cache_layouts()

    assert layouts
    assert all(layout.is_block_compact for layout in layouts)
