Files
enginex-ascend-910-vllm/tests/ut/distributed/ascend_store/test_backend.py
Sun Ruoxi 7f8a1b1f7a init v0.23.0
Signed-off-by: Sun Ruoxi <sunruoxi@4paradigm.com>
2026-08-27 15:11:51 +08:00

666 lines
25 KiB
Python

#
# 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()