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,485 @@
#
# Copyright (c) 2025 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.
#
import json
from unittest.mock import MagicMock, patch
import pytest
import torch
from torch import nn
from vllm_ascend.model_loader.netloader.netloader import DRAFT_PORT_OFFSET, ModelNetLoaderElastic
class DummyDeviceConfig:
device = "cuda"
device_type = "cuda"
class DummyParallelConfig:
tensor_parallel_size = 1
pipeline_parallel_size = 1
class DummyVllmConfig:
device_config = DummyDeviceConfig()
parallel_config = DummyParallelConfig()
additional_config = None
quant_config = None
speculative_config: object | None = None
class DummyModelConfig:
model = "dummy-model"
dtype = torch.float32
runner_type: str | None = None
class DummyDraftModelConfig(DummyModelConfig):
model = "draft-model"
runner_type = "draft"
@pytest.fixture
def default_load_config():
class DummyLoadConfig:
model_loader_extra_config = None
load_format = "default"
return DummyLoadConfig()
def make_loader_with_config(extra):
class DummyLoadConfig:
model_loader_extra_config = extra
load_format = "default"
return ModelNetLoaderElastic(DummyLoadConfig())
def test_init_with_extra_config_file(tmp_path, monkeypatch):
# Generate test JSON file
config_content = {
"SOURCE": [{"device_id": 0}],
"MODEL": "foo-model",
"LISTEN_PORT": 5001,
"INT8_CACHE": "hbm",
"OUTPUT_PREFIX": str(tmp_path),
}
config_file = tmp_path / "config.json"
config_file.write_text(json.dumps(config_content))
dummy_logger = MagicMock()
monkeypatch.setattr("vllm.logger.logger", dummy_logger)
monkeypatch.setattr("vllm_ascend.model_loader.netloader.utils.is_valid_path_prefix", lambda x: True)
extra = {"CONFIG_FILE": str(config_file)}
loader = make_loader_with_config(extra)
assert loader.model_path == "foo-model"
assert loader.source == [{"device_id": 0}]
assert loader.listen_port == 5001
assert loader.int8_cache == "hbm"
assert loader.output_prefix == str(tmp_path)
def test_init_with_extra_config(monkeypatch):
dummy_logger = MagicMock()
monkeypatch.setattr("vllm.logger.logger", dummy_logger)
monkeypatch.setattr("vllm_ascend.model_loader.netloader.utils.is_valid_path_prefix", lambda x: True)
extra = {
"SOURCE": [{"device_id": 0}],
"MODEL": "foo",
"LISTEN_PORT": "4000",
"INT8_CACHE": "dram",
"OUTPUT_PREFIX": "/tmp/",
}
loader = make_loader_with_config(extra)
assert loader.model_path == "foo"
assert loader.listen_port == 4000
assert loader.int8_cache == "dram"
assert loader.output_prefix == "/tmp/"
assert loader.source == [{"device_id": 0}]
def test_init_with_invalid_config(monkeypatch):
dummy_logger = MagicMock()
monkeypatch.setattr("vllm.logger.logger", dummy_logger)
monkeypatch.setattr("vllm_ascend.model_loader.netloader.utils.is_valid_path_prefix", lambda x: False)
# c
extra = {
"SOURCE": None,
"MODEL": None,
"LISTEN_PORT": None,
"INT8_CACHE": "something",
"OUTPUT_PREFIX": None,
}
loader = make_loader_with_config(extra)
assert loader.model_path is None
assert loader.listen_port is None
assert loader.int8_cache == "no"
assert loader.output_prefix is None
def test_clear_static_forward_context_clears_current_vllm_config(monkeypatch):
class DummyCompilationConfig:
def __init__(self):
self.static_forward_context = {}
class ConfigWithCompilation:
def __init__(self):
self.compilation_config = DummyCompilationConfig()
passed_config = ConfigWithCompilation()
current_config = ConfigWithCompilation()
passed_config.compilation_config.static_forward_context["passed.layer"] = object()
current_config.compilation_config.static_forward_context["current.layer"] = object()
monkeypatch.setattr(
"vllm_ascend.model_loader.netloader.netloader.get_current_vllm_config",
lambda: current_config,
)
ModelNetLoaderElastic._clear_static_forward_context(passed_config)
assert passed_config.compilation_config.static_forward_context == {}
assert current_config.compilation_config.static_forward_context == {}
@patch("vllm_ascend.model_loader.netloader.netloader.logger")
def test_load_model_elastic_success(mock_logger, monkeypatch, tmp_path):
monkeypatch.setattr("torch.distributed.get_rank", lambda: 0)
class FakeContext:
def __enter__(self):
pass
def __exit__(self, a, b, c):
pass
monkeypatch.setattr("torch.device", lambda d: FakeContext())
# patch deep copy
monkeypatch.setattr("vllm_ascend.model_loader.netloader.netloader.deepcopy", lambda x: x)
# patch set_default_torch_dtype
monkeypatch.setattr(
"vllm_ascend.model_loader.netloader.netloader.set_default_torch_dtype", lambda dtype: FakeContext()
)
# patch initialize_model
dummy_model = MagicMock(spec=nn.Module)
dummy_model.eval.return_value = dummy_model
monkeypatch.setattr("vllm_ascend.model_loader.netloader.netloader.initialize_model", lambda **kwargs: dummy_model)
# patch elastic_load
monkeypatch.setattr("vllm_ascend.model_loader.netloader.netloader.elastic_load", lambda **kwargs: dummy_model)
# patch process_weights_after_loading
monkeypatch.setattr(
"vllm_ascend.model_loader.netloader.netloader.process_weights_after_loading", lambda *a, **k: None
)
# patch get_ip
monkeypatch.setattr("vllm.utils.network_utils.get_ip", lambda: "127.0.0.1")
# patch find_free_port
monkeypatch.setattr("vllm_ascend.model_loader.netloader.netloader.find_free_port", lambda: 8888)
# patch ElasticServer
class DummyElasticServer:
def __init__(*a, **k):
pass
def start(self):
pass
monkeypatch.setattr("vllm_ascend.model_loader.netloader.netloader.ElasticServer", DummyElasticServer)
# write output_prefix to the temporary directory
extra = {
"SOURCE": [{"device_id": 0}],
"MODEL": "foo",
"LISTEN_PORT": 5555,
"OUTPUT_PREFIX": str(tmp_path) + "/output_",
"INT8_CACHE": "no",
}
loader = make_loader_with_config(extra)
vllm_config = DummyVllmConfig()
model_config = DummyModelConfig()
result = loader.load_model(vllm_config, model_config)
assert isinstance(result, nn.Module)
# Check file
written_file = tmp_path / "output_0.txt"
assert written_file.exists()
def _patch_loader_common(monkeypatch):
monkeypatch.setattr("torch.distributed.get_rank", lambda: 0)
class FakeContext:
def __enter__(self):
pass
def __exit__(self, a, b, c):
pass
monkeypatch.setattr("torch.device", lambda d: FakeContext())
monkeypatch.setattr("vllm_ascend.model_loader.netloader.netloader.deepcopy", lambda x: x)
monkeypatch.setattr(
"vllm_ascend.model_loader.netloader.netloader.set_default_torch_dtype", lambda dtype: FakeContext()
)
dummy_model = MagicMock(spec=nn.Module)
dummy_model.eval.return_value = dummy_model
monkeypatch.setattr("vllm_ascend.model_loader.netloader.netloader.initialize_model", lambda **kwargs: dummy_model)
monkeypatch.setattr(
"vllm_ascend.model_loader.netloader.netloader.process_weights_after_loading", lambda *a, **k: None
)
monkeypatch.setattr("vllm.utils.network_utils.get_ip", lambda: "127.0.0.1")
monkeypatch.setattr("vllm_ascend.model_loader.netloader.netloader.find_free_port", lambda: 8888)
return dummy_model
@patch("vllm_ascend.model_loader.netloader.netloader.logger")
def test_target_model_waits_for_all_netloader_ranks_before_draft(mock_logger, monkeypatch):
dummy_model = _patch_loader_common(monkeypatch)
monkeypatch.setattr("vllm_ascend.model_loader.netloader.netloader.elastic_load", lambda **kwargs: dummy_model)
class DummyElasticServer:
def __init__(*a, **k):
pass
def start(self):
pass
barrier_calls = []
monkeypatch.setattr("vllm_ascend.model_loader.netloader.netloader.ElasticServer", DummyElasticServer)
monkeypatch.setattr("torch.distributed.is_available", lambda: True)
monkeypatch.setattr("torch.distributed.is_initialized", lambda: True)
monkeypatch.setattr("torch.distributed.barrier", lambda: barrier_calls.append("barrier"))
extra = {
"SOURCE": [{"device_id": 0, "sources": ["127.0.0.1:5000"]}],
"MODEL": "dummy-model",
"LISTEN_PORT": 5555,
"INT8_CACHE": "dram",
}
loader = make_loader_with_config(extra)
vllm_config = DummyVllmConfig()
vllm_config.speculative_config = object()
loader.load_model(vllm_config, DummyModelConfig())
assert barrier_calls == ["barrier"]
@patch("vllm_ascend.model_loader.netloader.netloader.logger")
def test_failed_target_model_participates_in_barrier_before_error(mock_logger, monkeypatch):
_patch_loader_common(monkeypatch)
monkeypatch.setattr("vllm_ascend.model_loader.netloader.netloader.elastic_load", lambda **kwargs: None)
monkeypatch.setattr(
ModelNetLoaderElastic,
"revert_to_default",
lambda self, *args, **kwargs: (None, False),
)
barrier_calls = []
monkeypatch.setattr("torch.distributed.is_available", lambda: True)
monkeypatch.setattr("torch.distributed.is_initialized", lambda: True)
monkeypatch.setattr("torch.distributed.barrier", lambda: barrier_calls.append("barrier"))
extra = {
"SOURCE": [{"device_id": 0, "sources": ["127.0.0.1:5000"]}],
"MODEL": "dummy-model",
"LISTEN_PORT": 5555,
"INT8_CACHE": "dram",
}
loader = make_loader_with_config(extra)
vllm_config = DummyVllmConfig()
vllm_config.speculative_config = object()
with pytest.raises(RuntimeError, match="NetLoader elastic loads model fails"):
loader.load_model(vllm_config, DummyModelConfig())
assert barrier_calls == ["barrier"]
@patch("vllm_ascend.model_loader.netloader.netloader.logger")
def test_draft_model_does_not_wait_for_target_netloader_barrier(mock_logger, monkeypatch):
dummy_model = _patch_loader_common(monkeypatch)
monkeypatch.setattr("vllm_ascend.model_loader.netloader.netloader.elastic_load", lambda **kwargs: dummy_model)
class DummyElasticServer:
def __init__(*a, **k):
pass
def start(self):
pass
barrier_calls = []
monkeypatch.setattr("vllm_ascend.model_loader.netloader.netloader.ElasticServer", DummyElasticServer)
monkeypatch.setattr("torch.distributed.is_available", lambda: True)
monkeypatch.setattr("torch.distributed.is_initialized", lambda: True)
monkeypatch.setattr("torch.distributed.barrier", lambda: barrier_calls.append("barrier"))
extra = {
"SOURCE": [{"device_id": 0, "sources": ["127.0.0.1:5000"]}],
"MODEL": "draft-model",
"LISTEN_PORT": 5555,
"INT8_CACHE": "dram",
}
loader = make_loader_with_config(extra)
vllm_config = DummyVllmConfig()
vllm_config.speculative_config = object()
loader.load_model(vllm_config, DummyDraftModelConfig())
assert barrier_calls == []
@patch("vllm_ascend.model_loader.netloader.netloader.logger")
def test_load_draft_model_elastic_success(mock_logger, monkeypatch, tmp_path):
dummy_model = _patch_loader_common(monkeypatch)
monkeypatch.setattr("vllm_ascend.model_loader.netloader.netloader.elastic_load", lambda **kwargs: dummy_model)
elastic_server_instances = []
class DummyElasticServer:
def __init__(self, *args, **kwargs):
elastic_server_instances.append(self)
self.int8_cache = args[7]
self.group_name = kwargs.get("group_name")
def start(self):
pass
monkeypatch.setattr("vllm_ascend.model_loader.netloader.netloader.ElasticServer", DummyElasticServer)
extra = {
"SOURCE": [{"device_id": 0, "sources": ["127.0.0.1:5000"]}],
"MODEL": "draft-model",
"LISTEN_PORT": 5555,
"OUTPUT_PREFIX": str(tmp_path) + "/output_",
"INT8_CACHE": "no",
}
loader = make_loader_with_config(extra)
result = loader.load_model(DummyVllmConfig(), DummyDraftModelConfig())
assert isinstance(result, nn.Module)
assert loader._draft_elastic_server is elastic_server_instances[0]
assert loader._draft_elastic_server.int8_cache == "no"
assert loader._draft_elastic_server.group_name == "netloader_draft"
assert not (tmp_path / "output_0.txt").exists()
@patch("vllm_ascend.model_loader.netloader.netloader.logger")
def test_load_draft_model_uses_hbm_when_int8_cache_is_dram(mock_logger, monkeypatch):
dummy_model = _patch_loader_common(monkeypatch)
monkeypatch.setattr("vllm_ascend.model_loader.netloader.netloader.elastic_load", lambda **kwargs: dummy_model)
captured = {}
class DummyElasticServer:
def __init__(self, *args, **kwargs):
captured["int8_cache"] = args[7]
captured["group_name"] = kwargs.get("group_name")
def start(self):
pass
monkeypatch.setattr("vllm_ascend.model_loader.netloader.netloader.ElasticServer", DummyElasticServer)
extra = {
"SOURCE": [{"device_id": 0, "sources": ["127.0.0.1:5000"]}],
"MODEL": "draft-model",
"LISTEN_PORT": 5555,
"INT8_CACHE": "dram",
}
loader = make_loader_with_config(extra)
loader.load_model(DummyVllmConfig(), DummyDraftModelConfig())
assert captured["int8_cache"] == "hbm"
assert captured["group_name"] == "netloader_draft"
@patch("vllm_ascend.model_loader.netloader.netloader.logger")
def test_load_draft_model_port_offset_and_group_name(mock_logger, monkeypatch, tmp_path):
dummy_model = _patch_loader_common(monkeypatch)
captured = {}
def capture_elastic_load(**kwargs):
captured.update(kwargs)
return dummy_model
monkeypatch.setattr("vllm_ascend.model_loader.netloader.netloader.elastic_load", capture_elastic_load)
class DummyElasticServer:
def __init__(*a, **k):
pass
def start(self):
pass
monkeypatch.setattr("vllm_ascend.model_loader.netloader.netloader.ElasticServer", DummyElasticServer)
extra = {
"SOURCE": [{"device_id": 0, "sources": ["127.0.0.1:5000", "10.0.0.1:6000"]}],
"MODEL": "draft-model",
"LISTEN_PORT": 5555,
"INT8_CACHE": "no",
}
loader = make_loader_with_config(extra)
loader.load_model(DummyVllmConfig(), DummyDraftModelConfig())
assert captured["group_name"] == "netloader_draft"
assert captured["model_path"] == "draft-model"
assert captured["sources"] == [
{
"device_id": 0,
"sources": [
f"127.0.0.1:{5000 + DRAFT_PORT_OFFSET}",
f"10.0.0.1:{6000 + DRAFT_PORT_OFFSET}",
],
}
]
assert loader.listen_port == 5555 + DRAFT_PORT_OFFSET
@patch("vllm_ascend.model_loader.netloader.netloader.logger")
def test_load_draft_model_skips_invalid_source_addresses(mock_logger, monkeypatch):
dummy_model = _patch_loader_common(monkeypatch)
captured = {}
def capture_elastic_load(**kwargs):
captured.update(kwargs)
return dummy_model
monkeypatch.setattr("vllm_ascend.model_loader.netloader.netloader.elastic_load", capture_elastic_load)
class DummyElasticServer:
def __init__(*a, **k):
pass
def start(self):
pass
monkeypatch.setattr("vllm_ascend.model_loader.netloader.netloader.ElasticServer", DummyElasticServer)
extra = {
"SOURCE": [{"device_id": 0, "sources": ["127.0.0.1:5000", "invalid", "10.0.0.1:not_port"]}],
"MODEL": "draft-model",
"LISTEN_PORT": 5555,
"INT8_CACHE": "no",
}
loader = make_loader_with_config(extra)
loader.load_model(DummyVllmConfig(), DummyDraftModelConfig())
assert captured["sources"] == [
{"device_id": 0, "sources": [f"127.0.0.1:{5000 + DRAFT_PORT_OFFSET}"]},
]
if __name__ == "__main__":
pytest.main()

