[Utils] Add NetworkAddress abstraction for IPv6-safe address handling (#20306)
This commit is contained in:
@@ -4,11 +4,17 @@ import unittest
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
from sglang.srt.server_args import PortArgs, ServerArgs, prepare_server_args
|
||||
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
|
||||
from sglang.test.test_utils import CustomTestCase
|
||||
from sglang.test.ci.ci_register import register_cpu_ci
|
||||
from sglang.test.test_utils import (
|
||||
DEFAULT_SMALL_MODEL_NAME_FOR_TEST_QWEN,
|
||||
CustomTestCase,
|
||||
)
|
||||
|
||||
register_cuda_ci(est_time=9, suite="stage-b-test-small-1-gpu")
|
||||
register_amd_ci(est_time=1, suite="stage-b-test-small-1-gpu-amd")
|
||||
register_cpu_ci(est_time=1, suite="stage-a-cpu-only")
|
||||
|
||||
# Mock get_device() so all tests run on CPU-only CI runners
|
||||
_mock_device = patch("sglang.srt.server_args.get_device", return_value="cuda")
|
||||
_mock_device.start()
|
||||
|
||||
|
||||
class TestPrepareServerArgs(CustomTestCase):
|
||||
@@ -16,14 +22,12 @@ class TestPrepareServerArgs(CustomTestCase):
|
||||
server_args = prepare_server_args(
|
||||
[
|
||||
"--model-path",
|
||||
"meta-llama/Meta-Llama-3.1-8B-Instruct",
|
||||
DEFAULT_SMALL_MODEL_NAME_FOR_TEST_QWEN,
|
||||
"--json-model-override-args",
|
||||
'{"rope_scaling": {"factor": 2.0, "rope_type": "linear"}}',
|
||||
]
|
||||
)
|
||||
self.assertEqual(
|
||||
server_args.model_path, "meta-llama/Meta-Llama-3.1-8B-Instruct"
|
||||
)
|
||||
self.assertEqual(server_args.model_path, DEFAULT_SMALL_MODEL_NAME_FOR_TEST_QWEN)
|
||||
self.assertEqual(
|
||||
json.loads(server_args.json_model_override_args),
|
||||
{"rope_scaling": {"factor": 2.0, "rope_type": "linear"}},
|
||||
@@ -144,12 +148,10 @@ class TestPortArgs(unittest.TestCase):
|
||||
server_args.nnodes = 2
|
||||
server_args.dist_init_addr = "192.168.1.1"
|
||||
|
||||
with self.assertRaises(AssertionError) as context:
|
||||
with self.assertRaises(ValueError) as context:
|
||||
PortArgs.init_new(server_args)
|
||||
|
||||
self.assertIn(
|
||||
"please provide --dist-init-addr as host:port", str(context.exception)
|
||||
)
|
||||
self.assertIn("Missing port", str(context.exception))
|
||||
|
||||
def test_init_new_with_malformed_ipv4_address_invalid_port(self):
|
||||
server_args = ServerArgs(model_path="dummy")
|
||||
@@ -163,109 +165,6 @@ class TestPortArgs(unittest.TestCase):
|
||||
with self.assertRaises(ValueError):
|
||||
PortArgs.init_new(server_args)
|
||||
|
||||
@patch("sglang.srt.server_args.is_valid_ipv6_address", return_value=True)
|
||||
def test_init_new_with_ipv6_address(self, mock_is_valid_ipv6):
|
||||
server_args = ServerArgs(model_path="dummy")
|
||||
server_args.port = 30000
|
||||
server_args.nccl_port = None
|
||||
|
||||
server_args.enable_dp_attention = True
|
||||
server_args.nnodes = 2
|
||||
server_args.dist_init_addr = "[2001:db8::1]:25000"
|
||||
|
||||
port_args = PortArgs.init_new(server_args)
|
||||
|
||||
self.assertTrue(port_args.tokenizer_ipc_name.startswith("tcp://[2001:db8::1]:"))
|
||||
self.assertTrue(
|
||||
port_args.scheduler_input_ipc_name.startswith("tcp://[2001:db8::1]:")
|
||||
)
|
||||
self.assertTrue(
|
||||
port_args.detokenizer_ipc_name.startswith("tcp://[2001:db8::1]:")
|
||||
)
|
||||
self.assertIsInstance(port_args.nccl_port, int)
|
||||
|
||||
@patch("sglang.srt.server_args.is_valid_ipv6_address", return_value=False)
|
||||
def test_init_new_with_invalid_ipv6_address(self, mock_is_valid_ipv6):
|
||||
server_args = ServerArgs(model_path="dummy")
|
||||
server_args.port = 30000
|
||||
server_args.nccl_port = None
|
||||
|
||||
server_args.enable_dp_attention = True
|
||||
server_args.nnodes = 2
|
||||
server_args.dist_init_addr = "[invalid-ipv6]:25000"
|
||||
|
||||
with self.assertRaises(ValueError) as context:
|
||||
PortArgs.init_new(server_args)
|
||||
|
||||
self.assertIn("invalid IPv6 address", str(context.exception))
|
||||
|
||||
def test_init_new_with_malformed_ipv6_address_missing_bracket(self):
|
||||
server_args = ServerArgs(model_path="dummy")
|
||||
server_args.port = 30000
|
||||
server_args.nccl_port = None
|
||||
|
||||
server_args.enable_dp_attention = True
|
||||
server_args.nnodes = 2
|
||||
server_args.dist_init_addr = "[2001:db8::1:25000"
|
||||
|
||||
with self.assertRaises(ValueError) as context:
|
||||
PortArgs.init_new(server_args)
|
||||
|
||||
self.assertIn("invalid IPv6 address format", str(context.exception))
|
||||
|
||||
@patch("sglang.srt.server_args.is_valid_ipv6_address", return_value=True)
|
||||
def test_init_new_with_malformed_ipv6_address_missing_port(
|
||||
self, mock_is_valid_ipv6
|
||||
):
|
||||
server_args = ServerArgs(model_path="dummy")
|
||||
server_args.port = 30000
|
||||
server_args.nccl_port = None
|
||||
|
||||
server_args.enable_dp_attention = True
|
||||
server_args.nnodes = 2
|
||||
server_args.dist_init_addr = "[2001:db8::1]"
|
||||
|
||||
with self.assertRaises(ValueError) as context:
|
||||
PortArgs.init_new(server_args)
|
||||
|
||||
self.assertIn(
|
||||
"a port must be specified in IPv6 address", str(context.exception)
|
||||
)
|
||||
|
||||
@patch("sglang.srt.server_args.is_valid_ipv6_address", return_value=True)
|
||||
def test_init_new_with_malformed_ipv6_address_invalid_port(
|
||||
self, mock_is_valid_ipv6
|
||||
):
|
||||
server_args = ServerArgs(model_path="dummy")
|
||||
server_args.port = 30000
|
||||
server_args.nccl_port = None
|
||||
|
||||
server_args.enable_dp_attention = True
|
||||
server_args.nnodes = 2
|
||||
server_args.dist_init_addr = "[2001:db8::1]:abcde"
|
||||
|
||||
with self.assertRaises(ValueError) as context:
|
||||
PortArgs.init_new(server_args)
|
||||
|
||||
self.assertIn("invalid port in IPv6 address", str(context.exception))
|
||||
|
||||
@patch("sglang.srt.server_args.is_valid_ipv6_address", return_value=True)
|
||||
def test_init_new_with_malformed_ipv6_address_wrong_separator(
|
||||
self, mock_is_valid_ipv6
|
||||
):
|
||||
server_args = ServerArgs(model_path="dummy")
|
||||
server_args.port = 30000
|
||||
server_args.nccl_port = None
|
||||
|
||||
server_args.enable_dp_attention = True
|
||||
server_args.nnodes = 2
|
||||
server_args.dist_init_addr = "[2001:db8::1]#25000"
|
||||
|
||||
with self.assertRaises(ValueError) as context:
|
||||
PortArgs.init_new(server_args)
|
||||
|
||||
self.assertIn("expected ':' after ']'", str(context.exception))
|
||||
|
||||
|
||||
class TestSSLArgs(unittest.TestCase):
|
||||
def test_default_ssl_fields_are_none(self):
|
||||
@@ -305,10 +204,6 @@ class TestSSLArgs(unittest.TestCase):
|
||||
server_args = ServerArgs(model_path="dummy", host="")
|
||||
self.assertEqual(server_args.url(), "http://127.0.0.1:30000")
|
||||
|
||||
def test_url_rewrites_ipv6_all_interfaces_to_loopback(self):
|
||||
server_args = ServerArgs(model_path="dummy", host="::")
|
||||
self.assertEqual(server_args.url(), "http://[::1]:30000")
|
||||
|
||||
@patch("os.path.isfile", return_value=True)
|
||||
def test_url_returns_https_with_ssl(self, _mock_isfile):
|
||||
server_args = ServerArgs(
|
||||
@@ -401,16 +296,6 @@ class TestSSLArgs(unittest.TestCase):
|
||||
"SSL CA certificates file not found", str(context.exception)
|
||||
)
|
||||
|
||||
@patch("os.path.isfile", return_value=True)
|
||||
def test_url_returns_https_with_ssl_and_ipv6(self, _mock_isfile):
|
||||
server_args = ServerArgs(
|
||||
model_path="dummy",
|
||||
host="::1",
|
||||
ssl_keyfile="key.pem",
|
||||
ssl_certfile="cert.pem",
|
||||
)
|
||||
self.assertEqual(server_args.url(), "https://[::1]:30000")
|
||||
|
||||
def test_enable_ssl_refresh_default_false(self):
|
||||
server_args = ServerArgs(model_path="dummy")
|
||||
self.assertFalse(server_args.enable_ssl_refresh)
|
||||
|
||||
Reference in New Issue
Block a user