# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project

import contextlib
import inspect
import os
import queue
import tempfile
import textwrap
import threading
import time
import uuid
from collections import defaultdict
from concurrent.futures import Future
from types import SimpleNamespace
from typing import Any, cast
from unittest.mock import MagicMock, patch

import msgspec
import numpy as np
import pytest
import ray
import torch

from tests.v1.attention.utils import dense_kv_cache_views
from vllm import LLM
from vllm.config import KVTransferConfig, set_current_vllm_config
from vllm.distributed.kv_transfer.kv_connector.utils import (
    EngineTransferInfo,
    KVOutputAggregator,
    TransferTopology,
    get_current_attn_backend,
)
from vllm.distributed.kv_transfer.kv_connector.v1 import nixl
from vllm.distributed.kv_transfer.kv_connector.v1.base import (
    KVConnectorRole,
    KVConnectorTransferResults,
)
from vllm.distributed.kv_transfer.kv_connector.v1.metrics import KVConnectorStats
from vllm.distributed.kv_transfer.kv_connector.v1.multi_connector import (
    MultiKVConnectorStats,
)
from vllm.distributed.kv_transfer.kv_connector.v1.nixl import (
    NixlAgentMetadata,
    NixlConnector,
    NixlConnectorMetadata,
    NixlConnectorScheduler,
    NixlConnectorWorker,
    NixlHandshakePayload,
    NixlKVConnectorStats,
)
from vllm.distributed.kv_transfer.kv_connector.v1.nixl.metadata import (
    RemoteMeta,
    ReqMeta,
    compute_nixl_compatibility_hash,
)
from vllm.distributed.kv_transfer.kv_transfer_state import (
    ensure_kv_transfer_shutdown,
    has_kv_transfer_group,
)
from vllm.forward_context import ForwardContext
from vllm.outputs import RequestOutput
from vllm.platforms import current_platform
from vllm.platforms.interface import Platform
from vllm.sampling_params import SamplingParams
from vllm.v1.engine import EngineCoreRequest
from vllm.v1.engine.output_processor import OutputProcessor
from vllm.v1.kv_cache_interface import (
    AttentionSpec,
    FullAttentionSpec,
    KVCacheConfig,
    KVCacheGroupSpec,
    KVCacheLayout,
    MLAAttentionSpec,
    compute_layer_kv_cache_shape_bytes,
)
from vllm.v1.outputs import KVConnectorOutput, ModelRunnerOutput
from vllm.v1.request import RequestStatus

from .utils import (
    create_model_runner_output,
    create_request,
    create_scheduler,
    create_vllm_config,
    make_kv_cache_config,
)


@pytest.fixture(scope="module", autouse=True)
def clear_kv_transfer():
    """The test cases in this file use `VLLM_ENABLE_V1_MULTIPROCESSING=0`,
    causing the global variable `_KV_CONNECTOR_AGENT`
    to be assigned but never deleted.

    Since the current pytest process does not terminate and instead
    continues running tests from other files,
    this global variable remains in memory and interferes
    with test cases in other modules.

    So we use this fixture to ensure that the global variable
    `_KV_CONNECTOR_AGENT` is properly cleaned up after each test.
    """
    yield
    if has_kv_transfer_group():
        ensure_kv_transfer_shutdown()


def get_default_xfer_telemetry(
    xferDurationS: float = 1,
    postDurationS: float = 1,
    totalBytes: int = 1,
    descCount: int = 1,
) -> dict:
    class AttributeDict(dict):
        __slots__ = ()
        __getattr__ = dict.__getitem__
        __setattr__ = dict.__setitem__

    # We can't instantiate nixlXferTelemetry because it's read only and
    # ray env does not have NIXL, so we must fake it
    return AttributeDict(
        xferDuration=xferDurationS * 1e6,  # in us
        postDuration=postDurationS * 1e6,  # in us
        totalBytes=totalBytes,
        descCount=descCount,
    )


class FakeNixlWrapper:
    """Mock implementation of NixlWrapper for testing.

    We don't inherit from nixl._api.nixl_agent because nixl may not be
    installed.

    Note: The complete source of this class is also used in the
    `_make_fake_nixl_pkg` function to create a fake nixl package
    for Ray workers.
    """

    AGENT_METADATA = b"fake_agent_metadata"
    REMOTE_AGENT_NAME = "remote_agent"

    def __init__(self, agent_name: str, *args, **kwargs):
        self._cycles_before_xfer_done = 0
        self._check_xfer_state_cycles: defaultdict[int, int] = defaultdict(lambda: 0)

    def get_reg_descs(self, caches_data, memory_type: str) -> list:
        return [str(uuid.uuid4()) for _ in caches_data]

    def register_memory(self, descs, backends) -> None:
        pass

    def deregister_memory(self, descs) -> None:
        pass

    def get_xfer_descs(self, blocks_data, memory_type: str) -> list:
        return [str(uuid.uuid4()) for _ in blocks_data]

    def prep_xfer_dlist(self, agent_name: str, descs: list) -> int:
        return uuid.uuid4().int

    def get_agent_metadata(self) -> bytes:
        return self.AGENT_METADATA

    def add_remote_agent(self, agent_metadata: bytes) -> str:
        return self.REMOTE_AGENT_NAME

    def get_new_notifs(self) -> dict[str, list[bytes]]:
        # Used to collect done_sending, which we don't test yet.
        return {}

    def check_xfer_state(self, handle: int) -> str:
        if self._check_xfer_state_cycles[handle] >= self._cycles_before_xfer_done:
            return "DONE"
        self._check_xfer_state_cycles[handle] += 1
        return "PROC"

    def release_xfer_handle(self, handle: int) -> None:
        pass

    def release_dlist_handle(self, handle: int) -> None:
        pass

    def remove_remote_agent(self, agent: str) -> None:
        pass

    def send_notif(self, agent_name: str, notif_msg: bytes) -> None:
        pass

    def make_prepped_xfer(
        self,
        xfer_type: str,
        local_xfer_side_handle: int,
        local_block_descs_ids: list[int],
        remote_xfer_side_handle: int,
        remote_block_descs_ids: list[int],
        notif_msg: bytes | None = None,
    ) -> int:
        return uuid.uuid4().int

    def transfer(self, handle: int) -> str:
        return "PROC"

    def get_xfer_telemetry(self, handle: int) -> dict:
        return get_default_xfer_telemetry()


@contextlib.contextmanager
def _make_fake_nixl_pkg():
    """Context manager that creates a temporary package making
       `from nixl._api import nixl_agent` resolve to our FakeNixlWrapper.
       Also creates the ROCm NIXL packages.

    Automatically cleans up the temporary directory when done.
    """
    with tempfile.TemporaryDirectory() as td:
        for pkg_name in ["nixl", "nixl_rocm"]:
            pkg_root = os.path.join(td, pkg_name, "_api")
            os.makedirs(pkg_root, exist_ok=True)

            # Get the source code of FakeNixlWrapper class and dedent it
            fake_nixl_source = inspect.getsource(FakeNixlWrapper)
            fake_nixl_source = textwrap.dedent(fake_nixl_source)

            stub = f"""\
# Copy of FakeNixlWrapper implementation for Ray workers
import uuid
from collections import defaultdict

{fake_nixl_source}

# Export as nixl_agent
nixl_agent = FakeNixlWrapper
"""
            with open(os.path.join(pkg_root, "__init__.py"), "w") as f:
                f.write(stub)

            # Mock nixlXferTelemetry class
            pkg_root2 = os.path.join(td, pkg_name, "_bindings")
            os.makedirs(pkg_root2, exist_ok=True)
            with open(os.path.join(pkg_root2, "__init__.py"), "w") as f:
                f.write("class nixlXferTelemetry: pass")
            # touch parent package
            open(os.path.join(td, pkg_name, "__init__.py"), "w").close()

        yield td


def test_basic_interface():
    """Unit test for basic NixlConnector interface functionality."""
    vllm_config = create_vllm_config()
    scheduler = create_scheduler(vllm_config)

    # 2 Full Blocks and 1 Half Block.
    BLOCK_SIZE = vllm_config.cache_config.block_size
    NUM_EXTERNAL_FULL_BLOCKS = 2
    NUM_TOKENS = int(BLOCK_SIZE * (NUM_EXTERNAL_FULL_BLOCKS + 0.5))

    request = create_request(
        request_id=1,
        block_size=BLOCK_SIZE,
        num_tokens=NUM_TOKENS,
        do_remote_prefill=True,
    )
    request_id = request.request_id

    scheduler.add_request(request)

    # Remote Prefill, triggers NixlConnectorMetadata.
    scheduler_output = scheduler.schedule()
    kv_connector_metadata = scheduler_output.kv_connector_metadata
    assert kv_connector_metadata is not None
    assert isinstance(kv_connector_metadata, NixlConnectorMetadata)

    assert len(kv_connector_metadata.reqs_to_recv) == 1
    assert request_id in kv_connector_metadata.reqs_to_recv
    req_meta = kv_connector_metadata.reqs_to_recv[request_id]

    for block_id, block in zip(
        req_meta.local_block_ids[0],
        scheduler.kv_cache_manager.coordinator.single_type_managers[0].req_to_blocks[
            request_id
        ],
    ):
        assert block_id == block.block_id


def test_prompt_less_than_block_size():
    """Test that we can handle case where prompt is < block.

    In this case, the P worker will still send remote_block_ids of the
    partial block. The D worker should schedule an async read
    in this case.
    """
    vllm_config = create_vllm_config()
    scheduler = create_scheduler(vllm_config)

    # Half of a block.
    BLOCK_SIZE = vllm_config.cache_config.block_size
    NUM_TOKENS = int(BLOCK_SIZE * 0.5)

    # Request will have 1 partial remote block.
    request = create_request(
        request_id=1,
        block_size=BLOCK_SIZE,
        num_tokens=NUM_TOKENS,
        do_remote_prefill=True,
        num_remote_blocks=1,
    )
    scheduler.add_request(request)
    scheduler_output = scheduler.schedule()

    # This request will read async.
    kv_connector_metadata = scheduler_output.kv_connector_metadata
    assert kv_connector_metadata is not None
    assert isinstance(kv_connector_metadata, NixlConnectorMetadata)
    assert len(kv_connector_metadata.reqs_to_recv) == 1
    assert len(scheduler_output.scheduled_new_reqs) == 0


def test_abort_immediately_remote_prefill_enqueues_empty_recv():
    """A remote-prefill request added with abort_immediately=True should
    be added to the scheduler's waiting queue then immediately aborted, so the
    NIXL connector's request_finished hook enqueues an empty recv to notify
    the prefill instance to free its blocks."""
    from vllm.v1.request import RequestStatus

    scheduler = create_scheduler(create_vllm_config())

    request = create_request(request_id=42, num_tokens=10, do_remote_prefill=True)
    assert request.kv_transfer_params is not None
    assert request.kv_transfer_params["do_remote_prefill"] is True

    # Mimic the EngineCore.add_request path for an abort-immediately req.
    scheduler.add_request(request)
    scheduler.finish_requests([request.request_id], RequestStatus.FINISHED_ABORTED)

    scheduler_output = scheduler.schedule()
    meta = scheduler_output.kv_connector_metadata
    assert isinstance(meta, NixlConnectorMetadata)
    assert set(meta.reqs_to_recv) == {request.request_id}
    req_meta = meta.reqs_to_recv[request.request_id]
    assert req_meta.local_block_ids == []
    assert req_meta.remote.request_id == f"prefill-{42}"
    # do_remote_prefill is consumed by request_finished to prevent re-issuing.
    assert request.kv_transfer_params["do_remote_prefill"] is False
    # The scheduler is not waiting on this recv -- the request is already gone
    # from self.requests, so reporting it would trip `assert req_id in
    # self.requests` in _update_from_kv_xfer_finished.
    assert req_meta.awaiting_kvs is False


def test_prefill_exports_cached_tokens_in_kv_transfer_params():
    """The P worker reports its own prefix-cache hits in the returned
    kv_transfer_params so the D worker can surface them in
    prompt_tokens_details instead of the ~100% local hit it measures
    when pulling the KVs from the remote.
    """
    vllm_config = create_vllm_config()
    scheduler = create_scheduler(vllm_config)

    BLOCK_SIZE = vllm_config.cache_config.block_size
    NUM_TOKENS = BLOCK_SIZE * 3

    # Warm the prefix cache with a request sharing the full prompt.
    warmup = create_request(
        request_id=1,
        num_tokens=NUM_TOKENS,
        common_prefix_len=NUM_TOKENS,
        block_size=BLOCK_SIZE,
    )
    scheduler.add_request(warmup)
    scheduler_output = scheduler.schedule()
    model_runner_output = create_model_runner_output([warmup], use_eos=True)
    scheduler.update_from_output(scheduler_output, model_runner_output)

    # P-side request with the same prompt hits the local cache on all but
    # the last block (recomputed to obtain logits).
    request = create_request(
        request_id=2,
        num_tokens=NUM_TOKENS,
        common_prefix_len=NUM_TOKENS,
        block_size=BLOCK_SIZE,
        do_remote_decode=True,
    )
    scheduler.add_request(request)
    scheduler_output = scheduler.schedule()
    assert request.prefill_stats is not None
    assert request.prefill_stats.num_cached_tokens == NUM_TOKENS - BLOCK_SIZE

    # max_tokens=1, so the request finishes and returns kv_transfer_params.
    model_runner_output = create_model_runner_output([request])
    engine_core_outputs = scheduler.update_from_output(
        scheduler_output, model_runner_output
    )

    output = engine_core_outputs[0].outputs[0]
    assert output.finish_reason is not None
    assert output.kv_transfer_params is not None
    assert (
        output.kv_transfer_params["remote_prefill_cached_tokens"]
        == NUM_TOKENS - BLOCK_SIZE
    )


@patch(
    "vllm.distributed.kv_transfer.kv_connector.v1.nixl.base_worker.NixlWrapper",
    FakeNixlWrapper,
)
def test_kv_transfer_handshake(dist_init):
    """Unit test for basic NixlConnector interface functionality."""
    # Test setup, we creates a scheduler that contains a NixlConnector
    # of role SCHEDULER, and expect it to be serving NixlAgentMetadata from
    # all workers of the instance.
    vllm_config = create_vllm_config()
    vllm_config.cache_config.kv_cache_layout = "BLHNC"
    # in case the test runs on non-GPU machine
    vllm_config.kv_transfer_config.kv_buffer_device = "cpu"
    scheduler = create_scheduler(vllm_config)

    with set_current_vllm_config(vllm_config):
        # Create two NixlConnector of role WORKER, one is the worker of
        # the scheduler (prefill), the other is a worker of decode instance.

        # Prefill connector will register KV cache to populate proper handshake
        # metadata.
        kv_cache_groups = [
            KVCacheGroupSpec(
                ["layer0", "layer1", "layer2"],
                FullAttentionSpec(
                    block_size=16,
                    num_kv_heads=4,
                    head_size=16,
                    dtype=torch.float16,
                ),
            )
        ]
        kv_cache_config = KVCacheConfig(
            num_blocks=2, kv_cache_tensors=[], kv_cache_groups=kv_cache_groups
        )
        prefill_connector = NixlConnector(
            vllm_config, KVConnectorRole.WORKER, kv_cache_config
        )
        kv_cache_spec = cast(
            AttentionSpec, kv_cache_config.kv_cache_groups[0].kv_cache_spec
        )
        raw = torch.zeros(
            kv_cache_spec.page_size_bytes * kv_cache_config.num_blocks * 3,
            dtype=torch.int8,
        )
        caches = dense_kv_cache_views(
            raw,
            kv_cache_spec,
            kv_cache_config.num_blocks,
            num_layers=3,
            layout=KVCacheLayout.BLHNC,
        )
        kv_caches = {
            f"layer{layer_idx}": cache for layer_idx, cache in enumerate(caches)
        }
        prefill_connector.register_kv_caches(kv_caches)

        # Simulate EngineCore initialization that would gather connector
        # metadata from all workers
        metadata = prefill_connector.get_handshake_metadata()

        # metadata is a NixlHandshakePayload, decode it to get NixlAgentMetadata
        decoder = msgspec.msgpack.Decoder(NixlAgentMetadata)
        expected_agent_metadata = decoder.decode(metadata.agent_metadata_bytes)

        # The scheduler connector expects metadata keyed by
        # (pp_rank, tp_rank).
        scheduler_connector = scheduler.get_kv_connector()
        scheduler_connector.set_xfer_handshake_metadata_pp_aware({(0, 0): metadata})

        # Simulate a request that finishes prefill, which returns
        # corresponding NixlConnectorMetadata for decode instance.
        BLOCK_SIZE = vllm_config.cache_config.block_size
        NUM_EXTERNAL_FULL_BLOCKS = 2
        NUM_TOKENS = int(BLOCK_SIZE * (NUM_EXTERNAL_FULL_BLOCKS + 0.5))

        request = create_request(
            request_id=1,
            block_size=BLOCK_SIZE,
            num_tokens=NUM_TOKENS,
            do_remote_decode=True,
        )
        request.status = RequestStatus.FINISHED_LENGTH_CAPPED
        delay, kv_connector_metadata = (
            scheduler.get_kv_connector().request_finished_all_groups(
                request, ([0, 1, 2],)
            )
        )
        assert delay
        # Pull connector advertises its transfer mode in kv_transfer_params so
        # an external router can distinguish it from a push producer.
        assert kv_connector_metadata["transfer_mode"] == "pull"

        # Decode connector will be able to create handshake with the prefill connector.
        decode_connector = NixlConnector(
            vllm_config, KVConnectorRole.WORKER, kv_cache_config
        )
        decode_connector.register_kv_caches(kv_caches)

        # Here we are testing the retrieval of NIXLAgentMetadata.
        # Knowing the implementation detail, we override the add_remote_agent
        # to validate the metadata received is the same as the one in prefill_connector.
        with patch.object(
            decode_connector.connector_worker, "add_remote_agent"
        ) as mock_add_remote_agent:
            mock_add_remote_agent.return_type = "remote_agent"

            decode_connector.connector_worker._nixl_handshake(
                kv_connector_metadata["remote_host"],
                kv_connector_metadata["remote_port"],
                kv_connector_metadata["tp_size"],
                kv_connector_metadata["remote_engine_id"],
            )

            received_metadata = mock_add_remote_agent.call_args.args
            assert received_metadata[0] == expected_agent_metadata
            assert received_metadata[1] == 0  # remote_tp_rank
            assert received_metadata[2] == 1  # remote_tp_size

        # Need to shutdown the background thread to release NIXL side channel port
        scheduler_connector.shutdown()


