diff --git a/tests/v1/kv_connector/unit/test_kv_connector_lifecycle.py b/tests/v1/kv_connector/unit/test_kv_connector_lifecycle.py index a9a38a17b94..b61449d739b 100644 --- a/tests/v1/kv_connector/unit/test_kv_connector_lifecycle.py +++ b/tests/v1/kv_connector/unit/test_kv_connector_lifecycle.py @@ -1,6 +1,8 @@ # SPDX-License-Identifier: Apache-2.0 # SPDX-FileCopyrightText: Copyright contributors to the vLLM project +from unittest.mock import MagicMock, patch + from vllm.distributed.kv_transfer.kv_connector.v1.example_connector import ( # noqa: E501 ExampleConnectorMetadata, ) @@ -38,11 +40,20 @@ def test_kv_connector_mixin_clears_metadata(): vllm_config.kv_transfer_config.kv_role = "kv_both" vllm_config.kv_transfer_config.kv_connector_extra_config["name"] = "unit" - # Initialize the global connector instance kv_cache_config = KVCacheConfig( num_blocks=0, kv_cache_tensors=[], kv_cache_groups=[] ) - ensure_kv_transfer_initialized(vllm_config, kv_cache_config) + # Initialize the global connector instance. + # kv_transfer init now syncs engine_id across TP, so unit tests need + # a minimal mocked TP group. + mock_tp_group = MagicMock() + mock_tp_group.broadcast_object.side_effect = lambda value, src=0: value + + with patch( + "vllm.distributed.parallel_state.get_tp_group", + return_value=mock_tp_group, + ): + ensure_kv_transfer_initialized(vllm_config, kv_cache_config) try: # Minimal scheduler output with empty metadata; mixin should still diff --git a/vllm/distributed/kv_transfer/kv_transfer_state.py b/vllm/distributed/kv_transfer/kv_transfer_state.py index 4392d652077..67a6b4ca7a6 100644 --- a/vllm/distributed/kv_transfer/kv_transfer_state.py +++ b/vllm/distributed/kv_transfer/kv_transfer_state.py @@ -48,6 +48,20 @@ def is_v1_kv_transfer_group(connector: KVConnectorBaseType | None = None) -> boo return isinstance(connector, KVConnectorBase_V1) +def _sync_engine_id_across_tp(vllm_config: "VllmConfig") -> None: + """Broadcast engine_id from TP rank 0 so all workers in a + multi-node TP group share the same value.""" + from vllm.distributed.parallel_state import ( + get_tp_group, + ) + + assert vllm_config.kv_transfer_config is not None + synced_id = get_tp_group().broadcast_object( + vllm_config.kv_transfer_config.engine_id, src=0 + ) + vllm_config.kv_transfer_config.engine_id = synced_id + + def ensure_kv_transfer_initialized( vllm_config: "VllmConfig", kv_cache_config: "KVCacheConfig" ) -> None: @@ -64,6 +78,8 @@ def ensure_kv_transfer_initialized( vllm_config.kv_transfer_config.is_kv_transfer_instance and _KV_CONNECTOR_AGENT is None ): + _sync_engine_id_across_tp(vllm_config) + _KV_CONNECTOR_AGENT = KVConnectorFactory.create_connector( config=vllm_config, role=KVConnectorRole.WORKER,