# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
"""Unit tests for NixlConnectorScheduler with HMA and Mamba N-1 prefill."""

import gc
from unittest.mock import patch

import pytest
import torch

from tests.v1.attention.utils import MockMambaBuilder
from vllm import LLM, SamplingParams
from vllm.config import KVTransferConfig
from vllm.v1.core.single_type_kv_cache_manager import (
    FullAttentionManager,
    SlidingWindowManager,
)

from .utils import (
    create_request,
    create_vllm_config,
    make_kv_cache_config,
    make_nixl_scheduler,
)


@pytest.mark.cpu_test
@pytest.mark.parametrize(
    "swa_enabled,expected_sw_sizes",
    [
        # SWA enabled: FullAttentionSpec (0) + SlidingWindowSpec (2048/16=128)
        (True, [0, 128 + 1]),
        # SWA disabled: only FullAttentionSpec (0)
        (False, [0]),
    ],
)
@patch(
    "vllm.distributed.kv_transfer.kv_connector.v1.nixl.base_scheduler.current_platform"
)
def test_sw_sizes(mock_platform, swa_enabled, expected_sw_sizes):
    """Test sw_sizes is correctly computed based on SWA enabled/disabled."""
    from vllm.distributed.kv_transfer.kv_connector.v1.nixl.scheduler import (
        NixlConnectorScheduler,
    )

    mock_platform.device_type = "cpu"

    block_size = 16
    vllm_config = create_vllm_config(block_size=block_size)
    # SW 2048 tokens=>128 blocks
    kv_cache_config = make_kv_cache_config(
        block_size=block_size, swa_enabled=swa_enabled, sw_size=2048
    )

    scheduler = NixlConnectorScheduler(
        vllm_config=vllm_config,
        engine_id="test-engine",
        kv_cache_config=kv_cache_config,
    )
    # in number of blocks
    assert scheduler.blocks_per_sw == expected_sw_sizes, (
        f"Expected sw_sizes={expected_sw_sizes}, got {scheduler.blocks_per_sw}"
    )


@pytest.mark.cpu_test
def test_logical_to_kernel_block_ids_with_hma():
    """Test _logical_to_kernel_block_ids expands blocks when HMA is enabled.

    When HMA is enabled, the logical block size may differ from the kernel
    block size. Each logical block maps to multiple kernel blocks.
    """
    from vllm.distributed.kv_transfer.kv_connector.v1.nixl.worker import (
        NixlConnectorWorker,
    )

    # Create a mock worker with just the required attributes
    # (use __new__ to skip __init__)
    worker = object.__new__(NixlConnectorWorker)

    # Simulate HMA scenario: logical block size = 32, kernel block size = 16
    # So each logical block maps to 2 kernel blocks eg [0]->[0,1]
    worker._physical_blocks_per_logical_kv_block = 2
    # FA + SW groups (neither is MambaSpec, so both get expanded)
    worker.kv_cache_config = make_kv_cache_config(block_size=16, swa_enabled=True)

    # Test conversion: FA + SW group
    logical_block_ids = [[0, 1, 2], [3, 4]]
    kernel_block_ids = worker._logical_to_kernel_block_ids(
        logical_block_ids, worker._physical_blocks_per_logical_kv_block
    )

    expected_kernel_block_ids = [[0, 1, 2, 3, 4, 5], [6, 7, 8, 9]]
    assert kernel_block_ids == expected_kernel_block_ids, (
        f"Expected {expected_kernel_block_ids}, got {kernel_block_ids}"
    )


@pytest.mark.cpu_test
@pytest.mark.parametrize(
    "is_rocm,has_mamba,use_host_buffer,done_recving,failed_recving,expected_syncs",
    [
        (True, True, False, {"req"}, set(), 1),
        (False, True, False, {"req"}, set(), 0),
        (True, False, False, {"req"}, set(), 0),
        (True, True, True, {"req"}, set(), 0),
        (True, True, False, set(), set(), 0),
        (True, True, False, {"req"}, {"req"}, 0),
    ],
)
def test_sync_device_after_mamba_recv_gates(
    monkeypatch,
    is_rocm,
    has_mamba,
    use_host_buffer,
    done_recving,
    failed_recving,
    expected_syncs,
):
    """Only direct-GPU Mamba receives on ROCm need a device fence."""
    from vllm.distributed.kv_transfer.kv_connector.v1.nixl import base_worker
    from vllm.distributed.kv_transfer.kv_connector.v1.nixl.worker import (
        NixlConnectorWorker,
    )

    worker = object.__new__(NixlConnectorWorker)
    worker._has_mamba = has_mamba
    worker.use_host_buffer = use_host_buffer

    sync_calls = []
    monkeypatch.setattr(base_worker.current_platform, "is_rocm", lambda: is_rocm)
    monkeypatch.setattr(
        base_worker.torch.accelerator,
        "synchronize",
        lambda: sync_calls.append(True),
    )

    worker._sync_device_after_mamba_recv(done_recving, failed_recving)

    assert len(sync_calls) == expected_syncs


@pytest.mark.cpu_test
@pytest.mark.parametrize(
    "group_spec_types,remote_physical_per_logical,"
    "local_physical_per_logical,tp_ratio,remote_block_ids,"
    "expected_remote_block_ids",
    [
        pytest.param(
            ("FullAttentionSpec", "SlidingWindowSpec"),
            2,
            2,
            1,
            ([0, 1, 2], [3, 4]),
            [[0, 1, 2, 3, 4, 5], [6, 7, 8, 9]],
            id="dense_fa_swa",
        ),
        # Nemotron-3-Nano-30B-A3B 4p1d (P_TP=4, D_TP=1):
        # remote_physical_per_logical=34, local_physical_per_logical=66.
        # FA logical block 5 → kernel [170..203], block 6 → [204..237].
        # Mamba block unchanged.
        pytest.param(
            ("FullAttentionSpec", "MambaSpec"),
            34,
            66,
            -4,
            ([5, 6], [2]),
            [list(range(170, 238)), [2]],
            id="mamba_fa_ssm",
        ),
    ],
)
def test_read_blocks_for_req_expands_remote_ids(
    group_spec_types,
    remote_physical_per_logical,
    local_physical_per_logical,
    tp_ratio,
    remote_block_ids,
    expected_remote_block_ids,
):
    """_read_blocks_for_req must expand remote logical block IDs to kernel
    block IDs when kernel block size != logical block size.

    The hot path always calls _logical_to_kernel_block_ids with
    remote_info.remote_physical_blocks_per_logical (model-agnostic).
    """
    from unittest.mock import MagicMock

    from vllm.distributed.kv_transfer.kv_connector.v1.nixl.metadata import (
        NixlConnectorMetadata,
    )
    from vllm.distributed.kv_transfer.kv_connector.v1.nixl.tp_mapping import (
        TPMapping,
    )
    from vllm.distributed.kv_transfer.kv_connector.v1.nixl.worker import (
        NixlConnectorWorker,
    )
    from vllm.v1.kv_cache_interface import (
        FullAttentionSpec,
        MambaSpec,
        SlidingWindowSpec,
    )

    spec_name_to_type = {
        "FullAttentionSpec": FullAttentionSpec,
        "SlidingWindowSpec": SlidingWindowSpec,
        "MambaSpec": MambaSpec,
    }
    resolved_types = tuple(spec_name_to_type[n] for n in group_spec_types)

    worker = object.__new__(NixlConnectorWorker)
    worker._physical_blocks_per_logical_kv_block = local_physical_per_logical
    worker._engine_last_active = {}
    worker._bidirectional_kv_xfer_enabled = False

    has_mamba = any(t is MambaSpec for t in resolved_types)
    has_swa = any(t is SlidingWindowSpec for t in resolved_types)
    worker.kv_cache_config = make_kv_cache_config(
        block_size=16, swa_enabled=has_swa, mamba_enabled=has_mamba
    )

    remote_engine_id = "remote-engine"

    worker.transfer_topo = MagicMock()
    # tp_ratio not exercised (all_source_ranks is empty so no reads run),
    # but set for realism.
    worker.transfer_topo.tp_ratio.return_value = tp_ratio
    remote_info = MagicMock()
    remote_info.remote_physical_blocks_per_logical = remote_physical_per_logical
    worker.transfer_topo.get_engine_info.return_value = remote_info
    worker.use_mla = False

    mock_plan = MagicMock(spec=TPMapping)
    mock_plan.all_source_ranks = ()
    mock_plan.source_ranks_per_group = ()
    worker.tp_mappings = {remote_engine_id: mock_plan}

    metadata = NixlConnectorMetadata()
    metadata.add_new_req_to_recv(
        request_id="test-req",
        local_block_ids=([0, 1], [2, 3]),
        kv_transfer_params={
            "remote_block_ids": remote_block_ids,
            "remote_engine_id": remote_engine_id,
            "remote_request_id": "prefill-test-req",
            "remote_host": "localhost",
            "remote_port": 1234,
            "tp_size": 1,
        },
    )

    meta = metadata.reqs_to_recv["test-req"]
    worker._read_blocks_for_req("test-req", meta)

    assert meta.remote.block_ids == expected_remote_block_ids, (
        f"Expected {expected_remote_block_ids}, got {meta.remote.block_ids}"
    )