class FakeNixlConnectorWorker(NixlConnectorWorker):
    REMOTE_ENGINE_ID = "remote_engine"

    def __init__(
        self,
        *args,
        hand_shake_latency: float = 1.8,
        kv_cache_layout="LBHNC",
        kv_cache_config=None,
        **kwargs,
    ):
        if kv_cache_config is None:
            kv_cache_config = make_kv_cache_config(block_size=16)
        super().__init__(*args, kv_cache_config=kv_cache_config, **kwargs)
        self._hand_shake_latency = hand_shake_latency
        self.kv_cache_layout = kv_cache_layout
        # Mock register_kv_caches attributes needed for tests that do not call it.
        self.src_xfer_handles_by_block_size = {self.block_size: 1}
        self.src_blocks_data = np.empty((0, 3), dtype=np.uint64)
        rep_spec = self.kv_cache_config.kv_cache_groups[0].kv_cache_spec
        test_shape = compute_layer_kv_cache_shape_bytes(rep_spec, 1)
        self.transfer_topo = TransferTopology(
            tp_rank=self.tp_rank,
            tp_size=self.world_size,
            block_size=self.block_size,
            engine_id=self.engine_id,
            is_mla=self.use_mla,
            is_mamba=False,
            total_num_kv_heads=self.model_config.get_total_num_kv_heads(),
            attn_backends=self.attn_backends,
            tensor_shape=test_shape,
        )

        self.compat_hash = compute_nixl_compatibility_hash(
            self.vllm_config, self.backend_name
        )

    def _nixl_handshake(
        self,
        host: str,
        port: int,
        remote_tp_size: int,
        expected_engine_id: str,
        remote_dcp_size: int = 1,
        remote_pp_size: int = 1,
        notif_agents_only: bool = False,
    ) -> tuple[dict[tuple[int, int], str], float]:
        # Mimic slow _nixl_handshake, as well as bypass zmq communication.
        time.sleep(self._hand_shake_latency)
        # These should've been done in register_kv_caches(), called by
        # gpu_model_runner. Here we just hardcode some dummy values.
        slot_size_bytes = 4096
        self.slot_size_per_layer = [slot_size_bytes]
        self.block_len_per_layer = [slot_size_bytes * self.block_size]
        self.num_blocks = self.kv_cache_config.num_blocks
        self.num_regions = 1
        self.block_stride_per_layer = list(self.block_len_per_layer)
        self.region_num_blocks = [self.num_blocks]
        self.region_group_ids = [0]
        self.region_mem_types = [self.nixl_memory_type]
        self.dst_num_blocks[self.engine_id] = self.num_blocks
        self.dst_region_num_blocks[self.engine_id] = self.region_num_blocks
        self.dst_region_group_ids[self.engine_id] = self.region_group_ids
        self.dst_region_mem_types[self.engine_id] = self.region_mem_types

        assert expected_engine_id == self.REMOTE_ENGINE_ID

        # Adjust remote block length metadata to satisfy heterogeneous TP
        # invariants enforced during handshake validation.  Use per-rank
        # head ratio (not tp_ratio) to account for GQA replication capping.
        remote_block_lens = list(self.block_len_per_layer)
        tp_ratio = self.transfer_topo.tp_ratio(remote_tp_size)
        total_kv = self.transfer_topo.total_num_kv_heads
        local_heads = self.transfer_topo.local_physical_heads
        remote_heads = max(1, total_kv // remote_tp_size)
        if remote_tp_size != self.world_size:
            remote_block_lens = [
                block_len * remote_heads // local_heads
                for block_len in remote_block_lens
            ]

        # When remote tp_size > local tp_size, handshake with multiple
        # remote ranks.
        num_handshakes = 1 if tp_ratio > 0 else -tp_ratio
        remote_agents: dict[tuple[int, int], str] = {}
        for remote_tp_rank in range(num_handshakes):
            remote_agent_name = self.add_remote_agent(
                NixlAgentMetadata(
                    engine_id=self.REMOTE_ENGINE_ID,
                    agent_metadata=FakeNixlWrapper.AGENT_METADATA,
                    kv_caches_base_addr=[0],
                    device_id=remote_tp_rank,
                    num_blocks=self.num_blocks,
                    block_lens=remote_block_lens,
                    block_strides=remote_block_lens,
                    kv_cache_layout="LBHNC",
                    block_size=self.block_size,
                    ssm_sizes=(0, 0),
                    attn_backend_name=self.backend_name,
                    physical_blocks_per_logical_kv_block=1,
                    region_num_blocks=self.region_num_blocks,
                    region_group_ids=self.region_group_ids,
                    region_mem_types=self.region_mem_types,
                ),
                remote_tp_rank=remote_tp_rank,
                remote_tp_size=remote_tp_size,
            )
            remote_agents[(0, remote_tp_rank)] = remote_agent_name
        # Handshake bypasses zmq, so report a zero clock offset to the peer.
        return remote_agents, 0.0


class TestNixlHandshake:
    @pytest.mark.parametrize(
        ("pcp_rank", "pcp_size", "dcp_size", "expected_tracked"),
        [
            (0, 2, 1, True),
            (1, 2, 1, False),
            (0, 2, 2, True),
            (1, 2, 2, True),
            (0, 4, 4, True),
            (1, 4, 4, True),
            (2, 4, 4, True),
            (3, 4, 4, True),
        ],
    )
    @patch(
        "vllm.distributed.kv_transfer.kv_connector.v1.nixl.base_worker.NixlWrapper",
        FakeNixlWrapper,
    )
    def test_pcp_producer_exposes_dcp_shards_or_canonical_replica(
        self,
        default_vllm_config,
        dist_init,
        pcp_rank,
        pcp_size,
        dcp_size,
        expected_tracked,
    ):
        """Replicated PCP is canonicalized; PCP-DCP publishes every shard."""
        from vllm.v1.attention.backends.flash_attn import FlashAttentionBackend

        vllm_config = create_vllm_config(kv_role="kv_producer")
        vllm_config.parallel_config.prefill_context_parallel_size = pcp_size
        vllm_config.parallel_config.decode_context_parallel_size = dcp_size
        with (
            patch(
                "vllm.distributed.kv_transfer.kv_connector.v1.nixl."
                "base_worker.get_current_attn_backends",
                return_value=[FlashAttentionBackend],
            ),
            patch(
                "vllm.distributed.kv_transfer.kv_connector.v1.nixl."
                "base_worker.get_pcp_group"
            ) as mock_get_pcp_group,
        ):
            mock_get_pcp_group.return_value.rank_in_group = pcp_rank
            connector = NixlConnector(
                vllm_config,
                KVConnectorRole.WORKER,
                make_kv_cache_config(block_size=16),
            )

        worker = connector.connector_worker
        assert worker is not None
        assert worker.pcp_rank == pcp_rank
        assert worker.pcp_dcp_sharded is (dcp_size > 1)
        assert worker.transfer_tp_rank == (pcp_rank if dcp_size > 1 else 0)
        assert worker.transfer_tp_size == (pcp_size if dcp_size > 1 else 1)

        req_id = "req"
        metadata = NixlConnectorMetadata()
        metadata.reqs_in_batch.add(req_id)
        metadata.reqs_to_send[req_id] = time.perf_counter() + 10
        worker.start_load_kv(metadata)
        assert (req_id in worker._reqs_to_process) == expected_tracked
        assert (req_id in worker._reqs_to_send) == expected_tracked

        payload = MagicMock(spec=NixlHandshakePayload)
        worker.xfer_handshake_metadata = payload
        worker.transfer_topo = MagicMock()
        worker._get_new_notifs = MagicMock(
            side_effect=lambda: {"sent"} if expected_tracked else set()
        )

        expected_payload = payload if expected_tracked else None
        assert connector.get_handshake_metadata() is expected_payload
        done_sending, done_recving = connector.get_finished(set())
        assert done_sending == ({"sent"} if expected_tracked else {req_id})
        assert done_recving == set()
        if not expected_tracked:
            assert connector.get_finished(set()) == (set(), set())

        worker.get_transfer_results = MagicMock(
            return_value=KVConnectorTransferResults(finished_sending={"sent"})
        )
        results = connector.get_transfer_results(set())
        assert results.finished_sending == ({"sent"} if expected_tracked else set())

    @patch(
        "vllm.distributed.kv_transfer.kv_connector.v1.nixl.base_worker.NixlWrapper",
        FakeNixlWrapper,
    )
    def test_multi_xfer_one_engine(
        self,
        default_vllm_config,
        # dist_init is a fixture that initializes the distributed environment.
        dist_init,
    ):
        """Test case where multiple xfers are initiated to the same engine.

        This test triggers the connector to load remote KV for the same
        `request_id`.
        """
        vllm_config = create_vllm_config()

        request_id = "req_id"

        # Test worker role in decode server.
        kv_cache_config = make_kv_cache_config(block_size=16, num_blocks=10)
        connector = NixlConnector(vllm_config, KVConnectorRole.WORKER, kv_cache_config)
        connector.connector_worker = FakeNixlConnectorWorker(
            vllm_config,
            connector.engine_id,
            hand_shake_latency=0,
            kv_cache_config=kv_cache_config,
        )
        assert isinstance(connector.connector_worker.nixl_wrapper, FakeNixlWrapper)
        worker = connector.connector_worker
        # simulate handshake
        worker.dst_xfer_side_handles = {
            FakeNixlConnectorWorker.REMOTE_ENGINE_ID: {0: 1}
        }
        worker.kv_cache_layout = "LBHNC"
        num_xfers = 4
        while True:
            # For the same request_id, initiate multiple xfers across different
            # round of `execute_model` calls.
            metadata = NixlConnectorMetadata()
            if num_xfers > 0:
                num_xfers -= 1
                metadata.add_new_req_to_recv(
                    request_id=request_id,
                    local_block_ids=([num_xfers + 1, num_xfers + 2, num_xfers + 3],),
                    kv_transfer_params={
                        "remote_block_ids": (
                            [
                                num_xfers + 4,
                                num_xfers + 5,
                                num_xfers + 6,
                            ],
                        ),
                        "remote_engine_id": FakeNixlConnectorWorker.REMOTE_ENGINE_ID,
                        "remote_request_id": f"prefill-{request_id}",
                        "remote_host": "localhost",
                        "remote_port": 1234,
                        "remote_tp_size": 1,
                    },
                )
            connector.bind_connector_metadata(metadata)

            # Mimic logic in KVConnectorModelRunnerMixin._get_kv_connector_output.
            dummy_ctx = ForwardContext(
                no_compile_layers={},
                attn_metadata={},
                slot_mapping={},
            )
            _before_load = time.perf_counter()
            connector.start_load_kv(dummy_ctx)
            _after_load = time.perf_counter()
            assert _after_load - _before_load < 0.1, (
                f"start_load_kv took {_after_load - _before_load} seconds"
            )

            # Mimic logic in KVConnectorModelRunnerMixin._get_kv_connector_output.
            _, done_recving = connector.get_finished(finished_req_ids=set())
            if len(done_recving) > 0:
                assert request_id in done_recving
                break

            connector.clear_connector_metadata()

    @patch(
        "vllm.distributed.kv_transfer.kv_connector.v1.nixl.base_worker.NixlWrapper",
        FakeNixlWrapper,
    )
    @pytest.mark.parametrize(
        "decode_tp_size, prefill_tp_size",
        [
            (1, 1),
            (2, 1),
            (4, 2),
            (4, 4),
        ],
    )
    def test_async_load_kv(
        self,
        default_vllm_config,
        # Fixture that initializes the distributed environment.
        dist_init,
        # Simulate consumer-producer TP sizes.
        decode_tp_size,
        prefill_tp_size,
    ):
        """Test that NixlConnector's start_load_kv should be non-blocking."""
        vllm_config = create_vllm_config()
        vllm_config.parallel_config.tensor_parallel_size = decode_tp_size

        # Test worker role in decode server.
        connector = NixlConnector(
            vllm_config, KVConnectorRole.WORKER, make_kv_cache_config(block_size=16)
        )
        connector.connector_worker = FakeNixlConnectorWorker(
            vllm_config, connector.engine_id
        )
        metadata = NixlConnectorMetadata()
        metadata.add_new_req_to_recv(
            request_id="id",
            local_block_ids=([1, 2, 3],),
            kv_transfer_params={
                "remote_block_ids": ([4, 5, 6],),
                "remote_engine_id": FakeNixlConnectorWorker.REMOTE_ENGINE_ID,
                "remote_request_id": "prefill-id",
                "remote_host": "localhost",
                "remote_port": 1234,
                "remote_tp_size": prefill_tp_size,
            },
        )
        connector.bind_connector_metadata(metadata)

        timeout = 2.5
        start = time.perf_counter()
        while time.perf_counter() - start < timeout:
            dummy_ctx = ForwardContext(
                no_compile_layers={},
                attn_metadata={},
                slot_mapping={},
            )
            _before_load = time.perf_counter()
            connector.start_load_kv(dummy_ctx)
            _after_load = time.perf_counter()
            assert _after_load - _before_load < 0.1, (
                f"start_load_kv took {_after_load - _before_load} seconds"
            )
            time.sleep(0.5)  # backoff for the async handshake to complete.
            connector.bind_connector_metadata(NixlConnectorMetadata())
            _, done_recving = connector.get_finished(finished_req_ids=set())
            if len(done_recving) > 0:
                return
        raise TimeoutError("Took too long to complete async handshake.")

    @patch(
        "vllm.distributed.kv_transfer.kv_connector.v1.nixl.base_worker.NixlWrapper",
        FakeNixlWrapper,
    )
    @pytest.mark.parametrize("local_tp_size", [1, 2])
    def test_prefill_tp_size_greater_than_decode_tp_size(
        self, local_tp_size: int, default_vllm_config, dist_init, monkeypatch
    ):
        """Verify remote TP > local TP handshake succeeds with different
        remote configurations.
        """
        monkeypatch.setattr(
            "vllm.distributed.kv_transfer.kv_connector.v1.nixl.base_worker.get_tensor_model_parallel_world_size",
            lambda: local_tp_size,
        )

        vllm_config = create_vllm_config()

        connector = NixlConnector(
            vllm_config, KVConnectorRole.WORKER, make_kv_cache_config(block_size=16)
        )
        connector.connector_worker = FakeNixlConnectorWorker(
            vllm_config, connector.engine_id, hand_shake_latency=0
        )
        worker = connector.connector_worker

        # Minimal local registration params used by add_remote_agent
        worker.slot_size_per_layer = [4096]
        worker.block_len_per_layer = [4096 * worker.block_size]
        worker.num_blocks = 1
        worker.dst_num_blocks[worker.engine_id] = worker.num_blocks
        worker.src_blocks_data = np.array(
            [(0, worker.block_len_per_layer[0], worker.tp_rank)],
            dtype=np.uint64,
        )
        worker.num_descs = len(worker.src_blocks_data)

        def check_handshake(remote_tp_size: int):
            tp_ratio = remote_tp_size // local_tp_size
            assert set(remote_agents.keys()) == {(0, r) for r in range(tp_ratio)}

            remote_engine_id = worker.REMOTE_ENGINE_ID
            remote_info = worker.transfer_topo.get_engine_info(remote_engine_id)
            assert remote_info.remote_tp_size == remote_tp_size
            assert -tp_ratio == worker.transfer_topo.tp_ratio(remote_tp_size)
            # ensure src_xfer_handles_by_tp_ratio is populated with tpratio chunks
            split_key = (-tp_ratio, worker.block_size)
            assert split_key in worker.src_xfer_handles_by_tp_ratio
            assert len(worker.src_xfer_handles_by_tp_ratio[split_key]) == tp_ratio
            assert remote_engine_id in worker.dst_xfer_side_handles
            assert set(worker.dst_xfer_side_handles[remote_engine_id].keys()) == set(
                range(tp_ratio)
            )

        remote_agents, _ = worker._nixl_handshake(
            host="localhost",
            port=1234,
            remote_tp_size=4,
            expected_engine_id=worker.REMOTE_ENGINE_ID,
        )
        check_handshake(4)

        # NOTE flexibility: a second remote with higher number of ranks is
        # discovered. This is not a scenario we actively support right now, but
        # the connector allows it.
        worker.REMOTE_ENGINE_ID = "remote_engine_2"
        remote_agents, _ = worker._nixl_handshake(
            host="localhost",
            port=1234,
            remote_tp_size=6,
            expected_engine_id=worker.REMOTE_ENGINE_ID,
        )
        check_handshake(6)

    @patch(
        "vllm.distributed.kv_transfer.kv_connector.v1.nixl.base_worker.NixlWrapper",
        FakeNixlWrapper,
    )
    def test_prefill_tp_size_greater_than_decode_tp_size_mla(
        self, default_vllm_config, dist_init
    ):
        """Verify remote TP > local TP handshake succeeds with different
        remote configurations for an MLA model.
        """
        vllm_config = create_vllm_config()
        d_tp_size = 1
        p_tp_size = 2

        # Build two separate connectors/workers to emulate P TP=2 ranks.
        conn_p0 = NixlConnector(
            vllm_config, KVConnectorRole.WORKER, make_kv_cache_config(block_size=16)
        )
        conn_p1 = NixlConnector(
            vllm_config, KVConnectorRole.WORKER, make_kv_cache_config(block_size=16)
        )
        conn_p0.connector_worker = FakeNixlConnectorWorker(
            vllm_config, conn_p0.engine_id, hand_shake_latency=0
        )
        conn_p1.connector_worker = FakeNixlConnectorWorker(
            vllm_config, conn_p1.engine_id, hand_shake_latency=0
        )

        # Force P world size to 2 for both workers and emulate distinct tp_ranks.
        # Also enable MLA path so that expected_finished_count is updated.
        for rank, worker in enumerate(
            (conn_p0.connector_worker, conn_p1.connector_worker)
        ):
            worker.world_size = p_tp_size
            worker.transfer_topo.tp_size = p_tp_size
            worker.tp_rank = rank
            worker.use_mla = True

        req_id = "req-ep-dp2-p0"
        now = time.perf_counter()
        # Register a request on P that is waiting for consumers to read
        # (both workers track it).
        conn_p0.connector_worker._reqs_to_send[req_id] = now + 10.0
        conn_p0.connector_worker._reqs_to_process.add(req_id)
        conn_p1.connector_worker._reqs_to_send[req_id] = now + 10.0
        conn_p1.connector_worker._reqs_to_process.add(req_id)

        # Simulate a read notification coming from D with (tp=1, dp=2).
        notif = f"{req_id}:{d_tp_size}".encode()
        # D0-0->P0 notif
        conn_p0.connector_worker.nixl_wrapper.get_new_notifs = lambda: {
            "agent": [notif]
        }  # type: ignore[method-assign]
        conn_p1.connector_worker.nixl_wrapper.get_new_notifs = lambda: {
            "agent": [notif]
        }  # type: ignore[method-assign]

        # Trigger notification processing via get_finished().
        done_sending0, _ = conn_p0.get_finished(finished_req_ids=set())
        done_sending1, _ = conn_p1.get_finished(finished_req_ids=set())
        assert req_id in done_sending0 and req_id in done_sending1

        # E2E aggregation: ensure the aggregated output marks the request
        # as finished using the connector's expected_finished_count.
        from vllm.v1.outputs import KVConnectorOutput, ModelRunnerOutput

        aggregator = KVOutputAggregator.from_connector(conn_p0, world_size=2)

        out0 = ModelRunnerOutput(
            req_ids=[req_id],
            req_id_to_index={req_id: 0},
            sampled_token_ids=[[0]],
            logprobs=None,
            prompt_logprobs_dict={},
            pooler_output=[None],
            kv_connector_output=KVConnectorOutput(
                finished_sending=done_sending0,
                finished_recving=None,
            ),
        )
        out1 = ModelRunnerOutput(
            req_ids=[req_id],
            req_id_to_index={req_id: 0},
            sampled_token_ids=[[0]],
            logprobs=None,
            prompt_logprobs_dict={},
            pooler_output=[None],
            kv_connector_output=KVConnectorOutput(
                finished_sending=done_sending1,
                finished_recving=None,
            ),
        )
        aggregated = aggregator.aggregate([out0, out1], output_rank=0)
        assert aggregated.kv_connector_output is not None
        assert aggregated.kv_connector_output.finished_sending == {req_id}

        # Producers cleaned up state for the finished request.
        assert req_id not in conn_p0.connector_worker._reqs_to_send
        assert req_id not in conn_p0.connector_worker._reqs_to_process
        assert req_id not in conn_p1.connector_worker._reqs_to_send
        assert req_id not in conn_p1.connector_worker._reqs_to_process

    @patch(
        "vllm.distributed.kv_transfer.kv_connector.v1.nixl.base_worker.NixlWrapper",
        FakeNixlWrapper,
    )
    def test_concurrent_load_kv(
        self,
        default_vllm_config,
        # dist_init is a fixture that initializes the distributed environment.
        dist_init,
    ):
        """Test that multiple start_load_kv calls should occur concurrently."""
        vllm_config = create_vllm_config()

        # Test worker role in decode server.
        connector = NixlConnector(
            vllm_config, KVConnectorRole.WORKER, make_kv_cache_config(block_size=16)
        )
        connector.connector_worker = FakeNixlConnectorWorker(
            vllm_config, connector.engine_id
        )
        # Register (mocked) local xfer handler
        # worker = connector.connector_worker
        # worker.src_xfer_handles_by_block_size = {worker.block_size: 1}
        metadata = NixlConnectorMetadata()
        total_reqs = 5
        for i in range(total_reqs):
            metadata.add_new_req_to_recv(
                request_id=f"id_{i}",
                local_block_ids=([1, 2, 3],),
                kv_transfer_params={
                    "remote_block_ids": ([4, 5, 6],),
                    "remote_engine_id": FakeNixlConnectorWorker.REMOTE_ENGINE_ID,
                    "remote_request_id": f"prefill-id-{i}",
                    "remote_host": "localhost",
                    "remote_port": 1234,
                    "remote_tp_size": 1,
                },
            )
        connector.bind_connector_metadata(metadata)

        timeout = 2.5 * total_reqs
        cnt_finished_reqs = 0
        start = time.perf_counter()
        while time.perf_counter() - start < timeout:
            dummy_ctx = ForwardContext(
                no_compile_layers={},
                attn_metadata={},
                slot_mapping={},
            )
            _before_load = time.perf_counter()
            connector.start_load_kv(dummy_ctx)
            _after_load = time.perf_counter()
            assert _after_load - _before_load < 0.1, (
                f"start_load_kv took {_after_load - _before_load} seconds"
            )
            time.sleep(0.5)  # backoff for the async handshake to complete.
            connector.bind_connector_metadata(NixlConnectorMetadata())
            _, done_recving = connector.get_finished(finished_req_ids=set())
            if len(done_recving) > 0:
                cnt_finished_reqs += len(done_recving)
                if cnt_finished_reqs == total_reqs:
                    return
        raise TimeoutError("Took too long to complete async handshake.")

    @patch(
        "vllm.distributed.kv_transfer.kv_connector.v1.nixl.base_worker.NixlWrapper",
        FakeNixlWrapper,
    )
    def test_handshake_fails_on_kv_cache_layout_mismatch(
        self, default_vllm_config, dist_init
    ):
        """Verify that adding a remote agent fails if kv_cache_layout differs.
        This test is only relevant for heterogeneous TP.
        """
        vllm_config = create_vllm_config()

        # Mock TP world size to 2 to force heterogeneous TP when
        # remote_tp_size=1
        with patch(
            "vllm.distributed.kv_transfer.kv_connector.v1.nixl.base_worker.get_tensor_model_parallel_world_size",  # noqa: E501
            return_value=2,
        ):
            # Initialize connector and worker (with fake NIXL wrapper)
            connector = NixlConnector(
                vllm_config, KVConnectorRole.WORKER, make_kv_cache_config(block_size=16)
            )
            connector.connector_worker = FakeNixlConnectorWorker(
                vllm_config, connector.engine_id, hand_shake_latency=0
            )
            worker = connector.connector_worker

            # Minimal local registration params used by add_remote_agent
            worker.slot_size_per_layer = [4096]
            worker.block_len_per_layer = [4096 * worker.block_size]
            worker.num_blocks = 1
            worker.dst_num_blocks[worker.engine_id] = worker.num_blocks

            # Metadata with different kv_cache_layout than local worker
            mismatched_layout = (
                "LBHNC" if worker.kv_cache_layout != "LBHNC" else "LBNHC"
            )
            meta = NixlAgentMetadata(
                engine_id=FakeNixlConnectorWorker.REMOTE_ENGINE_ID,
                agent_metadata=FakeNixlWrapper.AGENT_METADATA,
                kv_caches_base_addr=[0],
                device_id=0,
                num_blocks=1,
                block_lens=worker.block_len_per_layer,
                block_strides=worker.block_len_per_layer,
                kv_cache_layout=mismatched_layout,
                block_size=worker.block_size,
                ssm_sizes=(0, 0),
                attn_backend_name=worker.backend_name,
                physical_blocks_per_logical_kv_block=1,
            )

            with pytest.raises(RuntimeError):
                # mismatched layout is expected to fail
                worker.add_remote_agent(meta, remote_tp_size=2)
                worker.add_remote_agent(meta, remote_tp_size=1)

    @patch(
        "vllm.distributed.kv_transfer.kv_connector.v1.nixl.base_worker.NixlWrapper",
        FakeNixlWrapper,
    )
    def test_handshake_succeed_on_kv_cache_layout_mismatch_with_experimental(
        self, default_vllm_config, dist_init
    ):
        """Verify that adding a remote agent fails if kv_cache_layout differs.
        This test is only relevant for heterogeneous TP.
        """
        vllm_config = create_vllm_config(enable_permute_local_kv=True)

        # Mock TP world size to 2 to force heterogeneous TP when
        # remote_tp_size=1
        with patch(
            "vllm.distributed.kv_transfer.kv_connector.v1.nixl.base_worker.get_tensor_model_parallel_world_size",  # noqa: E501
            return_value=2,
        ):
            # Initialize connector and worker (with fake NIXL wrapper)
            connector = NixlConnector(
                vllm_config, KVConnectorRole.WORKER, make_kv_cache_config(block_size=16)
            )
            connector.connector_worker = FakeNixlConnectorWorker(
                vllm_config,
                connector.engine_id,
                hand_shake_latency=0,
                kv_cache_layout="LBNHC",
            )
            worker = connector.connector_worker

            # Minimal local registration params used by add_remote_agent
            worker.slot_size_per_layer = [2048]
            worker.block_len_per_layer = [2048 * worker.block_size]
            worker.num_blocks = 1
            worker.dst_num_blocks[worker.engine_id] = worker.num_blocks

            # Metadata with different kv_cache_layout than local worker
            # prefill TP=1, decode TP=2, remote block_lens is double to local
            remote_block_lens = [i * 2 for i in worker.block_len_per_layer]
            meta = NixlAgentMetadata(
                engine_id=FakeNixlConnectorWorker.REMOTE_ENGINE_ID,
                agent_metadata=FakeNixlWrapper.AGENT_METADATA,
                kv_caches_base_addr=[0],
                device_id=0,
                num_blocks=1,
                block_lens=remote_block_lens,
                block_strides=remote_block_lens,
                kv_cache_layout="LBHNC",
                block_size=worker.block_size,
                ssm_sizes=(0, 0),
                attn_backend_name=worker.backend_name,
                physical_blocks_per_logical_kv_block=1,
            )

            # We don't check layout for homogeneous TP and MLA for now, as the
            # whole block is moved.
            worker.add_remote_agent(meta, remote_tp_size=1)

    @patch(
        "vllm.distributed.kv_transfer.kv_connector.v1.nixl.base_worker.NixlWrapper",
        FakeNixlWrapper,
    )
    def test_hybrid_mamba_attention_remote_descs_use_packed_head_slices(
        self, default_vllm_config, dist_init
    ):
        worker = FakeNixlConnectorWorker(
            create_vllm_config(), "engine", hand_shake_latency=0
        )

        remote_block_len = 2048
        local_block_len = remote_block_len // 2
        worker.block_len_per_layer = [local_block_len]
        worker._region_is_mla = [False]
        worker.num_blocks = 1
        worker.num_regions = 1
        worker._has_mamba = True
        worker._mamba_ssm_size = (128, 256)
        worker.transfer_topo = TransferTopology(
            tp_rank=1,
            tp_size=2,
            block_size=worker.block_size,
            engine_id=worker.engine_id,
            is_mla=False,
            is_mamba=True,
            total_num_kv_heads=2,
            attn_backends=worker.attn_backends,
            tensor_shape=None,
        )
        plan = MagicMock(
            source_ranks_per_group=((0,), (0,)),
            rank_offset_factor=1,
        )
        meta = NixlAgentMetadata(
            engine_id=FakeNixlConnectorWorker.REMOTE_ENGINE_ID,
            agent_metadata=FakeNixlWrapper.AGENT_METADATA,
            kv_caches_base_addr=[0x1000],
            device_id=0,
            num_blocks=1,
            block_lens=[remote_block_len],
            block_strides=[remote_block_len],
            kv_cache_layout="HND",
            block_size=worker.block_size,
            ssm_sizes=(0, 0),
            attn_backend_name=worker.backend_name,
            physical_blocks_per_logical_kv_block=1,
            region_num_blocks=None,
        )

        assert worker._build_fa_remote(plan, meta, block_size_ratio=1).tolist() == [
            [0x1000 + local_block_len, local_block_len, 0]
        ]

    @patch(
        "vllm.distributed.kv_transfer.kv_connector.v1.nixl.base_worker.NixlWrapper",
        FakeNixlWrapper,
    )
    def test_handshake_mixed_fa_mla_hetero_tp(self, default_vllm_config, dist_init):
        """Mixed full-attn (SPLIT) + MLA (REPLICATE) single KV group under
        heterogeneous TP must NOT raise (previously a NotImplementedError),
        and the per-region gate must still reject a wrong block_len.
        """
        vllm_config = create_vllm_config()
        with patch(
            "vllm.distributed.kv_transfer.kv_connector.v1.nixl.base_worker.get_tensor_model_parallel_world_size",  # noqa: E501
            return_value=2,
        ):
            connector = NixlConnector(
                vllm_config, KVConnectorRole.WORKER, make_kv_cache_config(block_size=16)
            )
            connector.connector_worker = FakeNixlConnectorWorker(
                vllm_config, connector.engine_id, hand_shake_latency=0
            )
            worker = connector.connector_worker

            # Region 0: full-attn (SPLIT). Region 1: MLA (REPLICATE).
            fa_len = 4096 * worker.block_size
            idx_len = 512 * worker.block_size
            worker.slot_size_per_layer = [4096, 512]
            worker.block_len_per_layer = [fa_len, idx_len]
            worker._region_is_mla = [False, True]
            worker.num_blocks = 1
            worker.dst_num_blocks[worker.engine_id] = worker.num_blocks
            worker.src_blocks_data = np.array(
                [
                    (0, fa_len, worker.tp_rank),
                    (0, idx_len, worker.tp_rank),
                ],
                dtype=np.uint64,
            )
            worker.num_descs = len(worker.src_blocks_data)

            # D_TP=2, P_TP=1 -> tp_ratio=2. SPLIT region scales by tp_ratio;
            # REPLICATE region is unchanged.
            tp_ratio = 2
            meta = NixlAgentMetadata(
                engine_id=FakeNixlConnectorWorker.REMOTE_ENGINE_ID,
                agent_metadata=FakeNixlWrapper.AGENT_METADATA,
                kv_caches_base_addr=[0, 0],
                device_id=0,
                num_blocks=1,
                block_lens=[fa_len * tp_ratio, idx_len],
                block_strides=[fa_len * tp_ratio, idx_len],
                kv_cache_layout=worker.kv_cache_layout,
                block_size=worker.block_size,
                ssm_sizes=(0, 0),
                attn_backend_name=worker.backend_name,
                physical_blocks_per_logical_kv_block=1,
            )
            worker.add_remote_agent(meta, remote_tp_size=1)
            assert (
                FakeNixlConnectorWorker.REMOTE_ENGINE_ID in worker.dst_xfer_side_handles
            )
            # Gate rejects an MLA region wrongly scaled by tp_ratio.
            worker2 = FakeNixlConnectorWorker(
                vllm_config, connector.engine_id, hand_shake_latency=0
            )
            worker2.block_len_per_layer = [fa_len, idx_len]
            worker2._region_is_mla = [False, True]
            worker2.num_blocks = 1
            worker2.dst_num_blocks[worker2.engine_id] = worker2.num_blocks
            bad_meta = NixlAgentMetadata(
                engine_id=FakeNixlConnectorWorker.REMOTE_ENGINE_ID,
                agent_metadata=FakeNixlWrapper.AGENT_METADATA,
                kv_caches_base_addr=[0, 0],
                device_id=0,
                num_blocks=1,
                # WRONG: MLA region scaled by tp_ratio (it should be replicated).
                block_lens=[fa_len * tp_ratio, idx_len * tp_ratio],
                block_strides=[fa_len * tp_ratio, idx_len * tp_ratio],
                kv_cache_layout=worker2.kv_cache_layout,
                block_size=worker2.block_size,
                ssm_sizes=(0, 0),
                attn_backend_name=worker2.backend_name,
                physical_blocks_per_logical_kv_block=1,
            )
            with pytest.raises(AssertionError):
                worker2.add_remote_agent(bad_meta, remote_tp_size=1)

    @patch(
        "vllm.distributed.kv_transfer.kv_connector.v1.nixl.base_worker.NixlWrapper",
        FakeNixlWrapper,
    )
    def test_handshake_validates_gqa_replicated_block_len(
        self, default_vllm_config, dist_init
    ):
        """Regression test for #45330.

        When tp_size > total_num_kv_heads, GQA replication caps per-rank
        KV heads at 1, so block_len stops scaling with 1/tp.  With 8 KV
        heads and D_TP=16 pulling from P_TP=8, both sides hold one head
        per rank and report the *same* block_len; the old validation
        expected local_block_len * tp_ratio and rejected the valid
        handshake.
        """
        vllm_config = create_vllm_config()

        with patch(
            "vllm.distributed.kv_transfer.kv_connector.v1.nixl.base_worker.get_tensor_model_parallel_world_size",  # noqa: E501
            return_value=16,
        ):
            connector = NixlConnector(
                vllm_config,
                KVConnectorRole.WORKER,
                make_kv_cache_config(block_size=16),
            )
            connector.connector_worker = FakeNixlConnectorWorker(
                vllm_config, connector.engine_id, hand_shake_latency=0
            )
            worker = connector.connector_worker

            worker.transfer_topo.total_num_kv_heads = 8
            worker.transfer_topo.local_physical_heads = 1
            worker.kv_cache_layout = "LBHNC"

            worker.slot_size_per_layer = [4096]
            worker.block_len_per_layer = [4096 * worker.block_size]
            worker.num_blocks = 1
            worker.dst_num_blocks[worker.engine_id] = worker.num_blocks

            # Remote P with TP=8 also has 1 head/rank -> identical
            # block_len despite tp_ratio == 2.
            meta = NixlAgentMetadata(
                engine_id=FakeNixlConnectorWorker.REMOTE_ENGINE_ID,
                agent_metadata=FakeNixlWrapper.AGENT_METADATA,
                kv_caches_base_addr=[0],
                device_id=0,
                num_blocks=1,
                block_lens=list(worker.block_len_per_layer),
                block_strides=list(worker.block_len_per_layer),
                kv_cache_layout="LBHNC",
                block_size=worker.block_size,
                ssm_sizes=(0, 0),
                attn_backend_name=worker.backend_name,
                physical_blocks_per_logical_kv_block=1,
            )

            # Must validate cleanly (used to raise AssertionError).
            worker.add_remote_agent(meta, remote_tp_size=8)

    @patch(
        "vllm.distributed.kv_transfer.kv_connector.v1.nixl.base_worker.NixlWrapper",
        FakeNixlWrapper,
    )
    def test_handshake_rejects_wrong_block_len_without_gqa_replication(
        self, default_vllm_config, dist_init
    ):
        """Ensure the head-ratio validation still rejects genuinely wrong
        block_lens when GQA replication is NOT in effect (32 KV heads,
        D_TP=4, P_TP=2: head_ratio=4, both sides have >1 head/rank).
        """
        vllm_config = create_vllm_config()

        with patch(
            "vllm.distributed.kv_transfer.kv_connector.v1.nixl.base_worker.get_tensor_model_parallel_world_size",  # noqa: E501
            return_value=4,
        ):
            connector = NixlConnector(
                vllm_config,
                KVConnectorRole.WORKER,
                make_kv_cache_config(block_size=16),
            )
            connector.connector_worker = FakeNixlConnectorWorker(
                vllm_config, connector.engine_id, hand_shake_latency=0
            )
            worker = connector.connector_worker

            worker.transfer_topo.total_num_kv_heads = 32
            worker.transfer_topo.local_physical_heads = 8  # 32 // 4
            worker.kv_cache_layout = "LBHNC"

            slot_size = 4096
            worker.slot_size_per_layer = [slot_size]
            worker.block_len_per_layer = [slot_size * worker.block_size]
            worker.num_blocks = 1
            worker.dst_num_blocks[worker.engine_id] = worker.num_blocks

            # Remote P_TP=2 has 16 heads/rank -> head_ratio = 16/8 = 2.
            # Correct remote block_len = local * 2.  Send local * 1
            # (wrong) to verify rejection.
            bad_meta = NixlAgentMetadata(
                engine_id=FakeNixlConnectorWorker.REMOTE_ENGINE_ID,
                agent_metadata=FakeNixlWrapper.AGENT_METADATA,
                kv_caches_base_addr=[0],
                device_id=0,
                num_blocks=1,
                block_lens=list(worker.block_len_per_layer),
                block_strides=list(worker.block_len_per_layer),
                kv_cache_layout="LBHNC",
                block_size=worker.block_size,
                ssm_sizes=(0, 0),
                attn_backend_name=worker.backend_name,
                physical_blocks_per_logical_kv_block=1,
            )

            with pytest.raises(AssertionError):
                worker.add_remote_agent(bad_meta, remote_tp_size=2)


# NOTE: resource cleanup in mp backend is a bit finicky, so the order in which
# we put here is important. First run ray, it will clean up the resources, then
# the rest of the tests.
@patch(
    "vllm.distributed.kv_transfer.kv_connector.v1.nixl.base_worker.NixlWrapper",
    FakeNixlWrapper,
)
def test_kv_connector_stats(default_vllm_config, dist_init):
    """Test that KV transfer stats are properly recorded and retrieved."""
    vllm_config = create_vllm_config()

    # Test worker role in decode server.
    connector = NixlConnector(
        vllm_config, KVConnectorRole.WORKER, make_kv_cache_config(block_size=16)
    )
    connector.connector_worker = FakeNixlConnectorWorker(
        vllm_config, connector.engine_id, hand_shake_latency=0
    )

    # Verify that xfer_stats starts empty
    initial_stats = connector.get_kv_connector_stats()
    assert initial_stats is None

    # Create transfer metadata
    request_id = "test_req_for_stats"
    metadata = NixlConnectorMetadata()
    metadata.add_new_req_to_recv(
        request_id=request_id,
        local_block_ids=([0],),
        kv_transfer_params={
            "remote_block_ids": ([0],),
            "remote_engine_id": FakeNixlConnectorWorker.REMOTE_ENGINE_ID,
            "remote_request_id": f"prefill-{request_id}",
            "remote_host": "localhost",
            "remote_port": 1234,
            "remote_tp_size": 1,
        },
    )
    connector.bind_connector_metadata(metadata)

    # Start the transfer
    dummy_ctx = ForwardContext(
        no_compile_layers={},
        attn_metadata={},
        slot_mapping={},
    )
    connector.start_load_kv(dummy_ctx)

    # Verify stats are recorded after transfer is complete
    max_iterations = 2
    # Clear metadata before start_load_kv to prevent reprocessing same request
    connector.bind_connector_metadata(NixlConnectorMetadata())
    for _ in range(max_iterations):
        # Need to call start_load_kv to process completed handshakes
        connector.start_load_kv(dummy_ctx)
        _, done_recving = connector.get_finished(finished_req_ids=set())
        if len(done_recving) > 0 and request_id in done_recving:
            break
        time.sleep(0.1)  # Small delay to allow background handshake to complete
    else:
        assert "Transfer did not complete within expected iterations"

    # Now check that stats were recorded
    stats_after_transfer = connector.get_kv_connector_stats()
    assert isinstance(stats_after_transfer, NixlKVConnectorStats)

    # Verify stats values are recorded
    assert not stats_after_transfer.is_empty()
    assert stats_after_transfer.num_successful_transfers == 1

    # Verify stats are reset after retrieval
    stats_after_reset = connector.get_kv_connector_stats()
    assert stats_after_reset is None


@patch(
    "vllm.distributed.kv_transfer.kv_connector.v1.nixl.base_worker.NixlWrapper",
    FakeNixlWrapper,
)
def test_reqs_to_send_deadline_rebased_to_worker_clock(default_vllm_config, dist_init):
    """reqs_to_send deadlines are stamped with the scheduler process's
    perf_counter, whose epoch differs across processes and (by boot-time
    deltas) across nodes. Without rebasing, a P worker on a node whose
    monotonic clock is ahead of the scheduler's by more than the TTL
    expires the lease on arrival and reports done_sending before D has
    read the blocks — the freed blocks can then be reallocated and the
    remote read pulls another request's data (silent accuracy corruption).
    The worker must anchor the remaining TTL to its own clock.
    """
    vllm_config = create_vllm_config()
    connector = NixlConnector(
        vllm_config, KVConnectorRole.WORKER, make_kv_cache_config(block_size=16)
    )
    connector.connector_worker = FakeNixlConnectorWorker(
        vllm_config, connector.engine_id, hand_shake_latency=0
    )
    worker = connector.connector_worker

    req_id = "req-lease-clock"
    ttl = 480.0
    # Simulate a scheduler whose monotonic clock is 10,000 s behind this
    # worker's (e.g. its node booted much later): the raw deadline is
    # then already far in the past in this worker's clock domain.
    scheduler_clock = time.perf_counter() - 10_000.0

    metadata = NixlConnectorMetadata()
    metadata.reqs_in_batch = {req_id}
    metadata.reqs_to_send = {req_id: scheduler_clock + ttl}
    metadata.scheduler_clock = scheduler_clock
    connector.bind_connector_metadata(metadata)
    dummy_ctx = ForwardContext(
        no_compile_layers={},
        attn_metadata={},
        slot_mapping={},
    )
    connector.start_load_kv(dummy_ctx)

    remaining = worker._reqs_to_send[req_id] - time.perf_counter()
    assert ttl - 5.0 < remaining <= ttl + 5.0

    # The expiry sweep must not release the request.
    done_sending, _ = worker.get_finished()
    assert req_id not in done_sending
    assert req_id in worker._reqs_to_process


def test_kv_connector_stats_aggregation():
    """Test KV transfer stats aggregation across TP ranks using
    KVOutputAggregator (used by MultiprocExecutor).
    """
    # Create KVOutputAggregator for 3 workers (simulating TP=3), same thing
    # done in MultiprocExecutor.execute_model
    aggregator = KVOutputAggregator(expected_finished_count=3)

    # Create stats for multiple workers with different transfer patterns
    worker1_stats = NixlKVConnectorStats()
    worker2_stats = NixlKVConnectorStats()
    worker3_stats = NixlKVConnectorStats()

    # Record different transfers on each worker
    # Worker 1: 2 transfers
    stats = get_default_xfer_telemetry()
    worker1_stats.record_transfer(stats)
    worker1_stats.record_transfer(stats)

    # Worker 2: 1 transfer
    worker2_stats.record_transfer(stats)

    # Worker 3: 3 transfers
    stats = get_default_xfer_telemetry(
        xferDurationS=2, postDurationS=2, totalBytes=2, descCount=2
    )
    worker3_stats.record_transfer(stats)
    worker3_stats.record_transfer(stats)
    worker3_stats.record_transfer(stats)

    # Create ModelRunnerOutput instances for each worker
    worker_outputs = []
    for i, worker_stats in enumerate([worker1_stats, worker2_stats, worker3_stats]):
        output = ModelRunnerOutput(
            req_ids=[f"req_{i}"],
            req_id_to_index={f"req_{i}": 0},
            sampled_token_ids=[[123]],  # dummy token
            logprobs=None,
            prompt_logprobs_dict={},
            pooler_output=[None],
            kv_connector_output=KVConnectorOutput(
                finished_sending=set([f"req_{i}_send"])
                if i < 2
                else None,  # Workers 0,1 finished sending
                finished_recving=set([f"req_{i}_recv"])
                if i > 0
                else None,  # Workers 1,2 finished receiving
                kv_connector_stats=worker_stats,
            ),
        )
        worker_outputs.append(output)

    # Use the real aggregation mechanism (like MultiprocExecutor.execute_model)
    aggregated_output = aggregator.aggregate(worker_outputs, output_rank=0)
    kv_connector_stats = aggregated_output.kv_connector_output.kv_connector_stats
    assert isinstance(kv_connector_stats, NixlKVConnectorStats)
    # Number of total transfers across all workers.
    assert kv_connector_stats.num_successful_transfers == 6
    # Logging proc, call reduce() to get CLI-friendly stats.
    cli_stats = kv_connector_stats.reduce()
    assert cli_stats["Avg xfer time (ms)"] == 1500.0
    assert cli_stats["Avg post time (ms)"] == 1500.0
    assert cli_stats["Avg number of descriptors"] == 1.5
    # Reduced values must be plain Python scalars so CLI logging renders
    # them without numpy reprs (eg np.float64(...)).
    assert all(not isinstance(v, np.generic) for v in cli_stats.values())


def test_kv_connector_stats_failure_grouping():
    """Transfer, handshake and notification failures are reported as one
    transport-failure count, while KV expiry is reported separately: the
    former are sporadic lower-transport-layer events, the latter an
    autoscaler signal."""
    stats = NixlKVConnectorStats()
    assert stats.is_empty()

    stats.record_failed_transfer()
    stats.record_failed_handshake()
    stats.record_failed_notification()
    stats.record_kv_expired_req()
    stats.record_notification_after_expiry()
    assert not stats.is_empty()

    # No successful transfers: latency stats are zero but the failure
    # counts still surface.
    reduced = stats.reduce()
    assert reduced["Num successful transfers"] == 0
    assert reduced["Num failed transfers"] == 3
    assert reduced["Num KV expired reqs"] == 1
    assert reduced["Num notifs after expiry"] == 1


def test_nixl_prom_metrics_group_handshake_with_transfer_failures():
    """vllm:nixl_num_failed_transfers counts handshake and notification
    failures too, while vllm:nixl_num_kv_expired_reqs stays a separate
    counter."""
    from prometheus_client import CollectorRegistry, Counter, Gauge, Histogram

    from vllm.distributed.kv_transfer.kv_connector.v1.nixl.stats import (
        NixlPromMetrics,
    )

    registry = CollectorRegistry()

    class RegistryGauge(Gauge):
        def __init__(self, *args, **kwargs):
            super().__init__(*args, registry=registry, **kwargs)

    class RegistryCounter(Counter):
        def __init__(self, *args, **kwargs):
            super().__init__(*args, registry=registry, **kwargs)

    class RegistryHistogram(Histogram):
        def __init__(self, *args, **kwargs):
            super().__init__(*args, registry=registry, **kwargs)

    vllm_config = create_vllm_config()
    metric_types = {
        Gauge: RegistryGauge,
        Counter: RegistryCounter,
        Histogram: RegistryHistogram,
    }
    prom = NixlPromMetrics(
        vllm_config,
        metric_types,
        labelnames=["engine"],
        per_engine_labelvalues={0: ["engine-0"]},
    )

    stats = NixlKVConnectorStats()
    stats.record_failed_transfer()
    stats.record_failed_handshake()
    stats.record_failed_notification()
    stats.record_kv_expired_req()
    stats.record_notification_after_expiry()
    prom.observe(stats.data, engine_idx=0)

    def counter_value(name: str) -> float:
        for metric in registry.collect():
            for sample in metric.samples:
                if sample.name == name:
                    return sample.value
        raise AssertionError(f"metric {name} not found in registry")

    assert counter_value("vllm:nixl_num_failed_transfers_total") == 3.0
    assert counter_value("vllm:nixl_num_kv_expired_reqs_total") == 1.0
    assert counter_value("vllm:nixl_num_notifications_after_expiry_total") == 1.0


@patch(
    "vllm.distributed.kv_transfer.kv_connector.v1.nixl.base_worker.NixlWrapper",
    FakeNixlWrapper,
)
def test_notification_after_expiry_is_counted(default_vllm_config, dist_init):
    vllm_config = create_vllm_config()
    connector = NixlConnector(
        vllm_config, KVConnectorRole.WORKER, make_kv_cache_config(block_size=16)
    )
    connector.connector_worker = FakeNixlConnectorWorker(
        vllm_config, connector.engine_id, hand_shake_latency=0
    )
    worker = connector.connector_worker
    worker._reqs_to_process.add("known")
    worker._reqs_to_send["known"] = time.perf_counter() + 10
    worker.nixl_wrapper.get_new_notifs = MagicMock(
        return_value={"decode-agent": [b"known:1", b"unknown:1"]}
    )

    assert worker._get_new_notifs() == {"known"}

    stats = connector.get_kv_connector_stats()
    assert isinstance(stats, NixlKVConnectorStats)
    assert stats.data["num_notifications_after_expiry"] == [1]


def test_multi_kv_connector_stats_aggregation():
    """Test MultiKVConnectorStats aggregation across TP ranks using
    KVOutputAggregator (used by MultiprocExecutor).
    """
    aggregator = KVOutputAggregator(expected_finished_count=3)

    from dataclasses import dataclass

    # Mock a KVConnectorStats class for testing aggregation over connectors.
    @dataclass
    class FooKVConnectorStats(KVConnectorStats):
        def reset(self):
            self.data = {"num_foo_transfers": 0}

        def record_transfer(self):
            if "num_foo_transfers" not in self.data:
                self.data["num_foo_transfers"] = 0
            self.data["num_foo_transfers"] += 1

        def is_empty(self) -> bool:
            return self.data["num_foo_transfers"] == 0

        def aggregate(self, other: "FooKVConnectorStats") -> "FooKVConnectorStats":
            if not other.is_empty():
                self.data["num_foo_transfers"] += other.data["num_foo_transfers"]
            return self

    def make_multi_stats(nixl_count: int, foo_count: int) -> MultiKVConnectorStats:
        data: dict[str, KVConnectorStats] = {}
        if nixl_count > 0:
            nixl_stats = NixlKVConnectorStats()
            for _ in range(nixl_count):
                nixl_stats.record_transfer(get_default_xfer_telemetry())
            data["NixlConnector"] = nixl_stats
        if foo_count > 0:
            foo_stats = FooKVConnectorStats()
            for _ in range(foo_count):
                foo_stats.record_transfer()
            data["FooConnector"] = foo_stats
        return MultiKVConnectorStats(data=data)

    # Create heterogeneous stats across 3 workers
    worker_patterns = [(2, 1), (3, 0), (0, 5)]  # (Nixl, Foo)

    worker_outputs: list[ModelRunnerOutput] = []
    for i, (nixl_count, foo) in enumerate(worker_patterns):
        stats = make_multi_stats(nixl_count, foo)
        output = ModelRunnerOutput(
            req_ids=[f"req_{i}"],
            req_id_to_index={f"req_{i}": 0},
            sampled_token_ids=[[123]],
            logprobs=None,
            prompt_logprobs_dict={},
            pooler_output=[None],
            kv_connector_output=KVConnectorOutput(
                finished_sending=set([f"req_{i}_send"]) if i < 2 else None,
                finished_recving=set([f"req_{i}_recv"]) if i > 0 else None,
                kv_connector_stats=stats,
            ),
        )
        worker_outputs.append(output)

    aggregated_output = aggregator.aggregate(worker_outputs, output_rank=0)
    kv_connector_stats = aggregated_output.kv_connector_output.kv_connector_stats
    assert isinstance(kv_connector_stats, MultiKVConnectorStats)

    # Validate per-connector totals across workers
    assert isinstance(kv_connector_stats["NixlConnector"], NixlKVConnectorStats)
    assert kv_connector_stats["NixlConnector"].num_successful_transfers == 5
    assert isinstance(kv_connector_stats["FooConnector"], FooKVConnectorStats)
    assert kv_connector_stats["FooConnector"].data["num_foo_transfers"] == 6


@patch(
    "vllm.distributed.kv_transfer.kv_connector.v1.nixl.base_worker.NixlWrapper",
    FakeNixlWrapper,
)
def test_scheduler_kv_connector_stats_aggregation():
    """Test scheduler and worker KV connector stats aggregation."""
    from vllm.v1.core.sched.output import SchedulerOutput

    scheduler = create_scheduler(create_vllm_config())

    # Worker stats with transfer metrics
    worker_stats = NixlKVConnectorStats()
    worker_stats.record_transfer(get_default_xfer_telemetry())

    # Scheduler stats with custom metric (needs dummy transfer to avoid being skipped)
    scheduler_stats = NixlKVConnectorStats()
    scheduler_stats.data.update(
        {  # dummy transfer just for testing, to bypass is_empty() check
            "transfer_duration": [0],
            "post_duration": [0],
            "bytes_transferred": [0],
            "num_descriptors": [0],
        }
    )

    # Mock the scheduler connector's stats method
    scheduler.connector.get_kv_connector_stats = lambda: MultiKVConnectorStats(
        data={"NixlConnector": scheduler_stats}
    )

    model_output = ModelRunnerOutput(
        req_ids=["req_0"],
        req_id_to_index={"req_0": 0},
        sampled_token_ids=[[123]],
        logprobs=None,
        prompt_logprobs_dict={},
        pooler_output=[None],
        kv_connector_output=KVConnectorOutput(
            kv_connector_stats=MultiKVConnectorStats(
                data={"NixlConnector": worker_stats}
            )
        ),
    )
    scheduler_output = SchedulerOutput(
        scheduled_new_reqs=[],
        scheduled_cached_reqs=None,
        num_scheduled_tokens={"req_0": 1},
        total_num_scheduled_tokens=1,
        scheduled_spec_decode_tokens={},
        scheduled_encoder_inputs={},
        num_common_prefix_blocks=[0],
        finished_req_ids=set(),
        free_encoder_mm_hashes=[],
    )

    engine_core_outputs = scheduler.update_from_output(scheduler_output, model_output)

    final_stats = next(
        iter(engine_core_outputs.values())
    ).scheduler_stats.kv_connector_stats
    # The scheduler stats payload carries serialized per-connector dicts.
    nixl_stats = final_stats["NixlConnector"]
    assert len(nixl_stats["transfer_duration"]) == 2


@pytest.mark.parametrize("distributed_executor_backend", ["ray", None])
@patch(
    "vllm.distributed.kv_transfer.kv_connector.v1.nixl.base_worker.NixlWrapper",
    FakeNixlWrapper,
)
def test_abort_timeout_on_prefiller(monkeypatch, distributed_executor_backend):
    """Test lifecycle of an aborted Remote Prefill request hitting the timeout.
    -----> P
            |  {process request}
     <-/--- |  {result is NOT delivered, eg proxy is down}
            |
            |
            |  {eventually free blocks}
    """
    model_name = "Qwen/Qwen3-0.6B"
    timeout = 6
    kv_transfer_config = KVTransferConfig(
        kv_connector="NixlConnector",
        kv_role="kv_consumer",
        kv_connector_extra_config={"kv_lease_duration": timeout},
    )
    llm_kwargs = {
        "model": model_name,
        "enforce_eager": True,
        "gpu_memory_utilization": 0.5,
        "kv_transfer_config": kv_transfer_config,
        "distributed_executor_backend": distributed_executor_backend,
    }

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

    def run_test_and_cleanup():
        llm = LLM(**llm_kwargs)
        try:
            _run_abort_timeout_test(llm, timeout)
        finally:
            llm.llm_engine.engine_core.shutdown()

    # Build runtime_env only if we're using Ray
    if distributed_executor_backend == "ray":
        with _make_fake_nixl_pkg() as working_dir:
            runtime_env = {
                "working_dir": working_dir,  # ship fake nixl package
                "env_vars": {
                    "NIXL_TELEMETRY_ENABLE": "1",
                },
            }
            # On XPU/ROCm, vLLM expects Ray's device key to be "GPU".
            # Explicitly reserving GPU resources here prevents false negatives
            # when Ray cannot auto-detect accelerator resources in test envs.
            ray_init_kwargs: dict[str, Any] = {"runtime_env": runtime_env}
            if not current_platform.is_cuda():
                ray_init_kwargs["num_gpus"] = 1
            ray.init(**ray_init_kwargs)
            try:
                run_test_and_cleanup()
            finally:
                ray.shutdown()
    else:
        run_test_and_cleanup()


class RequestIdMapper:
    """Helper class to map external request IDs to internal request IDs."""

    def __init__(self, output_processor: OutputProcessor):
        self.req_id_mapping: dict[str, str] = {}
        self.original_add_request = output_processor.add_request
        output_processor.add_request = self._add_request

    def _add_request(self, request: EngineCoreRequest, *args, **kwargs):
        self.req_id_mapping[request.external_req_id] = request.request_id
        return self.original_add_request(request, *args, **kwargs)

    def __call__(self, external_req_id: str) -> str:
        return self.req_id_mapping[external_req_id]


def _run_abort_timeout_test(llm: LLM, timeout: int):
    """Helper function to run the abort timeout test logic."""
    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
    req_to_blocks = scheduler.kv_cache_manager.coordinator.single_type_managers[
        0
    ].req_to_blocks

    id_mapper = RequestIdMapper(llm.llm_engine.output_processor)

    def req_id(outputs: list[RequestOutput]) -> str:
        assert len(outputs) == 1
        return id_mapper(outputs[0].request_id)

    padding = "Just making this request a little longer so that we're sure "
    "we're not hitting the small-request lower bound beneath which we don't "
    "actually trigger the whole kv transfer, but rather just recompute the "
    "blocks on D."
    req0_id = req_id(
        llm.generate([f"What is the capital of Japan? {padding}"], sampling_params)
    )

    # Request finished but not freed
    assert req0_id in scheduler.finished_req_ids and req0_id in req_to_blocks
    # Some other request, 0 still not freed
    req1_id = req_id(
        llm.generate([f"What is the capital of Italy? {padding}"], sampling_params)
    )
    assert req0_id in req_to_blocks
    assert req1_id in scheduler.finished_req_ids and req1_id in req_to_blocks

    # Wait for timeout and trigger another scheduler loop
    time.sleep(timeout)
    _ = llm.generate([f"What is the capital of France? {padding}"], sampling_params)
    # Request-0 times out and is cleared!
    assert req0_id not in req_to_blocks
    # Need to shutdown the background thread to release NIXL side channel port
    llm.llm_engine.engine_core.shutdown()


def test_mixed_memory_local_descriptors_split_by_memory_type():
    worker = object.__new__(NixlConnectorWorker)
    worker.transfer_topo = MagicMock()
    worker.block_size = 16
    worker.engine_id = "local"
    worker.tp_rank = 0
    worker.device_id = 3
    worker.kv_caches_base_addr = {"local": {0: [100, 200]}}
    worker._has_mamba = False
    worker._mixed_mem_types = True
    worker.region_mem_types = ["DRAM", "VRAM"]
    worker.region_num_blocks = [2, 2]
    worker._desc_is_dram_by_block_size = {}
    worker._desc_pos_by_block_size = {}
    worker._dram_src_handles_by_block_size = {}
    worker.nixl_memory_type = "VRAM"
    worker._build_fa_local = MagicMock(  # type: ignore[method-assign]
        return_value=np.array(
            [
                [100, 10, 3],
                [110, 10, 3],
                [200, 10, 3],
                [210, 10, 3],
            ]
        )
    )
    worker.nixl_wrapper = MagicMock()
    worker.nixl_wrapper.get_xfer_descs.side_effect = lambda blocks, memory_type: (
        memory_type,
        blocks,
    )
    worker.nixl_wrapper.prep_xfer_dlist.side_effect = [11, 22]

    handle, blocks = worker.register_local_xfer_handler(worker.block_size)

    assert handle == 22
    assert worker._dram_src_handles_by_block_size[worker.block_size] == 11
    assert [block[2] for block in blocks] == [0, 0, 3, 3]
    memory_types = [
        call.args[1] for call in worker.nixl_wrapper.get_xfer_descs.call_args_list
    ]
    assert memory_types == ["DRAM", "VRAM"]


@pytest.fixture
def recv_worker():
    """Receive lifecycle state without distributed or device initialization."""
    worker = object.__new__(NixlConnectorWorker)
    worker.transfer_topo = MagicMock()
    worker.transfer_topo.block_size_ratio.return_value = 1
    worker.transfer_topo.get_engine_info.return_value = SimpleNamespace(
        remote_block_size=16, remote_physical_blocks_per_logical=1
    )
    worker._physical_blocks_per_logical_kv_block = 1
    worker._recving_metadata = {"request": MagicMock(local_block_ids=([1, 2, 3],))}
    worker._recving_transfers = defaultdict(list)
    worker._failed_recv_reqs = queue.Queue()
    worker._recv_failures = set()
    worker._handshake_lock = threading.RLock()
    worker._handshake_futures = {}
    worker._remote_agents = {}
    worker._engine_by_address = {}
    worker._replicated_pcp_done_sending = set()
    worker._invalid_block_ids = queue.Queue()
    worker._pending_recv_notifs = {}
    worker._reqs_to_send = {}
    worker._replicated_pcp_done_sending = set()
    worker._is_hma_required = False
    worker._has_mamba = False
    worker.use_host_buffer = False
    worker.enable_permute_local_kv = False
    worker.enable_heterogeneous_attn_post_process = False
    worker.tp_rank = 0
    worker._log_failure = MagicMock()  # type: ignore[method-assign]
    worker.xfer_stats = NixlKVConnectorStats()
    worker.nixl_wrapper = MagicMock()
    worker.nixl_wrapper.get_xfer_telemetry.return_value = get_default_xfer_telemetry()
    return worker


def test_mixed_memory_read_notifies_after_both_transfers_finish(recv_worker):
    worker = recv_worker
    worker._desc_is_dram_by_block_size = {16: np.array([True, True, False, False])}
    worker._desc_pos_by_block_size = {16: np.array([0, 1, 0, 1])}
    worker._dram_src_handles_by_block_size = {16: 10}
    worker.nixl_wrapper.make_prepped_xfer.side_effect = [101, 102]
    worker.nixl_wrapper.check_xfer_state.return_value = "DONE"

    worker._read_blocks_mixed(
        request_id="request",
        local_block_size_key=16,
        local_device_handle=20,
        local_dram_handle=10,
        remote_xfer_side_handle=30,
        local_block_descs_ids=np.array([0, 2]),
        remote_block_descs_ids=np.array([5, 7]),
        notif_agent="prefill",
        notif_id=b"request:1",
    )

    dram_read, device_read = worker.nixl_wrapper.make_prepped_xfer.call_args_list
    assert dram_read.args[:2] == ("READ", 10)
    assert dram_read.args[3] == 30
    np.testing.assert_array_equal(dram_read.args[2], [0])
    np.testing.assert_array_equal(dram_read.args[4], [5])
    assert device_read.args[:2] == ("READ", 20)
    assert device_read.args[3] == 30
    np.testing.assert_array_equal(device_read.args[2], [0])
    np.testing.assert_array_equal(device_read.args[4], [7])
    worker.nixl_wrapper.send_notif.assert_not_called()

    assert worker.get_finished() == (set(), {"request"})
    worker.nixl_wrapper.send_notif.assert_called_once_with(
        "prefill", notif_msg=b"request:1"
    )


def test_mixed_memory_read_failure_does_not_notify_producer(recv_worker):
    """A failed half of a split READ must suppress its success notification."""
    worker = recv_worker
    worker._recving_transfers = {"request": [101, 102]}
    worker._pending_recv_notifs = {"request": [("prefill", b"request:1")]}
    worker.nixl_wrapper.check_xfer_state.side_effect = ["ERR", "DONE"]

    assert worker.get_finished() == (set(), {"request"})

    assert "request" not in worker._pending_recv_notifs
    worker.nixl_wrapper.send_notif.assert_not_called()


def element_byte_addrs(view: torch.Tensor) -> list[int]:
    """Absolute byte addresses of every element of a (possibly strided) view."""
    offsets = torch.zeros(view.shape, dtype=torch.int64)
    for dim, (size, stride) in enumerate(zip(view.shape, view.stride())):
        shape = [1] * view.ndim
        shape[dim] = size
        offsets = offsets + torch.arange(size, dtype=torch.int64).view(shape) * stride
    esize = view.element_size()
    byte_offsets = (offsets.flatten() * esize).tolist()
    base = view.data_ptr()
    return [base + offset + i for offset in byte_offsets for i in range(esize)]


@pytest.mark.parametrize(
    "attn_backend",
    [
        pytest.param(
            "FLASH_ATTN",
            marks=pytest.mark.skipif(
                current_platform.is_rocm(),
                reason="Attention backend FLASH_ATTN is not supported on ROCm",
            ),
        ),
        "TRITON_ATTN",
    ],
)
@pytest.mark.parametrize("layout", [layout.name for layout in KVCacheLayout])
@pytest.mark.parametrize("separate_kv_head_groups", [False, True])
def test_register_kv_caches(
    default_vllm_config,
    dist_init,
    attn_backend,
    layout,
    separate_kv_head_groups,
):
    """Test that register_kv_caches() properly calls nixl_wrapper methods with
    correct data.

    This test verifies:
    1. nixl_wrapper.get_reg_descs() is called with caches_data containing
       tensor metadata
    2. nixl_wrapper.get_xfer_descs() is called with blocks_data containing
       block layout info
    """
    vllm_config = create_vllm_config(attention_backend=attn_backend)
    vllm_config.cache_config.kv_cache_layout = layout

    # Import the appropriate backend based on the parameter
    if attn_backend == "FLASH_ATTN":
        from vllm.v1.attention.backends.flash_attn import FlashAttentionBackend

        backend_cls = FlashAttentionBackend
    elif attn_backend == "ROCM_ATTN":
        from vllm.v1.attention.backends.rocm_attn import RocmAttentionBackend

        backend_cls = RocmAttentionBackend
    else:  # TRITON_ATTN
        from vllm.v1.attention.backends.triton_attn import TritonAttentionBackend

        backend_cls = TritonAttentionBackend

    nixl_worker = "vllm.distributed.kv_transfer.kv_connector.v1.nixl.base_worker"
    with (
        patch(f"{nixl_worker}.NixlWrapper") as mock_nixl_wrapper,
        patch(f"{nixl_worker}.threading.Event"),
        patch(f"{nixl_worker}.threading.Thread") as mock_thread,
        patch(f"{nixl_worker}.get_current_attn_backends") as mock_get_attn_backends,
    ):
        mock_get_attn_backends.return_value = [backend_cls]
        block_size = 16
        num_blocks = 8
        num_heads = 4
        head_size = 16

        kv_cache_spec = FullAttentionSpec(
            block_size=block_size,
            num_kv_heads=num_heads,
            head_size=head_size,
            dtype=torch.float16,
            num_head_slots=2 if separate_kv_head_groups else None,
            state_content_bytes=num_heads * head_size * 2
            if separate_kv_head_groups
            else None,
        )
        kv_cache_config = KVCacheConfig(
            num_blocks=num_blocks,
            kv_cache_tensors=[],
            kv_cache_groups=[
                KVCacheGroupSpec(
                    ["layer0", "layer1", "layer2", "layer3"], kv_cache_spec
                )
            ],
        )
        # Create connector
        connector = NixlConnector(vllm_config, KVConnectorRole.WORKER, kv_cache_config)
        connector.connector_worker = FakeNixlConnectorWorker(
            vllm_config,
            connector.engine_id,
            hand_shake_latency=0,
            kv_cache_config=kv_cache_config,
        )

        # Get the mock instance
        mock_wrapper_instance = mock_nixl_wrapper.return_value
        connector.connector_worker.nixl_wrapper = mock_wrapper_instance

        # Appease NixlHandshakePayload encoding with some bytes
        mock_wrapper_instance.get_agent_metadata.return_value = b"fake_agent_metadata"

        # Reassure the shutdown() check that the thread is terminated
        mock_thread.return_value.is_alive.return_value = False

        raw0 = torch.zeros(
            kv_cache_spec.page_size_bytes * kv_cache_config.num_blocks * 2,
            dtype=torch.int8,
            device=current_platform.device_type,
        )
        raw1 = torch.zeros(
            kv_cache_spec.page_size_bytes * kv_cache_config.num_blocks,
            dtype=torch.int8,
            device=current_platform.device_type,
        )
        tensor0, tensor1 = dense_kv_cache_views(
            raw0,
            kv_cache_spec,
            kv_cache_config.num_blocks,
            num_layers=2,
            layout=KVCacheLayout[layout],
        )
        (tensor2,) = dense_kv_cache_views(
            raw1,
            kv_cache_spec,
            kv_cache_config.num_blocks,
            num_layers=1,
            layout=KVCacheLayout[layout],
        )
        kv_caches = {
            "layer0": tensor0,
            "layer1": tensor1,
            "layer2": tensor2,
            "layer3": tensor0,
        }

        # Execute register_kv_caches
        connector.register_kv_caches(kv_caches)

        # Verify get_reg_descs was called with caches_data
        assert mock_wrapper_instance.get_reg_descs.called
        caches_data, _ = mock_wrapper_instance.get_reg_descs.call_args[0]
        assert len(caches_data) == 2

        for cache_entry, raw in zip(caches_data, (raw0, raw1)):
            base_addr, size, _tp_rank, _ = cache_entry
            assert size == raw.nbytes
            assert base_addr == raw.data_ptr()

        # Verify get_xfer_descs was called with blocks_data
        assert mock_wrapper_instance.get_xfer_descs.called
        blocks_data, _ = mock_wrapper_instance.get_xfer_descs.call_args[0]

        # Layout-blind contract: whatever regions the worker carves out,
        # transferring "block b" must move exactly logical block b's bytes for
        # every layer. Map each registered byte to the block that owns it, then
        # require every descriptor window to hold bytes of a single block and
        # the windows of block b to cover exactly block b's bytes.
        owner: dict[int, int] = {}
        for cache in (tensor0, tensor1, tensor2):
            for blk in range(num_blocks):
                for addr in element_byte_addrs(cache[blk]):
                    assert owner.setdefault(addr, blk) == blk
        block_bytes = defaultdict(set)
        for addr, blk in owner.items():
            block_bytes[blk].add(addr)

        covered: defaultdict[int, set[int]] = defaultdict(set)
        for block_start_addr, block_len, _tp_rank in blocks_data:
            window = range(block_start_addr, block_start_addr + block_len)
            owners = {owner[addr] for addr in window}
            assert len(owners) == 1, "descriptor window spans logical blocks"
            covered[owners.pop()].update(window)
        assert covered == block_bytes

        # Region bases are exactly the block-0 window starts, one per region.
        base_addrs = connector.connector_worker.kv_caches_base_addr[
            connector.connector_worker.engine_id
        ][0]
        if layout == "BLHNC":
            assert len(base_addrs) == 3
        assert set(base_addrs) == {
            start for start, _len, _tp in blocks_data if owner[start] == 0
        }
        assert len(blocks_data) == num_blocks * len(base_addrs)

        assert connector.connector_worker.block_size == 16


def test_register_packed_dsv4_mla_cache_as_single_region(
    default_vllm_config, dist_init
):
    from vllm.v1.attention.backends.triton_attn import TritonAttentionBackend

    nixl_worker = "vllm.distributed.kv_transfer.kv_connector.v1.nixl.base_worker"
    with (
        patch(f"{nixl_worker}.NixlWrapper") as mock_nixl_wrapper,
        patch(f"{nixl_worker}.threading.Event"),
        patch(f"{nixl_worker}.threading.Thread") as mock_thread,
        patch(f"{nixl_worker}.get_current_attn_backends") as mock_backends,
    ):
        mock_backends.return_value = [TritonAttentionBackend]
        num_blocks = 2
        num_layers = 4
        block_size = 256
        spec = MLAAttentionSpec(
            block_size=block_size,
            num_kv_heads=1,
            head_size=512,
            dtype=torch.uint8,
            cache_dtype_str="fp8_ds_mla",
            tokens_per_state=4,
            alignment=576,
            model_version="deepseek_v4",
            state_content_bytes=584,
        )
        layer_names = [f"layer{idx}" for idx in range(num_layers)]
        kv_cache_config = KVCacheConfig(
            num_blocks=num_blocks,
            kv_cache_tensors=[],
            kv_cache_groups=[KVCacheGroupSpec(layer_names, spec)],
        )
        vllm_config = create_vllm_config(attention_backend="TRITON_ATTN")
        vllm_config.cache_config.block_size = block_size
        vllm_config.cache_config.kv_cache_layout = "BLHNC"
        connector = NixlConnector(vllm_config, KVConnectorRole.WORKER, kv_cache_config)
        connector.connector_worker = FakeNixlConnectorWorker(
            vllm_config,
            connector.engine_id,
            hand_shake_latency=0,
            kv_cache_layout="BLHNC",
            kv_cache_config=kv_cache_config,
        )
        wrapper = mock_nixl_wrapper.return_value
        connector.connector_worker.nixl_wrapper = wrapper
        wrapper.get_agent_metadata.return_value = b"fake_agent_metadata"
        mock_thread.return_value.is_alive.return_value = False

        raw = torch.zeros(
            spec.page_size_bytes * num_blocks * num_layers,
            dtype=torch.int8,
            device=current_platform.device_type,
        )
        views = dense_kv_cache_views(
            raw, spec, num_blocks, num_layers, KVCacheLayout.BLHNC
        )
        connector.register_kv_caches(dict(zip(layer_names, views)))

        blocks_data, _ = wrapper.get_xfer_descs.call_args[0]
        packed_block_len = num_layers * spec.page_size_bytes
        assert connector.connector_worker.kv_caches_base_addr[
            connector.connector_worker.engine_id
        ][0] == [raw.data_ptr()]
        assert blocks_data.tolist() == [
            [raw.data_ptr() + block_idx * packed_block_len, packed_block_len, 0]
            for block_idx in range(num_blocks)
        ]


class FakePlatform(Platform):
    device_type: str = "oot"

    @classmethod
    def get_nixl_supported_devices(cls) -> dict[str, tuple[str, ...]]:
        """Returns a mapping from device_type to a tuple of supported
        kv_buffer_device for nixl.
        """
        return {"oot": ("oot",)}

    @classmethod
    def get_nixl_memory_type(cls) -> str | None:
        """Returns the nixl memory type for the current platform."""
        return "VRAM"


@pytest.mark.parametrize(
    "kv_buffer_device, nixl_memory_type",
    [
        ("oot", "VRAM"),
    ],
)
def test_kv_buffer_to_nixl_memory_types(
    default_vllm_config, dist_init, kv_buffer_device, nixl_memory_type
):
    """Test that register_kv_caches() passes the correct memory types from the
    config to the nixl_wrapper.
    """
    vllm_config = create_vllm_config()
    # Override the default memory types in the config
    vllm_config.kv_transfer_config.kv_buffer_device = kv_buffer_device
    from vllm.distributed.kv_transfer.kv_connector.v1.nixl.utils import (
        _NIXL_SUPPORTED_DEVICE,
    )

    _NIXL_SUPPORTED_DEVICE.update(FakePlatform.get_nixl_supported_devices())

    with (
        patch(
            "vllm.distributed.kv_transfer.kv_connector.v1.nixl.base_worker.NixlWrapper"
        ),
        patch(
            "vllm.distributed.kv_transfer.kv_connector.v1.nixl.base_worker.threading.Event"
        ),
        patch(
            "vllm.distributed.kv_transfer.kv_connector.v1.nixl.base_worker.threading.Thread"
        ),
        patch(
            "vllm.distributed.kv_transfer.kv_connector.v1.nixl.base_worker.current_platform",
            FakePlatform,
        ),
        patch(
            "vllm.distributed.kv_transfer.kv_connector.v1.nixl.utils._NIXL_SUPPORTED_DEVICE",
            _NIXL_SUPPORTED_DEVICE,
        ),
    ):  # noqa: E501
        # Create connector and replace its worker with a fake one for isolation
        connector = NixlConnector(
            vllm_config, KVConnectorRole.WORKER, make_kv_cache_config(block_size=16)
        )

        # Verify get_reg_descs was called with the correct memory_type
        assert connector.connector_worker.kv_buffer_device == kv_buffer_device
        assert connector.connector_worker.nixl_memory_type == nixl_memory_type


@patch(
    "vllm.distributed.kv_transfer.kv_connector.v1.nixl.base_worker.NixlWrapper",
    FakeNixlWrapper,
)
def test_shutdown_cleans_up_resources(default_vllm_config, dist_init):
    """Test that shutdown() properly cleans up all resources."""
    vllm_config = create_vllm_config()

    scheduler = NixlConnectorScheduler(
        vllm_config,
        vllm_config.kv_transfer_config.engine_id,
        make_kv_cache_config(block_size=16),
    )
    worker = NixlConnectorWorker(
        vllm_config,
        vllm_config.kv_transfer_config.engine_id,
        make_kv_cache_config(block_size=16),
    )
    nixl_wrapper = worker.nixl_wrapper

    with (
        patch.object(worker, "_handshake_initiation_executor") as mock_exec,
        patch.object(scheduler, "_nixl_handshake_listener_t") as mock_listener,
        patch.object(nixl_wrapper, "release_xfer_handle") as mock_rel_xfer,
        patch.object(nixl_wrapper, "release_dlist_handle") as mock_rel_dlist,
        patch.object(nixl_wrapper, "remove_remote_agent") as mock_rem_agent,
        patch.object(nixl_wrapper, "deregister_memory") as mock_dereg,
    ):
        worker._recving_transfers = {"req1": [123]}
        # Mock register_kv_cache which registers local handle
        worker.src_xfer_handles_by_block_size = {worker.block_size: 455}
        # P TP = 2 * D TP case, we should register 2 local handles
        worker.src_xfer_handles_by_tp_ratio = {(-2, 16): [456, 457]}
        worker.dst_xfer_side_handles = {"engine1": {0: 789}}
        worker._remote_agents = {"engine1": {(0, 0): "agent1"}}
        # _cleanup_remote_engine (called by shutdown) also clears these:
        worker.kv_caches_base_addr["engine1"] = {0: [0xABC]}
        worker.dst_num_blocks["engine1"] = 50
        worker.tp_mappings["engine1"] = MagicMock()
        worker._engine_last_active["engine1"] = time.perf_counter()
        worker._registered_descs = ["desc1", "desc2"]

        mock_listener.is_alive.return_value = False

        worker.shutdown()

        # Test idempotency
        worker.shutdown()
        worker.shutdown()

        mock_exec.shutdown.assert_called_with(wait=False)

        # Same sequence on scheduler.shutdown()
        scheduler.shutdown()
        scheduler.shutdown()
        scheduler.shutdown()
        mock_listener.join.assert_called_once()

        mock_rel_xfer.assert_called_once_with(123)
        assert mock_rel_dlist.call_count == 4
        mock_rel_dlist.assert_any_call(455)  # src handle (whole region)
        mock_rel_dlist.assert_any_call(456)  # src handle (1st chunk)
        mock_rel_dlist.assert_any_call(457)  # src handle (2nd chunk)
        mock_rel_dlist.assert_any_call(789)  # dst handle
        mock_rem_agent.assert_called_once_with("agent1")
        assert mock_dereg.call_count == 2
        mock_dereg.assert_any_call("desc1")
        mock_dereg.assert_any_call("desc2")


# ── TTL-based remote engine eviction tests ──────────────────────────


def _setup_worker_with_remote_engine(
    engine_ttl: float = 10.0,
) -> tuple[Any, str]:
    """Create a worker with one remote engine registered."""
    vllm_config = create_vllm_config(
        kv_connector_extra_config={"engine_ttl": engine_ttl},
    )
    worker = NixlConnectorWorker(
        vllm_config,
        vllm_config.kv_transfer_config.engine_id,
        make_kv_cache_config(block_size=16),
    )

    engine_id = "remote-engine-1"
    worker._remote_agents[engine_id] = {(0, 0): "agent_0", (0, 1): "agent_1"}
    worker.dst_xfer_side_handles[engine_id] = {0: 100, 1: 200}
    worker.kv_caches_base_addr[engine_id] = {0: [0xABC]}
    worker.dst_num_blocks[engine_id] = 50
    worker.tp_mappings[engine_id] = MagicMock()
    worker._engine_last_active[engine_id] = time.perf_counter()

    worker.transfer_topo = MagicMock()

    return worker, engine_id


@patch(
    "vllm.distributed.kv_transfer.kv_connector.v1.nixl.base_worker.NixlWrapper",
    FakeNixlWrapper,
)
def test_engine_ttl_eviction(default_vllm_config, dist_init):
    """Stale engines are evicted when TTL expires."""
    worker, engine_id = _setup_worker_with_remote_engine(engine_ttl=10.0)
    nixl_wrapper = worker.nixl_wrapper

    with (
        patch.object(nixl_wrapper, "release_dlist_handle") as mock_rel,
        patch.object(nixl_wrapper, "remove_remote_agent") as mock_rem,
    ):
        # Make the engine stale.
        worker._engine_last_active[engine_id] = time.perf_counter() - 20.0

        worker._evict_stale_engines()

        assert engine_id not in worker._remote_agents
        assert engine_id not in worker.dst_xfer_side_handles
        assert engine_id not in worker.kv_caches_base_addr
        assert engine_id not in worker.dst_num_blocks
        assert engine_id not in worker.tp_mappings
        assert engine_id not in worker._engine_last_active
        worker.transfer_topo.unregister_remote_engine.assert_called_with(engine_id)

        assert mock_rel.call_count == 2
        mock_rel.assert_any_call(100)
        mock_rel.assert_any_call(200)

        assert mock_rem.call_count == 2
        mock_rem.assert_any_call("agent_0")
        mock_rem.assert_any_call("agent_1")


@patch(
    "vllm.distributed.kv_transfer.kv_connector.v1.nixl.base_worker.NixlWrapper",
    FakeNixlWrapper,
)
def test_engine_with_inflight_transfer_is_not_evicted(default_vllm_config, dist_init):
    """A transfer that outlives the TTL must keep its engine registered.

    _engine_last_active is stamped when a read is issued and not refreshed
    while it runs, so a stalled transfer -- a peer that has lost its NIC holds
    one indefinitely -- leaves its engine looking idle. Evicting it releases
    the dlist handle and remote agent the transfer is still reading through.
    """
    worker, engine_id = _setup_worker_with_remote_engine(engine_ttl=10.0)
    nixl_wrapper = worker.nixl_wrapper

    request_id = "req-still-reading"
    worker._recving_transfers[request_id] = [MagicMock()]
    worker._recving_metadata[request_id] = MagicMock(
        remote=MagicMock(engine_id=engine_id)
    )

    with (
        patch.object(nixl_wrapper, "release_dlist_handle") as mock_rel,
        patch.object(nixl_wrapper, "remove_remote_agent") as mock_rem,
    ):
        worker._engine_last_active[engine_id] = time.perf_counter() - 20.0

        worker._evict_stale_engines()

        assert engine_id in worker._remote_agents
        assert engine_id in worker.dst_xfer_side_handles
        mock_rel.assert_not_called()
        mock_rem.assert_not_called()

        # Once the transfer is done the engine is stale like any other.
        del worker._recving_transfers[request_id]
        worker._evict_stale_engines()

        assert engine_id not in worker._remote_agents
        assert mock_rem.call_count == 2


@patch(
    "vllm.distributed.kv_transfer.kv_connector.v1.nixl.base_worker.NixlWrapper",
    FakeNixlWrapper,
)
def test_engine_ttl_disabled(default_vllm_config, dist_init):
    """Eviction is disabled when engine_ttl <= 0."""
    worker, engine_id = _setup_worker_with_remote_engine(engine_ttl=0.0)

    # Make the engine stale.
    worker._engine_last_active[engine_id] = time.perf_counter() - 9999.0

    worker._evict_stale_engines()

    # Nothing should be evicted.
    assert engine_id in worker._remote_agents
    assert engine_id in worker.dst_xfer_side_handles


@pytest.mark.cpu_test
class TestPeerReplacement:
    @pytest.fixture(autouse=True)
    def setup(self, monkeypatch):
        from vllm.distributed.kv_transfer.kv_connector.v1.nixl import base_worker as bw
        from vllm.v1.attention.backends.flash_attn import FlashAttentionBackend

        config = create_vllm_config(kv_connector_extra_config={"engine_ttl": 0})
        config.kv_transfer_config.kv_buffer_device = "cpu"
        platform = SimpleNamespace(
            device_type="cpu",
            discover_numa_topology=lambda: [],
            get_nixl_memory_type=lambda: "DRAM",
            is_rocm=lambda: False,
        )
        with (
            patch.object(bw, "NixlWrapper", FakeNixlWrapper),
            patch.object(bw, "current_platform", platform),
            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=[FlashAttentionBackend]
            ),
        ):
            self.worker = FakeNixlConnectorWorker(config, "local", hand_shake_latency=0)
            self.nixl = self.worker.nixl_wrapper
            self.transport = MagicMock(wraps=self.nixl)
            self.worker.nixl_wrapper = self.transport
            monkeypatch.setattr(
                self.worker._handshake_initiation_executor,
                "submit",
                lambda *args: Future(),
            )
            self._connect("old")
            yield
            self.worker.shutdown()

    def _connect(self, engine_id="new", host="localhost", port=1234):
        future = self.worker._ensure_handshake(engine_id, host, port, 1)
        self.worker.REMOTE_ENGINE_ID = engine_id
        self.nixl.REMOTE_AGENT_NAME = engine_id
        future.set_result(self.worker._nixl_handshake(host, port, 1, engine_id))

    def _request(self, engine_id="new", blocks=(1,), awaiting_kvs=True):
        metadata = NixlConnectorMetadata()
        metadata.reqs_to_recv[f"{engine_id}-req"] = ReqMeta(
            local_block_ids=(blocks,),
            local_physical_block_ids=(blocks,),
            tp_size=1,
            remote=RemoteMeta(
                ([2],), "localhost", 1234, engine_id, f"{engine_id}-prefill"
            ),
            awaiting_kvs=awaiting_kvs,
        )
        return metadata

    def test_cleans_old_peer_after_confirmed_handshake(self):
        self._connect("healthy", host="other-host")
        old_handle = self.worker.dst_xfer_side_handles["old"][0]
        self.worker.start_load_kv(self._request())
        self.worker.get_transfer_results()
        self.transport.remove_remote_agent.assert_not_called()

        self._connect()
        self.worker.get_transfer_results()
        self.transport.release_dlist_handle.assert_called_once_with(old_handle)
        self.transport.remove_remote_agent.assert_called_once_with("old")
        assert set(self.worker._remote_agents) == {"new", "healthy"}
        assert self.worker._engine_by_address == {
            ("localhost", 1234): "new",
            ("other-host", 1234): "healthy",
        }
        with pytest.raises(KeyError):
            self.worker.transfer_topo.get_engine_info("old")

        self.worker.start_load_kv(NixlConnectorMetadata())
        assert self.worker.get_transfer_results().finished_recving == {"new-req"}
        self.transport.remove_remote_agent.assert_called_once_with("old")

    @pytest.mark.parametrize("host,port", [("other-host", 1234), ("localhost", 5678)])
    def test_preserves_other_addresses(self, host, port):
        self._connect(host=host, port=port)
        self.worker.get_transfer_results()
        assert "old" in self.worker._remote_agents
        self.transport.remove_remote_agent.assert_not_called()

    def test_preserves_old_peer_when_handshake_fails(self):
        self.worker.start_load_kv(self._request())
        self.worker._handshake_futures["new"].set_exception(
            RuntimeError("handshake failed")
        )
        assert self.worker.get_transfer_results().failed_recving == {"new-req"}
        assert "old" in self.worker._remote_agents
        assert self.worker._engine_by_address == {("localhost", 1234): "old"}
        self.transport.remove_remote_agent.assert_not_called()

    @pytest.mark.parametrize(
        "states,failed",
        [
            (["PROC", "DONE"], False),
            (["ERR", "PROC", "DONE"], True),
            (["ERR", "PROC", "ERR"], True),
        ],
        ids=["healthy", "failed-sibling-completes", "failed-sibling-fails"],
    )
    def test_waits_for_reads_before_cleanup_and_failure_reporting(self, states, failed):
        self.worker._recving_metadata.update(self._request("old").reqs_to_recv)
        self.worker._recving_transfers["old-req"] = [1, 2] if failed else [1]
        self.transport.check_xfer_state.side_effect = states
        self._connect()

        result = self.worker.get_transfer_results()
        assert result.finished_recving == result.failed_recving == set()
        assert self.worker.get_block_ids_with_load_errors() == set()
        self.transport.remove_remote_agent.assert_not_called()

        result = self.worker.get_transfer_results()
        assert result.finished_recving == {"old-req"}
        assert result.failed_recving == ({"old-req"} if failed else set())
        assert self.worker.get_block_ids_with_load_errors() == (
            {1} if failed else set()
        )
        assert self.transport.release_xfer_handle.call_count == (2 if failed else 1)
        self.transport.remove_remote_agent.assert_called_once_with("old")

    def test_waits_for_queued_notification(self):
        metadata = self._request("old", blocks=())
        self.worker._recving_metadata.update(metadata.reqs_to_recv)
        self.worker._background_nixl_handshake(
            "old-req", "old", metadata.reqs_to_recv["old-req"]
        )
        self._connect()
        self.worker.get_transfer_results()
        self.transport.remove_remote_agent.assert_not_called()

        self.worker.start_load_kv(NixlConnectorMetadata())
        self.transport.send_notif.assert_called_once_with(
            "old", notif_msg=b"old-prefill:1"
        )
        result = self.worker.get_transfer_results()
        assert result.finished_recving == {"old-req"}
        assert result.failed_recving == set()
        self.transport.remove_remote_agent.assert_called_once_with("old")

    def test_waits_for_other_handshakes(self):
        self.worker._ensure_handshake("other", "other-host", 1234, 1)
        self._connect()
        self.worker.get_transfer_results()
        self.transport.remove_remote_agent.assert_not_called()

        self._connect("other", host="other-host")
        self.worker.get_transfer_results()
        self.transport.remove_remote_agent.assert_called_once_with("old")

    def test_rehandshakes_queued_request_for_released_engine(self):
        # A request queued after its engine's handshake, but drained only after
        # that engine was replaced and released, must not read from it.
        metadata = self._request("old", blocks=(), awaiting_kvs=False)
        self.worker._ready_requests.put(("old-req", metadata.reqs_to_recv["old-req"]))
        self._connect()
        self.worker.get_transfer_results()
        self.transport.remove_remote_agent.assert_called_once_with("old")

        self.worker.start_load_kv(NixlConnectorMetadata())
        assert "old" in self.worker._handshake_futures
        self.transport.send_notif.assert_not_called()


