init v0.23.0

Signed-off-by: Sun Ruoxi <sunruoxi@4paradigm.com>
This commit is contained in:
2026-08-27 15:11:51 +08:00
parent b582a8e7d1
commit 7f8a1b1f7a
2849 changed files with 712887 additions and 22001 deletions

View File

@@ -0,0 +1,463 @@
#
# Copyright (c) 2026 Huawei Technologies Co., Ltd. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
# This file is a part of the vllm-ascend project.
#
"""Mock heavy dependencies (torch, vllm, etc.) for ascend_store unit tests.
IMPORTANT: This module MUST be imported before any vllm_ascend or vllm
imports in each test file.
Usage at the top of each test file:
import tests.ut.distributed.ascend_store._mock_deps # noqa: F401, E402
"""
import importlib.util
import logging
import os
import sys
import types
from typing import Any
from unittest.mock import MagicMock
# ---------------------------------------------------------------------------
# Mock torch / torch_npu
# ---------------------------------------------------------------------------
if "torch" not in sys.modules and importlib.util.find_spec("torch") is None:
_torch = types.ModuleType("torch")
_torch.Tensor = MagicMock # type: ignore[attr-defined]
_torch.bool = "bool" # type: ignore[attr-defined]
_torch.float16 = "float16" # type: ignore[attr-defined]
_torch.float32 = "float32" # type: ignore[attr-defined]
_torch.zeros = MagicMock(return_value=MagicMock()) # type: ignore[attr-defined]
_torch.sum = MagicMock(return_value=0) # type: ignore[attr-defined]
_torch.device = MagicMock() # type: ignore[attr-defined]
_torch.distributed = MagicMock() # type: ignore[attr-defined]
_npu = MagicMock()
_npu.Event = MagicMock
_npu.current_device = MagicMock(return_value=0)
_npu.set_device = MagicMock()
_torch.npu = _npu # type: ignore[attr-defined]
sys.modules["torch"] = _torch
sys.modules["torch.distributed"] = _torch.distributed # type: ignore[attr-defined]
if "torch_npu" not in sys.modules:
sys.modules["torch_npu"] = MagicMock()
sys.modules["torch_npu._inductor"] = MagicMock()
# ---------------------------------------------------------------------------
# Mock vllm modules
# ---------------------------------------------------------------------------
_MOCK_VLLM_DEPS = importlib.util.find_spec("vllm") is None
_vllm_mock_modules = [
"vllm",
"vllm.config",
"vllm.distributed",
"vllm.distributed.kv_events",
"vllm.distributed.kv_transfer",
"vllm.distributed.kv_transfer.kv_connector",
"vllm.distributed.kv_transfer.kv_connector.factory",
"vllm.distributed.kv_transfer.kv_connector.v1",
"vllm.distributed.kv_transfer.kv_connector.v1.base",
"vllm.distributed.parallel_state",
"vllm.envs",
"vllm.forward_context",
"vllm.logger",
"vllm.model_executor",
"vllm.model_executor.layers",
"vllm.model_executor.layers.linear",
"vllm.model_executor.layers.quantization",
"vllm.platforms",
"vllm.utils",
"vllm.utils.hashing",
"vllm.utils.math_utils",
"vllm.utils.network_utils",
"vllm.v1",
"vllm.v1.attention",
"vllm.v1.attention.backend",
"vllm.v1.core",
"vllm.v1.core.block_pool",
"vllm.v1.core.kv_cache_manager",
"vllm.v1.core.kv_cache_utils",
"vllm.v1.core.sched",
"vllm.v1.core.sched.output",
"vllm.v1.core.single_type_kv_cache_manager",
"vllm.v1.kv_cache_interface",
"vllm.v1.kv_cache_spec_registry",
"vllm.v1.outputs",
"vllm.v1.request",
"vllm.v1.serial_utils",
]
if _MOCK_VLLM_DEPS:
for _mod_name in _vllm_mock_modules:
if _mod_name not in sys.modules:
sys.modules[_mod_name] = MagicMock()
if _MOCK_VLLM_DEPS:
sys.modules["vllm.utils.math_utils"].cdiv = lambda a, b: -(-a // b) # type: ignore[attr-defined]
sys.modules["vllm.logger"].logger = logging.getLogger("vllm") # type: ignore[attr-defined]
_base_mod: Any = (
sys.modules["vllm.distributed.kv_transfer.kv_connector.v1.base"] if _MOCK_VLLM_DEPS else types.SimpleNamespace()
)
_base_mod.KVConnectorBase_V1 = type("KVConnectorBase_V1", (), {"__init__": lambda self, **kw: None}) # type: ignore[attr-defined]
_base_mod.KVConnectorMetadata = type("KVConnectorMetadata", (), {}) # type: ignore[attr-defined]
_base_mod.KVConnectorWorkerMetadata = type("KVConnectorWorkerMetadata", (), {}) # type: ignore[attr-defined]
_base_mod.KVConnectorRole = MagicMock() # type: ignore[attr-defined]
_base_mod.KVConnectorRole.SCHEDULER = "SCHEDULER"
_base_mod.KVConnectorRole.WORKER = "WORKER"
_base_mod.SupportsHMA = type("SupportsHMA", (), {}) # type: ignore[attr-defined]
_events_mod: Any = sys.modules["vllm.distributed.kv_events"] if _MOCK_VLLM_DEPS else types.SimpleNamespace()
_events_mod.KVCacheEvent = type("KVCacheEvent", (), {}) # type: ignore[attr-defined]
_events_mod.KVConnectorKVEvents = type("KVConnectorKVEvents", (), {}) # type: ignore[attr-defined]
class _FakeAggregator:
def __init__(self, *args, **kwargs):
self._mock = MagicMock()
def __getattr__(self, name):
return getattr(self._mock, name)
_events_mod.KVEventAggregator = _FakeAggregator # type: ignore[attr-defined]
_events_mod.BlockStored = type( # type: ignore[attr-defined]
"BlockStored",
(),
{"__init__": lambda self, **kwargs: self.__dict__.update(kwargs)},
)
_kv_cache_utils_mod: Any = sys.modules["vllm.v1.core.kv_cache_utils"] if _MOCK_VLLM_DEPS else types.SimpleNamespace()
_kv_cache_utils_mod.BlockHash = bytes # type: ignore[attr-defined]
_kv_cache_utils_mod.maybe_convert_block_hash = lambda x: x # type: ignore[attr-defined]
class _FakeKVCacheBlock:
def __init__(self, block_id=0, **kwargs):
self.block_id = block_id
self.__dict__.update(kwargs)
class _FakeKVCacheSpec:
def __init__(self, block_size=16, **kwargs):
self.block_size = block_size
for key, value in kwargs.items():
setattr(self, key, value)
def __eq__(self, other):
return type(self) is type(other) and self.__dict__ == getattr(other, "__dict__", {})
def copy_with_new_block_size(self, block_size):
kwargs = self.__dict__.copy()
kwargs["block_size"] = block_size
return type(self)(**kwargs)
@property
def page_size_bytes(self):
num_kv_heads = getattr(self, "num_kv_heads", 1)
head_size = getattr(self, "head_size", 1)
dtype = getattr(self, "dtype", None)
dtype_size = getattr(dtype, "itemsize", None)
if dtype_size is None and dtype is not None and hasattr(dtype, "element_size"):
dtype_size = dtype.element_size()
return self.block_size * num_kv_heads * head_size * int(dtype_size or 1) * 2
class _FakeFullAttentionSpec(_FakeKVCacheSpec):
pass
class _FakeSlidingWindowSpec(_FakeKVCacheSpec):
def __init__(self, block_size=16, sliding_window=32, **kwargs):
super().__init__(block_size=block_size, sliding_window=sliding_window, **kwargs)
class _FakeMambaSpec(_FakeKVCacheSpec):
def __init__(self, block_size=16, **kwargs):
super().__init__(block_size=block_size, **kwargs)
self.num_speculative_blocks = getattr(self, "num_speculative_blocks", 0)
class _FakeUniformTypeKVCacheSpecs(_FakeKVCacheSpec):
def __init__(self, block_size=16, kv_cache_specs=None, **kwargs):
super().__init__(block_size=block_size, **kwargs)
self.kv_cache_specs = kv_cache_specs or {}
@classmethod
def from_specs(cls, kv_cache_specs):
if not kv_cache_specs:
return None
first_spec = next(iter(kv_cache_specs.values()))
return cls(
block_size=getattr(first_spec, "block_size", 16),
kv_cache_specs=kv_cache_specs,
)
class _FakeKVCacheGroupSpec:
def __init__(self, layer_names=None, kv_cache_spec=None, is_eagle_group=False):
self.layer_names = layer_names or []
self.kv_cache_spec = kv_cache_spec or _FakeFullAttentionSpec()
self.is_eagle_group = is_eagle_group
class _FakeKVCacheConfig:
def __init__(self, num_blocks=1, kv_cache_tensors=None, kv_cache_groups=None):
self.num_blocks = num_blocks
self.kv_cache_tensors = kv_cache_tensors or []
self.kv_cache_groups = kv_cache_groups or []
_kv_cache_utils_mod.KVCacheBlock = _FakeKVCacheBlock # type: ignore[attr-defined]
_kv_cache_utils_mod.BlockHashList = list # type: ignore[attr-defined]
class _FakeBlockPool:
def __init__(self, *args, **kwargs):
self.null_block = _FakeKVCacheBlock(block_id=0)
self._next_block_id = 1
def get_new_blocks(self, num_blocks):
blocks = []
for _ in range(num_blocks):
blocks.append(_FakeKVCacheBlock(block_id=self._next_block_id))
self._next_block_id += 1
return blocks
if _MOCK_VLLM_DEPS:
sys.modules["vllm.v1.core.block_pool"].BlockPool = _FakeBlockPool # type: ignore[attr-defined]
class _FakeSingleTypeKVCacheManager:
def __init__(self, *args, **kwargs):
self._mock = MagicMock()
def __getattr__(self, name):
return getattr(self._mock, name)
@classmethod
def reachable_block_mask(
cls,
start_block,
end_block,
alignment_tokens,
kv_cache_spec,
use_eagle,
retention_interval=None,
num_prompt_tokens=None,
):
return None
@classmethod
def find_longest_cache_hit(
cls,
block_hashes,
max_length,
kv_cache_group_ids,
block_pool,
kv_cache_spec,
drop_eagle_block=False,
alignment_tokens=16,
dcp_world_size=1,
pcp_world_size=1,
):
computed: tuple[list[object], ...] = tuple([] for _ in kv_cache_group_ids)
max_blocks = max_length // kv_cache_spec.block_size
for block_hash in list(block_hashes)[:max_blocks]:
cached = block_pool.get_cached_block(block_hash, kv_cache_group_ids)
if not cached:
break
for blocks, block in zip(computed, cached):
blocks.append(block)
if drop_eagle_block and computed and computed[0]:
for blocks in computed:
blocks.pop()
return computed
class _FakeSlidingWindowManager(_FakeSingleTypeKVCacheManager):
@classmethod
def reachable_block_mask(
cls,
start_block,
end_block,
alignment_tokens,
kv_cache_spec,
use_eagle,
retention_interval=None,
num_prompt_tokens=None,
):
if alignment_tokens is None:
return None
per_segment = max(alignment_tokens // kv_cache_spec.block_size, 1)
return [(idx + 1) % per_segment == 0 for idx in range(start_block, end_block)]
_single_type_mod: Any = (
sys.modules["vllm.v1.core.single_type_kv_cache_manager"] if _MOCK_VLLM_DEPS else types.SimpleNamespace()
)
_single_type_mod.SingleTypeKVCacheManager = _FakeSingleTypeKVCacheManager # type: ignore[attr-defined]
_single_type_mod.FullAttentionManager = _FakeSingleTypeKVCacheManager # type: ignore[attr-defined]
_single_type_mod.SlidingWindowManager = _FakeSlidingWindowManager # type: ignore[attr-defined]
_single_type_mod.MambaManager = _FakeSingleTypeKVCacheManager # type: ignore[attr-defined]
_single_type_mod.spec_manager_map = { # type: ignore[attr-defined]
_FakeFullAttentionSpec: _FakeSingleTypeKVCacheManager,
_FakeSlidingWindowSpec: _FakeSlidingWindowManager,
_FakeMambaSpec: _FakeSingleTypeKVCacheManager,
}
_kv_interface_mod: Any = sys.modules["vllm.v1.kv_cache_interface"] if _MOCK_VLLM_DEPS else types.SimpleNamespace()
_kv_interface_mod.KVCacheSpec = _FakeKVCacheSpec # type: ignore[attr-defined]
_kv_interface_mod.FullAttentionSpec = _FakeFullAttentionSpec # type: ignore[attr-defined]
_kv_interface_mod.SlidingWindowSpec = _FakeSlidingWindowSpec # type: ignore[attr-defined]
_kv_interface_mod.MambaSpec = _FakeMambaSpec # type: ignore[attr-defined]
_kv_interface_mod.UniformTypeKVCacheSpecs = _FakeUniformTypeKVCacheSpecs # type: ignore[attr-defined]
_kv_interface_mod.KVCacheGroupSpec = _FakeKVCacheGroupSpec # type: ignore[attr-defined]
_kv_interface_mod.KVCacheConfig = _FakeKVCacheConfig # type: ignore[attr-defined]
class _FakeKVCacheSpecRegistry:
@classmethod
def get_manager_class(cls, kv_cache_spec):
if isinstance(kv_cache_spec, _FakeSlidingWindowSpec):
return _FakeSlidingWindowManager
return _FakeSingleTypeKVCacheManager
if _MOCK_VLLM_DEPS:
sys.modules["vllm.v1.kv_cache_spec_registry"].KVCacheSpecRegistry = _FakeKVCacheSpecRegistry # type: ignore[attr-defined]
_sched_output_mod: Any = sys.modules["vllm.v1.core.sched.output"] if _MOCK_VLLM_DEPS else types.SimpleNamespace()
_sched_output_mod.NewRequestData = MagicMock # type: ignore[attr-defined]
if _MOCK_VLLM_DEPS:
sys.modules["vllm.envs"].VLLM_RPC_BASE_PATH = "/tmp/vllm_rpc" # type: ignore[attr-defined]
# ---------------------------------------------------------------------------
# Mock external backends
# ---------------------------------------------------------------------------
for _mod_name in [
"mooncake",
"mooncake.engine",
"mooncake.store",
"memcache_hybrid",
"yr",
"yr.datasystem",
"yr.datasystem.hetero_client",
"yr.datasystem.kv_client",
"yr.datasystem.object_client",
"zmq",
]:
if _mod_name not in sys.modules:
sys.modules[_mod_name] = MagicMock()
# ---------------------------------------------------------------------------
# Mock vllm_ascend transitive imports
# ---------------------------------------------------------------------------
def _make_pkg(name, path=""):
mod = types.ModuleType(name)
mod.__path__ = [path] # type: ignore[attr-defined]
mod.__package__ = name # type: ignore[attr-defined]
return mod
_vllm_ascend_real_path = os.path.abspath(os.path.join(os.path.dirname(__file__), "..", "..", "..", "..", "vllm_ascend"))
_vllm_ascend_package_paths = {
"vllm_ascend": _vllm_ascend_real_path,
"vllm_ascend.distributed": os.path.join(_vllm_ascend_real_path, "distributed"),
}
for _pkg, _path in _vllm_ascend_package_paths.items():
if _pkg not in sys.modules:
sys.modules[_pkg] = _make_pkg(_pkg, _path)
_distributed_utils = types.ModuleType("vllm_ascend.distributed.utils")
_distributed_utils.get_decode_context_model_parallel_rank = MagicMock( # type: ignore[attr-defined]
return_value=0
)
_distributed_utils.get_decode_context_model_parallel_world_size = MagicMock( # type: ignore[attr-defined]
return_value=1
)
sys.modules["vllm_ascend.distributed.utils"] = _distributed_utils
_kv_transfer_init = _make_pkg("vllm_ascend.distributed.kv_transfer")
_kv_transfer_init.register_connector = MagicMock() # type: ignore[attr-defined]
sys.modules["vllm_ascend.distributed.kv_transfer"] = _kv_transfer_init
_kv_utils_pkg = _make_pkg("vllm_ascend.distributed.kv_transfer.utils")
sys.modules["vllm_ascend.distributed.kv_transfer.utils"] = _kv_utils_pkg
sys.modules["vllm_ascend.distributed.kv_transfer.utils.mooncake_transfer_engine"] = MagicMock()
_kv_pool_pkg = _make_pkg("vllm_ascend.distributed.kv_transfer.kv_pool")
sys.modules["vllm_ascend.distributed.kv_transfer.kv_pool"] = _kv_pool_pkg
_ascend_store_real_path = os.path.join(
os.path.dirname(__file__),
"..",
"..",
"..",
"..",
"vllm_ascend",
"distributed",
"kv_transfer",
"kv_pool",
"ascend_store",
)
_ascend_store_pkg = _make_pkg(
"vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store",
os.path.abspath(_ascend_store_real_path),
)
sys.modules["vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store"] = _ascend_store_pkg
_backend_pkg = _make_pkg(
"vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.backend",
os.path.join(os.path.abspath(_ascend_store_real_path), "backend"),
)
sys.modules["vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.backend"] = _backend_pkg
# Mirror the real backend/__init__.py entry points. The scheduler/worker resolve
# the backend class dynamically via ``importlib.import_module(path)``; tests that
# exercise those paths patch ``<module>.importlib`` locally (see
# test_pool_scheduler.py / test_pool_worker.py) so the backend resolves to a
# MagicMock. Do NOT register the backends in sys.modules or globally wrap
# import_module here: test_backend.py imports the real backend classes and also
# relies on ``mock.patch`` (which itself calls importlib.import_module) resolving
# those real modules.
_backend_module_paths = {
"mooncake": "vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.backend.mooncake_backend",
"memcache": "vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.backend.memcache_backend",
"yuanrong": "vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.backend.yuanrong_backend",
}
_backend_pkg.backend_map = { # type: ignore[attr-defined]
"mooncake": {"name": "MooncakeBackend", "path": _backend_module_paths["mooncake"]},
"memcache": {"name": "MemcacheBackend", "path": _backend_module_paths["memcache"]},
"yuanrong": {"name": "YuanrongBackend", "path": _backend_module_paths["yuanrong"]},
}
if "vllm_ascend.utils" not in sys.modules or not hasattr(sys.modules["vllm_ascend.utils"], "AscendDeviceType"):
_ascend_utils = MagicMock()
_ascend_utils.AscendDeviceType = MagicMock()
_ascend_utils.get_ascend_device_type = MagicMock()
sys.modules["vllm_ascend.utils"] = _ascend_utils
# NOTE: vllm_ascend.{ascend_config, memcache_comm_fence} and their helpers
# (get_ascend_config, AttentionComputeStartGate, ...) are intentionally NOT
# mocked here. Doing so by mutating these real modules leaks into every other
# UT in the same pytest session (breaking test_ascend_config / test_platform,
# which collect after ascend_store and bind the polluted symbols at import).
# These helpers are mocked per-test, scoped to the ascend_store tests only,
# via the autouse fixture in tests/ut/conftest.py.

View File

@@ -0,0 +1,459 @@
#
# Copyright (c) 2026 Huawei Technologies Co., Ltd. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
# This file is a part of the vllm-ascend project.
#
import types
import unittest
from unittest.mock import MagicMock, patch
# isort: off
import tests.ut.distributed.ascend_store._mock_deps # noqa: F401, E402
from vllm.distributed.kv_events import KVCacheEvent
from vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.ascend_store_connector import (
AscendStoreConnector,
AscendStoreKVEvents,
)
# isort: on
def _mock_events(num_workers=1):
events = AscendStoreKVEvents(num_workers=num_workers)
events._aggregator = MagicMock()
return events
class TestAscendStoreKVEvents(unittest.TestCase):
def _make_events(self, num_workers=1):
return _mock_events(num_workers=num_workers)
def test_add_and_get_events(self):
ev = self._make_events()
mock_events = [MagicMock(spec=KVCacheEvent), MagicMock(spec=KVCacheEvent)]
ev.add_events(mock_events)
ev._aggregator.get_all_events.return_value = mock_events
result = ev.get_all_events()
self.assertEqual(result, mock_events)
def test_aggregate(self):
ev = self._make_events()
common = [MagicMock()]
ev._aggregator.get_common_events.return_value = common
result = ev.aggregate()
self.assertIs(result, ev)
ev._aggregator.clear_events.assert_called()
ev._aggregator.add_events.assert_called_with(common)
ev._aggregator.reset_workers.assert_called()
def test_increment_workers(self):
ev = self._make_events()
ev.increment_workers(3)
ev._aggregator.increment_workers.assert_called_with(3)
def test_get_number_of_workers(self):
ev = self._make_events()
ev._aggregator.get_number_of_workers.return_value = 5
self.assertEqual(ev.get_number_of_workers(), 5)
def test_clear_events(self):
ev = self._make_events()
ev.clear_events()
ev._aggregator.clear_events.assert_called()
ev._aggregator.reset_workers.assert_called()
def test_repr(self):
ev = self._make_events()
ev._aggregator.get_all_events.return_value = []
s = repr(ev)
self.assertIn("AscendStoreKVEvents", s)
class TestAscendStoreConnector(unittest.TestCase):
def _make_vllm_config(self, kv_role="kv_producer", extra_config=None):
config = MagicMock()
config.kv_transfer_config.kv_role = kv_role
config.kv_transfer_config.kv_connector = "AscendStoreConnector"
config.kv_transfer_config.kv_connector_extra_config = extra_config or {}
config.parallel_config.rank = 0
return config
def test_pp_handshake_metadata_is_ignored(self):
connector = AscendStoreConnector.__new__(AscendStoreConnector)
metadata = {
(0, 0): MagicMock(),
(1, 0): MagicMock(),
}
original_metadata = metadata.copy()
result = connector.set_xfer_handshake_metadata_pp_aware(metadata)
self.assertIsNone(result)
self.assertEqual(metadata, original_metadata)
@patch("vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.ascend_store_connector.KVPoolScheduler")
def test_init_scheduler_role(self, mock_scheduler_cls):
config = self._make_vllm_config()
from vllm.distributed.kv_transfer.kv_connector.v1.base import KVConnectorRole
_connector = AscendStoreConnector(
vllm_config=config,
role=KVConnectorRole.SCHEDULER,
kv_cache_config=MagicMock(),
)
mock_scheduler_cls.assert_called_once()
@patch("vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.ascend_store_connector.LookupKeyServer")
@patch("vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.ascend_store_connector.KVPoolWorker")
def test_init_worker_role(self, mock_worker_cls, mock_lookup_cls):
config = self._make_vllm_config()
from vllm.distributed.kv_transfer.kv_connector.v1.base import KVConnectorRole
_connector = AscendStoreConnector(
vllm_config=config,
role=KVConnectorRole.WORKER,
kv_cache_config=None,
)
mock_worker_cls.assert_called_once()
mock_lookup_cls.assert_called_once()
@patch("vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.ascend_store_connector.KVPoolScheduler")
def test_scheduler_methods_delegate(self, mock_scheduler_cls):
config = self._make_vllm_config()
from vllm.distributed.kv_transfer.kv_connector.v1.base import KVConnectorRole
connector = AscendStoreConnector(
vllm_config=config,
role=KVConnectorRole.SCHEDULER,
kv_cache_config=MagicMock(),
)
mock_sched = mock_scheduler_cls.return_value
# get_num_new_matched_tokens
mock_sched.get_num_new_matched_tokens.return_value = (10, False)
result = connector.get_num_new_matched_tokens(MagicMock(), 5)
self.assertEqual(result, (10, False))
# update_state_after_alloc
connector.update_state_after_alloc(MagicMock(), MagicMock(), 10)
mock_sched.update_state_after_alloc.assert_called_once()
# build_connector_meta
connector.build_connector_meta(MagicMock())
mock_sched.build_connector_meta.assert_called_once()
# request_finished
mock_sched.request_finished.return_value = (True, None)
result = connector.request_finished(MagicMock(), [1, 2])
self.assertEqual(result, (True, None))
@patch("vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.ascend_store_connector.KVPoolScheduler")
def test_update_connector_output_no_events(self, mock_scheduler_cls):
config = self._make_vllm_config()
from vllm.distributed.kv_transfer.kv_connector.v1.base import KVConnectorRole
connector = AscendStoreConnector(
vllm_config=config,
role=KVConnectorRole.SCHEDULER,
kv_cache_config=MagicMock(),
)
output = MagicMock()
output.kv_cache_events = None
connector.update_connector_output(output)
self.assertIsNone(connector._kv_cache_events)
@patch("vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.ascend_store_connector.KVPoolScheduler")
def test_update_connector_output_with_events(self, mock_scheduler_cls):
config = self._make_vllm_config()
from vllm.distributed.kv_transfer.kv_connector.v1.base import KVConnectorRole
connector = AscendStoreConnector(
vllm_config=config,
role=KVConnectorRole.SCHEDULER,
kv_cache_config=MagicMock(),
)
events = _mock_events(num_workers=1)
mock_kv_events = [MagicMock()]
events._aggregator.get_all_events.return_value = mock_kv_events
events._aggregator.get_number_of_workers.return_value = 1
output = MagicMock()
output.kv_cache_events = events
connector.update_connector_output(output)
self.assertIsNotNone(connector._kv_cache_events)
@patch("vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.ascend_store_connector.KVPoolScheduler")
def test_update_connector_output_accumulate(self, mock_scheduler_cls):
config = self._make_vllm_config()
from vllm.distributed.kv_transfer.kv_connector.v1.base import KVConnectorRole
connector = AscendStoreConnector(
vllm_config=config,
role=KVConnectorRole.SCHEDULER,
kv_cache_config=MagicMock(),
)
# First update
events1 = _mock_events(num_workers=1)
events1._aggregator.get_all_events.return_value = [MagicMock()]
events1._aggregator.get_number_of_workers.return_value = 1
output1 = MagicMock()
output1.kv_cache_events = events1
connector.update_connector_output(output1)
# Second update
events2 = _mock_events(num_workers=1)
events2._aggregator.get_all_events.return_value = [MagicMock()]
events2._aggregator.get_number_of_workers.return_value = 1
output2 = MagicMock()
output2.kv_cache_events = events2
connector.update_connector_output(output2)
self.assertIsNotNone(connector._kv_cache_events)
@patch("vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.ascend_store_connector.KVPoolScheduler")
def test_take_events(self, mock_scheduler_cls):
config = self._make_vllm_config()
from vllm.distributed.kv_transfer.kv_connector.v1.base import KVConnectorRole
connector = AscendStoreConnector(
vllm_config=config,
role=KVConnectorRole.SCHEDULER,
kv_cache_config=MagicMock(),
)
# No events
result = list(connector.take_events())
self.assertEqual(result, [])
# With events
events = _mock_events(num_workers=1)
mock_event = MagicMock()
events._aggregator.get_common_events.return_value = [mock_event]
events._aggregator.get_all_events.return_value = [mock_event]
connector._kv_cache_events = events
result = list(connector.take_events())
self.assertEqual(len(result), 1)
self.assertIsNone(connector._kv_cache_events)
@patch("vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.ascend_store_connector.LookupKeyServer")
@patch("vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.ascend_store_connector.KVPoolWorker")
def test_worker_methods(self, mock_worker_cls, mock_lookup_cls):
config = self._make_vllm_config()
from vllm.distributed.kv_transfer.kv_connector.v1.base import KVConnectorRole
connector = AscendStoreConnector(
vllm_config=config,
role=KVConnectorRole.WORKER,
kv_cache_config=None,
)
mock_worker = mock_worker_cls.return_value
# register_kv_caches
connector.register_kv_caches({"layer1": MagicMock()})
mock_worker.register_kv_caches.assert_called_once()
# start_load_kv
connector._get_connector_metadata = MagicMock(return_value=MagicMock())
connector.start_load_kv(MagicMock())
mock_worker.start_load_kv.assert_called_once()
# wait_for_save (non-consumer)
connector.kv_role = "kv_producer"
connector.use_layerwise = False
connector.wait_for_save()
mock_worker.wait_for_save.assert_called_once()
# get_finished
mock_worker.get_finished.return_value = ({"r1"}, {"r2"})
done_s, done_r = connector.get_finished({"r1"})
self.assertEqual(done_s, {"r1"})
@patch("vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.ascend_store_connector.LookupKeyServer")
@patch("vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.ascend_store_connector.KVPoolWorker")
def test_wait_for_layer_load_not_layerwise(self, mock_worker_cls, mock_lookup_cls):
config = self._make_vllm_config(extra_config={"use_layerwise": False})
from vllm.distributed.kv_transfer.kv_connector.v1.base import KVConnectorRole
connector = AscendStoreConnector(
vllm_config=config,
role=KVConnectorRole.WORKER,
kv_cache_config=None,
)
# Should return immediately without calling worker
connector.wait_for_layer_load("layer_0")
@patch("vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.ascend_store_connector.LookupKeyServer")
@patch("vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.ascend_store_connector.KVPoolWorker")
def test_save_kv_layer_not_layerwise(self, mock_worker_cls, mock_lookup_cls):
config = self._make_vllm_config(extra_config={"use_layerwise": False})
from vllm.distributed.kv_transfer.kv_connector.v1.base import KVConnectorRole
connector = AscendStoreConnector(
vllm_config=config,
role=KVConnectorRole.WORKER,
kv_cache_config=None,
)
connector.save_kv_layer("layer_0", MagicMock(), MagicMock())
# Should return immediately
@patch("vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.ascend_store_connector.LookupKeyServer")
@patch("vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.ascend_store_connector.KVPoolWorker")
def test_save_kv_layer_consumer(self, mock_worker_cls, mock_lookup_cls):
config = self._make_vllm_config(kv_role="kv_consumer", extra_config={"use_layerwise": True})
from vllm.distributed.kv_transfer.kv_connector.v1.base import KVConnectorRole
connector = AscendStoreConnector(
vllm_config=config,
role=KVConnectorRole.WORKER,
kv_cache_config=None,
)
connector.save_kv_layer("layer_0", MagicMock(), MagicMock())
# Consumer should not save
@patch("vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.ascend_store_connector.LookupKeyServer")
@patch("vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.ascend_store_connector.KVPoolWorker")
def test_wait_for_save_consumer(self, mock_worker_cls, mock_lookup_cls):
config = self._make_vllm_config(kv_role="kv_consumer")
from vllm.distributed.kv_transfer.kv_connector.v1.base import KVConnectorRole
connector = AscendStoreConnector(
vllm_config=config,
role=KVConnectorRole.WORKER,
kv_cache_config=None,
)
connector.wait_for_save()
mock_worker_cls.return_value.wait_for_save.assert_not_called()
@patch("vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.ascend_store_connector.LookupKeyServer")
@patch("vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.ascend_store_connector.KVPoolWorker")
def test_get_kv_connector_kv_cache_events_empty(self, mock_worker_cls, mock_lookup_cls):
config = self._make_vllm_config()
from vllm.distributed.kv_transfer.kv_connector.v1.base import KVConnectorRole
connector = AscendStoreConnector(
vllm_config=config,
role=KVConnectorRole.WORKER,
kv_cache_config=None,
)
mock_worker_cls.return_value.get_kv_events.return_value = []
result = connector.get_kv_connector_kv_cache_events()
self.assertIsNone(result)
@patch("vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.ascend_store_connector.LookupKeyServer")
@patch("vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.ascend_store_connector.KVPoolWorker")
def test_get_kv_connector_kv_cache_events_with_events(self, mock_worker_cls, mock_lookup_cls):
config = self._make_vllm_config()
from vllm.distributed.kv_transfer.kv_connector.v1.base import KVConnectorRole
connector = AscendStoreConnector(
vllm_config=config,
role=KVConnectorRole.WORKER,
kv_cache_config=None,
)
mock_worker_cls.return_value.get_kv_events.return_value = [MagicMock()]
result = connector.get_kv_connector_kv_cache_events()
self.assertIsNotNone(result)
self.assertIsInstance(result, AscendStoreKVEvents)
class TestAscendStoreConnectorLayerwise(unittest.TestCase):
"""Test connector methods that are specific to layerwise mode."""
connector_mod: types.ModuleType
@classmethod
def setUpClass(cls):
from vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store import ascend_store_connector
cls.connector_mod = ascend_store_connector
def test_requires_piecewise_for_cudagraph_enabled(self):
self.assertTrue(
self.connector_mod.AscendStoreConnector.requires_piecewise_for_cudagraph({"use_layerwise": True})
)
def test_requires_piecewise_for_cudagraph_disabled(self):
self.assertFalse(
self.connector_mod.AscendStoreConnector.requires_piecewise_for_cudagraph({"use_layerwise": False})
)
def test_requires_piecewise_for_cudagraph_missing(self):
self.assertFalse(self.connector_mod.AscendStoreConnector.requires_piecewise_for_cudagraph({}))
def test_wait_for_save_layerwise_returns_early(self):
from vllm.distributed.kv_transfer.kv_connector.v1.base import KVConnectorRole
with (
patch.object(self.connector_mod, "KVPoolWorker") as mock_worker_cls,
patch.object(self.connector_mod, "LookupKeyServer") as _mock_lookup_cls,
):
config = MagicMock()
config.kv_transfer_config.kv_role = "kv_producer"
config.kv_transfer_config.kv_connector = "AscendStoreConnector"
config.kv_transfer_config.kv_connector_extra_config = {"use_layerwise": True}
config.parallel_config.rank = 0
connector = self.connector_mod.AscendStoreConnector(
vllm_config=config,
role=KVConnectorRole.WORKER,
kv_cache_config=None,
)
connector.wait_for_save()
mock_worker_cls.return_value.wait_for_save.assert_not_called()
def test_save_kv_layer_layerwise_producer(self):
from vllm.distributed.kv_transfer.kv_connector.v1.base import KVConnectorRole
with (
patch.object(self.connector_mod, "KVPoolWorker") as mock_worker_cls,
patch.object(self.connector_mod, "LookupKeyServer") as _mock_lookup_cls,
):
config = MagicMock()
config.kv_transfer_config.kv_role = "kv_producer"
config.kv_transfer_config.kv_connector = "AscendStoreConnector"
config.kv_transfer_config.kv_connector_extra_config = {"use_layerwise": True}
config.parallel_config.rank = 0
connector = self.connector_mod.AscendStoreConnector(
vllm_config=config,
role=KVConnectorRole.WORKER,
kv_cache_config=None,
)
connector._get_connector_metadata = MagicMock(return_value=MagicMock())
connector.save_kv_layer("layer_0", MagicMock(), MagicMock())
mock_worker_cls.return_value.save_kv_layer.assert_called_once()
def test_wait_for_layer_load_layerwise(self):
from vllm.distributed.kv_transfer.kv_connector.v1.base import KVConnectorRole
with (
patch.object(self.connector_mod, "KVPoolWorker") as mock_worker_cls,
patch.object(self.connector_mod, "LookupKeyServer") as _mock_lookup_cls,
):
config = MagicMock()
config.kv_transfer_config.kv_role = "kv_consumer"
config.kv_transfer_config.kv_connector = "AscendStoreConnector"
config.kv_transfer_config.kv_connector_extra_config = {"use_layerwise": True}
config.parallel_config.rank = 0
connector = self.connector_mod.AscendStoreConnector(
vllm_config=config,
role=KVConnectorRole.WORKER,
kv_cache_config=None,
)
connector.wait_for_layer_load("layer_0")
mock_worker_cls.return_value.wait_for_layer_load.assert_called_once()
if __name__ == "__main__":
unittest.main()

View File

@@ -0,0 +1,665 @@
#
# Copyright (c) 2026 Huawei Technologies Co., Ltd. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
# This file is a part of the vllm-ascend project.
#
import json
import os
import tempfile
import unittest
from unittest.mock import MagicMock, patch
import tests.ut.distributed.ascend_store._mock_deps # noqa: F401, E402
from vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.backend.backend import Backend
from vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.backend.mooncake_backend import (
MooncakeStoreConfig,
_convert_to_bytes,
_parse_global_segment_size,
_ssd_setup_kwargs,
)
from vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.backend.yuanrong_backend import (
YuanrongConfig,
YuanrongHelper,
)
def _format_log_call(call):
args = call.args
return args[0] % args[1:]
# =========================================================================
# Backend ABC
# =========================================================================
class TestBackendABC(unittest.TestCase):
def test_cannot_instantiate(self):
with self.assertRaises(TypeError):
Backend(MagicMock()) # type: ignore[abstract]
def _make_mooncake_store_config(**overrides) -> MooncakeStoreConfig:
"""Build MooncakeStoreConfig via from_file(); inherits from_file() defaults."""
config = dict(overrides)
with tempfile.NamedTemporaryFile(mode="w", suffix=".json", delete=False) as f:
json.dump(config, f)
f.flush()
path = f.name
try:
return MooncakeStoreConfig.from_file(path)
finally:
os.unlink(path)
# =========================================================================
# MooncakeStoreConfig
# =========================================================================
class TestMooncakeStoreConfig(unittest.TestCase):
def test_from_file(self):
config = {
"metadata_server": "127.0.0.1:2379",
"global_segment_size": "2GB",
"local_buffer_size": "1GB",
"protocol": "ascend",
"device_name": "npu0",
"master_server_address": "127.0.0.1:8080",
}
with tempfile.NamedTemporaryFile(mode="w", suffix=".json", delete=False) as f:
json.dump(config, f)
f.flush()
path = f.name
try:
cfg = MooncakeStoreConfig.from_file(path)
self.assertEqual(cfg.metadata_server, "127.0.0.1:2379")
self.assertEqual(cfg.global_segment_size, 2 * 1024**3)
self.assertEqual(cfg.local_buffer_size, 1 * 1024**3)
self.assertEqual(cfg.protocol, "ascend")
self.assertEqual(cfg.device_name, "npu0")
finally:
os.unlink(path)
def test_from_file_defaults(self):
config = {
"metadata_server": "localhost:2379",
"master_server_address": "localhost:8080",
}
with tempfile.NamedTemporaryFile(mode="w", suffix=".json", delete=False) as f:
json.dump(config, f)
f.flush()
path = f.name
try:
cfg = MooncakeStoreConfig.from_file(path)
self.assertEqual(cfg.protocol, "ascend")
self.assertEqual(cfg.device_name, "")
self.assertFalse(cfg.enable_ssd_offload)
self.assertEqual(cfg.ssd_offload_path, "")
finally:
os.unlink(path)
def test_from_file_ssd_offload(self):
ssd_path = TestMooncakeStoreConfig._writable_ssd_path()
self.addCleanup(lambda: os.rmdir(ssd_path))
cfg = _make_mooncake_store_config(
enable_ssd_offload=True,
ssd_offload_path=ssd_path,
)
self.assertTrue(cfg.enable_ssd_offload)
self.assertEqual(cfg.ssd_offload_path, ssd_path)
def test_ssd_offload_requires_absolute_path(self):
with self.assertRaises(ValueError):
_make_mooncake_store_config(
enable_ssd_offload=True,
ssd_offload_path="relative/path",
)
def test_ssd_offload_requires_path_in_json(self):
with self.assertRaises(ValueError):
_make_mooncake_store_config(enable_ssd_offload=True)
@staticmethod
def _writable_ssd_path() -> str:
return tempfile.mkdtemp(prefix="mooncake_ssd_ut_")
@patch(
"vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.backend."
"mooncake_backend._mooncake_setup_supports_ssd_offload",
return_value=False,
)
def test_ssd_setup_kwargs_off_when_disabled(self, _mock_supports):
cfg = _make_mooncake_store_config()
self.assertEqual(_ssd_setup_kwargs(cfg), {})
@patch(
"vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.backend."
"mooncake_backend._mooncake_setup_supports_ssd_offload",
return_value=False,
)
def test_ssd_setup_kwargs_raises_on_old_mooncake(self, _mock_supports):
ssd_path = TestMooncakeStoreConfig._writable_ssd_path()
self.addCleanup(lambda: os.rmdir(ssd_path))
cfg = _make_mooncake_store_config(
enable_ssd_offload=True,
ssd_offload_path=ssd_path,
)
with self.assertRaises(RuntimeError):
_ssd_setup_kwargs(cfg)
@patch(
"vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.backend."
"mooncake_backend._mooncake_setup_supports_ssd_offload",
return_value=True,
)
def test_ssd_setup_kwargs_when_supported(self, _mock_supports):
ssd_path = TestMooncakeStoreConfig._writable_ssd_path()
self.addCleanup(lambda: os.rmdir(ssd_path))
cfg = _make_mooncake_store_config(
enable_ssd_offload=True,
ssd_offload_path=ssd_path,
)
self.assertEqual(
_ssd_setup_kwargs(cfg),
{
"enable_ssd_offload": cfg.enable_ssd_offload,
"ssd_offload_path": cfg.ssd_offload_path,
},
)
def test_load_from_env_missing(self):
with patch.dict(os.environ, {}, clear=True):
os.environ.pop("MOONCAKE_CONFIG_PATH", None)
with self.assertRaises(ValueError):
MooncakeStoreConfig.load_from_env()
def test_load_from_env(self):
config = {
"metadata_server": "host:1234",
"master_server_address": "host:5678",
}
with tempfile.NamedTemporaryFile(mode="w", suffix=".json", delete=False) as f:
json.dump(config, f)
f.flush()
path = f.name
try:
with patch.dict(os.environ, {"MOONCAKE_CONFIG_PATH": path}):
cfg = MooncakeStoreConfig.load_from_env()
self.assertEqual(cfg.metadata_server, "host:1234")
finally:
os.unlink(path)
class TestParseGlobalSegmentSize(unittest.TestCase):
def test_int(self):
self.assertEqual(_parse_global_segment_size(1024), 1024)
def test_gb(self):
self.assertEqual(_parse_global_segment_size("2GB"), 2 * 1024**3)
def test_mb(self):
self.assertEqual(_parse_global_segment_size("512MB"), 512 * 1024**2)
def test_kb(self):
self.assertEqual(_parse_global_segment_size("256KB"), 256 * 1024)
def test_b(self):
self.assertEqual(_parse_global_segment_size("4096B"), 4096)
def test_no_unit(self):
self.assertEqual(_parse_global_segment_size("2048"), 2048)
def test_float_input(self):
self.assertEqual(_parse_global_segment_size(2048.0), 2048)
def test_empty_string(self):
with self.assertRaises(ValueError):
_parse_global_segment_size("")
def test_invalid_format(self):
with self.assertRaises(ValueError):
_parse_global_segment_size("abcGB")
def test_unsupported_type(self):
with self.assertRaises(TypeError):
_parse_global_segment_size(None) # type: ignore[arg-type]
class TestConvertToBytes(unittest.TestCase):
def test_valid(self):
self.assertEqual(_convert_to_bytes("10", 1, "10"), 10)
self.assertEqual(_convert_to_bytes("1.5", 1024, "1.5KB"), int(1.5 * 1024))
def test_invalid_number(self):
with self.assertRaises(ValueError):
_convert_to_bytes("abc", 1, "abc")
# =========================================================================
# YuanrongConfig
# =========================================================================
class TestYuanrongConfig(unittest.TestCase):
def test_load_from_env(self):
with patch.dict(
os.environ,
{
"DS_WORKER_ADDR": "host:1234",
"DS_ENABLE_EXCLUSIVE_CONNECTION": "1",
"DS_ENABLE_REMOTE_H2D": "0",
},
):
cfg = YuanrongConfig.load_from_env()
self.assertEqual(cfg.worker_addr, "host:1234")
self.assertTrue(cfg.enable_exclusive_connection)
self.assertFalse(cfg.enable_remote_h2d)
def test_load_from_env_missing(self):
with patch.dict(os.environ, {}, clear=True):
os.environ.pop("DS_WORKER_ADDR", None)
with self.assertRaises(ValueError):
YuanrongConfig.load_from_env()
def test_load_from_env_defaults(self):
with patch.dict(os.environ, {"DS_WORKER_ADDR": "h:1"}):
cfg = YuanrongConfig.load_from_env()
self.assertFalse(cfg.enable_exclusive_connection)
self.assertFalse(cfg.enable_remote_h2d)
# =========================================================================
# YuanrongHelper
# =========================================================================
class TestYuanrongHelper(unittest.TestCase):
def setUp(self):
self.blob_cls = MagicMock()
self.blob_list_cls = MagicMock()
self.helper = YuanrongHelper(self.blob_cls, self.blob_list_cls)
def test_normalize_keys_short_valid(self):
keys = ["abc-123", "key_2"]
result = self.helper.normalize_keys(keys)
self.assertEqual(result, keys)
def test_normalize_keys_with_invalid_chars(self):
keys = ["key with spaces/and.dots"]
result = self.helper.normalize_keys(keys)
self.assertEqual(len(result), 1)
# Should not contain the original invalid chars
self.assertNotIn(" ", result[0])
self.assertNotIn("/", result[0])
# Should have hash suffix
self.assertIn("__", result[0])
def test_normalize_keys_at_max_length(self):
max_length_key = "a" * 1024
result = self.helper.normalize_keys([max_length_key])
self.assertEqual(result, [max_length_key])
def test_normalize_keys_over_max_length(self):
long_key = "a" * 1025
result = self.helper.normalize_keys([long_key])
self.assertEqual(len(result), 1)
self.assertEqual(len(result[0]), 1024)
self.assertIn("__", result[0])
def test_make_blob_lists(self):
self.helper._device_id = 0
addrs = [[100, 200], [300, 400]]
sizes = [[10, 20], [30, 40]]
result = self.helper.make_blob_lists(addrs, sizes)
self.assertEqual(len(result), 2)
self.assertEqual(self.blob_cls.call_count, 4)
def test_make_blob_lists_length_mismatch(self):
self.helper._device_id = 0
with self.assertRaises(ValueError):
self.helper.make_blob_lists([[1]], [[1, 2], [3, 4]])
def test_make_blob_lists_inner_length_mismatch(self):
self.helper._device_id = 0
with self.assertRaises(ValueError):
self.helper.make_blob_lists([[1, 2]], [[1]])
def test_make_blob_lists_no_device(self):
self.helper._device_id = None
with self.assertRaises(RuntimeError):
self.helper.make_blob_lists([[1]], [[1]])
# =========================================================================
# MooncakeBackend (mocked store)
# =========================================================================
class TestMooncakeBackendMethods(unittest.TestCase):
def _make_backend(self):
from vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.backend.mooncake_backend import MooncakeBackend
with (
patch.dict(os.environ, {"MOONCAKE_CONFIG_PATH": "/dev/null"}),
patch.object(MooncakeBackend, "__init__", lambda self, pc: None),
):
backend = MooncakeBackend.__new__(MooncakeBackend)
backend.store = MagicMock()
backend.config = MagicMock()
backend.local_seg = "127.0.0.1:1234"
backend._lazy_init = False
backend._store_initialized = True
backend._use_fabric_mem = False
backend._store_init_lock = MagicMock()
backend.local_seg = None
return backend
def test_exists(self):
b = self._make_backend()
b.store.batch_is_exist.return_value = [1, 0]
result = b.exists(["k1", "k2"])
self.assertEqual(result, [1, 0])
def test_put(self):
b = self._make_backend()
b.store.batch_put_from_multi_buffers.return_value = [0, 0]
b.put(["k1"], [[100]], [[10]])
b.store.batch_put_from_multi_buffers.assert_called_once()
def test_put_error(self):
b = self._make_backend()
b.store.batch_put_from_multi_buffers.return_value = [-1]
b.put(["k1"], [[100]], [[10]]) # Should log error but not raise
def test_put_exception(self):
b = self._make_backend()
b.store.batch_put_from_multi_buffers.side_effect = RuntimeError("backend fail")
with patch(
"vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.backend.mooncake_backend.logger"
) as mock_logger:
b.put(["k1"], [[100]], [[10]]) # Should log error but not raise
error_log = _format_log_call(mock_logger.error.call_args)
self.assertIn("RuntimeError", error_log)
self.assertIn("backend fail", error_log)
def test_get(self):
b = self._make_backend()
b.store.batch_get_into_multi_buffers.return_value = [0]
b.get(["k1"], [[100]], [[10]])
b.store.batch_get_into_multi_buffers.assert_called_once()
def test_get_error(self):
b = self._make_backend()
b.store.batch_get_into_multi_buffers.return_value = [-1]
b.get(["k1"], [[100]], [[10]])
def test_get_exception(self):
b = self._make_backend()
b.store.batch_get_into_multi_buffers.side_effect = RuntimeError("backend fail")
with patch(
"vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.backend.mooncake_backend.logger"
) as mock_logger:
b.get(["k1"], [[100]], [[10]])
error_log = _format_log_call(mock_logger.error.call_args)
self.assertIn("RuntimeError", error_log)
self.assertIn("backend fail", error_log)
def test_register_buffer(self):
b = self._make_backend()
with (
patch(
"vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.backend.mooncake_backend.global_te"
) as mock_te,
patch("vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.backend.mooncake_backend.get_ip"),
):
b.register_buffer([100], [200])
mock_te.register_buffer.assert_called_once()
# =========================================================================
# YuanrongBackend (mocked store)
# =========================================================================
class TestYuanrongBackendMethods(unittest.TestCase):
def _make_backend(self):
from vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.backend.yuanrong_backend import YuanrongBackend
with patch.object(YuanrongBackend, "__init__", lambda self, pc: None):
backend = YuanrongBackend.__new__(YuanrongBackend)
backend._helper = MagicMock()
backend._helper._device_id = 0
backend._helper.normalize_keys = lambda keys: keys
backend._helper.make_blob_lists = lambda a, s: [MagicMock() for _ in a]
backend._hetero_client = MagicMock()
backend._ds_set_param = MagicMock()
backend._is_a2 = False
backend._registered_buffers = None
backend._buffers_registered = False
backend.config = YuanrongConfig(
worker_addr="127.0.0.1:0",
enable_exclusive_connection=False,
enable_remote_h2d=False,
)
backend.rank = 0
return backend
def test_exists_empty(self):
b = self._make_backend()
result = b.exists([])
self.assertEqual(result, [])
def test_exists(self):
b = self._make_backend()
b._hetero_client.exist.return_value = [True, False]
result = b.exists(["k1", "k2"])
self.assertEqual(result, [1, 0])
def test_exists_exception(self):
b = self._make_backend()
b._hetero_client.exist.side_effect = Exception("fail")
result = b.exists(["k1"])
self.assertEqual(result, [0])
def test_get_empty(self):
b = self._make_backend()
result = b.get([], [], [])
self.assertEqual(result, [])
b._hetero_client.mget_h2d.assert_not_called()
def test_get(self):
b = self._make_backend()
b._hetero_client.mget_h2d.return_value = []
result = b.get(["k1"], [[100]], [[10]])
self.assertEqual(result, [0])
b._hetero_client.mget_h2d.assert_called_once()
def test_get_partial_failure(self):
b = self._make_backend()
b._hetero_client.mget_h2d.return_value = ["k2"]
result = b.get(["k1", "k2", "k3"], [[100], [200], [300]], [[10], [20], [30]])
self.assertEqual(result, [0, 1, 0])
def test_get_failed_keys(self):
b = self._make_backend()
b._hetero_client.mget_h2d.return_value = ["k1"]
result = b.get(["k1"], [[100]], [[10]]) # Should log error
self.assertEqual(result, [1])
def test_get_exception(self):
b = self._make_backend()
b._hetero_client.mget_h2d.side_effect = RuntimeError("backend fail")
with patch(
"vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.backend.yuanrong_backend.logger"
) as mock_logger:
result = b.get(["k1"], [[100]], [[10]])
error_log = _format_log_call(mock_logger.error.call_args)
self.assertIsNone(result)
self.assertIn("RuntimeError", error_log)
self.assertIn("backend fail", error_log)
def test_put_empty(self):
b = self._make_backend()
b.put([], [], [])
b._hetero_client.mset_d2h.assert_not_called()
def test_put(self):
b = self._make_backend()
b.put(["k1"], [[100]], [[10]])
b._hetero_client.mset_d2h.assert_called_once()
def test_put_exception(self):
b = self._make_backend()
b._hetero_client.mset_d2h.side_effect = RuntimeError("backend fail")
with patch(
"vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.backend.yuanrong_backend.logger"
) as mock_logger:
b.put(["k1"], [[100]], [[10]])
error_log = _format_log_call(mock_logger.error.call_args)
self.assertIn("RuntimeError", error_log)
self.assertIn("backend fail", error_log)
def test_register_buffer_noop_when_remote_h2d_disabled(self):
b = self._make_backend()
b.register_buffer([100], [200])
b._hetero_client.pre_register_device_memory.assert_not_called()
def test_register_buffer_when_remote_h2d_enabled(self):
b = self._make_backend()
b.config.enable_remote_h2d = True
b.register_buffer([100], [200])
b._hetero_client.pre_register_device_memory.assert_called_once_with([100], [200])
def test_register_buffer_noop_on_a2(self):
# A2 must not register (opposite of memcache_backend's _is_a2 gating).
b = self._make_backend()
b._is_a2 = True
b.config.enable_remote_h2d = True
b.register_buffer([100], [200])
b._hetero_client.pre_register_device_memory.assert_not_called()
def test_register_buffer_idempotent(self):
b = self._make_backend()
b.config.enable_remote_h2d = True
b.register_buffer([100], [200])
b.register_buffer([300], [400])
b._hetero_client.pre_register_device_memory.assert_called_once_with([100], [200])
def test_register_buffers_if_needed_no_buffers(self):
b = self._make_backend()
b.config.enable_remote_h2d = True
b._registered_buffers = None
b._register_buffers_if_needed()
b._hetero_client.pre_register_device_memory.assert_not_called()
def test_register_buffers_if_needed_already_registered(self):
b = self._make_backend()
b.config.enable_remote_h2d = True
b._registered_buffers = ([100], [200])
b._buffers_registered = True
b._register_buffers_if_needed()
b._hetero_client.pre_register_device_memory.assert_not_called()
def test_register_buffers_if_needed_disabled(self):
b = self._make_backend()
b.config.enable_remote_h2d = False
b._registered_buffers = ([100], [200])
b._register_buffers_if_needed()
b._hetero_client.pre_register_device_memory.assert_not_called()
def test_ensure_device_ready(self):
b = self._make_backend()
b._helper._device_id = None
b.set_device = MagicMock()
b._ensure_device_ready()
b.set_device.assert_called_once()
def test_ensure_device_ready_already_set(self):
b = self._make_backend()
b._helper._device_id = 0
b.set_device = MagicMock()
b._ensure_device_ready()
b.set_device.assert_not_called()
# =========================================================================
# MemcacheBackend (mocked store)
# =========================================================================
class TestMemcacheBackendMethods(unittest.TestCase):
def _make_backend(self):
from vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.backend.memcache_backend import MemcacheBackend
with patch.object(MemcacheBackend, "__init__", lambda self, pc: None):
backend = MemcacheBackend.__new__(MemcacheBackend)
backend.store = MagicMock()
backend.local_rank = 0
# Set internal state to avoid lazy init logic during tests
backend._lazy_init = False
backend._store_initialized = True
backend._is_a2 = False
backend._registered_buffers = None
backend._buffers_registered = False
return backend
def test_exists(self):
b = self._make_backend()
b.store.batch_is_exist.return_value = [1]
self.assertEqual(b.exists(["k1"]), [1])
def test_register_buffer(self):
b = self._make_backend()
b._is_a2 = True
b.register_buffer([100], [200])
b.store.register_buffer.assert_called_once()
def test_get(self):
b = self._make_backend()
b.store.batch_get_into_layers.return_value = [0]
b.get(["k1"], [[100]], [[10]])
b.store.batch_get_into_layers.assert_called_once()
def test_get_error(self):
b = self._make_backend()
b.store.batch_get_into_layers.return_value = [1] # non-zero = error
b.get(["k1"], [[100]], [[10]])
def test_get_exception(self):
b = self._make_backend()
b.store.batch_get_into_layers.side_effect = RuntimeError("backend fail")
with patch(
"vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.backend.memcache_backend.logger"
) as mock_logger:
b.get(["k1"], [[100]], [[10]])
error_log = _format_log_call(mock_logger.error.call_args)
self.assertIn("RuntimeError", error_log)
self.assertIn("backend fail", error_log)
def test_put(self):
b = self._make_backend()
b.store.batch_put_from_layers.return_value = [0]
b.put(["k1"], [[100]], [[10]])
b.store.batch_put_from_layers.assert_called_once()
def test_put_error(self):
b = self._make_backend()
b.store.batch_put_from_layers.return_value = [1]
b.put(["k1"], [[100]], [[10]])
def test_put_exception(self):
b = self._make_backend()
b.store.batch_put_from_layers.side_effect = RuntimeError("backend fail")
with patch(
"vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.backend.memcache_backend.logger"
) as mock_logger:
b.put(["k1"], [[100]], [[10]])
error_log = _format_log_call(mock_logger.error.call_args)
self.assertIn("RuntimeError", error_log)
self.assertIn("backend fail", error_log)
if __name__ == "__main__":
unittest.main()

View File

@@ -0,0 +1,581 @@
#
# Copyright (c) 2026 Huawei Technologies Co., Ltd. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
# This file is a part of the vllm-ascend project.
#
import hashlib
import unittest
from unittest.mock import MagicMock
import tests.ut.distributed.ascend_store._mock_deps # noqa: F401, E402
from vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.config_data import (
AscendConnectorMetadata,
ChunkedTokenDatabase,
KeyMetadata,
LayerMultiBlockReqMeta,
LayerPoolKey,
LoadSpec,
PoolKey,
ReqMeta,
RequestTracker,
get_block_hashes,
)
_GROUPED_BLOCK_HASH_DOMAIN = b"vllm-ascend-grouped-block-hash-v1\0"
_GROUPED_BLOCK_HASH_LENGTH_PREFIX_BYTES = 4
def _expected_grouped_hash(*block_hashes):
hasher = hashlib.sha256()
hasher.update(_GROUPED_BLOCK_HASH_DOMAIN)
hasher.update(len(block_hashes).to_bytes(_GROUPED_BLOCK_HASH_LENGTH_PREFIX_BYTES, "big"))
for block_hash in block_hashes:
hash_bytes = block_hash.encode("utf-8") if isinstance(block_hash, str) else bytes(block_hash)
hasher.update(len(hash_bytes).to_bytes(_GROUPED_BLOCK_HASH_LENGTH_PREFIX_BYTES, "big"))
hasher.update(hash_bytes)
return hasher.digest()
class TestKeyMetadata(unittest.TestCase):
def test_fields(self):
meta = KeyMetadata(
model_name="llama",
head_or_tp_rank=0,
pcp_rank=0,
dcp_rank=0,
pp_rank=0,
)
self.assertEqual(meta.model_name, "llama")
self.assertEqual(meta.head_or_tp_rank, 0)
self.assertEqual(meta.pcp_rank, 0)
self.assertEqual(meta.dcp_rank, 0)
self.assertEqual(meta.pp_rank, 0)
class TestPoolKey(unittest.TestCase):
def setUp(self):
self.meta = KeyMetadata("llama", 1, 2, 3, 0)
def test_hash_equal(self):
k1 = PoolKey(self.meta, "abc123")
k2 = PoolKey(self.meta, "abc123")
self.assertEqual(hash(k1), hash(k2))
def test_hash_diff(self):
k1 = PoolKey(self.meta, "abc123")
k2 = PoolKey(self.meta, "def456")
self.assertNotEqual(hash(k1), hash(k2))
def test_to_string(self):
k = PoolKey(self.meta, "hash1")
s = k.to_string()
self.assertIn("llama", s)
self.assertIn("@pcp2", s)
self.assertIn("@dcp3", s)
self.assertIn("@head_or_tp_rank:1", s)
self.assertIn("@pp_rank:0", s)
self.assertIn("hash1", s)
def test_pp_ranks_use_distinct_keys(self):
other_pp_meta = KeyMetadata("llama", 1, 2, 3, 1)
pp0_key = PoolKey(self.meta, "hash1")
pp1_key = PoolKey(other_pp_meta, "hash1")
self.assertNotEqual(pp0_key.to_string(), pp1_key.to_string())
self.assertIn("@pp_rank:0", pp0_key.to_string())
self.assertIn("@pp_rank:1", pp1_key.to_string())
def test_split_layers(self):
k = PoolKey(self.meta, "hash1")
layers = k.split_layers(3)
self.assertEqual(len(layers), 3)
for i, lk in enumerate(layers):
self.assertIsInstance(lk, LayerPoolKey)
self.assertEqual(lk.layer_id, i)
self.assertEqual(lk.chunk_hash, "hash1")
class TestLayerPoolKey(unittest.TestCase):
def test_hash(self):
meta = KeyMetadata("model", 0, 0, 0, 0)
k1 = LayerPoolKey(meta, "h1", 0)
k2 = LayerPoolKey(meta, "h1", 1)
self.assertNotEqual(hash(k1), hash(k2))
def test_to_string_contains_layer_id(self):
meta = KeyMetadata("model", 0, 0, 0, 0)
k = LayerPoolKey(meta, "h1", 5)
s = k.to_string()
self.assertIn("@layer_id:5", s)
self.assertIn("model", s)
self.assertTrue(s.endswith("@h1"))
class TestChunkedTokenDatabase(unittest.TestCase):
def setUp(self):
self.meta = KeyMetadata("llama", 0, 0, 0, 0)
self.db = ChunkedTokenDatabase([self.meta], block_size=[16], partitions=None)
self.db.set_group_buffers({0: [1000, 2000]}, {0: [160, 320]}, group_num_layers={0: 1})
def test_make_key_by_hash(self):
key = self.db._make_key_by_hash("abc")
self.assertIsInstance(key, PoolKey)
self.assertEqual(key.chunk_hash, "abc")
def test_process_tokens_empty(self):
result = list(self.db.process_tokens(32, []))
self.assertEqual(result, [])
def test_process_tokens_with_str_hashes(self):
hashes = ["aaa", "bbb"]
result = list(self.db.process_tokens(32, hashes))
self.assertEqual(len(result), 2)
self.assertEqual(result[0][0], 0) # start
self.assertEqual(result[0][1], 16) # end
self.assertEqual(result[1][0], 16)
self.assertEqual(result[1][1], 32)
def test_process_tokens_with_bytes_hashes(self):
hashes = [b"\xaa\xbb", b"\xcc\xdd"]
result = list(self.db.process_tokens(32, hashes))
self.assertEqual(len(result), 2)
def test_process_tokens_with_mask(self):
hashes = ["a", "b", "c"]
result = list(self.db.process_tokens(48, hashes, mask_num=16))
# first chunk (start=0 < mask_num=16) should be skipped
self.assertEqual(len(result), 2)
self.assertEqual(result[0][0], 16)
def test_process_tokens_with_tail_clipped_block_ids_maps_tail_chunks(self):
db = ChunkedTokenDatabase([self.meta], block_size=[128], partitions=None)
hashes = [bytes([idx % 251]) * 32 for idx in range(128)]
result = list(
db.process_token_key_strings_with_block_ids(
128 * 128,
hashes,
[1000, 1001, 1002, 1003],
)
)
self.assertEqual(
[start for start, _, _, _, _ in result],
[124 * 128, 125 * 128, 126 * 128, 127 * 128],
)
self.assertEqual(
[block_id for _, _, _, _, block_id in result],
[1000, 1001, 1002, 1003],
)
def test_process_tokens_token_len_shorter_than_all_blocks(self):
hashes = ["a", "b", "c", "d"]
# token_len=32 means only first 2 blocks valid
result = list(self.db.process_tokens(32, hashes))
self.assertEqual(len(result), 2)
def test_process_tokens_rehashes_grouped_hashes(self):
db = ChunkedTokenDatabase([self.meta], block_size=[16], partitions=None, hash_block_size=8)
result = list(db.process_tokens(32, ["a", "b", "c", "d"]))
self.assertEqual(len(result), 2)
self.assertEqual(result[0][2].chunk_hash, _expected_grouped_hash("a", "b").hex())
self.assertEqual(len(result[0][2].chunk_hash), 64)
def test_key_strings_match_pool_keys(self):
hashes = ["aaa", "bbb", "ccc"]
pool_keys = list(self.db.process_tokens(40, hashes))
self.assertEqual(
list(self.db.process_token_key_strings(40, hashes)),
[
(start, end, key.to_string(), hash_val)
for (start, end, key), hash_val in zip(pool_keys, hashes, strict=True)
],
)
block_ids = [5, 6]
self.assertEqual(
list(self.db.process_token_key_strings_with_block_ids(32, hashes, block_ids)),
[
(start, end, key.to_string(), hash_val, block_id)
for (start, end, key), hash_val, block_id in zip(pool_keys[:2], hashes[:2], block_ids, strict=True)
],
)
def test_direct_keys_preserve_multigroup_layerwise_key_semantics(self):
group_metadata = [
KeyMetadata("llama", 0, 0, 0, 0),
KeyMetadata("llama", 1, 0, 0, 0),
]
db = ChunkedTokenDatabase(group_metadata, block_size=[16, 32], partitions=None, hash_block_size=16)
db.set_group_buffers(
{0: [1000], 1: [2000]},
{0: [160], 1: [320]},
group_cache_families={0: "c1", 1: "c2"},
group_num_layers={0: 2, 1: 2},
)
hashes = ["a", "b", "c", "d"]
pool_key_result = list(db.process_tokens(64, hashes, kv_cache_group_id=1))
direct_key_result = list(db.process_token_key_strings(64, hashes, kv_cache_group_id=1))
self.assertEqual(len(pool_key_result), 1)
self.assertEqual(
direct_key_result[0][:3],
(pool_key_result[0][0], pool_key_result[0][1], pool_key_result[0][2].to_string()),
)
layer_key = pool_key_result[0][2].split_layers(2)[1]
self.assertIn("@group:1@cache_role:kv@cache_family:c2@layer_id:1", layer_key.to_string())
def test_key_strings_pre_shard_after_filtering(self):
hashes = ["a", "b", "c", "d"]
store_mask = [True, False, True, True]
result = list(
self.db.process_token_key_strings_with_block_ids(
64,
hashes,
[10, 11, 12, 13],
chunk_filter=lambda start: store_mask[start // 16],
shard_rank=1,
shard_size=2,
)
)
self.assertEqual([(start, end, block_id) for start, end, _, _, block_id in result], [(32, 48, 12)])
def test_get_block_hashes_rehashes_groups(self):
for hashes in (["a", "b", "c", "d"], [b"a", b"b", b"c", b"d"]):
with self.subTest(hash_type=type(hashes[0])):
result = get_block_hashes(hashes, group_block_size=32, hash_block_size=16)
expected = [
_expected_grouped_hash(hashes[0], hashes[1]),
_expected_grouped_hash(hashes[2], hashes[3]),
]
self.assertEqual(list(result), expected)
def test_prepare_value(self):
addr, size, block_id = self.db.prepare_value(0, 16, [5, 6, 7])
self.assertEqual(block_id, 5)
self.assertEqual(len(addr), 2)
self.assertEqual(addr[0], 1000 + 5 * 160)
self.assertEqual(addr[1], 2000 + 5 * 320)
self.assertEqual(size[0], 160)
self.assertEqual(size[1], 320)
def test_prepare_value_partial_block(self):
addr, size, block_id = self.db.prepare_value(0, 8, [5])
self.assertEqual(size[0], 80) # 160/16*8
self.assertEqual(size[1], 160) # 320/16*8
def test_prepare_value_uses_block_id_override(self):
addr, size, block_id = self.db.prepare_value(64, 80, [5], block_id=99)
self.assertEqual(block_id, 99)
self.assertEqual(addr[0], 1000 + 99 * 160)
self.assertEqual(addr[1], 2000 + 99 * 320)
self.assertEqual(size[0], 160)
self.assertEqual(size[1], 320)
def test_prepare_value_layer(self):
addr, size, block_id = self.db.prepare_value_layer(0, 16, [5, 6], layer_id=0)
self.assertEqual(block_id, 5)
self.assertEqual(len(addr), 2)
# layer_id=0, entries_per_layers=2 => group_addrs[0] and group_addrs[1]
self.assertEqual(addr[0], 1000 + 5 * 160)
self.assertEqual(addr[1], 2000 + 5 * 320)
def test_decode_adaptor_prefill_pp_no_partitions(self):
key, addr, size = self.db.decode_adaptor_prefill_pp(["k1"], [[1, 2]], [[10, 20]])
self.assertEqual(key, ["k1"])
def test_decode_adaptor_prefill_pp_single_partition(self):
db = ChunkedTokenDatabase([self.meta], [16], partitions=[4])
key, addr, size = db.decode_adaptor_prefill_pp(["k1"], [[1, 2]], [[10, 20]])
self.assertEqual(key, ["k1"])
def test_decode_adaptor_prefill_pp_multi_partition(self):
db = ChunkedTokenDatabase([self.meta], [16], partitions=[2, 2])
db.set_group_buffers({0: [1000, 2000]}, {0: [160, 320]})
keys = ["k1@pp_rank:0"]
addrs = [[1, 2, 3, 4, 5, 6, 7, 8]]
sizes = [[10, 20, 30, 40, 50, 60, 70, 80]]
new_keys, new_addrs, new_sizes = db.decode_adaptor_prefill_pp(keys, addrs, sizes)
self.assertEqual(len(new_keys), 2)
self.assertIn("@pp_rank:0", new_keys[0])
self.assertIn("@pp_rank:1", new_keys[1])
class TestLoadSpec(unittest.TestCase):
def test_fields(self):
spec = LoadSpec(vllm_cached_tokens=10, kvpool_cached_tokens=20, can_load=True)
self.assertEqual(spec.vllm_cached_tokens, 10)
self.assertEqual(spec.kvpool_cached_tokens, 20)
self.assertTrue(spec.can_load)
self.assertEqual(spec.token_len, 0)
def test_token_len_default(self):
spec = LoadSpec(0, 0, False, token_len=128)
self.assertEqual(spec.token_len, 128)
class TestRequestTracker(unittest.TestCase):
def test_from_new_request(self):
new_req = MagicMock()
new_req.req_id = "req-1"
new_req.block_ids = [10, 20, 30]
new_req.prompt_token_ids = list(range(100))
tracker = RequestTracker.from_new_request(new_req, num_tokens_to_compute=48)
self.assertEqual(tracker.req_id, "req-1")
self.assertEqual(tracker.token_len, 48)
self.assertEqual(tracker.allocated_block_ids, [10, 20, 30])
self.assertEqual(len(tracker.token_ids), 48)
self.assertEqual(tracker.num_saved_tokens, 0)
def test_from_new_request_nested_block_ids(self):
new_req = MagicMock()
new_req.req_id = "req-2"
new_req.block_ids = [[10, 20], [30, 40]]
new_req.prompt_token_ids = list(range(32))
tracker = RequestTracker.from_new_request(new_req, num_tokens_to_compute=32)
self.assertEqual(tracker.allocated_block_ids, [10, 20])
def test_update_with_list(self):
tracker = RequestTracker(req_id="r1", token_len=16, allocated_block_ids=[1, 2])
tracker.update([3, 4])
self.assertEqual(tracker.allocated_block_ids, [1, 2, 3, 4])
def test_update_with_tuple(self):
tracker = RequestTracker(req_id="r1", token_len=16, allocated_block_ids=[1])
tracker.update(([5, 6], [7, 8]))
self.assertEqual(tracker.allocated_block_ids, [1, 5, 6])
def test_update_with_empty(self):
tracker = RequestTracker(req_id="r1", token_len=16, allocated_block_ids=[1])
tracker.update([])
self.assertEqual(tracker.allocated_block_ids, [1])
def test_update_invalid_type(self):
tracker = RequestTracker(req_id="r1", token_len=16, allocated_block_ids=[1])
with self.assertRaises(ValueError):
tracker.update("invalid") # type: ignore[arg-type]
def test_update_mamba_with_tuple(self):
tracker = RequestTracker(
req_id="r1", token_len=16, allocated_block_ids_by_group=[[1], [2], [3], [4]], block_sizes=[16] * 4
)
tracker.update(([5, 6], [0, 7], [0, 8], [0, 9]))
self.assertEqual(tracker.allocated_block_ids_by_group[0], [1, 5, 6])
self.assertEqual(tracker.allocated_block_ids_by_group[1], [2, 0, 7])
self.assertEqual(tracker.allocated_block_ids_by_group[2], [3, 0, 8])
self.assertEqual(tracker.allocated_block_ids_by_group[3], [4, 0, 9])
def test_update_mamba_mtp_with_tuple_chunk2(self):
tracker = RequestTracker(
req_id="r1",
token_len=32,
allocated_block_ids_by_group=[
[1, 2],
[0, 3, 4, 5, 6],
[0, 7, 8, 9, 10],
[0, 11, 12, 13, 14],
],
mamba_group_ids=[1, 2, 3],
num_speculative_blocks=3,
block_sizes=[16] * 4,
)
tracker.update(([15, 16], [4, 17], [8, 18], [12, 19]), 32)
self.assertEqual(tracker.allocated_block_ids_by_group[0], [1, 2, 15, 16])
self.assertEqual(tracker.allocated_block_ids_by_group[1], [0, 3, 0, 5, 6, 4, 17])
self.assertEqual(tracker.allocated_block_ids_by_group[2], [0, 7, 0, 9, 10, 8, 18])
self.assertEqual(tracker.allocated_block_ids_by_group[3], [0, 11, 0, 13, 14, 12, 19])
def test_update_mamba_mtp_with_tuple_chunk8(self):
tracker = RequestTracker(
req_id="r1",
token_len=128,
allocated_block_ids_by_group=[
[1, 2, 3, 4, 5, 6, 7, 8],
[0, 0, 0, 0, 0, 0, 0, 9, 10, 11, 12],
[0, 0, 0, 0, 0, 0, 0, 13, 14, 15, 16],
[0, 0, 0, 0, 0, 0, 0, 17, 18, 19, 20],
],
mamba_group_ids=[1, 2, 3],
num_speculative_blocks=3,
block_sizes=[16] * 4,
)
tracker.update(
(
[21, 22, 23, 24, 25, 26, 27, 28],
[0, 0, 0, 0, 10, 11, 12, 29],
[0, 0, 0, 0, 14, 15, 16, 30],
[0, 0, 0, 0, 18, 19, 20, 31],
),
128,
)
self.assertEqual(
tracker.allocated_block_ids_by_group[0], [1, 2, 3, 4, 5, 6, 7, 8, 21, 22, 23, 24, 25, 26, 27, 28]
)
self.assertEqual(
tracker.allocated_block_ids_by_group[1], [0, 0, 0, 0, 0, 0, 0, 9, 0, 0, 0, 0, 0, 0, 0, 10, 11, 12, 29]
)
self.assertEqual(
tracker.allocated_block_ids_by_group[2], [0, 0, 0, 0, 0, 0, 0, 13, 0, 0, 0, 0, 0, 0, 0, 14, 15, 16, 30]
)
self.assertEqual(
tracker.allocated_block_ids_by_group[3], [0, 0, 0, 0, 0, 0, 0, 17, 0, 0, 0, 0, 0, 0, 0, 18, 19, 20, 31]
)
class TestReqMeta(unittest.TestCase):
def test_from_request_tracker_basic_save(self):
tracker = RequestTracker(
req_id="r1",
token_len=32,
allocated_block_ids=[0, 1],
num_saved_tokens=0,
token_ids=list(range(32)),
)
meta = ReqMeta.from_request_tracker(tracker, cache_transfer_granularity=16, block_hashes=[b"h1", b"h2"])
self.assertIsNotNone(meta)
self.assertEqual(meta.req_id, "r1")
self.assertTrue(meta.can_save)
self.assertEqual(meta.token_len_chunk, 32)
self.assertIsNone(meta.load_spec)
def test_from_request_tracker_skip_save(self):
tracker = RequestTracker(
req_id="r1",
token_len=32,
allocated_block_ids=[0, 1],
num_saved_tokens=0,
)
meta = ReqMeta.from_request_tracker(tracker, cache_transfer_granularity=16, skip_save=True)
self.assertIsNone(meta)
def test_from_request_tracker_with_load_spec(self):
tracker = RequestTracker(
req_id="r1",
token_len=32,
allocated_block_ids=[0, 1],
num_saved_tokens=0,
)
load_spec = LoadSpec(vllm_cached_tokens=0, kvpool_cached_tokens=32, can_load=True)
meta = ReqMeta.from_request_tracker(tracker, cache_transfer_granularity=16, load_spec=load_spec, skip_save=True)
self.assertIsNotNone(meta)
self.assertIsNotNone(meta.load_spec)
def test_from_request_tracker_load_spec_cannot_load(self):
tracker = RequestTracker(
req_id="r1",
token_len=32,
allocated_block_ids=[0, 1],
num_saved_tokens=32,
)
load_spec = LoadSpec(vllm_cached_tokens=0, kvpool_cached_tokens=32, can_load=False)
meta = ReqMeta.from_request_tracker(tracker, cache_transfer_granularity=16, load_spec=load_spec, skip_save=True)
# can_load=False => load_spec set to None in meta,
# but skip_save+load_spec input is not None, so meta is still created
self.assertIsNotNone(meta)
self.assertIsNone(meta.load_spec)
self.assertFalse(meta.can_save)
def test_from_request_tracker_partial_tokens_discarded(self):
tracker = RequestTracker(
req_id="r1",
token_len=20,
allocated_block_ids=[0, 1],
num_saved_tokens=0,
)
meta = ReqMeta.from_request_tracker(tracker, cache_transfer_granularity=16, discard_partial_chunks=True)
self.assertIsNotNone(meta)
self.assertEqual(meta.token_len_chunk, 16)
def test_from_request_tracker_no_discard(self):
tracker = RequestTracker(
req_id="r1",
token_len=20,
allocated_block_ids=[0, 1],
num_saved_tokens=0,
)
meta = ReqMeta.from_request_tracker(tracker, cache_transfer_granularity=16, discard_partial_chunks=False)
self.assertIsNotNone(meta)
self.assertEqual(meta.token_len_chunk, 20)
def test_from_request_tracker_already_saved(self):
tracker = RequestTracker(
req_id="r1",
token_len=32,
allocated_block_ids=[0, 1],
num_saved_tokens=32,
)
meta = ReqMeta.from_request_tracker(tracker, cache_transfer_granularity=16)
# num_saved_tokens=32, chunk_boundary=ceil(33/16)*16=48 > 32
# so skip_save, and no load_spec => None
self.assertIsNone(meta)
def test_from_request_tracker_with_original_block_size(self):
tracker = RequestTracker(
req_id="r1",
token_len=32,
allocated_block_ids=[0, 1],
num_saved_tokens=0,
)
# Provide block_hashes (2 full blocks for token_len=32 / granularity=16)
# so the boundary_without_hash short-circuit does not zero out the save
# length and skip; this exercises the original_block_size propagation.
meta = ReqMeta.from_request_tracker(
tracker,
cache_transfer_granularity=16,
original_block_size=8,
block_hashes=[b"h0", b"h1"],
)
self.assertIsNotNone(meta)
self.assertEqual(meta.original_block_size, 8)
class TestAscendConnectorMetadata(unittest.TestCase):
def test_add_request(self):
meta = AscendConnectorMetadata(unfinished_request_ids=set(), preempted_req_ids=set())
req = ReqMeta(
req_id="r1",
token_len_chunk=16,
block_ids=[0],
block_hashes=[],
)
meta.add_request(req)
self.assertEqual(len(meta.requests), 1)
self.assertEqual(meta.requests[0].req_id, "r1")
class TestLayerMultiBlockReqMeta(unittest.TestCase):
def test_fields(self):
meta = LayerMultiBlockReqMeta(
req_id="r1",
keys=[],
starts=[0, 16],
ends=[16, 32],
block_ids=[0, 1],
layer_id=2,
)
self.assertEqual(meta.req_id, "r1")
self.assertEqual(meta.layer_id, 2)
self.assertTrue(meta.is_last_chunk)
self.assertIsNone(meta.current_event)
if __name__ == "__main__":
unittest.main()

View File

@@ -0,0 +1,291 @@
#
# Copyright (c) 2026 Huawei Technologies Co., Ltd. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
# This file is a part of the vllm-ascend project.
#
import unittest
from dataclasses import dataclass, replace
from unittest.mock import patch
# isort: off
import tests.ut.distributed.ascend_store._mock_deps # noqa: F401, E402
import torch
from vllm.v1.kv_cache_interface import FullAttentionSpec, KVCacheGroupSpec, SlidingWindowSpec
from vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store import config_data
from vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.config_data import (
ChunkedTokenDatabase,
KeyMetadata,
get_block_hashes,
)
from vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.coordinator import (
AscendStoreCoordinator,
ExternalCachedBlockPool,
)
# isort: on
def _hashes(num_blocks: int) -> list[bytes]:
return [bytes([idx % 251]) * 32 for idx in range(num_blocks)]
def _full_spec(block_size: int) -> FullAttentionSpec:
return FullAttentionSpec(
block_size=block_size,
num_kv_heads=1,
head_size=1,
dtype=torch.float32,
)
def _sliding_spec(block_size: int, sliding_window: int) -> SlidingWindowSpec:
return SlidingWindowSpec(
block_size=block_size,
num_kv_heads=1,
head_size=1,
dtype=torch.float32,
sliding_window=sliding_window,
)
@dataclass(frozen=True)
class _FakeCompressedSpec:
block_size: int
compress_ratio: int
def copy_with_new_block_size(self, block_size):
return replace(self, block_size=block_size)
class _FakeCompressedManager:
@classmethod
def find_longest_cache_hit(
cls,
block_hashes,
max_length,
kv_cache_group_ids,
block_pool,
kv_cache_spec,
drop_eagle_block=False,
alignment_tokens=16,
**kwargs,
):
computed: tuple[list[object], ...] = tuple([] for _ in kv_cache_group_ids)
logical_block_size = kv_cache_spec.block_size * kv_cache_spec.compress_ratio
max_blocks = max_length // logical_block_size
for block_hash in list(block_hashes)[:max_blocks]:
cached = block_pool.get_cached_block(block_hash, kv_cache_group_ids)
if not cached:
break
for blocks, block in zip(computed, cached):
blocks.append(block)
return computed
class TestAscendStoreCoordinator(unittest.TestCase):
def test_load_mask_grouped_hashes_are_reused_by_key_build(self):
block_hashes = _hashes(4)
coord = AscendStoreCoordinator(
[KVCacheGroupSpec(["layer.0"], _full_spec(16))],
scheduler_block_size=16,
hash_block_size=8,
group_block_sizes=[16],
group_cache_families=["c1"],
)
db = ChunkedTokenDatabase(
[KeyMetadata("model", 0, 0, 0, 0)],
block_size=[16],
partitions=None,
hash_block_size=8,
)
db.cache_coordinator = coord
grouped_hash_cache: config_data.GroupedBlockHashCache = {}
with patch.object(
config_data,
"_rehash_block_hash_group",
wraps=config_data._rehash_block_hash_group,
) as rehash:
self.assertEqual(
db.load_mask(
block_hashes,
32,
grouped_hash_cache=grouped_hash_cache,
),
([True, True],),
)
self.assertEqual(rehash.call_count, 2)
keys = list(
db.process_token_key_strings_with_block_ids(
32,
block_hashes,
[10, 11],
grouped_hash_cache=grouped_hash_cache,
)
)
self.assertEqual(len(keys), 2)
self.assertEqual(rehash.call_count, 2)
def test_compressed_group_hits_on_effective_granularity(self):
block_hashes = _hashes(128)
grouped_hash = get_block_hashes(block_hashes, group_block_size=128 * 128, hash_block_size=128)[0]
coord = AscendStoreCoordinator(
[KVCacheGroupSpec(["layer.0"], _full_spec(128))],
scheduler_block_size=128 * 128,
hash_block_size=128,
group_block_sizes=[128],
group_cache_families=["c128"],
)
_, hit_length = coord.find_longest_cache_hit(
block_hashes,
128 * 128,
ExternalCachedBlockPool({(0, bytes(grouped_hash))}),
)
self.assertEqual(hit_length, 128 * 128)
def test_compressed_spec_does_not_apply_ratio_twice(self):
block_hashes = _hashes(128)
grouped_hash = get_block_hashes(block_hashes, group_block_size=128 * 128, hash_block_size=128)[0]
with patch(
"vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.coordinator._get_manager_class",
return_value=_FakeCompressedManager,
):
coord = AscendStoreCoordinator(
[KVCacheGroupSpec(["layer.0"], _FakeCompressedSpec(block_size=128, compress_ratio=128))],
scheduler_block_size=128 * 128,
hash_block_size=128,
group_block_sizes=[128],
group_cache_families=["c128"],
)
_, hit_length = coord.find_longest_cache_hit(
block_hashes,
128 * 128,
ExternalCachedBlockPool({(0, bytes(grouped_hash))}),
)
self.assertEqual(coord.group_effective_specs[0].compress_ratio, 1)
self.assertEqual(hit_length, 128 * 128)
def test_missing_required_group_returns_zero(self):
block_hashes = _hashes(128)
c1_exists = {(0, block_hash) for block_hash in block_hashes}
coord = AscendStoreCoordinator(
[
KVCacheGroupSpec(["layer.0"], _full_spec(128)),
KVCacheGroupSpec(["layer.1"], _full_spec(128)),
],
scheduler_block_size=128 * 128,
hash_block_size=128,
group_block_sizes=[128, 128],
group_cache_families=["c1", "c128"],
)
_, hit_length = coord.find_longest_cache_hit(
block_hashes,
128 * 128,
ExternalCachedBlockPool(c1_exists),
)
self.assertEqual(hit_length, 0)
def test_store_mask_uses_manager_reachability(self):
coord = AscendStoreCoordinator(
[KVCacheGroupSpec(["layer.0"], _sliding_spec(block_size=128, sliding_window=256))],
scheduler_block_size=512,
hash_block_size=128,
group_block_sizes=[128],
group_cache_families=["c1"],
)
with patch(
"vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.coordinator._reachable_block_mask",
return_value=[False, False, False, True],
):
masks = coord.store_mask(512)
self.assertEqual(masks, ([False, False, False, True],))
def test_lookup_mask_uses_reachability_without_retention(self):
coord = AscendStoreCoordinator(
[KVCacheGroupSpec(["layer.0"], _sliding_spec(block_size=128, sliding_window=256))],
scheduler_block_size=512,
hash_block_size=128,
group_block_sizes=[128],
group_cache_families=["c1"],
retention_interval=256,
)
with patch(
"vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.coordinator._reachable_block_mask",
return_value=[False, False, False, True],
) as reachable:
masks = coord.lookup_mask(512)
self.assertEqual(masks, ([False, False, False, True],))
self.assertIsNone(reachable.call_args.kwargs["retention_interval"])
def test_store_mask_propagates_eagle_to_same_spec_siblings(self):
calls = []
def fake_reachable_block_mask(*args, **kwargs):
calls.append(kwargs["use_eagle"])
return [True, False, True, False]
shared_spec = _sliding_spec(block_size=128, sliding_window=256)
coord = AscendStoreCoordinator(
[
KVCacheGroupSpec(["layer.0"], shared_spec),
KVCacheGroupSpec(["layer.mtp"], shared_spec, is_eagle_group=True),
],
scheduler_block_size=512,
hash_block_size=128,
group_block_sizes=[128, 128],
group_cache_families=["c1", "c1"],
)
with patch(
"vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.coordinator._reachable_block_mask",
side_effect=fake_reachable_block_mask,
):
masks = coord.store_mask(512)
self.assertEqual(calls, [True, True])
self.assertEqual(masks, ([True, False, True, False], [True, False, True, False]))
def test_compressed_masks_stay_unmasked(self):
coord = AscendStoreCoordinator(
[KVCacheGroupSpec(["layer.0"], _sliding_spec(block_size=128, sliding_window=512))],
scheduler_block_size=2048,
hash_block_size=128,
group_block_sizes=[128],
group_cache_families=["c4"],
)
self.assertEqual(coord.store_mask(2048, num_prompt_tokens=2048), ([True] * 4,))
with patch.object(
coord,
"find_longest_cache_hit",
return_value=(([False, False, False, True],), 2048),
):
self.assertEqual(coord.load_mask(_hashes(16), 2048), ([True] * 4,))
if __name__ == "__main__":
unittest.main()

View File

@@ -0,0 +1,664 @@
#
# Copyright (c) 2026 Huawei Technologies Co., Ltd. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
# This file is a part of the vllm-ascend project.
#
import threading
import unittest
from unittest.mock import MagicMock, patch
# isort: off
import tests.ut.distributed.ascend_store._mock_deps # noqa: F401, E402
from vllm.distributed.kv_events import BlockStored
from vllm.v1.core.kv_cache_utils import maybe_convert_block_hash
from vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store import config_data
from vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.config_data import (
ChunkedTokenDatabase,
KeyMetadata,
LayerMultiBlockReqMeta,
LayerPoolKey,
LoadSpec,
ReqMeta,
)
# isort: on
from vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.kv_transfer import (
KVCacheStoreLayerRecvingThread,
KVCacheStoreLayerSendingThread,
KVCacheStoreRecvingThread,
KVCacheStoreSendingThread,
KVTransferThread,
)
class FakeStore:
def __init__(self, exists_result=None):
self.exists_result = exists_result or []
self.put_calls = []
self.get_calls = []
def set_device(self):
pass
def exists(self, keys):
return self.exists_result[: len(keys)]
def put(self, keys, addrs, sizes):
self.put_calls.append((list(keys), list(addrs), list(sizes)))
def get(self, keys, addrs, sizes):
self.get_calls.append((list(keys), list(addrs), list(sizes)))
class FakeTokenDatabase(ChunkedTokenDatabase):
def __init__(self, block_size=16):
super().__init__([KeyMetadata("m", 0, 0, 0, 0)], [block_size], None)
self.set_group_buffers({0: [1000]}, {0: [block_size]}, {0: [1]}, group_num_layers={0: 1})
class MaskedFakeTokenDatabase(FakeTokenDatabase):
def __init__(self, block_size=16, masks=([True],)):
super().__init__(block_size)
self.masks = masks
def store_mask(self, token_len, num_prompt_tokens=None):
return self.masks
def load_mask(self, block_hashes, token_len, grouped_hash_cache=None):
return self.masks
def mask_allows_chunk(self, masks, kv_cache_group_id, start):
if masks is None:
return True
block_idx = start // self.get_block_size(kv_cache_group_id)
return block_idx < len(masks[kv_cache_group_id]) and masks[kv_cache_group_id][block_idx]
class TestKVTransferThread(unittest.TestCase):
def _make_thread(self, exists_result=None):
store = FakeStore(exists_result or [])
db = FakeTokenDatabase()
t = KVTransferThread(
m_store=store,
token_database=db,
block_size=16,
tp_rank=0,
dcp_size=1,
ready_event=threading.Event(),
name="test",
)
return t, store
def test_add_request(self):
t, _ = self._make_thread()
req = MagicMock()
t.add_request(req)
self.assertFalse(t.request_queue.empty())
def test_get_and_clear_finished_requests(self):
t, _ = self._make_thread()
t.set_finished_request("r1")
t.set_finished_request("r2")
finished = t.get_and_clear_finished_requests()
self.assertEqual(finished, {"r1", "r2"})
self.assertEqual(t.get_and_clear_finished_requests(), set())
def test_lookup_all_exist(self):
t, _ = self._make_thread([1, 1, 1])
result = t.lookup(["k1", "k2", "k3"])
self.assertEqual(result, [True, True, True])
def test_lookup_partial(self):
t, _ = self._make_thread([1, 0, 1])
result = t.lookup(["k1", "k2", "k3"])
self.assertEqual(result, [True, False, True])
def test_lookup_exception(self):
t, store = self._make_thread()
store.exists = MagicMock(side_effect=Exception("conn fail"))
result = t.lookup(["k1"])
self.assertEqual(result, [False])
def test_update_and_get_kv_events(self):
t, _ = self._make_thread()
event1 = BlockStored(
block_hashes=["h1"],
parent_block_hash=None,
token_ids=[1, 2, 3],
block_size=16,
lora_id=None,
medium="cpu",
lora_name=None,
)
event2 = BlockStored(
block_hashes=["h2"],
parent_block_hash="h1",
token_ids=[4, 5, 6],
block_size=16,
lora_id=None,
medium="cpu",
lora_name=None,
)
t.update_kv_event([event1, event2])
events = t.get_kv_events()
self.assertEqual(len(events), 2)
# After get, events should be cleared
self.assertEqual(len(t.get_kv_events()), 0)
def test_handle_request_base_noop(self):
t, _ = self._make_thread()
# Base class _handle_request does nothing
t._handle_request(MagicMock())
class TestKVCacheStoreSendingThread(unittest.TestCase):
def _make_thread(self, exists_result=None, kv_role="kv_producer", enable_kv_event=False):
store = FakeStore(exists_result or [0, 0, 0, 0])
db = FakeTokenDatabase()
t = KVCacheStoreSendingThread(
m_store=store,
token_database=db,
block_size=16,
tp_rank=0,
dcp_size=1,
put_step=1,
kv_role=kv_role,
ready_event=threading.Event(),
group_uses_align_state=[False],
enable_kv_event=enable_kv_event,
)
return t, store
def test_handle_request_puts_missing_keys(self):
t, store = self._make_thread([1, 0, 1, 0])
req = ReqMeta(
req_id="r1",
token_len_chunk=64,
block_ids=[0, 1, 2, 3],
block_hashes=[b"h0", b"h1", b"h2", b"h3"], # type: ignore[arg-type]
current_event=None,
)
t.add_stored_request("r1")
t.request_queue.put(req)
t._handle_request(req)
self.assertEqual(len(store.put_calls), 1)
keys, _, _ = store.put_calls[0]
self.assertEqual(len(keys), 2)
def test_handle_request_all_exist_no_put(self):
t, store = self._make_thread([1, 1])
req = ReqMeta(
req_id="r1",
token_len_chunk=32,
block_ids=[0, 1],
block_hashes=[b"h0", b"h1"], # type: ignore[arg-type]
current_event=None,
)
t.add_stored_request("r1")
t.request_queue.put(req)
t._handle_request(req)
self.assertEqual(len(store.put_calls), 0)
def test_handle_request_not_in_stored(self):
t, store = self._make_thread([0])
req = ReqMeta(
req_id="r1",
token_len_chunk=16,
block_ids=[0],
block_hashes=[b"h0"], # type: ignore[arg-type]
current_event=None,
)
t.request_queue.put(req)
t._handle_request(req)
self.assertEqual(len(store.put_calls), 0)
def test_handle_request_with_kv_event(self):
t, store = self._make_thread([1, 0, 1], enable_kv_event=True)
req = ReqMeta(
req_id="r1",
token_len_chunk=48,
block_ids=[0, 1, 2],
block_hashes=[b"h0", b"h1", b"h2"], # type: ignore[arg-type]
current_event=None,
token_ids=list(range(48)),
original_block_size=16,
)
t.add_stored_request("r1")
t.request_queue.put(req)
t._handle_request(req)
events = t.get_kv_events()
self.assertEqual(len(events), 1)
self.assertEqual(events[0].block_hashes, [maybe_convert_block_hash(b"h1")])
self.assertEqual(events[0].parent_block_hash, maybe_convert_block_hash(b"h0"))
def test_save_reuses_grouped_hashes_for_kv_events(self):
store = FakeStore([0, 0])
db = ChunkedTokenDatabase(
[KeyMetadata("m", 0, 0, 0, 0)],
[16],
None,
hash_block_size=8,
)
db.set_group_buffers({0: [1000]}, {0: [16]}, {0: [1]}, group_num_layers={0: 1})
thread = KVCacheStoreSendingThread(
m_store=store,
token_database=db,
block_size=16,
tp_rank=0,
dcp_size=1,
put_step=1,
kv_role="kv_producer",
ready_event=threading.Event(),
group_uses_align_state=[False],
enable_kv_event=True,
)
request = ReqMeta(
req_id="r1",
token_len_chunk=32,
block_ids=[0, 1],
block_hashes=[b"h0", b"h1", b"h2", b"h3"], # type: ignore[arg-type]
current_event=None,
token_ids=list(range(32)),
original_block_size=16,
)
thread.add_stored_request("r1")
thread.request_queue.put(request)
with patch.object(
config_data,
"_rehash_block_hash_group",
wraps=config_data._rehash_block_hash_group,
) as rehash:
thread._handle_request(request)
self.assertEqual(rehash.call_count, 2)
self.assertEqual(len(thread.get_kv_events()), 2)
def test_handle_request_consumer_role(self):
t, store = self._make_thread([0], kv_role="kv_consumer")
req = ReqMeta(
req_id="r1",
token_len_chunk=16,
block_ids=[0],
block_hashes=[b"h0"], # type: ignore[arg-type]
current_event=None,
)
t.add_stored_request("r1")
t.request_queue.put(req)
t._handle_request(req)
self.assertEqual(len(store.put_calls), 1)
def test_add_dec_delete_stored_request(self):
t, _ = self._make_thread()
t.add_stored_request("r1")
t.add_stored_request("r1")
self.assertEqual(t.stored_requests["r1"], 2)
t.dec_stored_request("r1")
self.assertEqual(t.stored_requests["r1"], 1)
t.delete_finished_stored_request("r1")
self.assertNotIn("r1", t.stored_requests)
def test_dec_nonexistent_request(self):
t, _ = self._make_thread()
t.dec_stored_request("nonexist") # should not raise
def test_delete_nonexistent_request(self):
t, _ = self._make_thread()
t.delete_finished_stored_request("nonexist") # should not raise
def test_handle_request_with_current_event(self):
t, store = self._make_thread([0])
event = MagicMock()
req = ReqMeta(
req_id="r1",
token_len_chunk=16,
block_ids=[0],
block_hashes=[b"h0"], # type: ignore[arg-type]
current_event=event,
)
t.add_stored_request("r1")
t.request_queue.put(req)
t._handle_request(req)
event.synchronize.assert_called_once()
def test_handle_request_dcp_size_gt_1(self):
store = FakeStore([0, 0])
db = FakeTokenDatabase()
t = KVCacheStoreSendingThread(
m_store=store,
token_database=db,
block_size=16,
tp_rank=0,
dcp_size=2,
put_step=1,
kv_role="kv_producer",
ready_event=threading.Event(),
group_uses_align_state=[False],
)
req = ReqMeta(
req_id="r1",
token_len_chunk=32,
block_ids=[0, 1],
block_hashes=[b"h0", b"h1"], # type: ignore[arg-type]
current_event=None,
)
t.add_stored_request("r1")
t.request_queue.put(req)
t._handle_request(req)
# dcp_size > 1 means no slicing
self.assertEqual(len(store.put_calls), 1)
def test_handle_request_applies_store_mask(self):
store = FakeStore([0, 0])
db = MaskedFakeTokenDatabase(masks=([True, False],))
t = KVCacheStoreSendingThread(
m_store=store,
token_database=db,
block_size=16,
tp_rank=0,
dcp_size=1,
put_step=1,
kv_role="kv_producer",
ready_event=threading.Event(),
group_uses_align_state=[False],
)
req = ReqMeta(
req_id="r1",
token_len_chunk=32,
block_ids=[0, 1],
block_hashes=[b"h0", b"h1"], # type: ignore[arg-type]
current_event=None,
)
t.add_stored_request("r1")
t.request_queue.put(req)
t._handle_request(req)
keys, _, _ = store.put_calls[0]
self.assertEqual(len(keys), 1)
def test_handle_request_skips_compressed_hit_in_raw_token_domain(self):
t, store = self._make_thread([0, 0])
t.token_database.group_cache_families["kv"][0] = "c4"
req = ReqMeta(
req_id="r1",
token_len_chunk=128,
block_ids=[0, 1],
block_hashes=[f"h{i}" for i in range(8)],
load_spec=LoadSpec(
vllm_cached_tokens=0,
kvpool_cached_tokens=63,
kvpool_store_skip_tokens=64,
can_load=True,
),
)
t.add_stored_request("r1")
t.request_queue.put(req)
t._handle_request(req)
keys, addrs, _ = store.put_calls[0]
self.assertEqual(len(keys), 1)
self.assertEqual(addrs, [[1001]])
def test_save_exception_cleans_queue_lifecycle(self):
t, store = self._make_thread([0])
store.put = MagicMock(side_effect=RuntimeError("put failed"))
req = ReqMeta(
req_id="r1",
token_len_chunk=16,
block_ids=[0],
block_hashes=[b"h0"], # type: ignore[arg-type]
)
t.add_stored_request("r1")
t.request_queue.put(req)
t._handle_request(req)
self.assertEqual(t.request_queue.unfinished_tasks, 0)
self.assertNotIn("r1", t.stored_requests)
class TestKVCacheStoreRecvingThread(unittest.TestCase):
def test_handle_request(self):
store = FakeStore()
db = FakeTokenDatabase()
t = KVCacheStoreRecvingThread(
m_store=store,
token_database=db,
block_size=16,
tp_rank=0,
dcp_size=1,
ready_event=threading.Event(),
invalid_block_ids=set(),
invalid_block_ids_lock=threading.Lock(),
)
load_spec = LoadSpec(vllm_cached_tokens=0, kvpool_cached_tokens=32, can_load=True, token_len=32)
req = ReqMeta(
req_id="r1",
token_len_chunk=32,
block_ids=[0, 1],
block_hashes=[b"h0", b"h1"], # type: ignore[arg-type]
load_spec=load_spec,
)
t.request_queue.put(req)
t._handle_request(req)
self.assertEqual(len(store.get_calls), 1)
finished = t.get_and_clear_finished_requests()
self.assertIn("r1", finished)
def test_handle_request_applies_load_mask(self):
store = FakeStore()
db = MaskedFakeTokenDatabase(masks=([True, False],))
t = KVCacheStoreRecvingThread(
m_store=store,
token_database=db,
block_size=16,
tp_rank=0,
dcp_size=1,
ready_event=threading.Event(),
invalid_block_ids=set(),
invalid_block_ids_lock=threading.Lock(),
)
load_spec = LoadSpec(vllm_cached_tokens=0, kvpool_cached_tokens=32, can_load=True, token_len=32)
req = ReqMeta(
req_id="r1",
token_len_chunk=32,
block_ids=[0, 1],
block_hashes=[b"h0", b"h1"], # type: ignore[arg-type]
load_spec=load_spec,
)
t.request_queue.put(req)
t._handle_request(req)
keys, _, _ = store.get_calls[0]
self.assertEqual(len(keys), 1)
@unittest.skip("LayerMultiBlockReqMeta API is deprecated, tests need update for LayerTransferTask")
class TestKVCacheStoreLayerSendingThread(unittest.TestCase):
def _make_thread(self, exists_result=None, num_layers=2):
store = FakeStore(exists_result or [0, 0])
db = FakeTokenDatabase()
t = KVCacheStoreLayerSendingThread(
m_store=store,
token_database=db,
block_size=16,
tp_rank=0,
tp_size=1,
dcp_size=1,
put_step=1,
my_key_index=0,
num_ranks_per_layer=1,
page_size_bytes=32,
ready_event=threading.Event(),
num_layers=num_layers,
layer_save_finished_events=[threading.Event() for _ in range(num_layers)],
sync_save_events=[],
)
return t, store
def _make_layer_req(self, layer_id=0, is_last_chunk=False, num_keys=2):
meta = KeyMetadata("m", 0, 0, 0, 0)
keys = [LayerPoolKey(meta, f"h{i}", layer_id) for i in range(num_keys)]
return LayerMultiBlockReqMeta(
req_id="r1",
keys=keys,
starts=[i * 16 for i in range(num_keys)],
ends=[(i + 1) * 16 for i in range(num_keys)],
block_ids=list(range(num_keys)),
layer_id=layer_id,
is_last_chunk=is_last_chunk,
current_event=None,
token_ids=list(range(num_keys * 16)),
original_block_size=16,
block_hashes=[f"h{i}".encode() for i in range(num_keys)],
)
def test_handle_request_puts_missing(self):
t, store = self._make_thread([1, 0])
req = self._make_layer_req(layer_id=0)
t.add_stored_request(req.req_id)
t.request_queue.put(req)
t._handle_request(req)
self.assertEqual(len(store.put_calls), 1)
keys, _, _ = store.put_calls[0]
self.assertEqual(len(keys), 1)
def test_handle_request_all_exist_not_last(self):
t, store = self._make_thread([1, 1])
req = self._make_layer_req(layer_id=0, is_last_chunk=False)
t.add_stored_request(req.req_id)
t.request_queue.put(req)
t._handle_request(req)
self.assertEqual(len(store.put_calls), 0)
def test_handle_request_all_exist_last_chunk_final_layer(self):
t, store = self._make_thread([1, 1], num_layers=2)
req = self._make_layer_req(layer_id=1, is_last_chunk=True)
t.add_stored_request(req.req_id)
t.request_queue.put(req)
t._handle_request(req)
finished = t.get_and_clear_finished_requests()
self.assertIn("r1", finished)
def test_handle_request_empty_keys(self):
t, store = self._make_thread()
_meta = KeyMetadata("m", 0, 0, 0, 0)
req = LayerMultiBlockReqMeta(
req_id="r1",
keys=[],
starts=[],
ends=[],
block_ids=[],
layer_id=0,
is_last_chunk=True,
)
t.add_stored_request(req.req_id)
t.request_queue.put(req)
t._handle_request(req)
finished = t.get_and_clear_finished_requests()
self.assertNotIn("r1", finished)
def test_handle_request_with_current_event(self):
t, store = self._make_thread([0])
event = MagicMock()
meta = KeyMetadata("m", 0, 0, 0, 0)
req = LayerMultiBlockReqMeta(
req_id="r1",
keys=[LayerPoolKey(meta, "h0", 0)],
starts=[0],
ends=[16],
block_ids=[0],
layer_id=0,
is_last_chunk=False,
current_event=event,
)
t.add_stored_request(req.req_id)
t.request_queue.put(req)
t._handle_request(req)
event.synchronize.assert_called_once()
def test_handle_request_last_chunk_final_layer_with_missing(self):
t, store = self._make_thread([0], num_layers=2)
req = self._make_layer_req(layer_id=1, is_last_chunk=True, num_keys=1)
t.add_stored_request(req.req_id)
t.request_queue.put(req)
t._handle_request(req)
finished = t.get_and_clear_finished_requests()
self.assertIn("r1", finished)
def test_layerwise_kv_event_published_on_final_layer(self):
t, store = self._make_thread([0], num_layers=2)
req = self._make_layer_req(layer_id=1, is_last_chunk=True, num_keys=1)
t.add_stored_request(req.req_id)
t.request_queue.put(req)
t._handle_request(req)
events = t.get_kv_events()
self.assertEqual(len(events), 1)
self.assertEqual(events[0].block_hashes, [maybe_convert_block_hash(b"h0")])
self.assertEqual(events[0].token_ids, list(range(16)))
self.assertEqual(events[0].block_size, 16)
def test_layerwise_kv_event_not_published_before_final_layer(self):
t, store = self._make_thread([0], num_layers=2)
req = self._make_layer_req(layer_id=0, is_last_chunk=False, num_keys=1)
t.add_stored_request(req.req_id)
t.request_queue.put(req)
t._handle_request(req)
self.assertEqual(t.get_kv_events(), [])
def test_layerwise_kv_event_uses_missing_blocks_from_previous_layers(self):
t, store = self._make_thread([0], num_layers=2)
first_layer_req = self._make_layer_req(layer_id=0, is_last_chunk=True, num_keys=1)
t.add_stored_request(first_layer_req.req_id)
t.request_queue.put(first_layer_req)
t._handle_request(first_layer_req)
t.m_store.exists_result = [1]
final_layer_req = self._make_layer_req(layer_id=1, is_last_chunk=True, num_keys=1)
t.request_queue.put(final_layer_req)
t._handle_request(final_layer_req)
events = t.get_kv_events()
self.assertEqual(len(events), 1)
self.assertEqual(events[0].block_hashes, [maybe_convert_block_hash(b"h0")])
@unittest.skip("LayerMultiBlockReqMeta API is deprecated, tests need update for LayerTransferTask")
class TestKVCacheStoreLayerRecvingThread(unittest.TestCase):
def test_handle_request(self):
store = FakeStore()
db = FakeTokenDatabase()
get_event = threading.Event()
t = KVCacheStoreLayerRecvingThread(
m_store=store,
token_database=db,
block_size=16,
tp_rank=0,
dcp_size=1,
ready_event=threading.Event(),
get_event=get_event,
invalid_block_ids=set(),
invalid_block_ids_lock=threading.Lock(),
)
meta = KeyMetadata("m", 0, 0, 0, 0)
req = LayerMultiBlockReqMeta(
req_id="r1",
keys=[LayerPoolKey(meta, "h0", 0)],
starts=[0],
ends=[16],
block_ids=[0],
layer_id=0,
)
t.request_queue.put(req)
t._handle_request(req)
self.assertEqual(len(store.get_calls), 1)
self.assertTrue(get_event.is_set())
if __name__ == "__main__":
unittest.main()

File diff suppressed because it is too large Load Diff

File diff suppressed because it is too large Load Diff