@pytest.mark.cpu_test
@pytest.mark.parametrize(
    "local_physical_per_logical,remote_physical_per_logical,"
    "local_block_ids,remote_block_ids,"
    "expected_local,expected_remote",
    [
        # 10 kernel blocks of data, local has more logical blocks.
        # remote physical_per_logical=10 → 1 logical → 10 kernel blocks
        # local  physical_per_logical=6  → 2 logical → 12 kernel blocks
        # Trim local from 12 to 10.
        pytest.param(
            6,
            10,
            [list(range(12)), [42]],
            [list(range(10)), [42]],
            [list(range(10)), [42]],
            [list(range(10)), [42]],
            id="align_local6_remote10",
        ),
        # 10 kernel blocks of data, remote has more logical blocks.
        # remote physical_per_logical=6  → 2 logical → 12 kernel blocks
        # local  physical_per_logical=10 → 1 logical → 10 kernel blocks
        # Trim remote from 12 to 10.
        pytest.param(
            10,
            6,
            [list(range(10)), [42]],
            [list(range(12)), [42]],
            [list(range(10)), [42]],
            [list(range(10)), [42]],
            id="align_local10_remote6",
        ),
    ],
)
def test_apply_prefix_caching_mamba_hybrid(
    local_physical_per_logical,
    remote_physical_per_logical,
    local_block_ids,
    remote_block_ids,
    expected_local,
    expected_remote,
):
    """_apply_prefix_caching front-trims FA groups to
    min(local, remote) for Mamba hybrid models with heterogeneous TP.
    """
    from vllm.distributed.kv_transfer.kv_connector.v1.nixl.worker import (
        NixlConnectorWorker,
    )
    from vllm.v1.kv_cache_interface import FullAttentionSpec, MambaSpec

    worker = object.__new__(NixlConnectorWorker)
    worker._has_mamba = True
    worker._physical_blocks_per_logical_kv_block = local_physical_per_logical
    worker._group_spec_types = (FullAttentionSpec, MambaSpec)
    worker.kv_cache_config = make_kv_cache_config(block_size=16, mamba_enabled=True)

    aligned_local, aligned_remote = worker._apply_prefix_caching(
        local_block_ids,
        remote_block_ids,
        local_physical_per_logical,
        remote_physical_per_logical,
    )

    assert aligned_local == expected_local, (
        f"Expected local {expected_local}, got {aligned_local}"
    )
    assert aligned_remote == expected_remote, (
        f"Expected remote {expected_remote}, got {aligned_remote}"
    )


@pytest.mark.cpu_test
@pytest.mark.parametrize(
    "local_physical_per_logical,remote_physical_per_logical,"
    "local_block_ids,remote_block_ids,"
    "expected_local,expected_remote",
    [
        # SSM prefix caching: remote has 3 placeholder + 1 real block,
        # local has only the 1 real block. FA blocks are equal (no trim).
        pytest.param(
            10,
            10,
            [list(range(10)), [42]],
            [list(range(10)), [40, 41, 42, 43]],
            [list(range(10)), [42]],
            [list(range(10)), [43]],
            id="ssm_prefix_trim_only",
        ),
        # FA partial prefix cache hit with homogeneous TP: local has 4 FA
        # blocks (prefix cached), remote has full 10. SSM equal (no trim).
        pytest.param(
            10,
            10,
            [list(range(6, 10)), [42]],
            [list(range(10)), [42]],
            [list(range(6, 10)), [42]],
            [list(range(6, 10)), [42]],
            id="fa_prefix_hit_homo_tp",
        ),
        # Both: FA partial prefix hit + SSM placeholder trim.
        # local FA=[6..9] (4 blocks, prefix cached), remote FA=[0..9]
        # local SSM=[99], remote SSM=[10, 20, 99] (2 placeholders + real)
        pytest.param(
            10,
            10,
            [[6, 7, 8, 9], [99]],
            [list(range(10)), [10, 20, 99]],
            [[6, 7, 8, 9], [99]],
            [[6, 7, 8, 9], [99]],
            id="fa_prefix_hit_and_ssm_trim",
        ),
        # Multi-slot SSM ("all" mode): a local prefix hit leaves fewer local
        # slots; the earlier remote slots are covered locally → remote tail.
        pytest.param(
            10,
            10,
            [list(range(10)), [5, 6]],
            [list(range(10)), [1, 2, 3]],
            [list(range(10)), [5, 6]],
            [list(range(10)), [2, 3]],
            id="ssm_multi_block_local_hit_tail",
        ),
        # Multi-slot SSM ("all" mode): the one trailing local position holds
        # the token D recomputes itself → local head-clip.
        pytest.param(
            10,
            10,
            [list(range(10)), [4, 5, 6]],
            [list(range(10)), [8, 9]],
            [list(range(10)), [4, 5]],
            [list(range(10)), [8, 9]],
            id="ssm_multi_block_local_extra_head_clip",
        ),
    ],
)
def test_apply_prefix_caching_ssm_prefix_cache_hit(
    local_physical_per_logical,
    remote_physical_per_logical,
    local_block_ids,
    remote_block_ids,
    expected_local,
    expected_remote,
):
    """_apply_prefix_caching end-trims SSM remote blocks to match the single
    local block (placeholders dropped) and end-trims FA remote blocks on
    partial prefix cache hits when physical_per_logical matches.
    """
    from vllm.distributed.kv_transfer.kv_connector.v1.nixl.worker import (
        NixlConnectorWorker,
    )
    from vllm.v1.kv_cache_interface import FullAttentionSpec, MambaSpec

    worker = object.__new__(NixlConnectorWorker)
    worker._has_mamba = True
    worker._physical_blocks_per_logical_kv_block = local_physical_per_logical
    worker._group_spec_types = (FullAttentionSpec, MambaSpec)
    worker.kv_cache_config = make_kv_cache_config(block_size=16, mamba_enabled=True)

    aligned_local, aligned_remote = worker._apply_prefix_caching(
        local_block_ids,
        remote_block_ids,
        local_physical_per_logical,
        remote_physical_per_logical,
    )

    assert aligned_local == expected_local, (
        f"Expected local {expected_local}, got {aligned_local}"
    )
    assert aligned_remote == expected_remote, (
        f"Expected remote {expected_remote}, got {aligned_remote}"
    )


@pytest.mark.cpu_test
def test_apply_prefix_caching_ssm_unpairable_slots_rejected():
    """Local SSM slots can only exceed the remote ones by the position D
    recomputes itself. A larger excess means the lists aren't
    position-aligned: fail loudly rather than transfer into wrong slots."""
    from vllm.distributed.kv_transfer.kv_connector.v1.nixl.worker import (
        NixlConnectorWorker,
    )
    from vllm.v1.kv_cache_interface import FullAttentionSpec, MambaSpec

    worker = object.__new__(NixlConnectorWorker)
    worker._has_mamba = True
    worker._physical_blocks_per_logical_kv_block = 10
    worker._group_spec_types = (FullAttentionSpec, MambaSpec)
    worker.kv_cache_config = make_kv_cache_config(block_size=16, mamba_enabled=True)

    with pytest.raises(AssertionError, match="unpairable SSM state slots"):
        worker._apply_prefix_caching(
            [list(range(10)), [4, 5, 6, 7]], [list(range(10)), [8, 9]], 10, 10
        )


