[P/D] [NixlConnector] kv load recovery integration (#26171)

Signed-off-by: Will Eaton <weaton@redhat.com>
This commit is contained in:
Will Eaton
2025-10-13 11:48:04 -04:00
committed by GitHub
parent 0d21b9b51e
commit 53c9a7cee2
3 changed files with 252 additions and 26 deletions

View File

@@ -190,7 +190,6 @@ def _make_fake_nixl_pkg():
# Copy of FakeNixlWrapper implementation for Ray workers
import uuid
from collections import defaultdict
from typing import Optional
{fake_nixl_source}
@@ -1143,3 +1142,145 @@ def test_aborted_request_removed_from_worker_in_batch(dist_init):
# 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
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)
@patch(
"vllm.distributed.kv_transfer.kv_connector.v1.nixl_connector.NixlWrapper",
FailingNixlWrapper,
)
def test_handshake_failure_returns_finished(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)
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(
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_host": "localhost",
"remote_port": 1234,
"remote_tp_size": 1,
},
)
connector.bind_connector_metadata(metadata)
dummy_ctx = ForwardContext(
no_compile_layers={},
attn_metadata={},
virtual_engine=0,
)
connector.start_load_kv(dummy_ctx)
# Wait for handshake to fail
time.sleep(0.3)
# Check that blocks were marked invalid
invalid_blocks = connector.get_block_ids_with_load_errors()
assert invalid_blocks == {1, 2, 3}
# Check that request appears in get_finished
_, done_recving = connector.get_finished(finished_req_ids=set())
assert request_id in done_recving
@patch(
"vllm.distributed.kv_transfer.kv_connector.v1.nixl_connector.NixlWrapper",
FailingNixlWrapper,
)
def test_transfer_setup_failure_returns_finished(dist_init):
"""Test that transfer setup failures mark blocks invalid
and return via get_finished."""
vllm_config = create_vllm_config()
connector = NixlConnector(vllm_config, KVConnectorRole.WORKER)
connector.connector_worker = FakeNixlConnectorWorker(
vllm_config, connector.engine_id, hand_shake_latency=0
)
connector.connector_worker.nixl_wrapper.fail_transfer_setup = True
request_id = "test_transfer_fail"
metadata = NixlConnectorMetadata()
metadata.add_new_req(
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_host": "localhost",
"remote_port": 1234,
"remote_tp_size": 1,
},
)
connector.bind_connector_metadata(metadata)
dummy_ctx = ForwardContext(
no_compile_layers={},
attn_metadata={},
virtual_engine=0,
)
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)
# check that blocks were marked invalid
invalid_blocks = connector.get_block_ids_with_load_errors()
assert invalid_blocks == {7, 8, 9}
# ensure request appears in get_finished
_, done_recving = connector.get_finished(finished_req_ids=set())
assert request_id in done_recving