View File

@@ -0,0 +1,437 @@
#
# Copyright (c) 2025 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.
#
import io
import json
import logging
import socket
from contextlib import contextmanager
from unittest.mock import MagicMock, patch
import pytest
import torch
from vllm_ascend.model_loader.netloader.interaction import elastic
from vllm_ascend.model_loader.netloader.interaction.elastic import ElasticClient, ElasticServer
# Simulate server's normal response
def mock_server_response(data):
return json.dumps({"label": "JOIN_ACK", "content": {"name": "mocked_name"}}).encode("utf-8")
# Simulate server's error response
def mock_server_error_response(data):
return json.dumps({"label": "JOIN_ACK", "content": None}).encode("utf-8")
# Simulated server's abnormal response
def mock_server_exception_response(data):
raise Exception("Mocked server exception")
@contextmanager
def capture_elastic_logs(level=logging.DEBUG):
log_capture_string = io.StringIO()
handler = logging.StreamHandler(log_capture_string)
handler.setLevel(level)
original_level = elastic.logger.level
elastic.logger.setLevel(level)
elastic.logger.addHandler(handler)
try:
yield log_capture_string
finally:
elastic.logger.removeHandler(handler)
elastic.logger.setLevel(original_level)
log_capture_string.close()
# Test the initialization of ElasticClient
def test_elastic_client_init():
sources = ["127.0.0.1:12345"]
device_id = 0
model_path = "mocked_model_path"
tp = 1
pp = 1
with patch("socket.socket") as mock_socket:
mock_socket_instance = MagicMock()
mock_socket.return_value = mock_socket_instance
mock_socket_instance.recv.return_value = mock_server_response(None)
mock_socket_instance.getsockname.return_value = ("127.0.0.1", 12346)
mock_socket_instance.__enter__.return_value = mock_socket_instance
with ElasticClient(sources, device_id, model_path, tp, pp) as client:
assert client.server_addr == "127.0.0.1"
assert client.server_port == 12345
assert client.ack == ("mocked_name", 12346)
mock_socket_instance.close.assert_called_once()
# Test the register method of ElasticClient
def test_elastic_client_register():
sources = ["127.0.0.1:12345"]
device_id = 0
model_path = "mocked_model_path"
tp = 1
pp = 1
with patch("socket.socket") as mock_socket:
mock_socket_instance = MagicMock()
mock_socket.return_value = mock_socket_instance
mock_socket_instance.connect.return_value = None
mock_socket_instance.recv.return_value = mock_server_response(None)
mock_socket_instance.getsockname.return_value = ("127.0.0.1", 12346)
mock_socket_instance.__enter__.return_value = mock_socket_instance
client = ElasticClient(sources, device_id, model_path, tp, pp)
assert client.register(device_id, model_path, tp, pp) == ("mocked_name", 12346)
# Test the behavior of the `register` method of ElasticClient when the server returns an error response.
def test_elastic_client_register_error_response():
sources = ["127.0.0.1:12345"]
device_id = 0
model_path = "mocked_model_path"
tp = 1
pp = 1
with patch("socket.socket") as mock_socket:
mock_socket_instance = MagicMock()
mock_socket.return_value = mock_socket_instance
mock_socket_instance.connect.return_value = None
mock_socket_instance.recv.return_value = mock_server_error_response(None)
with ElasticClient(sources, device_id, model_path, tp, pp) as client, pytest.raises(RuntimeError):
client.register(device_id, model_path, tp, pp)
mock_socket_instance.close.assert_called_once()
# Test the behavior of the `register` method of ElasticClient when an exception is thrown on the server.
def test_elastic_client_register_exception():
sources = ["127.0.0.1:12345"]
device_id = 0
model_path = "mocked_model_path"
tp = 1
pp = 1
with patch("socket.socket") as mock_socket:
mock_socket_instance = MagicMock()
mock_socket.return_value = mock_socket_instance
mock_socket_instance.connect.return_value = None
mock_socket_instance.recv.side_effect = mock_server_exception_response
mock_socket_instance.__enter__.return_value = mock_socket_instance
mock_socket_instance.__exit__.return_value = None
with ElasticClient(sources, device_id, model_path, tp, pp) as client, pytest.raises(RuntimeError):
client.register(device_id, model_path, tp, pp)
mock_socket_instance.close.assert_called_once()
class FakeInt8Param:
def __init__(self, name="param", device="npu", dtype=torch.int8):
self.dtype = dtype
self.device = torch.device(device)
@property
def data(self):
return self # Simulate .data returning self so .cpu() etc. can be chained
def clone(self):
return self
def detach(self):
return self
def cpu(self):
self.device = torch.device("cpu")
return self
class FakeModel:
def __init__(self):
self.params = {
"param1": MagicMock(dtype=torch.float32), # This will be ignored
"param2": FakeInt8Param(), # This simulates a real int8 param
}
def named_parameters(self):
return self.params.items()
@pytest.fixture
def mock_model():
return FakeModel()
@pytest.fixture
def server_config():
return {
"addr": "127.0.0.1",
"port": 8080,
"model": MagicMock(),
"device_id": 0,
"model_path": "/test/model",
"tp": 1,
"pp": 1,
"int8_cache": "dram",
"int8_cache_name": None,
}
# Test server initialization
def test_server_initialization(server_config, mock_model):
server_config["model"] = mock_model
with patch("socket.socket") as mock_socket, capture_elastic_logs() as log_capture_string:
server = ElasticServer(**server_config)
# Check the socket configuration
mock_socket.assert_called_with(socket.AF_INET, socket.SOCK_STREAM)
mock_socket.return_value.setsockopt.assert_called_with(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1)
mock_socket.return_value.bind.assert_called_with(("127.0.0.1", 8080))
mock_socket.return_value.listen.assert_called_with(256)
# Check int8 cache
assert "param2" in server.original_int8
assert server.original_int8["param2"].device.type == "cpu" # Verifying DRAM Cache
assert server.addr == server_config["addr"]
assert server.port == server_config["port"]
assert server.device_id == server_config["device_id"]
assert server.model_path == server_config["model_path"]
assert server.tp == server_config["tp"]
assert server.pp == server_config["pp"]
# Get captured logs
log_output = log_capture_string.getvalue()
# Check output
assert "Server 127.0.0.1:8080 starts" in log_output
# Test the int8 cache option
@pytest.mark.parametrize("cache_option,expected_device", [("dram", "cpu"), ("no", None), ("invalid", None)])
def test_int8_cache_handling(server_config, mock_model, cache_option, expected_device):
server_config["int8_cache"] = cache_option
server_config["model"] = mock_model
with patch("socket.socket"), capture_elastic_logs() as log_capture_string:
server = ElasticServer(**server_config)
log_output = log_capture_string.getvalue()
if cache_option == "invalid":
assert "int8_cache should be selected in [HBM, DRAM]" in log_output
if expected_device is None:
assert len(server.original_int8) == 0
else:
assert server.original_int8["param2"].device.type == expected_device
# Test client processing
def test_client_handler_valid_join(server_config, mock_model):
server_config["model"] = mock_model
with patch("vllm_ascend.model_loader.netloader.interaction.elastic.P2PSend") as mock_p2p_send:
# Create a simulated connection
mock_conn = MagicMock()
mock_addr = ("192.168.1.1", 12345)
# Configuring Client Data
valid_data = {
"label": "JOIN",
"content": {"device_id": 0, "model_path": "/test/model", "tp": 1, "pp": 1, "port": 9090},
}
mock_conn.recv.return_value = json.dumps(valid_data).encode("utf-8")
# Start the server
server = ElasticServer(**server_config)
server.register_handler(mock_conn, mock_addr)
# Verify response
expected_ack = {"label": "JOIN_ACK", "content": {"name": "192.168.1.1:12345"}}
mock_conn.send.assert_called_once_with(json.dumps(expected_ack).encode("utf-8"))
mock_p2p_send.assert_called_once_with("127.0.0.1", 9090, "192.168.1.1:12345", "netloader")
mock_conn.close.assert_called_once()
# Test mismatched JOIN requests
def test_client_handler_mismatch(server_config):
with patch("socket.socket"):
server = ElasticServer(**server_config)
mock_conn = MagicMock()
mock_addr = ("192.168.1.1", 12345)
# Send mismatched data
mismatch_data = {
"label": "JOIN",
"content": {
"device_id": 1, # 不匹配的ID
"model_path": "/wrong/model",
"tp": 2,
"pp": 2,
"port": 9090,
},
}
mock_conn.recv.return_value = json.dumps(mismatch_data).encode("utf-8")
server.register_handler(mock_conn, mock_addr)
assert isinstance(mismatch_data["content"], dict)
# Verify response
mismatch_tuple = (
mismatch_data["content"]["device_id"],
mismatch_data["content"]["model_path"],
mismatch_data["content"]["tp"],
mismatch_data["content"]["pp"],
)
server_tuple = (
server_config["device_id"],
server_config["model_path"],
server_config["tp"],
server_config["pp"],
)
expected_ack = {
"label": "JOIN_NACK",
"content": (f"Received data {mismatch_tuple} does not consist with this server {server_tuple}"),
}
mock_conn.send.assert_called_once_with(json.dumps(expected_ack).encode("utf-8"))
mock_conn.close.assert_called_once()
# Test Invalid Request
@pytest.mark.parametrize(
"invalid_data,should_send",
[
({"label": "WRONG_LABEL"}, True), # Incorrect label, can be decoded as JSON, but the content is invalid.
(
{"content": {"missing_fields": True}},
True,
), # Missing field, can be decoded as JSON, but the content is invalid.
("plain text", False), # Non-JSON data, json.loads failed
(b"invalid_bytes", False), # Invalid byte, decode or json.loads failed
],
)
def test_client_handler_invalid_requests(server_config, invalid_data, should_send):
with patch("socket.socket"), capture_elastic_logs() as log_capture_string:
server = ElasticServer(**server_config)
mock_conn = MagicMock()
mock_addr = ("192.168.1.1", 12345)
if isinstance(invalid_data, (str, bytes)):
mock_conn.recv.return_value = invalid_data if isinstance(invalid_data, bytes) else invalid_data.encode()
else:
mock_conn.recv.return_value = json.dumps(invalid_data).encode("utf-8")
server.register_handler(mock_conn, mock_addr)
if should_send:
expected_ack = {
"label": "JOIN_NACK",
"content": f"Received data does not contain required fields: {invalid_data}",
}
mock_conn.send.assert_called_once_with(json.dumps(expected_ack).encode("utf-8"))
else:
mock_conn.send.assert_not_called()
log_output = log_capture_string.getvalue()
# Any warning in the log is acceptable
assert "Failed to load" in log_output or "does not contain" in log_output
mock_conn.close.assert_called_once()
# Test the thread startup.
def test_server_start(server_config):
with patch("socket.socket"), patch("threading.Thread") as mock_thread:
handler_thread_instance = mock_thread.return_value
server = ElasticServer(**server_config)
server.start()
# Assert that the correct target parameter was passed when instantiating the Thread instance.
mock_thread.assert_called_once()
args, kwargs = mock_thread.call_args
assert kwargs["target"] == server.elastic_client_handler
# Verify the daemon attribute is set to True (the attribute value will be recorded after MagicMock assignment).
assert handler_thread_instance.daemon is True
# Check if the start() method is called.
handler_thread_instance.start.assert_called_once()
# Test resource clearing
def test_server_cleanup(server_config):
with patch("socket.socket") as mock_socket:
server = ElasticServer(**server_config)
del server
mock_socket.return_value.close.assert_called_once()
def test_draft_group_name_in_client_register():
sent_payloads = []
with (
patch("socket.socket") as mock_socket,
patch("vllm_ascend.model_loader.netloader.interaction.elastic.find_free_port", return_value=12346),
):
mock_socket_instance = MagicMock()
mock_socket.return_value = mock_socket_instance
mock_socket_instance.recv.return_value = mock_server_response(None)
mock_socket_instance.send.side_effect = lambda data: sent_payloads.append(json.loads(data.decode()))
client = ElasticClient(["127.0.0.1:12345"], 0, "draft-model", 1, 1, "netloader_draft")
client.s = mock_socket_instance
client.register(0, "draft-model", 1, 1)
assert client.group_name == "netloader_draft"
assert sent_payloads[0]["content"]["group_name"] == "netloader_draft"
def test_draft_group_name_in_server_p2p_send(server_config, mock_model):
server_config["model"] = mock_model
join_data = {
"label": "JOIN",
"content": {
"device_id": 0,
"model_path": "/test/model",
"tp": 1,
"pp": 1,
"port": 9090,
"group_name": "netloader_draft",
},
}
with (
patch("socket.socket"),
patch("vllm_ascend.model_loader.netloader.interaction.elastic.P2PSend") as mock_p2p_send,
):
mock_conn = MagicMock()
mock_conn.recv.return_value = json.dumps(join_data).encode("utf-8")
ElasticServer(**server_config, group_name="netloader_draft").register_handler(mock_conn, ("192.168.1.1", 12345))
mock_p2p_send.assert_called_once_with("127.0.0.1", 9090, "192.168.1.1:12345", "netloader_draft")
if __name__ == "__main__":
pytest.main()