@pytest.mark.cpu_test
@pytest.mark.parametrize(
    "local_physical_per_logical,remote_physical_per_logical,"
    "remote_fa_blocks,local_fa_blocks,ssm_blocks,"
    "correct_remote_fa,correct_local_fa",
    [
        # 10 kernel blocks of data (640 tokens).
        # remote physical_per_logical=10 → 1 logical → 10 kernel [0..9]
        # local  physical_per_logical=6  → 2 logical → 12 kernel [0..11]
        # 1st local logical block cached → suffix [6..11]
        # Correct: transfer only uncached suffix tokens (384-639)
        #   = remote [6,7,8,9] → local [6,7,8,9].
        # Actual (front-trim): remote[:6]=[0..5] → local [6..11]. Wrong.
        pytest.param(
            6,
            10,
            [0, 1, 2, 3, 4, 5, 6, 7, 8, 9],
            [6, 7, 8, 9, 10, 11],
            [42],
            [6, 7, 8, 9],
            [6, 7, 8, 9],
            id="local6_remote10_fail",
        ),
        # 15 kernel blocks of data (960 tokens).
        # remote physical_per_logical=6  → 3 logical → 18 kernel [0..17]
        # local  physical_per_logical=10 → 2 logical → 20 kernel [0..19]
        # 1st local logical block cached → suffix [10..19]
        # Correct: transfer only uncached suffix tokens (640-959)
        #   = remote [10,11,12,13,14] → local [10,11,12,13,14].
        # Actual (front-trim): remote[:10]=[0..9] → local [10..19]. Wrong.
        pytest.param(
            10,
            6,
            list(range(18)),
            list(range(10, 20)),
            [42],
            [10, 11, 12, 13, 14],
            [10, 11, 12, 13, 14],
            id="local10_remote6_fail",
        ),
    ],
)
def test_mismatched_physical_per_logical_fails_with_prefix_caching(
    local_physical_per_logical,
    remote_physical_per_logical,
    remote_fa_blocks,
    local_fa_blocks,
    ssm_blocks,
    correct_remote_fa,
    correct_local_fa,
):
    """Demonstrate that _apply_prefix_caching front-trims ([:N])
    in the Mamba hybrid path, which fails when prefix caching produces
    suffix-only local blocks.

    Prefix caching operates at logical block granularity. When a logical
    block is cached locally, the decode side only allocates kernel blocks
    for the uncached suffix. The front-trim pairs remote prefix blocks
    with local suffix slots — a silent data corruption.
    """
    from vllm.distributed.kv_transfer.kv_connector.v1.nixl.worker import (
        NixlConnectorWorker,
    )

    worker = object.__new__(NixlConnectorWorker)
    worker._physical_blocks_per_logical_kv_block = local_physical_per_logical
    worker.kv_cache_config = make_kv_cache_config(
        block_size=16,
        mamba_enabled=True,
    )
    worker._has_mamba = True
    worker._group_spec_types = tuple(
        type(g.kv_cache_spec) for g in worker.kv_cache_config.kv_cache_groups
    )

    local_block_ids = (local_fa_blocks, ssm_blocks)
    remote_block_ids = (remote_fa_blocks, ssm_blocks)

    aligned_local, aligned_remote = worker._apply_prefix_caching(
        local_block_ids,
        remote_block_ids,
        local_physical_per_logical,
        remote_physical_per_logical,
    )

    assert (
        aligned_remote[0] != correct_remote_fa or aligned_local[0] != correct_local_fa
    ), (
        f"Prefix caching with mismatched physical_per_logical should not "
        f"produce correct transfer ids: "
        f"remote={aligned_remote[0]}, local={aligned_local[0]}, "
        f"correct_remote={correct_remote_fa}, correct_local={correct_local_fa}"
    )


@pytest.mark.parametrize("model_name, sw_size", [("google/gemma-3-1b-it", 512)])
def test_fewer_blocks_with_hma(monkeypatch, model_name, sw_size):
    """Test that a prefill instance returns fewer "remote blocks" for the SWA groups
    when sequence exceeds the sliding window.
    """
    kv_transfer_config = KVTransferConfig(
        kv_connector="NixlConnector",
        kv_role="kv_consumer",
    )
    block_size = 16
    llm_kwargs = {
        "model": model_name,
        "enforce_eager": True,
        "gpu_memory_utilization": 0.3,
        "kv_transfer_config": kv_transfer_config,
        "max_model_len": 2048,
        "max_num_seqs": 1,
        "max_num_batched_tokens": 2048,
        "enable_prefix_caching": False,
        "block_size": block_size,
    }

    monkeypatch.setenv("VLLM_ENABLE_V1_MULTIPROCESSING", "0")

    def run_hma_test(llm: LLM):
        remote_prefill_opts = {
            "do_remote_decode": True,
            "do_remote_prefill": False,
            "remote_engine_id": None,
            "remote_block_ids": None,
            "remote_host": None,
            "remote_port": None,
        }
        # Simulate sidecar request
        sampling_params = SamplingParams(
            temperature=0.0,
            max_tokens=1,
            extra_args={"kv_transfer_params": remote_prefill_opts},
        )
        scheduler = llm.llm_engine.engine_core.engine_core.scheduler
        kv_managers = scheduler.kv_cache_manager.coordinator.single_type_managers
        # HMA enabled with FA + SWA groups
        assert len(kv_managers) > 2
        for kv_manager in kv_managers:
            assert isinstance(kv_manager, (SlidingWindowManager, FullAttentionManager))
        req_to_blocks = kv_managers[0].req_to_blocks
        assert len(req_to_blocks) == 0

        # Process some request with length exceeding the sliding window
        outputs = llm.generate(["hi" * 1401], sampling_params)
        kv_params = outputs[0].kv_transfer_params

        # +1 to account for overlapping window across blocks.
        expected_num_remote_blocks = sw_size // block_size + 1
        remote_block_ids = kv_params["remote_block_ids"]
        assert (
            len(remote_block_ids[0])
            == expected_num_remote_blocks
            < len(remote_block_ids[-1])
        )
        for group_block_ids in remote_block_ids[:-1]:
            assert len(group_block_ids) == expected_num_remote_blocks

    def run_test_and_cleanup():
        gc.collect()
        torch.accelerator.empty_cache()
        llm = LLM(**llm_kwargs)
        try:
            run_hma_test(llm)
        finally:
            llm.llm_engine.engine_core.shutdown()

    run_test_and_cleanup()


@pytest.mark.cpu_test
def test_nixl_metadata_hma_block_ids_structure():
    """
    Test that NixlConnectorMetadata correctly stores block IDs for multiple
    KV cache groups when HMA is enabled.
    """
    from vllm.distributed.kv_transfer.kv_connector.v1.nixl.metadata import (
        NixlConnectorMetadata,
    )

    metadata = NixlConnectorMetadata()

    # Add request with block IDs for 2 groups (FA + SW)
    fa_blocks = [0, 1, 2, 3, 4, 5, 6, 7]  # 8 blocks for FA
    sw_blocks = [8, 9, 10, 11]  # 4 blocks for SW (clipped)

    metadata.add_new_req_to_recv(
        request_id="test-req-hma",
        local_block_ids=(fa_blocks, sw_blocks),
        kv_transfer_params={
            "remote_block_ids": ([10, 11, 12, 13, 14, 15, 16, 17], [18, 19, 20, 21]),
            "remote_engine_id": "remote-engine",
            "remote_request_id": "prefill-test-req-hma",
            "remote_host": "localhost",
            "remote_port": 1234,
            "tp_size": 1,
        },
    )

    assert "test-req-hma" in metadata.reqs_to_recv
    req_meta = metadata.reqs_to_recv["test-req-hma"]

    # Verify local block IDs structure
    assert len(req_meta.local_block_ids) == 2
    assert list(req_meta.local_block_ids[0]) == fa_blocks
    assert list(req_meta.local_block_ids[1]) == sw_blocks

    # Verify remote block IDs structure
    assert req_meta.remote is not None
    assert len(req_meta.remote.block_ids) == 2
    assert list(req_meta.remote.block_ids[0]) == [10, 11, 12, 13, 14, 15, 16, 17]
    assert list(req_meta.remote.block_ids[1]) == [18, 19, 20, 21]