def test_transfer_topology_unregister():
    """TransferTopology.unregister_remote_engine removes the engine."""
    from vllm.v1.attention.backends.flash_attn import FlashAttentionBackend

    topo = TransferTopology(
        tp_rank=0,
        tp_size=1,
        block_size=16,
        engine_id="local",
        is_mla=False,
        is_mamba=False,
        total_num_kv_heads=4,
        attn_backends=[FlashAttentionBackend],
    )

    info = EngineTransferInfo(
        remote_tp_size=1,
        remote_block_size=16,
        remote_block_len=64,
        remote_physical_blocks_per_logical=1,
    )
    topo.register_remote_engine("remote-1", info)
    assert topo.get_engine_info("remote-1") is info

    topo.unregister_remote_engine("remote-1")
    with pytest.raises(KeyError):
        topo.get_engine_info("remote-1")

    # Idempotent: no error on double-unregister
    topo.unregister_remote_engine("remote-1")


@patch(
    "vllm.distributed.kv_transfer.kv_connector.v1.nixl.base_worker.NixlWrapper",
    FakeNixlWrapper,
)
def test_aborted_request_removed_from_worker_in_batch(default_vllm_config, dist_init):
    """Create and schedule a request so that P adds it to in-batch tracking via
    the real scheduler, then simulate an abort (request not in next scheduler
    iteration) and verify the worker no longer tracks it as in-batch.
    """
    vllm_config = create_vllm_config()

    scheduler = create_scheduler(vllm_config)
    # KVConnector Worker in P
    connector = NixlConnector(
        vllm_config, KVConnectorRole.WORKER, make_kv_cache_config(block_size=16)
    )
    connector.connector_worker = FakeNixlConnectorWorker(
        vllm_config, connector.engine_id, hand_shake_latency=0
    )

    # Create a request that triggers do_remote_decode so that
    # the scheduler adds it to reqs_in_batch
    req = create_request(request_id=1, do_remote_decode=True, max_tokens=1)
    scheduler.add_request(req)

    # First scheduling pass - examine build_connector_meta output
    sched_out = scheduler.schedule()
    kv_meta = sched_out.kv_connector_metadata
    assert kv_meta is not None
    assert isinstance(kv_meta, NixlConnectorMetadata)
    assert req.request_id in kv_meta.reqs_in_batch

    #### Model Runner start ####
    # Bind scheduler-produced metadata and start worker processing.
    connector.bind_connector_metadata(kv_meta)

    dummy_ctx = ForwardContext(
        no_compile_layers={},
        attn_metadata={},
        slot_mapping={},
    )
    connector.start_load_kv(dummy_ctx)

    # Ensure it was tracked by the worker
    assert req.request_id in connector.connector_worker._reqs_to_process

    #### Model Runner end ####

    # Abort request - request_finished call in connector scheduler
    scheduler.finish_requests(req.request_id, RequestStatus.FINISHED_ABORTED)
    # Second scheduling pass - build metadata with aborted request
    sched_out2 = scheduler.schedule()
    kv_meta2 = sched_out2.kv_connector_metadata
    assert kv_meta2 is not None
    assert isinstance(kv_meta2, NixlConnectorMetadata)
    assert req.request_id not in kv_meta2.reqs_in_batch

    # Bind empty/abort metadata and run worker step
    #### Model Runner start ####
    connector.bind_connector_metadata(kv_meta2)
    connector.start_load_kv(dummy_ctx)

    # After abort, the worker should not keep tracking it as "in-batch"
    assert req.request_id not in connector.connector_worker._reqs_to_process
    #### Model Runner end ####