View File

@@ -0,0 +1,129 @@
#
# Copyright (c) 2025 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.
#
from unittest.mock import MagicMock, patch
import pytest
from vllm_ascend.model_loader.netloader.load import elastic_load
@pytest.fixture
def mock_sources():
return [
{"device_id": 0, "sources": ["a", "b"]},
{"device_id": 1, "sources": ["c"]},
]
@patch("vllm_ascend.model_loader.netloader.interaction.elastic.ElasticClient")
@patch("vllm_ascend.model_loader.netloader.executor.elastic_load.P2PLoad")
def test_sources_this_device_empty(mock_p2p, mock_client):
sources = [{"device_id": 1, "sources": ["c"]}]
result = elastic_load("model", 0, "model_path", sources, 1, 1)
assert result is None
mock_client.assert_not_called()
mock_p2p.assert_not_called()
@patch("vllm_ascend.model_loader.netloader.interaction.elastic.ElasticClient")
@patch("vllm_ascend.model_loader.netloader.executor.elastic_load.P2PLoad")
def test_client_s_none(mock_p2p, mock_client, mock_sources):
# Simulate ElasticClient.s as None
mock_instance = MagicMock()
mock_instance.s = None
mock_client.return_value = mock_instance
result = elastic_load("model", 0, "model_path", mock_sources, 1, 1)
assert result is None
@patch("vllm_ascend.model_loader.netloader.interaction.elastic.ElasticClient")
@patch("vllm_ascend.model_loader.netloader.executor.elastic_load.P2PLoad")
def test_client_ack_none(mock_p2p, mock_client, mock_sources):
# Simulate ElasticClient.ack as None
mock_instance = MagicMock()
mock_instance.s = True
mock_instance.ack = None
mock_client.return_value = mock_instance
result = elastic_load("model", 0, "model_path", mock_sources, 1, 1)
assert result is None
@patch("vllm_ascend.model_loader.netloader.load.P2PLoad")
@patch("vllm_ascend.model_loader.netloader.load.logger")
def test_model_load_fail(mock_logger, mock_p2p):
mock_client = MagicMock()
mock_client.s = True
mock_client.ack = ["foo", "bar"]
mock_client.server_addr = "addr"
with patch("vllm_ascend.model_loader.netloader.load.ElasticClient", return_value=mock_client):
# P2PLoad.load returns None
mock_p2p_instance = MagicMock()
mock_p2p_instance.load.return_value = None
mock_p2p.return_value = mock_p2p_instance
sources = [{"device_id": 0, "sources": ["whatever"]}]
result = elastic_load("model", 0, "model_path", sources, 1, 1)
assert result is None
mock_logger.error.assert_called_once()
@patch("vllm_ascend.model_loader.netloader.load.P2PLoad")
@patch("vllm_ascend.model_loader.netloader.load.logger")
def test_model_load_success(mock_logger, mock_p2p):
mock_client = MagicMock()
mock_client.s = True
mock_client.ack = ["foo", "bar"]
mock_client.server_addr = "addr"
with patch("vllm_ascend.model_loader.netloader.load.ElasticClient", return_value=mock_client):
expected_model = object()
mock_p2p_instance = MagicMock()
mock_p2p_instance.load.return_value = expected_model
mock_p2p.return_value = mock_p2p_instance
sources = [{"device_id": 0, "sources": ["whatever"]}]
result = elastic_load("model", 0, "model_path", sources, 1, 1)
assert result is expected_model
mock_logger.info.assert_called_once()
@patch("vllm_ascend.model_loader.netloader.load.P2PLoad")
@patch("vllm_ascend.model_loader.netloader.load.ElasticClient")
def test_elastic_load_passes_draft_group_name(mock_client, mock_p2p):
mock_client_instance = MagicMock()
mock_client_instance.s = True
mock_client_instance.ack = ["foo", "bar"]
mock_client_instance.server_addr = "addr"
mock_client_instance.__enter__.return_value = mock_client_instance
mock_client.return_value = mock_client_instance
expected_model = object()
mock_p2p_instance = MagicMock()
mock_p2p_instance.load.return_value = expected_model
mock_p2p.return_value = mock_p2p_instance
sources = [{"device_id": 0, "sources": ["127.0.0.1:15000"]}]
result = elastic_load("model", 0, "draft-model", sources, 1, 1, group_name="netloader_draft")
assert result is expected_model
mock_client.assert_called_once_with(["127.0.0.1:15000"], 0, "draft-model", 1, 1, "netloader_draft")
mock_p2p.assert_called_once_with("foo", "addr", "bar", "netloader_draft")
if __name__ == "__main__":
pytest.main()

View File

@@ -0,0 +1,60 @@
#
# Copyright (c) 2025 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.
#
import os
import tempfile
from unittest.mock import patch
import pytest
from vllm_ascend.model_loader.netloader.utils import find_free_port, is_valid_path_prefix
def test_find_free_port():
port = find_free_port()
assert isinstance(port, int)
assert port > 0
def test_is_valid_path_prefix_empty():
assert not is_valid_path_prefix("")
def test_is_valid_path_prefixIllegal_characters():
assert not is_valid_path_prefix('test<>:"|?*')
def test_is_valid_path_prefixRelative_path():
assert is_valid_path_prefix("test")
def test_is_valid_path_prefixAbsolute_path():
with tempfile.TemporaryDirectory() as tmpdir:
assert is_valid_path_prefix(os.path.join(tmpdir, "test"))
@patch("os.path.exists", return_value=False)
def test_is_valid_path_prefix_no_directory(mock_exists):
assert not is_valid_path_prefix("/nonexistent_dir/test")
@patch("os.path.exists", return_value=True)
def test_is_valid_path_prefix_directory_exists(mock_exists):
assert is_valid_path_prefix("/existing_dir/test")
if __name__ == "__main__":
pytest.main()