def _make_mock_worker_for_desc_ids(
    num_regions: int,
    has_mamba: bool,
    group_spec_types: tuple,
    block_len_per_layer: list[int] | None = None,
):
    """Build a mock NixlConnectorWorker with attrs needed by _compute_desc_ids."""
    from unittest.mock import MagicMock

    from vllm.distributed.kv_transfer.kv_connector.v1.nixl.worker import (
        NixlConnectorWorker,
    )

    worker = MagicMock(spec=NixlConnectorWorker)
    worker.num_regions = num_regions
    worker._has_mamba = has_mamba
    worker._group_spec_types = group_spec_types
    worker.block_len_per_layer = block_len_per_layer or [100]
    worker._conv_decomp = None
    if has_mamba:
        from vllm.distributed.kv_transfer.kv_connector.v1.ssm_conv_transfer_utils import (  # noqa: E501
            MambaConvSplitInfo,
        )

        # Mamba2/GDN layout: 3 conv sub-projections -> 4 NIXL regions per layer.
        worker._conv_decomp = MambaConvSplitInfo(
            conv_rows=3,
            local_proj_dims=(1, 1, 1),
            conv_dtype_size=2,
            ssm_sizes=(0, 0),
        )
    worker._compute_desc_ids = NixlConnectorWorker._compute_desc_ids.__get__(
        worker, NixlConnectorWorker
    )
    return worker


@pytest.mark.cpu_test
def test_get_block_descs_ids_hybrid_ssm():
    """Test _compute_desc_ids uses per-group strides for hybrid
    FA+SSM when ratio=1 (no kernel block size mismatch)."""
    from vllm.v1.kv_cache_interface import FullAttentionSpec, MambaSpec

    worker = _make_mock_worker_for_desc_ids(
        num_regions=2,
        has_mamba=True,
        group_spec_types=(FullAttentionSpec, MambaSpec),
        block_len_per_layer=[100],
    )

    fa_blocks = [3, 5]
    ssm_blocks = [1, 2]
    result = worker._compute_desc_ids(
        block_ids=(fa_blocks, ssm_blocks),
        dst_num_blocks=100,
        block_size_ratio=None,
        physical_blocks_per_logical=1,
    )

    expected = [3, 5, 103, 105, 201, 202, 301, 302, 401, 402, 501, 502]
    assert list(result) == expected, f"Expected {expected}, got {list(result)}"


@pytest.mark.cpu_test
def test_get_block_descs_ids_kernel_block_mismatch():
    """Test _compute_desc_ids uses different strides for FA
    (kernel blocks) vs SSM (logical blocks) when ratio > 1."""
    from vllm.v1.kv_cache_interface import FullAttentionSpec, MambaSpec

    ratio = 4
    logical_blocks = 100
    num_blocks = logical_blocks * ratio  # 400 kernel blocks

    worker = _make_mock_worker_for_desc_ids(
        num_regions=2,
        has_mamba=True,
        group_spec_types=(FullAttentionSpec, MambaSpec),
        block_len_per_layer=[100],
    )

    fa_blocks = [3, 7]
    ssm_blocks = [1, 2]
    result = worker._compute_desc_ids(
        block_ids=(fa_blocks, ssm_blocks),
        dst_num_blocks=num_blocks,
        block_size_ratio=None,
        physical_blocks_per_logical=ratio,
    )

    expected = [3, 7, 403, 407, 801, 802, 901, 902, 1001, 1002, 1101, 1102]
    assert list(result) == expected, f"Expected {expected}, got {list(result)}"


@pytest.mark.cpu_test
def test_get_block_descs_ids_hetero_block_size_hybrid():
    """With a block-size ratio, FA desc ids are ratio-expanded while SSM
    desc ids keep the unexpanded logical stride (state blocks are never
    sub-split)."""
    from vllm.v1.kv_cache_interface import FullAttentionSpec, MambaSpec

    worker = _make_mock_worker_for_desc_ids(
        num_regions=2,
        has_mamba=True,
        group_spec_types=(FullAttentionSpec, MambaSpec),
        block_len_per_layer=[100],
    )

    ratio = 4
    # FA ids are already remote-granularity (expanded) sub-block ids.
    fa_sub_blocks = [3, 5]
    ssm_blocks = [1]
    result = worker._compute_desc_ids(
        block_ids=(fa_sub_blocks, ssm_blocks),
        dst_num_blocks=100,
        block_size_ratio=ratio,
        physical_blocks_per_logical=1,
    )

    # FA regions have 100*4 entries each; SSM regions (4 per layer) start at
    # 2*400 and stride by the unexpanded 100 logical blocks.
    expected = [3, 5, 403, 405, 801, 901, 1001, 1101]
    assert list(result) == expected, f"Expected {expected}, got {list(result)}"


def _bind_worker_method(worker, name):
    from vllm.distributed.kv_transfer.kv_connector.v1.nixl.worker import (
        NixlConnectorWorker,
    )

    method = getattr(NixlConnectorWorker, name)
    setattr(worker, name, method.__get__(worker, NixlConnectorWorker))


@pytest.mark.cpu_test
def test_map_block_ids_for_block_size_ratio_hybrid():
    """Attention groups expand to remote granularity and clip to the remote
    coverage; mamba state blocks pass through 1:1."""
    from unittest.mock import MagicMock

    from vllm.distributed.kv_transfer.kv_connector.v1.nixl.worker import (
        NixlConnectorWorker,
    )
    from vllm.v1.kv_cache_interface import FullAttentionSpec, MambaSpec

    worker = MagicMock(spec=NixlConnectorWorker)
    worker._group_spec_types = (FullAttentionSpec, MambaSpec)
    _bind_worker_method(worker, "get_mapped_blocks")
    _bind_worker_method(worker, "_map_block_ids_for_block_size_ratio")

    local, remote = worker._map_block_ids_for_block_size_ratio(
        [[1, 2, 3], [7]],
        [list(range(30, 40)), [42]],
        4,
    )
    # [1, 2, 3] expand to sub-blocks [4..15], clipped to the 10 remote blocks.
    assert local == [list(range(4, 14)), [7]]
    assert remote == [list(range(30, 40)), [42]]

    # Attention-only full prefix hit: empty local list is preserved.
    worker._group_spec_types = (FullAttentionSpec,)
    local, remote = worker._map_block_ids_for_block_size_ratio([[]], [[30, 31]], 4)
    assert local == []


@pytest.mark.cpu_test
def test_post_process_zeroes_untransferred_tail():
    """The untransferred sub-blocks of the last local block are zeroed on
    receive; mamba state caches are untouched by the attention permute."""
    from unittest.mock import MagicMock

    from vllm.distributed.kv_transfer.kv_connector.v1.nixl.worker import (
        NixlConnectorWorker,
    )
    from vllm.v1.kv_cache_interface import FullAttentionSpec, MambaSpec

    ratio = 4
    block_tokens = 8  # 2 tokens per remote sub-block

    worker = MagicMock(spec=NixlConnectorWorker)
    worker._group_spec_types = (FullAttentionSpec, MambaSpec)
    worker.transfer_topo = MagicMock()
    worker.device_type = "cpu"
    worker.enable_permute_local_kv = False
    attn_cache = torch.ones(6, block_tokens, 2, 4)
    mamba_cache = torch.ones(6, 16)
    worker.device_kv_caches = {"attn.0": attn_cache, "mamba.0": mamba_cache}
    fa_group = MagicMock(layer_names=["attn.0"])
    ssm_group = MagicMock(layer_names=["mamba.0"])
    worker.kv_cache_config = MagicMock(kv_cache_groups=[fa_group, ssm_group])
    # The cached property filters mamba layers out of the permuted caches.
    attn_caches = NixlConnectorWorker._attention_kv_caches.func(worker)
    assert len(attn_caches) == 1 and attn_caches[0] is attn_cache
    worker._attention_kv_caches = attn_caches
    _bind_worker_method(worker, "post_process_device_kv_on_receive")

    # Request occupies blocks [2, 3]; only 6 of 8 sub-blocks were received.
    worker.post_process_device_kv_on_receive(ratio, [([2, 3], 6)])

    # Block 2 fully covered; block 3 covered for 2 sub-blocks (4 tokens).
    assert torch.all(attn_cache[2] == 1)
    assert torch.all(attn_cache[3, :4] == 1)
    assert torch.all(attn_cache[3, 4:] == 0)
    # Untouched blocks and the mamba cache keep their content.
    assert torch.all(attn_cache[4] == 1)
    assert torch.all(mamba_cache == 1)