class FailingNixlWrapper(FakeNixlWrapper):
    """Mock NixlWrapper that fails on specific operations."""

    def __init__(self, *args, **kwargs):
        super().__init__(*args, **kwargs)
        self.fail_handshake = False
        self.fail_transfer_setup = False
        self.fail_send_notif = False
        self.fail_transfer_state = False  # Returns "ERR" state
        self.fail_transfer_exception = False  # Raises exception in check_xfer_state

    def add_remote_agent(self, agent_metadata: bytes) -> str:
        if self.fail_handshake:
            from zmq.error import Again

            raise Again("Simulated timeout failure")
        return super().add_remote_agent(agent_metadata)

    def make_prepped_xfer(
        self,
        xfer_type: str,
        local_xfer_side_handle: int,
        local_block_descs_ids: list[int],
        remote_xfer_side_handle: int,
        remote_block_descs_ids: list[int],
        notif_msg: bytes | None = None,
    ) -> int:
        if self.fail_transfer_setup:
            # classic RuntimeError to simulate failure
            raise RuntimeError("BAD STATUS")
        return super().make_prepped_xfer(
            xfer_type,
            local_xfer_side_handle,
            local_block_descs_ids,
            remote_xfer_side_handle,
            remote_block_descs_ids,
            notif_msg,
        )

    def send_notif(self, agent_name: str, notif_msg: bytes) -> None:
        if self.fail_send_notif:
            raise RuntimeError("Simulated send_notif failure")
        return super().send_notif(agent_name, notif_msg)

    def check_xfer_state(self, handle: int) -> str:
        if self.fail_transfer_exception:
            raise RuntimeError("Simulated check_xfer_state exception")
        if self.fail_transfer_state:
            return "ERR"  # Bad transfer state
        return super().check_xfer_state(handle)