@pytest.mark.cpu_test
def test_nixl_metadata_hybrid_ssm_block_ids():
    """Test NixlConnectorMetadata correctly stores block IDs for FA + SSM
    groups with different block counts (kernel mismatch active)."""
    from vllm.distributed.kv_transfer.kv_connector.v1.nixl.metadata import (
        NixlConnectorMetadata,
    )

    metadata = NixlConnectorMetadata()

    # FA: 8 kernel blocks (2 logical * ratio=4), SSM: 2 logical blocks
    fa_blocks = [0, 1, 2, 3, 4, 5, 6, 7]
    ssm_blocks = [0, 1]

    metadata.add_new_req_to_recv(
        request_id="test-req-hybrid",
        local_block_ids=(fa_blocks, ssm_blocks),
        kv_transfer_params={
            "remote_block_ids": ([10, 11, 12, 13, 14, 15, 16, 17], [20, 21]),
            "remote_engine_id": "remote-engine",
            "remote_request_id": "prefill-test-req-hybrid",
            "remote_host": "localhost",
            "remote_port": 1234,
            "tp_size": 1,
        },
    )

    assert "test-req-hybrid" in metadata.reqs_to_recv
    req_meta = metadata.reqs_to_recv["test-req-hybrid"]

    # Verify local block IDs: different lengths per group
    assert len(req_meta.local_block_ids) == 2
    assert list(req_meta.local_block_ids[0]) == fa_blocks
    assert list(req_meta.local_block_ids[1]) == ssm_blocks
    assert len(req_meta.local_block_ids[0]) != len(req_meta.local_block_ids[1])

    # Verify remote block IDs: same asymmetry preserved
    assert req_meta.remote is not None
    assert len(req_meta.remote.block_ids) == 2
    assert list(req_meta.remote.block_ids[0]) == [10, 11, 12, 13, 14, 15, 16, 17]
    assert list(req_meta.remote.block_ids[1]) == [20, 21]
    assert len(req_meta.remote.block_ids[0]) != len(req_meta.remote.block_ids[1])


class _FakeBlock:
    def __init__(self, block_id):
        self.block_id = block_id


class _FakeSingleTypeManager:
    def __init__(self, records, block_size, block_ids):
        self.records_new_block_ids = records
        self.block_size = block_size
        self.req_to_blocks = {"req-1": [_FakeBlock(b) for b in block_ids]}
        self.new_block_ids: list[int] = []

    def take_new_block_ids(self):
        ids = self.new_block_ids
        self.new_block_ids = []
        return ids


def _make_fake_kv_cache_manager():
    from unittest.mock import MagicMock

    from vllm.v1.core.kv_cache_manager import KVCacheManager

    manager = object.__new__(KVCacheManager)
    manager.coordinator = MagicMock()
    manager.coordinator.single_type_managers = (
        _FakeSingleTypeManager(True, 16, [10, 11, 12, 13, 14, 15]),  # attention
        _FakeSingleTypeManager(False, 16, [20, 21, 22, 23, 24, 25]),  # mamba
    )
    return manager


@pytest.mark.cpu_test
def test_zeroing_block_ids_cover_only_loaded_attention_blocks():
    """Only zero-recorded (attention) groups contribute, sliced to the
    externally-loaded token range; Mamba state blocks are never zeroed."""
    manager = _make_fake_kv_cache_manager()

    # Tokens [0, 16) are locally cached; the load covers tokens [16, 56).
    assert manager.get_zeroing_block_ids_in_range("req-1", 16, 56) == [11, 12, 13]


@pytest.mark.cpu_test
def test_scheduler_filters_connector_loaded_blocks_from_zeroing():
    """Blocks that will be loaded by the connector must not be zeroed."""
    from vllm.v1.core.sched.scheduler import Scheduler

    class FakeKVCacheManager:
        def take_new_block_ids(self):
            return [9, 10, 11, 12]

    scheduler = object.__new__(Scheduler)
    scheduler.needs_kv_cache_zeroing = True
    scheduler.kv_cache_manager = FakeKVCacheManager()
    scheduler._skip_zero_block_ids = {10, 12}

    assert scheduler._get_new_block_ids_to_zero() == [9, 11]
    assert not scheduler._skip_zero_block_ids


@pytest.mark.cpu_test
def test_failed_load_rezeroes_unwritten_skipped_blocks():
    """A failed async load leaves zeroing-skipped blocks unwritten beyond
    the valid prefix; they must be zeroed before local recompute."""
    from unittest.mock import MagicMock

    from vllm.v1.core.sched.scheduler import Scheduler

    scheduler = object.__new__(Scheduler)
    scheduler.connector = MagicMock()
    scheduler.needs_kv_cache_zeroing = True
    scheduler.kv_cache_manager = _make_fake_kv_cache_manager()
    scheduler.kv_cache_manager.cache_blocks = MagicMock()
    scheduler.failed_recving_kv_req_ids = {"req-1"}
    scheduler.finished_recving_kv_req_ids = {"req-1"}

    request = MagicMock()
    request.request_id = "req-1"
    request.num_computed_tokens = 48  # Truncated at the first invalid block.

    scheduler._update_waiting_for_remote_kv(request)

    # Attention blocks covering tokens >= 48 are re-recorded for zeroing
    # and flow into the next step's zero list; Mamba blocks are not.
    scheduler._skip_zero_block_ids = set()
    assert scheduler._get_new_block_ids_to_zero() == [13, 14, 15]


# ── Mamba N-1 prefill tests ──────────────────────────────────────────────


@pytest.mark.cpu_test
@pytest.mark.parametrize(
    "has_mamba,is_hma_required,expected_count",
    [
        (True, True, 9),
        (False, False, 10),
        (False, True, 10),
    ],
    ids=["mamba", "fa_only", "swa_only"],
)
def test_mamba_n1_d_side(has_mamba, is_hma_required, expected_count):
    """D-side: Mamba gets N-1 matched tokens, non-Mamba gets N."""
    sched = make_nixl_scheduler(has_mamba=has_mamba, is_hma_required=is_hma_required)
    req = create_request(num_tokens=10, do_remote_prefill=True)

    count, is_async = sched.get_num_new_matched_tokens(req, num_computed_tokens=0)
    assert count == expected_count
    assert is_async is True


@pytest.mark.cpu_test
def test_mamba_n1_d_side_builds_decode_metadata():
    req = create_request(num_tokens=10, do_remote_prefill=True)
    sched = make_nixl_scheduler(has_mamba=True, is_hma_required=True)

    num_computed_tokens, is_async = sched.get_num_new_matched_tokens(
        req, num_computed_tokens=0
    )

    assert num_computed_tokens == req.num_prompt_tokens - 1
    assert is_async is True

    vllm_config = create_vllm_config()
    metadata = MockMambaBuilder.build_mamba_metadata(
        vllm_config,
        seq_lens=[req.num_prompt_tokens],
        query_lens=[1],
        is_prefilling=[True],
    )

    assert metadata.num_decodes == 1
    assert metadata.num_prefills == 0


@pytest.mark.cpu_test
def test_mamba_n1_p_side_truncation():
    """P-side: Mamba truncates prompt to N-1, sets max_tokens=1.

    Also verifies idempotency (calling again is a no-op) which is
    needed for preemption safety via the _p_side_truncated guard,
    and that non-Mamba models skip truncation entirely.
    """
    sched = make_nixl_scheduler(has_mamba=True, is_hma_required=True)
    req = create_request(num_tokens=10, do_remote_decode=True)
    req.max_tokens = 128
    original_len = len(req.prompt_token_ids)

    count, is_async = sched.get_num_new_matched_tokens(req, num_computed_tokens=0)

    assert count == 0
    assert is_async is False
    assert len(req.prompt_token_ids) == original_len - 1
    assert req.num_prompt_tokens == original_len - 1
    assert req.max_tokens == 1
    assert req.kv_transfer_params["_p_side_truncated"] is True

    # Idempotency: second call must not truncate further
    sched.get_num_new_matched_tokens(req, num_computed_tokens=0)
    assert len(req.prompt_token_ids) == original_len - 1

    # Non-Mamba: truncation is skipped
    fa_sched = make_nixl_scheduler(has_mamba=False, is_hma_required=False)
    fa_req = create_request(num_tokens=10, do_remote_decode=True)
    fa_original = len(fa_req.prompt_token_ids)

    fa_sched.get_num_new_matched_tokens(fa_req, num_computed_tokens=0)
    assert len(fa_req.prompt_token_ids) == fa_original


@pytest.mark.cpu_test
@pytest.mark.parametrize(
    "swa_enabled,mamba_enabled,expected_has_mamba,expected_is_hma",
    [
        (True, True, True, True),
        (True, False, False, True),
        (False, False, False, False),
    ],
    ids=["fa_swa_mamba", "fa_swa_only", "fa_only"],
)
@patch(
    "vllm.distributed.kv_transfer.kv_connector.v1.nixl.base_scheduler.current_platform"
)
def test_has_mamba_init(
    mock_platform,
    swa_enabled,
    mamba_enabled,
    expected_has_mamba,
    expected_is_hma,
):
    """Test _has_mamba / _is_hma_required derived from kv_cache_groups."""
    from vllm.distributed.kv_transfer.kv_connector.v1.nixl.scheduler import (
        NixlConnectorScheduler,
    )

    mock_platform.device_type = "cpu"

    block_size = 16
    vllm_config = create_vllm_config(block_size=block_size)
    # Explicitly enable HMA so we can test the scheduler's own derivation.
    vllm_config.scheduler_config.disable_hybrid_kv_cache_manager = False
    kv_cache_config = make_kv_cache_config(
        block_size=block_size,
        swa_enabled=swa_enabled,
        mamba_enabled=mamba_enabled,
    )

    scheduler = NixlConnectorScheduler(
        vllm_config=vllm_config,
        engine_id="test-engine",
        kv_cache_config=kv_cache_config,
    )
    assert scheduler._has_mamba is expected_has_mamba
    assert scheduler._is_hma_required is expected_is_hma


@pytest.mark.cpu_test
@pytest.mark.parametrize(
    "ssm_sizes,block_len,expected_ratio",
    [
        # Nemotron 30B TP=1: ceil((36864 + 2097152) / 8192) = 261
        ((36864, 2097152), 8192, 261),
        # Nemotron 30B TP=2: ceil((18432 + 1048576) / 4096) = 261
        ((18432, 1048576), 4096, 261),
        # Nemotron 30B TP=4: ceil((9216 + 524288) / 4096) = 131
        ((9216, 524288), 4096, 131),
    ],
)
def test_compute_physical_blocks_per_logical(ssm_sizes, block_len, expected_ratio):
    """Verify that compute_physical_blocks_per_logical is TP-dependent.

    With dimension-sharded Mamba state, the ratio differs across TP sizes
    (e.g. TP=1 → 261, TP=4 → 131 for Nemotron 30B). This is why
    _physical_blocks_per_logical must be stored per-engine.
    """
    from vllm.distributed.kv_transfer.kv_connector.v1.ssm_conv_transfer_utils import (
        compute_physical_blocks_per_logical,
    )

    assert compute_physical_blocks_per_logical(ssm_sizes, block_len) == expected_ratio


@pytest.mark.cpu_test
@pytest.mark.parametrize(
    "mamba_type,local_tp,conv_dim_local,conv_rows,temporal_shape,expected_proj_dims",
    [
        # nvidia/Nemotron-H-8B-Base-8K (Mamba2)
        # mamba_num_heads=128, head_dim=64, n_groups=8, ssm_state_size=128
        pytest.param(
            "mamba2",
            1,
            10240,
            3,
            (128, 64, 128),
            (8192, 1024, 1024),
            id="nemotron_h_8b_tp1",
        ),
        pytest.param(
            "mamba2",
            4,
            2560,
            3,
            (32, 64, 128),
            (2048, 256, 256),
            id="nemotron_h_8b_tp4",
        ),
        # nvidia/NVIDIA-Nemotron-3-Nano-30B-A3B (Mamba2)
        # mamba_num_heads=64, head_dim=64, n_groups=8, ssm_state_size=128
        pytest.param(
            "mamba2",
            1,
            6144,
            3,
            (64, 64, 128),
            (4096, 1024, 1024),
            id="nemotron_nano_30b_tp1",
        ),
        # Qwen/Qwen3.5-0.8B (GDN, symmetric: num_v=num_k=16)
        # key_dim=2048, value_dim=2048, conv_dim=6144
        pytest.param(
            "gdn_attention",
            1,
            6144,
            3,
            (16, 128, 128),
            (2048, 2048, 2048),
            id="qwen35_08b_tp1",
        ),
        pytest.param(
            "gdn_attention",
            4,
            1536,
            3,
            (4, 128, 128),
            (512, 512, 512),
            id="qwen35_08b_tp4",
        ),
        # Qwen/Qwen3.5-4B (GDN, asymmetric: num_v=32, num_k=16, K:V=1:2)
        # key_dim=2048, value_dim=4096, conv_dim=8192
        pytest.param(
            "gdn_attention",
            1,
            8192,
            3,
            (32, 128, 128),
            (2048, 2048, 4096),
            id="qwen35_4b_tp1",
        ),
        # Qwen/Qwen3.5-27B (GDN, asymmetric: num_v=48, num_k=16, K:V=1:3)
        # key_dim=2048, value_dim=6144, conv_dim=10240
        pytest.param(
            "gdn_attention",
            1,
            10240,
            3,
            (48, 128, 128),
            (2048, 2048, 6144),
            id="qwen35_27b_tp1",
        ),
        pytest.param(
            "gdn_attention",
            8,
            1280,
            3,
            (6, 128, 128),
            (256, 256, 768),
            id="qwen35_27b_tp8",
        ),
        # ai21labs/AI21-Jamba2-Mini (Mamba1)
        # mamba d_inner = mamba_expand(2) * hidden_size(4096) = 8192
        # mamba_d_state=16, mamba_d_conv=4 → conv_rows=3.
        # Conv state holds only x: a single contiguous sub-projection.
        pytest.param(
            "mamba1",
            1,
            8192,
            3,
            (8192, 16),
            (8192,),
            id="jamba_mini_tp1",
        ),
        pytest.param(
            "mamba1",
            4,
            2048,
            3,
            (2048, 16),
            (2048,),
            id="jamba_mini_tp4",
        ),
        pytest.param(
            "mamba1",
            8,
            1024,
            3,
            (1024, 16),
            (1024,),
            id="jamba_mini_tp8",
        ),
    ],
)
def test_derive_mamba_conv_split(
    monkeypatch,
    mamba_type,
    local_tp,
    conv_dim_local,
    conv_rows,
    temporal_shape,
    expected_proj_dims,
):
    """Parametrized test for derive_mamba_conv_split with real model configs.

    Values generated by verify_conv_split.py which loads HuggingFace configs
    and calls vLLM's derive_mamba_conv_split directly.
    """
    from vllm.distributed.kv_transfer.kv_connector.v1.ssm_conv_transfer_utils import (
        derive_mamba_conv_split,
    )
    from vllm.v1.attention.backends.registry import MambaAttentionBackendEnum
    from vllm.v1.kv_cache_interface import MambaSpec

    _TYPE_MAP = {
        "mamba1": MambaAttentionBackendEnum.MAMBA1,
        "mamba2": MambaAttentionBackendEnum.MAMBA2,
        "gdn_attention": MambaAttentionBackendEnum.GDN_ATTN,
    }
    mamba_type_enum = _TYPE_MAP[mamba_type]

    monkeypatch.setenv("VLLM_SSM_CONV_STATE_LAYOUT", "DS")
    spec = MambaSpec(
        block_size=64,
        shapes=((conv_dim_local, conv_rows), temporal_shape),
        dtypes=(torch.bfloat16, torch.bfloat16),
        mamba_type=mamba_type_enum,
    )
    out = derive_mamba_conv_split(spec, local_tp=local_tp)
    assert out.local_proj_dims == expected_proj_dims
    assert out.conv_rows == conv_rows