@patch(
    "vllm.distributed.kv_transfer.kv_connector.v1.nixl.base_worker.NixlWrapper",
    FailingNixlWrapper,
)
@pytest.mark.parametrize("awaiting_kvs", [True, False])
@pytest.mark.parametrize("notification_fails", [False, True])
def test_empty_recv_is_reported_only_when_awaited(
    default_vllm_config, dist_init, awaiting_kvs, notification_fails
):
    """An empty recv is reported iff the scheduler parked the request on it.

    _read_blocks returns early when the local block list is empty -- D already
    holds the KV and only the producer needs telling. Two kinds of caller reach
    that path, and they need opposite handling:

    Parked (awaiting_kvs=True, a full prefix cache hit). Without a report the
    request is named in neither _recving_transfers nor _failed_recv_reqs and
    never reaches finished_recving. The scheduler has no other way to release a
    WAITING_FOR_REMOTE_KVS request, and there is no timeout, so it sits in
    kv_holding_waiting holding its blocks for the life of the process.

    Notify-only (awaiting_kvs=False): request_finished seeding an empty recv to
    free P's blocks for a request aborted before it was scheduled, or a readback
    on a request that stays RUNNING. Reporting either one crashes the engine
    core -- _update_from_kv_xfer_finished asserts the id is still in
    self.requests and that the request is parked or finished.

    A failed notification changes neither: the KV is local either way, and the
    producer frees its own blocks on a timeout.
    """
    vllm_config = create_vllm_config()
    connector = NixlConnector(
        vllm_config, KVConnectorRole.WORKER, make_kv_cache_config(block_size=16)
    )
    connector.connector_worker = FakeNixlConnectorWorker(
        vllm_config,
        connector.engine_id,
        hand_shake_latency=0.0,
        kv_cache_config=connector._kv_cache_config,
    )
    connector.connector_worker.nixl_wrapper.fail_send_notif = notification_fails

    request_id = f"test_empty_recv_awaited_{awaiting_kvs}"
    metadata = NixlConnectorMetadata()
    metadata.add_new_req_to_recv(
        request_id=request_id,
        local_block_ids=(),  # empty: the whole prompt is already cached locally
        kv_transfer_params={
            "remote_block_ids": [[20, 21, 22]],
            "remote_engine_id": FakeNixlConnectorWorker.REMOTE_ENGINE_ID,
            "remote_request_id": f"prefill-{request_id}",
            "remote_host": "localhost",
            "remote_port": 1234,
            "remote_tp_size": 1,
        },
        awaiting_kvs=awaiting_kvs,
    )
    connector.bind_connector_metadata(metadata)
    dummy_ctx = ForwardContext(no_compile_layers={}, attn_metadata={}, slot_mapping={})
    connector.start_load_kv(dummy_ctx)
    connector.bind_connector_metadata(NixlConnectorMetadata())
    time.sleep(0.1)
    connector.start_load_kv(dummy_ctx)

    _, done_recving = connector.get_finished(finished_req_ids=set())
    assert (request_id in done_recving) is awaiting_kvs