@pytest.mark.cpu_test
@pytest.mark.parametrize(
    "mamba_enabled,swa_enabled,"
    "local_physical_per_logical,remote_physical_per_logical,"
    "logical_block_ids,expected_kernel_block_ids",
    [
        # Qwen3.5-0.8B 4P2D (kernel_block_size=64):
        #   prefill TP=4: logical_block_size=384 → physical_per_logical=6
        #   decode  TP=2: logical_block_size=640 → physical_per_logical=10
        # FA logical [0] → remote kernel [0..9] (1 * 10)
        # SSM logical [10] → unchanged [10]
        pytest.param(
            True,
            False,
            6,
            10,
            ([0], [10]),
            [[0, 1, 2, 3, 4, 5, 6, 7, 8, 9], [10]],
            id="qwen35_4p2d",
        ),
        # Qwen3.5-0.8B 2P4D (kernel_block_size=64):
        #   prefill TP=2: logical_block_size=640 → physical_per_logical=10
        #   decode  TP=4: logical_block_size=384 → physical_per_logical=6
        # FA logical [0, 1] → remote kernel [0..5, 6..11] (2 * 6)
        # SSM logical [10] → unchanged [10]
        pytest.param(
            True,
            False,
            10,
            6,
            ([0, 1], [10]),
            [[0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11], [10]],
            id="qwen35_2p4d",
        ),
        # Homogeneous TP (kernel_block_size=64):
        #   both sides: logical_block_size=640 → physical_per_logical=10
        # FA logical [0] → kernel [0..9], SSM unchanged
        pytest.param(
            True,
            False,
            10,
            10,
            ([0], [10]),
            [[0, 1, 2, 3, 4, 5, 6, 7, 8, 9], [10]],
            id="homo_tp",
        ),
        # remote physical_per_logical=1: early return, no expansion
        pytest.param(
            True,
            False,
            10,
            1,
            ([0, 1, 2], [5]),
            [[0, 1, 2], [5]],
            id="mamba_remote_physical_per_logical_1",
        ),
        # Pure FA (no mamba): single group expanded with remote stride
        pytest.param(
            False,
            False,
            2,
            4,
            ([0, 1],),
            [[0, 1, 2, 3, 4, 5, 6, 7]],
            id="pure_fa",
        ),
        # FA + SWA (no mamba): both groups expanded
        pytest.param(
            False,
            True,
            2,
            3,
            ([0, 1], [2, 3]),
            [[0, 1, 2, 3, 4, 5], [6, 7, 8, 9, 10, 11]],
            id="fa_swa",
        ),
    ],
)
def test_logical_to_kernel_block_ids_with_remote_ratio(
    mamba_enabled,
    swa_enabled,
    local_physical_per_logical,
    remote_physical_per_logical,
    logical_block_ids,
    expected_kernel_block_ids,
):
    """Verify _logical_to_kernel_block_ids uses the remote
    physical_per_logical for FA expansion, not the local one.

    This was the root cause of silent accuracy corruption in Qwen3.5
    heterogeneous TP (e.g. 4P2D): the old code used local physical_per_logical
    for the expansion arange, producing wrong kernel block indices.

    Qwen3.5-0.8B values verified by verify_conv_split.py (issue #13).
    """
    from vllm.distributed.kv_transfer.kv_connector.v1.nixl.worker import (
        NixlConnectorWorker,
    )

    worker = object.__new__(NixlConnectorWorker)
    worker._physical_blocks_per_logical_kv_block = local_physical_per_logical
    worker.kv_cache_config = make_kv_cache_config(
        block_size=16,
        mamba_enabled=mamba_enabled,
        swa_enabled=swa_enabled,
    )

    result = worker._logical_to_kernel_block_ids(
        logical_block_ids,
        remote_physical_per_logical,
    )
    assert list(result) == expected_kernel_block_ids, (
        f"Expected {expected_kernel_block_ids}, got {result}"
    )


@pytest.mark.cpu_test
def test_exchange_clipped_blocks_ssm_single_state():
    """In single-state cache modes, SSM lists are reduced to the running
    state slot: speculative scratch slots, null placeholders and the previous
    step's state carry nothing. Attention groups pass through untouched."""
    sched = make_nixl_scheduler(has_mamba=True, is_hma_required=True)
    sched.blocks_per_sw = [0, 0]
    sched._ssm_spec_blocks = [None, 2]
    sched._ssm_state_slots_are_positional = False

    # Align-mode list: null placeholders, state block, 2 speculative slots.
    clipped = sched.get_exchange_clipped_blocks(([1, 2, 3], [0, 0, 7, 8, 9]))
    assert clipped == ([1, 2, 3], [7])

    # Same, still holding the previous step's state block (freed a step later).
    assert sched.get_exchange_clipped_blocks(([1], [0, 6, 7, 8, 9]))[1] == [7]

    # Default (mamba_block_size=max_model_len): state block, 2 scratch slots.
    assert sched.get_exchange_clipped_blocks(([1], [7, 8, 9]))[1] == [7]

    # Scratch slots not allocated: the state slot still survives.
    assert sched.get_exchange_clipped_blocks(([1], [5]))[1] == [5]

    # Non-mamba models pass through unchanged.
    fa_sched = make_nixl_scheduler(has_mamba=False)
    assert fa_sched.get_exchange_clipped_blocks(([1, 2],)) == ([1, 2],)


@pytest.mark.cpu_test
def test_exchange_clipped_blocks_ssm_positional_states():
    """In "all" mode every position holds a state, so only the speculative
    slots go; placeholders stay to keep the list position-indexed."""
    sched = make_nixl_scheduler(has_mamba=True, is_hma_required=True)
    sched.blocks_per_sw = [0, 0]
    sched._ssm_spec_blocks = [None, 2]
    sched._ssm_state_slots_are_positional = True

    clipped = sched.get_exchange_clipped_blocks(([1, 2, 3], [0, 5, 6, 7, 8, 9]))
    assert clipped == ([1, 2, 3], [0, 5, 6, 7])


# ── Hybrid MLA+SSM (KimiLinear-shaped KDA+MLA) tests ─────────────────────