@pytest.mark.parametrize("is_hma", [False, True])
@pytest.mark.parametrize(
    "local_block_ids,awaiting_kvs",
    [((), False), (([],), False), (([],), True), (([1, 2, 3],), True)],
)
def test_handshake_failure_reports_only_awaited_recvs(
    recv_worker, is_hma, local_block_ids, awaiting_kvs
):
    """A failed cleanup handshake must not complete an already-aborted request."""
    worker = recv_worker
    worker._is_hma_required = is_hma
    worker._recving_metadata.clear()
    worker._remote_agents = {}
    worker._handshake_lock = contextlib.nullcontext()
    worker._ready_requests = queue.Queue()
    worker._reqs_to_process = set()
    worker.pcp_rank = 0
    metadata = NixlConnectorMetadata()
    metadata.add_new_req_to_recv(
        request_id="request",
        local_block_ids=local_block_ids,
        kv_transfer_params={
            "remote_block_ids": ([4, 5, 6],),
            "remote_engine_id": "prefill",
            "remote_request_id": "prefill-request",
            "remote_host": "localhost",
            "remote_port": 1234,
        },
        awaiting_kvs=awaiting_kvs,
    )
    handshake = Future[None]()
    handshake.set_exception(RuntimeError("handshake failed"))
    with patch.object(worker, "_ensure_handshake", return_value=handshake):
        worker.start_load_kv(metadata)

    results = worker.get_transfer_results()
    expected = {"request"} if awaiting_kvs else set()
    assert results.finished_recving == results.failed_recving == expected
    assert worker.get_block_ids_with_load_errors() == (
        set(local_block_ids[0]) if awaiting_kvs and not is_hma else set()
    )
    assert not worker._recving_metadata
    assert not worker._recv_failures
    assert not worker._pending_recv_notifs


@patch(
    "vllm.distributed.kv_transfer.kv_connector.v1.nixl.base_worker.NixlWrapper",
    FailingNixlWrapper,
)
@pytest.mark.parametrize(
    "failure_type,wrapper_config,needs_get_finished",
    [
        ("transfer_setup_failed", {"fail_transfer_setup": True}, False),
        ("handshake_failed", {"fail_handshake": True}, False),
        ("notification_failed", {"fail_send_notif": True}, False),
        ("transfer_failed", {"fail_transfer_state": True}, True),
        ("transfer_exception", {"fail_transfer_exception": True}, True),
    ],
)
@pytest.mark.parametrize("enable_hma", [False, True])
def test_transfer_failure_logging(
    default_vllm_config,
    dist_init,
    failure_type,
    wrapper_config,
    needs_get_finished,
    enable_hma,
):
    """Test that transfer failures are logged with structured context.

    Run with `pytest -sv` to see the log output.

    Covers failure types:
    - transfer_setup_failed: make_prepped_xfer fails
    - handshake_failed: add_remote_agent fails during request handshake
    - notification_failed: send_notif fails
    - transfer_failed: check_xfer_state returns bad state (e.g., "ERR")
    - transfer_exception: check_xfer_state raises exception
    """
    import logging

    vllm_config = create_vllm_config()

    connector = NixlConnector(
        vllm_config,
        KVConnectorRole.WORKER,
        make_kv_cache_config(block_size=16, swa_enabled=enable_hma),
    )
    connector.connector_worker = FakeNixlConnectorWorker(
        vllm_config,
        connector.engine_id,
        hand_shake_latency=0.0,
        kv_cache_config=connector._kv_cache_config,
    )

    # Configure FailingNixlWrapper to fail in the specified way
    for key, value in wrapper_config.items():
        setattr(connector.connector_worker.nixl_wrapper, key, value)

    request_id = f"test_{failure_type}_req"

    # For notification_failed, we need empty local blocks
    # (full cache hit path to trigger send_notif)
    local_blocks: tuple[()] | tuple[list[int], ...]
    if enable_hma:
        # HMA enabled: multiple groups (FA + SW)
        local_blocks = (
            () if failure_type == "notification_failed" else ([10, 11, 12], [13, 14])
        )
        remote_blocks = [[20, 21, 22], [23, 24]]
    else:
        # HMA disabled: single group
        local_blocks = () if failure_type == "notification_failed" else ([10, 11, 12],)
        remote_blocks = [[20, 21, 22]]

    metadata = NixlConnectorMetadata()
    metadata.add_new_req_to_recv(
        request_id=request_id,
        local_block_ids=local_blocks,
        kv_transfer_params={
            "remote_block_ids": remote_blocks,
            "remote_engine_id": FakeNixlConnectorWorker.REMOTE_ENGINE_ID,
            "remote_request_id": f"prefill-{request_id}",
            "remote_host": "localhost",
            "remote_port": 1234,
            "remote_tp_size": 1,
        },
    )
    connector.bind_connector_metadata(metadata)

    dummy_ctx = ForwardContext(
        no_compile_layers={},
        attn_metadata={},
        slot_mapping={},
    )

    # Capture logs from the nixl connector loggers
    # vLLM loggers have propagate=False, so we need to capture directly
    nixl_logger = logging.getLogger(
        "vllm.distributed.kv_transfer.kv_connector.v1.nixl.base_worker"
    )
    pull_logger = logging.getLogger(
        "vllm.distributed.kv_transfer.kv_connector.v1.nixl.pull_worker"
    )
    captured_logs: list[logging.LogRecord] = []

    class LogCapture(logging.Handler):
        def emit(self, record):
            captured_logs.append(record)

    handler = LogCapture()
    handler.setLevel(logging.ERROR)
    nixl_logger.addHandler(handler)
    pull_logger.addHandler(handler)

    try:
        connector.start_load_kv(dummy_ctx)
        # Process the ready_requests queue (for async handshake)
        connector.bind_connector_metadata(NixlConnectorMetadata())
        # Wait for async handshake to complete
        time.sleep(0.2)
        connector.start_load_kv(dummy_ctx)

        # For transfer_failed/transfer_exception, the error happens in
        # get_finished() when checking transfer state
        if needs_get_finished:
            connector.get_finished(finished_req_ids=set())
    finally:
        nixl_logger.removeHandler(handler)
        pull_logger.removeHandler(handler)

    # Print logs for manual comparison between commits
    error_logs = [r for r in captured_logs if r.levelno >= logging.ERROR]
    print("\n" + "=" * 60)
    print(f"CAPTURED ERROR LOGS for {failure_type}:")
    print("=" * 60)
    for i, record in enumerate(error_logs):
        print(f"\n--- Log {i + 1} ---")
        print(f"Message: {record.message}")
    print("=" * 60 + "\n")

    assert len(error_logs) >= 1, f"Expected at least one error log for {failure_type}"

    # Verify structured logging output (new format)
    # Check that at least one log matches the expected format
    all_messages = [r.message for r in error_logs]
    combined_logs = "\n".join(all_messages)

    assert any("NIXL transfer failure" in msg for msg in all_messages), (
        f"Expected structured log format with 'NIXL transfer failure' prefix "
        f"for {failure_type}. Got: {all_messages}"
    )
    assert any("failure_type" in msg for msg in all_messages), (
        f"Expected 'failure_type' in logs. Got: {all_messages}"
    )
    assert any("Context:" in msg for msg in all_messages), (
        f"Expected 'Context:' in logs. Got: {all_messages}"
    )
    # Check that the expected failure_type appears in at least one log
    # Note: handshake_failed also triggers handshake_setup_failed
    assert failure_type in combined_logs or (
        failure_type == "handshake_failed" and "handshake_setup_failed" in combined_logs
    ), f"Expected '{failure_type}' in logs. Got: {all_messages}"


@patch(
    "vllm.distributed.kv_transfer.kv_connector.v1.nixl.base_worker.NixlWrapper",
    FailingNixlWrapper,
)
def test_handshake_failure_returns_finished(default_vllm_config, dist_init):
    """Test that handshake failures mark blocks invalid and return via get_finished."""
    vllm_config = create_vllm_config()

    connector = NixlConnector(
        vllm_config, KVConnectorRole.WORKER, make_kv_cache_config(block_size=16)
    )
    connector.connector_worker = FakeNixlConnectorWorker(
        vllm_config, connector.engine_id, hand_shake_latency=0.1
    )
    connector.connector_worker.nixl_wrapper.fail_handshake = True

    request_id = "test_handshake_fail"
    metadata = NixlConnectorMetadata()
    metadata.add_new_req_to_recv(
        request_id=request_id,
        local_block_ids=([1, 2, 3],),
        kv_transfer_params={
            "remote_block_ids": ([4, 5, 6],),
            "remote_engine_id": FakeNixlConnectorWorker.REMOTE_ENGINE_ID,
            "remote_request_id": f"prefill-{request_id}",
            "remote_host": "localhost",
            "remote_port": 1234,
            "remote_tp_size": 1,
        },
    )
    connector.bind_connector_metadata(metadata)

    dummy_ctx = ForwardContext(
        no_compile_layers={},
        attn_metadata={},
        slot_mapping={},
    )
    connector.start_load_kv(dummy_ctx)

    # Wait for handshake to fail
    time.sleep(0.3)

    assert connector.get_block_ids_with_load_errors() == set()
    _, done_recving = connector.get_finished(finished_req_ids=set())
    assert request_id in done_recving
    assert connector.get_block_ids_with_load_errors() == {1, 2, 3}

    # Handshake failures are recorded as transport failures, separately
    # from KV expiry.
    assert connector.connector_worker.xfer_stats.data["num_failed_handshakes"]
    assert connector.connector_worker.xfer_stats.data["num_failed_transfers"] == []


@patch(
    "vllm.distributed.kv_transfer.kv_connector.v1.nixl.base_worker.NixlWrapper",
    FailingNixlWrapper,
)
@pytest.mark.parametrize("is_hma", [False, True])
def test_transfer_setup_failure_returns_finished(
    default_vllm_config, dist_init, is_hma
):
    """Setup failures report the request; only non-HMA reports block IDs."""
    vllm_config = create_vllm_config()

    connector = NixlConnector(
        vllm_config, KVConnectorRole.WORKER, make_kv_cache_config(block_size=16)
    )
    connector.connector_worker = FakeNixlConnectorWorker(
        vllm_config, connector.engine_id, hand_shake_latency=0
    )
    connector.connector_worker.nixl_wrapper.fail_transfer_setup = True
    connector.connector_worker._is_hma_required = is_hma

    request_id = "test_transfer_fail"
    metadata = NixlConnectorMetadata()
    metadata.add_new_req_to_recv(
        request_id=request_id,
        local_block_ids=([7, 8, 9],),
        kv_transfer_params={
            "remote_block_ids": ([10, 11, 12],),
            "remote_engine_id": FakeNixlConnectorWorker.REMOTE_ENGINE_ID,
            "remote_request_id": f"prefill-{request_id}",
            "remote_host": "localhost",
            "remote_port": 1234,
            "remote_tp_size": 1,
        },
    )
    connector.bind_connector_metadata(metadata)

    dummy_ctx = ForwardContext(
        no_compile_layers={},
        attn_metadata={},
        slot_mapping={},
    )
    connector.start_load_kv(dummy_ctx)

    # Wait for handshake to complete and process ready_requests
    connector.bind_connector_metadata(NixlConnectorMetadata())
    time.sleep(0.1)
    connector.start_load_kv(dummy_ctx)

    assert connector.get_block_ids_with_load_errors() == set()
    results = connector.get_transfer_results(finished_req_ids=set())
    assert request_id in results.finished_recving
    assert results.failed_recving == {request_id}
    invalid_blocks = connector.get_block_ids_with_load_errors()
    assert invalid_blocks == (set() if is_hma else {7, 8, 9})


class _ScriptedXferWrapper(FakeNixlWrapper):
    """Scripts per-handle xfer states; forbids releasing in-flight handles."""

    def __init__(self, *args, **kwargs):
        super().__init__(*args, **kwargs)
        # handle -> successive states; the last one repeats.
        self.xfer_states: dict[int, list[str]] = {}
        self.released: list[int] = []

    def check_xfer_state(self, handle: int) -> str:
        states = self.xfer_states[handle]
        return states.pop(0) if len(states) > 1 else states[0]

    def release_xfer_handle(self, handle: int) -> None:
        # A posted-but-unfinished transfer cannot be aborted: releasing it
        # leaves the RDMA READ armed. Production code must never do this.
        assert self.xfer_states.get(handle, ["DONE"])[0] != "PROC", (
            f"released in-flight handle {handle}"
        )
        self.released.append(handle)


def _make_split_read_connector(vllm_config, request_id, states):
    """Seed a request with scripted in-flight xfer handles + recv metadata."""
    connector = NixlConnector(
        vllm_config, KVConnectorRole.WORKER, make_kv_cache_config(block_size=16)
    )
    connector.connector_worker = FakeNixlConnectorWorker(
        vllm_config, connector.engine_id, hand_shake_latency=0
    )
    worker = connector.connector_worker
    wrapper = _ScriptedXferWrapper("agent")
    worker.nixl_wrapper = wrapper
    metadata = NixlConnectorMetadata()
    metadata.add_new_req_to_recv(
        request_id=request_id,
        local_block_ids=([7, 8, 9],),
        kv_transfer_params={
            "remote_block_ids": ([10, 11, 12],),
            "remote_engine_id": FakeNixlConnectorWorker.REMOTE_ENGINE_ID,
            "remote_request_id": f"prefill-{request_id}",
            "remote_host": "localhost",
            "remote_port": 1234,
            "remote_tp_size": 1,
        },
    )
    worker._recving_metadata[request_id] = metadata.reqs_to_recv[request_id]
    wrapper.xfer_states = dict(states)
    worker._recving_transfers[request_id] = list(states)
    return connector, worker, wrapper


@patch(
    "vllm.distributed.kv_transfer.kv_connector.v1.nixl.base_worker.NixlWrapper",
    FakeNixlWrapper,
)
def test_split_read_failure_defers_report_until_last_handle(
    default_vllm_config, dist_init
):
    """One half of a split (mixed DRAM/VRAM) read failing must not report the
    request — nor invalidate its blocks — while the sibling xfer is still in
    flight: a posted READ cannot be aborted and would DMA into blocks the
    scheduler could free and reuse. The report happens exactly once, when the
    last handle is terminal."""
    request_id = "split_read_partial_failure"
    err_handle, live_handle = 11, 22
    connector, worker, wrapper = _make_split_read_connector(
        create_vllm_config(),
        request_id,
        {err_handle: ["ERR"], live_handle: ["PROC", "PROC", "DONE"]},
    )

    # Poll 1: one half fails; the sibling is in flight -> nothing reported.
    _, done_recving = connector.get_finished(finished_req_ids=set())
    assert done_recving == set()
    assert connector.get_block_ids_with_load_errors() == set()
    assert request_id in worker._recving_metadata
    assert wrapper.released == [err_handle]

    # Poll 2: sibling still in flight.
    _, done_recving = connector.get_finished(finished_req_ids=set())
    assert done_recving == set()

    # Poll 3: sibling terminal -> reported exactly once, blocks invalidated.
    _, done_recving = connector.get_finished(finished_req_ids=set())
    assert done_recving == {request_id}
    assert connector.get_block_ids_with_load_errors() == {7, 8, 9}
    assert request_id not in worker._recving_metadata
    assert wrapper.released == [err_handle, live_handle]

    # Poll 4: nothing left; no double report.
    _, done_recving = connector.get_finished(finished_req_ids=set())
    assert done_recving == set()


@patch(
    "vllm.distributed.kv_transfer.kv_connector.v1.nixl.base_worker.NixlWrapper",
    FailingNixlWrapper,
)
@pytest.mark.parametrize(
    "failure_mode",
    [
        "handshake",
        "transfer_setup",
        "transfer_failed",
        "transfer_exception",
    ],
)
def test_failed_request_skips_kv_postprocessing(
    default_vllm_config, dist_init, failure_mode
):
    """Test that failed requests skip KV sync and post-processing in
    get_transfer_results().

    This is the core safety behavior: when a KV transfer fails at any stage,
    the request must still appear in done_recving (so the scheduler can apply
    kv_load_failure_policy), but sync_recved_kv_to_device and post-processing
    must NOT be called since no valid KV data was received.

    Covers all failure paths that involve an actual (attempted) KV transfer:
    - handshake: add_remote_agent raises during async handshake
    - transfer_setup: make_prepped_xfer raises before handle is in transfers
    - transfer_failed: check_xfer_state returns bad state ("ERR") in
      _pop_done_transfers — this is the path that previously had the bug
      where post-processing was NOT skipped
    - transfer_exception: check_xfer_state raises in _pop_done_transfers

    Note: notification_failed (send_notif raises on the full-cache-hit path)
    is intentionally excluded. That path is a best-effort D→P courtesy
    notification; the blocks are already in D's cache, so no KV transfer
    was attempted and done_recving is correctly empty.
    """
    # Map each failure mode to the FailingNixlWrapper attribute to set.
    _WRAPPER_CONFIG: dict[str, str] = {
        "handshake": "fail_handshake",
        "transfer_setup": "fail_transfer_setup",
        "transfer_failed": "fail_transfer_state",
        "transfer_exception": "fail_transfer_exception",
    }

    # Use enable_permute_local_kv=True so that
    # post_process_device_kv_on_receive would be called on the success path,
    # making the assertion meaningful (not trivially true).
    vllm_config = create_vllm_config(enable_permute_local_kv=True)

    connector = NixlConnector(
        vllm_config, KVConnectorRole.WORKER, make_kv_cache_config(block_size=16)
    )
    connector.connector_worker = FakeNixlConnectorWorker(
        vllm_config,
        connector.engine_id,
        hand_shake_latency=0.1 if failure_mode == "handshake" else 0,
    )
    worker = connector.connector_worker
    setattr(worker.nixl_wrapper, _WRAPPER_CONFIG[failure_mode], True)

    request_id = f"test_{failure_mode}_skip_postprocess"
    metadata = NixlConnectorMetadata()
    metadata.add_new_req_to_recv(
        request_id=request_id,
        local_block_ids=([1, 2, 3],),
        kv_transfer_params={
            "remote_block_ids": ([4, 5, 6],),
            "remote_engine_id": FakeNixlConnectorWorker.REMOTE_ENGINE_ID,
            "remote_request_id": f"prefill-{request_id}",
            "remote_host": "localhost",
            "remote_port": 1234,
            "remote_tp_size": 1,
        },
    )
    connector.bind_connector_metadata(metadata)

    dummy_ctx = ForwardContext(
        no_compile_layers={},
        attn_metadata={},
        slot_mapping={},
    )
    connector.start_load_kv(dummy_ctx)

    if failure_mode == "handshake":
        # Wait for async handshake to fail.
        time.sleep(0.3)
    else:
        # All other modes: let the handshake complete, then process the
        # ready_requests queue. For transfer_failed / transfer_exception the
        # handle ends up in _recving_transfers; the failure surfaces in
        # get_finished() via _pop_done_transfers below.
        connector.bind_connector_metadata(NixlConnectorMetadata())
        time.sleep(0.1)
        connector.start_load_kv(dummy_ctx)

    # Spy on sync_recved_kv_to_device and post_process_device_kv_on_receive
    # to verify they are NOT called for the failed request.
    with (
        patch.object(worker, "sync_recved_kv_to_device") as mock_sync,
        patch.object(worker, "post_process_device_kv_on_receive") as mock_postprocess,
    ):
        results = connector.get_transfer_results(finished_req_ids=set())

    assert request_id in results.finished_recving
    assert results.failed_recving == {request_id}

    # Critical: KV sync and post-processing must NOT have been called
    # since no valid KV data was received for the failed request.
    mock_sync.assert_not_called()
    mock_postprocess.assert_not_called()

    # Metadata for the request should have been cleaned up.
    assert request_id not in worker._recving_metadata

    # Blocks should have been marked as invalid.
    invalid_blocks = connector.get_block_ids_with_load_errors()
    assert invalid_blocks == {1, 2, 3}


@pytest.mark.parametrize("sibling_state", ["ERR", "DONE"])
def test_recv_failure_waits_for_sibling_transfer(recv_worker, sibling_state):
    """Failure must not release blocks while a sibling can still write to them."""
    worker = recv_worker
    worker._recving_transfers["request"] = [101, 102]
    worker._pending_recv_notifs = {"request": [("prefill", b"request:1")]}
    worker.nixl_wrapper.check_xfer_state.side_effect = ["ERR", "PROC", sibling_state]

    results = worker.get_transfer_results()
    assert results.finished_recving == results.failed_recving == set()
    assert worker.get_block_ids_with_load_errors() == set()
    assert "request" in worker._recving_metadata
    worker.nixl_wrapper.release_xfer_handle.assert_called_once_with(101)

    results = worker.get_transfer_results()
    assert results.finished_recving == results.failed_recving == {"request"}
    assert worker.get_block_ids_with_load_errors() == {1, 2, 3}
    worker.nixl_wrapper.send_notif.assert_not_called()

    results = worker.get_transfer_results()
    assert results.finished_recving == results.failed_recving == set()
    assert worker.get_block_ids_with_load_errors() == set()


def test_recv_failure_waits_for_unpollable_handle_to_be_released(recv_worker):
    """A polling exception and failed release do not prove that DMA has stopped."""
    worker = recv_worker
    worker._recving_transfers["request"] = [101]
    worker.nixl_wrapper.check_xfer_state.side_effect = [RuntimeError("poll"), "DONE"]
    worker.nixl_wrapper.release_xfer_handle.side_effect = [
        RuntimeError("release"),
        None,
    ]

    results = worker.get_transfer_results()
    assert results.finished_recving == results.failed_recving == set()
    assert worker.get_block_ids_with_load_errors() == set()
    assert "request" in worker._recving_metadata

    results = worker.get_transfer_results()
    assert results.finished_recving == results.failed_recving == {"request"}
    assert worker.get_block_ids_with_load_errors() == {1, 2, 3}
    assert worker.nixl_wrapper.release_xfer_handle.call_count == 2


def _set_test_speculative_config(
    vllm_config,
    *,
    method: str = "eagle3",
    model: str = "test/eagle3-drafter",
    revision: str | None = None,
    code_revision: str | None = None,
    num_speculative_tokens: int = 1,
    parallel_drafting: bool = False,
    kv_cache_dtype: str | None = None,
    attention_backend: str | None = None,
    auxiliary_layer_ids: tuple[int, ...] = (2, 16, 29),
) -> None:
    draft_model_config = SimpleNamespace(
        model=model,
        revision=revision,
        code_revision=code_revision,
        hf_config=SimpleNamespace(eagle_aux_hidden_state_layer_ids=auxiliary_layer_ids),
    )
    vllm_config.speculative_config = SimpleNamespace(
        method=method,
        draft_model_config=draft_model_config,
        num_speculative_tokens=num_speculative_tokens,
        parallel_drafting=parallel_drafting,
        kv_cache_dtype=kv_cache_dtype,
        attention_backend=attention_backend,
        use_eagle=lambda: True,
    )


@pytest.mark.parametrize(
    "remote_overrides,should_match",
    [
        ({}, True),
        ({"num_speculative_tokens": 2}, True),
        ({"method": "mtp"}, False),
        ({"model": "test/different-drafter"}, False),
        ({"revision": "different-revision"}, False),
        ({"parallel_drafting": True}, False),
        ({"kv_cache_dtype": "fp8"}, False),
        # attention_backend is intentionally not part of the compat hash
        # (see _get_speculative_compatibility_factors); overriding it must
        # not change the hash.
        ({"attention_backend": "FLASHINFER"}, True),
        ({"auxiliary_layer_ids": (2, 16, 30)}, False),
    ],
)
@pytest.mark.skip_global_cleanup
def test_speculative_config_compatibility_hash(
    remote_overrides: dict[str, Any], should_match: bool
):
    local_config = create_vllm_config()
    remote_config = create_vllm_config()
    _set_test_speculative_config(local_config)
    _set_test_speculative_config(remote_config, **remote_overrides)

    local_hash = compute_nixl_compatibility_hash(local_config, "FLASH_ATTN")
    remote_hash = compute_nixl_compatibility_hash(remote_config, "FLASH_ATTN")

    assert (local_hash == remote_hash) is should_match


@pytest.mark.skip_global_cleanup
def test_missing_speculative_config_changes_compatibility_hash():
    regular_config = create_vllm_config()
    speculative_config = create_vllm_config()
    _set_test_speculative_config(speculative_config)

    regular_hash = compute_nixl_compatibility_hash(regular_config, "FLASH_ATTN")
    speculative_hash = compute_nixl_compatibility_hash(speculative_config, "FLASH_ATTN")

    assert regular_hash != speculative_hash