def _make_hybrid_mla_kv_cache_config(num_blocks: int = 4):
    """KimiLinear-shaped config: one MLA group and two KDA (GDN-typed
    MambaSpec) groups whose layers share the same HMA tensors, with a
    mamba-aligned unified page and an MLA kernel block smaller than the
    logical block."""
    from vllm.v1.attention.backends.registry import MambaAttentionBackendEnum
    from vllm.v1.kv_cache_interface import (
        KVCacheConfig,
        KVCacheGroupSpec,
        KVCacheTensor,
        MambaSpec,
        MLAAttentionSpec,
    )

    # 12-token logical blocks over a 4-token MLA kernel block.
    mla_spec = MLAAttentionSpec(
        block_size=12, num_kv_heads=1, head_size=6, dtype=torch.float16
    )
    unified_page = mla_spec.page_size_bytes
    kda_spec = MambaSpec(
        block_size=12,
        # GDN-decomposable conv (Q|K|V = 2|2|4 cols x 3 rows) + fp32 temporal.
        shapes=((8, 3), (1, 4, 4)),
        dtypes=(torch.float16, torch.float32),
        page_size_padded=unified_page,
        mamba_type=MambaAttentionBackendEnum.GDN_ATTN,
    )
    assert kda_spec.page_size_bytes == unified_page
    return KVCacheConfig(
        num_blocks=num_blocks,
        kv_cache_tensors=[
            KVCacheTensor(
                size=num_blocks * unified_page,
                shared_by=[f"mla.{i}", f"kda_a.{i}", f"kda_b.{i}"],
            )
            for i in range(2)
        ],
        kv_cache_groups=[
            KVCacheGroupSpec(["mla.0", "mla.1"], mla_spec),
            KVCacheGroupSpec(["kda_a.0", "kda_a.1"], kda_spec),
            KVCacheGroupSpec(["kda_b.0", "kda_b.1"], kda_spec),
        ],
    )


@pytest.mark.cpu_test
def test_register_kv_caches_hybrid_mla_dual_purpose_regions():
    """Hybrid MLA+KDA registration: HMA tensors shared by both layer types
    must be flagged as MLA regions even when a KDA layer registers them
    first, expose TP-independent kernel-granularity block lens, and build
    FA + mamba descriptors for every region."""
    from unittest.mock import MagicMock

    from vllm.config import set_current_vllm_config
    from vllm.distributed.kv_transfer.kv_connector.v1.nixl import base_worker as bw
    from vllm.distributed.kv_transfer.kv_connector.v1.nixl.worker import (
        NixlConnectorWorker,
    )

    kv_cache_config = _make_hybrid_mla_kv_cache_config()
    unified_page = kv_cache_config.kv_cache_groups[0].kv_cache_spec.page_size_bytes
    vllm_config = create_vllm_config(block_size=12)
    # kv_buffer_device defaults to the *real* platform's device type, which on
    # a CPU-only test host would make this a host-buffer worker: host xfer
    # buffers are per-layer, so the HMA shared tensors would not be
    # deduplicated. Pin it to the faked device type.
    vllm_config.kv_transfer_config.kv_buffer_device = "cuda"

    fake_backend = MagicMock()
    fake_backend.get_supported_kernel_block_sizes.return_value = [4]
    fake_backend.get_name.return_value = "FLASHMLA"
    fake_backend.full_cls_name.return_value = "fake.FLASHMLA"
    fake_platform = MagicMock()
    fake_platform.device_type = "cuda"
    fake_platform.get_nixl_memory_type.return_value = "VRAM"

    with (
        patch.object(bw, "NixlWrapper"),
        patch.object(bw, "get_tensor_model_parallel_rank", return_value=0),
        patch.object(bw, "get_tensor_model_parallel_world_size", return_value=1),
        patch.object(bw, "get_current_attn_backends", return_value=[fake_backend]),
        patch.object(bw, "current_platform", fake_platform),
        patch(
            "vllm.model_executor.layers.mamba.mamba_utils.get_conv_state_layout",
            return_value="DS",
        ),
        set_current_vllm_config(vllm_config),
    ):
        worker = NixlConnectorWorker(vllm_config, "test-engine", kv_cache_config)
        worker.use_mla = True  # opt-125m test config is not MLA; force the flag
        worker.nixl_wrapper.get_agent_metadata.return_value = b"fake-agent-metadata"

        tensors = [torch.zeros(4 * unified_page, dtype=torch.uint8) for _ in range(2)]
        # KDA layer first per tensor: exercises the dual-purpose flag merge.
        worker.register_kv_caches(
            {
                "kda_a.0": tensors[0],
                "mla.0": tensors[0],
                "kda_b.0": tensors[0],
                "kda_a.1": tensors[1],
                "mla.1": tensors[1],
                "kda_b.1": tensors[1],
            }
        )

    # 12-token logical blocks over the 4-token MLA kernel block.
    assert worker._physical_blocks_per_logical_kv_block == 3
    assert worker.block_size == 4 and worker.num_blocks == 12
    # Both shared tensors are dual-purpose: their FA view is MLA even though
    # a KDA layer registered them first.
    assert worker._region_is_mla == [True, True]
    assert worker.num_regions == 2 and worker.num_descs == 24
    # Kernel-granularity block lens; TP-independent for MLA hybrids.
    assert worker.block_len_per_layer == [unified_page // 3] * 2
    # Split handles must replicate every FA descriptor (MLA isn't head-sharded).
    assert worker._fa_desc_replicated(worker.num_descs) == [True] * 24
    # FA descs: 2 regions x 12 kernel blocks, page stride = kernel page.
    # Mamba descs: 2 regions x (3 conv sub-projections + 1 ssm) x 4 blocks.
    assert worker.src_blocks_data.shape == (24 + 32, 3)
    fa_descs = worker.src_blocks_data[:24]
    assert fa_descs[1][0] - fa_descs[0][0] == unified_page // 3
    assert all(size == unified_page // 3 for size in fa_descs[:, 1])


@pytest.mark.cpu_test
def test_push_write_hybrid_mla_replicates_attention():
    """Hybrid MLA+SSM push with P_TP < D_TP: attention blocks must be
    written to every covered D rank (replicated MLA latent) while SSM state
    is written per-rank through the split handles."""
    import threading
    from collections import defaultdict
    from unittest.mock import MagicMock

    from vllm.distributed.kv_transfer.kv_connector.v1.nixl.push_worker import (
        NixlPushConnectorWorker,
    )
    from vllm.distributed.kv_transfer.kv_connector.v1.nixl.tp_mapping import (
        TPMapping,
    )
    from vllm.v1.kv_cache_interface import MambaSpec, MLAAttentionSpec

    worker = object.__new__(NixlPushConnectorWorker)
    worker.shutdown = lambda: None  # skeleton worker: silence __del__
    worker.use_mla = True
    worker._has_mamba = True
    worker._group_spec_types = (MLAAttentionSpec, MambaSpec)
    worker.transfer_topo = MagicMock()
    worker.transfer_topo.tp_ratio.return_value = -2
    remote_info = MagicMock()
    remote_info.remote_physical_blocks_per_logical = 1
    remote_info.remote_block_size = 4
    worker.transfer_topo.get_engine_info.return_value = remote_info

    engine_id = "remote-engine"
    # Read-oriented mapping collapses the replicated attention group to one
    # source rank; the SSM state is sharded across both covered D ranks.
    worker.tp_mappings = {
        engine_id: TPMapping(
            source_ranks_per_group=((0,), (0, 1)),
            all_source_ranks=(0, 1),
            rank_to_attention_slot={0: 0, 1: 0},
            rank_offset_factor=0,
        )
    }
    worker.dst_xfer_side_handles = {engine_id: {0: 100, 1: 101}}
    worker.src_xfer_handles_by_tp_ratio = {(-2, 4): [200, 201]}
    worker.src_xfer_handles_by_block_size = {4: 300}
    worker._sending_transfers = defaultdict(list)
    worker._sending_transfers_lock = threading.Lock()
    worker.kv_cache_config = _make_hybrid_mla_kv_cache_config()
    worker._xfer_blocks = MagicMock(return_value=1)

    meta = MagicMock()
    meta.remote.engine_id = engine_id
    meta.remote.block_ids = [[7, 8], [3]]
    meta.local_physical_block_ids = [[1, 2], [5]]

    worker._xfer_blocks_for_req("req-1", meta)

    calls = worker._xfer_blocks.call_args_list
    assert len(calls) == 2
    for call, rank, local_handle, remote_handle in zip(
        calls, (0, 1), (200, 201), (100, 101)
    ):
        spec = call.kwargs["read_spec"]
        assert spec.remote_rank == rank
        # Attention group replicated to every rank, SSM by membership.
        assert spec.local_block_ids == [[1, 2], [5]]
        assert spec.remote_block_ids == [[7, 8], [3]]
        assert call.kwargs["local_xfer_side_handle"] == local_handle
        assert call.kwargs["remote_xfer_side_handle"] == remote_handle