@pytest.mark.skip_global_cleanup
def test_speculative_kv_cache_dtype_resolves_to_target():
    # The draft kv_cache_dtype override defaults to None ("inherit the target's
    # --kv-cache-dtype"). An explicit setting on one side that matches the
    # other side's inherited (resolved) dtype must not spuriously mismatch.
    local_config = create_vllm_config(cache_dtype="fp8")
    remote_config = create_vllm_config(cache_dtype="fp8")
    _set_test_speculative_config(local_config, kv_cache_dtype="fp8")  # explicit
    _set_test_speculative_config(remote_config, kv_cache_dtype=None)  # inherits

    local_hash = compute_nixl_compatibility_hash(local_config, "FLASH_ATTN")
    remote_hash = compute_nixl_compatibility_hash(remote_config, "FLASH_ATTN")

    assert local_hash == remote_hash


@pytest.mark.skip_global_cleanup
def test_speculative_attention_backend_not_in_compatibility_hash():
    # The draft attention_backend is intentionally excluded from the hash: the
    # connector only has the raw override (auto-select), and its transfer-
    # relevant effect is validated per region at runtime. Differing overrides
    # must not change the hash.
    local_config = create_vllm_config()
    remote_config = create_vllm_config()
    _set_test_speculative_config(local_config, attention_backend=None)
    _set_test_speculative_config(remote_config, attention_backend="FLASHINFER")

    local_hash = compute_nixl_compatibility_hash(local_config, "FLASH_ATTN")
    remote_hash = compute_nixl_compatibility_hash(remote_config, "FLASH_ATTN")

    assert local_hash == remote_hash


@pytest.mark.skip_global_cleanup
def test_transfer_mode_changes_compatibility_hash():
    # push (WRITE) and pull (READ) connectors use incompatible transfer
    # protocols, so their compatibility hashes must differ; identical modes
    # must match. The default mode is pull.
    config = create_vllm_config()

    pull_hash = compute_nixl_compatibility_hash(
        config, "FLASH_ATTN", transfer_mode="pull"
    )
    push_hash = compute_nixl_compatibility_hash(
        config, "FLASH_ATTN", transfer_mode="push"
    )

    assert pull_hash != push_hash
    assert pull_hash == compute_nixl_compatibility_hash(
        config, "FLASH_ATTN", transfer_mode="pull"
    )
    assert compute_nixl_compatibility_hash(config, "FLASH_ATTN") == pull_hash


@pytest.mark.skip_global_cleanup
@pytest.mark.parametrize("mode", ["pull", "push"])
@pytest.mark.parametrize(
    "tp_size,pcp_size,dcp_size,transfer_tp_size",
    [(1, 1, 1, 1), (4, 1, 1, 4), (4, 1, 4, 4), (1, 4, 1, 1), (1, 4, 4, 4)],
)
def test_scheduler_advertises_transfer_topology(
    mode, tp_size, pcp_size, dcp_size, transfer_tp_size
):
    """The consumer must address every distinct producer KV shard."""
    from vllm.distributed.kv_transfer.kv_connector.v1.nixl.pull_scheduler import (
        NixlPullConnectorScheduler,
    )
    from vllm.distributed.kv_transfer.kv_connector.v1.nixl.push_scheduler import (
        NixlPushConnectorScheduler,
    )

    config = create_vllm_config(kv_role="kv_producer")
    config.parallel_config.tensor_parallel_size = tp_size
    config.parallel_config.prefill_context_parallel_size = pcp_size
    config.parallel_config.decode_context_parallel_size = dcp_size
    cls = NixlPullConnectorScheduler if mode == "pull" else NixlPushConnectorScheduler
    scheduler = cls(config, "prefiller", make_kv_cache_config(block_size=16))
    request = create_request(request_id=1, num_tokens=32, do_remote_decode=True)
    request.status = RequestStatus.FINISHED_LENGTH_CAPPED
    try:
        delay, params = scheduler.request_finished(request, ([0, 1],))
        assert delay
        assert params["transfer_mode"] == mode
        assert params["tp_size"] == transfer_tp_size
        assert params["dcp_size"] == dcp_size
    finally:
        scheduler.shutdown()


@pytest.mark.parametrize(
    "mismatch_type,config_overrides,version_override,should_fail,enforce_handshake_compat",
    [
        ("vllm_version", {}, {"vllm_version": "0.6.1"}, True, True),
        ("nixl_connector_version", {}, {"connector_version": 37}, True, True),
        ("model_name", {"model": "facebook/opt-350m"}, {}, True, True),
        ("dtype", {"dtype": "bfloat16"}, {}, True, True),
        ("cache_dtype", {"cache_dtype": "fp8"}, {}, True, True),
        ("num_kv_heads", {"hf_overrides": {"num_key_value_heads": 8}}, {}, True, True),
        (
            "num_hidden_layers",
            {"hf_overrides": {"num_hidden_layers": 24}},
            {},
            True,
            True,
        ),
        ("hidden_size", {"hf_overrides": {"hidden_size": 1536}}, {}, True, True),
        ("block_size", {"block_size": 8}, {}, False, True),
        ("matching_config", {}, {}, False, True),
        ("escape_hatch", {"model": "facebook/opt-350m"}, {}, False, False),
    ],
)
@patch(
    "vllm.distributed.kv_transfer.kv_connector.v1.nixl.base_worker.NixlWrapper",
    FakeNixlWrapper,
)
def test_compatibility_hash_validation(
    default_vllm_config,
    dist_init,
    mismatch_type,
    config_overrides,
    version_override,
    should_fail,
    enforce_handshake_compat,
):
    """Test NIXL compatibility hash validation during handshake.

    Parameters
    ----------
        mismatch_type: description of what is being tested
        config_overrides: dict of config to override for the remote instance
        version_override: version dict e.g. {"vllm_version": "0.6.1"}
        should_fail: whether the handshake should fail
        enforce_handshake_compat: whether to enforce compatibility checking

    """
    local_vllm_config = create_vllm_config(
        model="facebook/opt-125m",
        block_size=16,
        kv_connector_extra_config={
            "enforce_handshake_compat": enforce_handshake_compat
        },
    )
    kv_cache_config = make_kv_cache_config(block_size=16, num_blocks=2)
    decode_connector = NixlConnector(
        local_vllm_config, KVConnectorRole.WORKER, kv_cache_config
    )
    decode_worker = decode_connector.connector_worker
    kv_cache_spec = cast(
        AttentionSpec, kv_cache_config.kv_cache_groups[0].kv_cache_spec
    )
    shape = compute_layer_kv_cache_shape_bytes(
        kv_cache_spec, kv_cache_config.num_blocks
    )
    shared_tensor = torch.zeros(*shape, dtype=torch.int8).view(kv_cache_spec.dtype)
    unique_tensor = torch.zeros(*shape, dtype=torch.int8).view(kv_cache_spec.dtype)
    # Build kv_caches from the actual layer names in kv_cache_config so that
    # _layer_specs lookups in register_kv_caches always find a matching key.
    layer_names = [
        name for group in kv_cache_config.kv_cache_groups for name in group.layer_names
    ]
    kv_caches = {
        name: shared_tensor if i % 2 == 0 else unique_tensor
        for i, name in enumerate(layer_names)
    }
    decode_connector.register_kv_caches(kv_caches)

    remote_config_params: dict[str, Any] = {
        "model": "facebook/opt-125m",
        "block_size": 16,
        **config_overrides,
    }
    remote_vllm_config = create_vllm_config(**remote_config_params)

    with contextlib.ExitStack() as stack:
        if "vllm_version" in version_override:
            stack.enter_context(
                patch("vllm.__version__", version_override["vllm_version"])
            )
        elif "connector_version" in version_override:
            stack.enter_context(
                patch.object(
                    nixl.metadata,
                    "NIXL_CONNECTOR_VERSION",
                    version_override["connector_version"],
                )
            )
        remote_hash = compute_nixl_compatibility_hash(
            remote_vllm_config,
            decode_worker.backend_name,
        )

    prefill_block_size = config_overrides.get("block_size", 16)
    prefill_block_lens = [4096 * prefill_block_size]
    prefill_metadata = NixlAgentMetadata(
        engine_id=FakeNixlConnectorWorker.REMOTE_ENGINE_ID,
        agent_metadata=FakeNixlWrapper.AGENT_METADATA,
        kv_caches_base_addr=[0],
        device_id=0,
        num_blocks=1,
        block_lens=prefill_block_lens,
        block_strides=prefill_block_lens,
        kv_cache_layout="LBHNC",
        block_size=prefill_block_size,
        ssm_sizes=(0, 0),
        attn_backend_name=decode_worker.backend_name,
        physical_blocks_per_logical_kv_block=1,
    )
    handshake_payload = NixlHandshakePayload(
        compatibility_hash=remote_hash,
        agent_metadata_bytes=msgspec.msgpack.encode(prefill_metadata),
    )

    # Mock ZMQ socket to return our handshake payload
    mock_socket = MagicMock()
    mock_socket.recv_multipart.return_value = [
        msgspec.msgpack.encode(handshake_payload),
        msgspec.msgpack.encode(time.perf_counter()),
    ]

    # Mock add_remote_agent to avoid actual NIXL operations
    # Patch zmq_ctx to return our mock socket
    with (
        patch.object(decode_worker, "add_remote_agent", return_value="fake_agent"),
        patch.object(nixl.base_worker, "zmq_ctx") as mock_zmq_ctx,
    ):
        mock_zmq_ctx.return_value.__enter__.return_value = mock_socket

        if should_fail:
            with pytest.raises(RuntimeError, match="compatibility hash mismatch"):
                decode_worker._nixl_handshake(
                    host="localhost",
                    port=1234,
                    remote_tp_size=1,
                    expected_engine_id=FakeNixlConnectorWorker.REMOTE_ENGINE_ID,
                )
        else:
            result, _ = decode_worker._nixl_handshake(
                host="localhost",
                port=1234,
                remote_tp_size=1,
                expected_engine_id=FakeNixlConnectorWorker.REMOTE_ENGINE_ID,
            )
            # Verify handshake returned agent mapping
            assert isinstance(result, dict)
            assert len(result) == 1


@pytest.mark.parametrize(
    "error_scenario",
    [
        "handshake_decode_error",
        "handshake_validation_error",
        "metadata_decode_error",
        "metadata_validation_error",
    ],
)
@patch(
    "vllm.distributed.kv_transfer.kv_connector.v1.nixl.base_worker.NixlWrapper",
    FakeNixlWrapper,
)
def test_handshake_decode_errors(default_vllm_config, dist_init, error_scenario):
    """Test that msgspec decode errors are properly handled during handshake.

    Tests both DecodeError and ValidationError for both decoders:
    - NixlHandshakePayload decoder
    - NixlAgentMetadata decoder
    """
    local_vllm_config = create_vllm_config(
        model="facebook/opt-125m",
        block_size=16,
    )
    decode_connector = NixlConnector(
        local_vllm_config, KVConnectorRole.WORKER, make_kv_cache_config(block_size=16)
    )
    decode_worker = decode_connector.connector_worker

    backend = get_current_attn_backend(local_vllm_config)
    probe_spec = decode_worker.kv_cache_config.kv_cache_groups[0].kv_cache_spec
    test_shape = compute_layer_kv_cache_shape_bytes(probe_spec, 1)
    decode_worker.transfer_topo = TransferTopology(
        tp_rank=decode_worker.tp_rank,
        tp_size=decode_worker.world_size,
        block_size=decode_worker.block_size,
        engine_id=decode_worker.engine_id,
        is_mla=decode_worker.use_mla,
        is_mamba=False,
        total_num_kv_heads=decode_worker.model_config.get_total_num_kv_heads(),
        attn_backends=[backend],
        tensor_shape=test_shape,
    )

    decode_worker.compat_hash = compute_nixl_compatibility_hash(
        decode_worker.vllm_config,
        decode_worker.backend_name,
    )

    if error_scenario == "handshake_decode_error":
        msg_bytes = b"this is not valid msgpack data"
    elif error_scenario == "handshake_validation_error":
        msg_bytes = msgspec.msgpack.encode({"wrong_field": "value"})
    elif error_scenario == "metadata_decode_error":
        valid_handshake = NixlHandshakePayload(
            compatibility_hash=decode_worker.compat_hash,
            agent_metadata_bytes=b"invalid msgpack for metadata",
        )
        msg_bytes = msgspec.msgpack.encode(valid_handshake)

    elif error_scenario == "metadata_validation_error":
        valid_handshake = NixlHandshakePayload(
            compatibility_hash=decode_worker.compat_hash,
            agent_metadata_bytes=msgspec.msgpack.encode({"missing": "fields"}),
        )
        msg_bytes = msgspec.msgpack.encode(valid_handshake)
    else:
        raise AssertionError(f"{error_scenario} not a valid scenario")

    mock_socket = MagicMock()
    mock_socket.recv_multipart.return_value = [
        msg_bytes,
        msgspec.msgpack.encode(time.perf_counter()),
    ]
    with (
        patch.object(decode_worker, "add_remote_agent", return_value="fake_agent"),
        patch.object(nixl.base_worker, "zmq_ctx") as mock_zmq_ctx,
    ):
        mock_zmq_ctx.return_value.__enter__.return_value = mock_socket

        with pytest.raises(RuntimeError):
            decode_worker._nixl_handshake(
                host="localhost",
                port=1234,
                remote_tp_size=1,
                expected_engine_id=FakeNixlConnectorWorker.REMOTE_ENGINE_ID,
            )

    @patch(
        "vllm.distributed.kv_transfer.kv_connector.v1.nixl.base_worker.NixlWrapper",
        FakeNixlWrapper,
    )
    def test_mla_broadcast_notif_uses_remote_request_id(
        self, default_vllm_config, dist_init
    ):
        """MLA + remote TP > local TP: the broadcast notification sent to
        non-read prefill ranks must be keyed by the prefill-side request
        id (``meta.remote.request_id``), not the local decode request id.

        Prefill ranks key ``_reqs_to_send`` by their own request id, so a
        broadcast keyed by the decode id is rejected in
        ``_get_new_notifs`` with "Potentially invalid KV blocks for
        unrecognized request" and the blocks only release via the abort
        timeout. See ``_read_blocks_for_req`` in
        ``vllm/distributed/kv_transfer/kv_connector/v1/nixl/worker.py``.
        """
        decode_tp_size = 1
        prefill_tp_size = 4

        vllm_config = create_vllm_config()
        vllm_config.parallel_config.tensor_parallel_size = decode_tp_size

        connector = NixlConnector(
            vllm_config, KVConnectorRole.WORKER, make_kv_cache_config(block_size=16)
        )
        connector.connector_worker = FakeNixlConnectorWorker(
            vllm_config, connector.engine_id, hand_shake_latency=0
        )
        worker = connector.connector_worker

        # Force the MLA path; only `self.use_mla` gates the branches we
        # exercise inside `_read_blocks_for_req`.
        worker.use_mla = True

        # Manually register the remote (P) engine and pre-populate the
        # per-rank state the handshake would normally fill in. The real
        # `_nixl_handshake` is unnecessary here — we only need
        # `transfer_topo` to know `remote_tp_size`, and `_remote_agents`
        # / `dst_xfer_side_handles` to be keyed by remote rank.
        remote_engine_id = "remote_engine"
        worker.transfer_topo.register_remote_engine(
            remote_engine_id=remote_engine_id,
            remote_tp_size=prefill_tp_size,
            remote_block_size=worker.block_size,
            remote_block_len=worker.block_size * 4096,
            remote_physical_blocks_per_logical=1,
            local_block_len=worker.block_size * 4096,
        )
        worker._remote_agents[remote_engine_id] = {
            (0, rank): f"agent_p{rank}" for rank in range(prefill_tp_size)
        }
        worker.dst_xfer_side_handles = {
            remote_engine_id: {rank: 100 + rank for rank in range(prefill_tp_size)}
        }
        # Sanity: D TP=1, P TP=4 => tp_ratio = -4 (P > D).
        assert worker.transfer_topo.tp_ratio(prefill_tp_size) == -prefill_tp_size

        # Distinct ids on each side — that's the whole point of the bug.
        decode_req_id = "decode-req-AAAA"
        prefill_req_id = "prefill-req-BBBB"
        assert decode_req_id != prefill_req_id

        metadata = NixlConnectorMetadata()
        metadata.add_new_req_to_recv(
            request_id=decode_req_id,
            local_block_ids=([0, 1, 2],),
            kv_transfer_params={
                "remote_block_ids": ([10, 11, 12],),
                "remote_engine_id": remote_engine_id,
                "remote_request_id": prefill_req_id,
                "remote_host": "localhost",
                "remote_port": 1234,
                "remote_tp_size": prefill_tp_size,
            },
        )
        meta = metadata.reqs_to_recv[decode_req_id]

        # Capture broadcast send_notif calls; stub `_read_blocks` so we
        # don't need a working xfer path. Real `_read_blocks` emits its
        # auto-notif via `make_prepped_xfer`, not via `send_notif`, so
        # any captured `send_notif` here is a broadcast.
        send_notif_calls: list[tuple[str, bytes]] = []
        worker.nixl_wrapper.send_notif = (  # type: ignore[method-assign]
            lambda agent_name, notif_msg: send_notif_calls.append(
                (agent_name, notif_msg)
            )
        )
        worker._read_blocks = MagicMock()  # type: ignore[method-assign]

        worker._read_blocks_for_req(decode_req_id, meta)

        # MLA: read once from rank 0 and broadcast to the other ranks.
        worker._read_blocks.assert_called_once()
        assert worker._read_blocks.call_args.kwargs["remote_rank"] == 0
        assert (
            worker._read_blocks.call_args.kwargs["remote_request_id"] == prefill_req_id
        )

        # Broadcast goes to ranks {1, 2, 3} only, never to the read target.
        expected_recipients = {
            worker._remote_agents[remote_engine_id][(0, r)]
            for r in range(1, prefill_tp_size)
        }
        assert {agent for agent, _ in send_notif_calls} == expected_recipients

        # Every broadcast notif must be keyed by the prefill request id.
        # Pre-fix this used the *decode* request id, which prefill ranks
        # didn't recognize.
        expected_notif = f"{prefill_req_id}:{decode_tp_size}".encode()
        bad_notif = f"{decode_req_id}:{decode_tp_size}".encode()
        for agent, notif in send_notif_calls:
            assert notif == expected_notif, (
                f"Broadcast notif to {agent!r} must use prefill_req_id; "
                f"got {notif!r} (expected {expected_notif!r}, "
                f"buggy form would be {bad_notif!r})"
            )


@patch(
    "vllm.distributed.kv_transfer.kv_connector.v1.nixl.base_worker.NixlWrapper",
    FakeNixlWrapper,
)
def test_kv_both_deprecation_warning(default_vllm_config, dist_init):
    """kv_role='kv_both' should emit a deprecation log warning."""
    from vllm.logger import _print_warning_once

    _print_warning_once.cache_clear()

    vllm_config = create_vllm_config(kv_role="kv_both")

    with patch(
        "vllm.distributed.kv_transfer.kv_connector.v1.nixl.connector.logger"
    ) as mock_logger:
        mock_logger.warning_once = mock_logger.warning_once
        NixlConnector(
            vllm_config,
            KVConnectorRole.WORKER,
            make_kv_cache_config(block_size=16),
        )

    mock_logger.warning_once.assert_called_once()
    msg = mock_logger.warning_once.call_args[0][0]
    assert "kv_role='kv_both'" in msg
    assert "deprecated" in msg


@patch(
    "vllm.distributed.kv_transfer.kv_connector.v1.nixl.base_worker.NixlWrapper",
    FakeNixlWrapper,
)
def test_explicit_kv_role_no_deprecation_warning(default_vllm_config, dist_init):
    """kv_role='kv_consumer' or 'kv_producer' should NOT emit a warning."""
    for role in ("kv_consumer", "kv_producer"):
        vllm_config = create_vllm_config(kv_role=role)
        with patch(
            "vllm.distributed.kv_transfer.kv_connector.v1.nixl.connector.logger"
        ) as mock_logger:
            NixlConnector(
                vllm_config,
                KVConnectorRole.WORKER,
                make_kv_cache_config(block_size=16),
            )

        (
            mock_logger.warning_once.assert_not_called(),
            (f"kv_role={role!r} should not emit deprecation warning"),
        )
