Files
sglang/test/registered/unit/server_args/test_server_args.py
T

1502 lines
59 KiB
Python

import importlib
import json
import os
import socket
import tempfile
import unittest
from types import SimpleNamespace
from unittest.mock import MagicMock, patch
import sglang.srt.server_args as server_args_module
from sglang.srt.arg_groups.speculative_hook import handle_speculative_decoding
from sglang.srt.environ import envs
from sglang.srt.layers.cp.base import is_cp_enabled, is_interleave
from sglang.srt.model_executor.cuda_graph_config import (
Backend,
CudaGraphConfig,
Phase,
PhaseConfig,
)
from sglang.srt.server_args import PortArgs, ServerArgs, prepare_server_args
from sglang.srt.server_args_config_parser import ConfigArgumentMerger
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_cpu_ci(est_time=10, suite="base-a-test-cpu")
register_cpu_ci(est_time=12, suite="base-c-test-cpu")
# 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):
def test_config_nested_dict_args_are_json(self):
with tempfile.NamedTemporaryFile(mode="w", suffix=".yaml", delete=False) as f:
f.write("mm-process-config:\n image:\n resize: 128\n")
config_file = f.name
try:
parser = server_args_module.argparse.ArgumentParser()
ServerArgs.add_cli_args(parser)
merged = ConfigArgumentMerger(parser).merge_config_with_args(
[
"--config",
config_file,
"--model-path",
DEFAULT_SMALL_MODEL_NAME_FOR_TEST_QWEN,
]
)
value = merged[merged.index("--mm-process-config") + 1]
parsed = parser.parse_args(merged)
self.assertEqual(json.loads(value), {"image": {"resize": 128}})
self.assertEqual(parsed.mm_process_config, {"image": {"resize": 128}})
finally:
os.unlink(config_file)
class TestMambaCacheStochasticRounding(unittest.TestCase):
def test_rejects_fp32_ssm_cache(self):
server_args = ServerArgs(
model_path="dummy",
mamba_ssm_dtype="float32",
enable_mamba_cache_stochastic_rounding=True,
)
with self.assertRaisesRegex(ValueError, "--mamba-ssm-dtype float16"):
server_args._handle_mamba_backend()
@patch("sglang.srt.server_args.is_cuda", return_value=False)
def test_rejects_non_cuda(self, _mock_is_cuda):
server_args = ServerArgs(
model_path="dummy",
mamba_ssm_dtype="float16",
enable_mamba_cache_stochastic_rounding=True,
)
with self.assertRaisesRegex(ValueError, "NVIDIA CUDA"):
server_args._handle_mamba_backend()
@patch("sglang.srt.server_args.is_cuda", return_value=True)
@patch("sglang.srt.server_args.is_sm100_supported", return_value=False)
def test_rejects_triton_without_sm100(self, _mock_sm100, _mock_is_cuda):
server_args = ServerArgs(
model_path="dummy",
mamba_ssm_dtype="float16",
mamba_backend="triton",
enable_mamba_cache_stochastic_rounding=True,
)
with self.assertRaisesRegex(ValueError, "requires SM100"):
server_args._handle_mamba_backend()
class TestLoadBalanceMethod(unittest.TestCase):
def _load_balance_args(self, **kwargs):
server_args = ServerArgs(model_path="dummy", **kwargs)
server_args._handle_pd_disaggregation()
server_args._handle_load_balance_method()
return server_args
def test_non_pd_defaults_to_round_robin(self):
server_args = self._load_balance_args(disaggregation_mode="null")
self.assertEqual(server_args.load_balance_method, "round_robin")
def test_pd_prefill_defaults_to_follow_bootstrap_room(self):
server_args = self._load_balance_args(disaggregation_mode="prefill")
self.assertEqual(server_args.load_balance_method, "follow_bootstrap_room")
def test_pd_decode_defaults_to_round_robin(self):
server_args = self._load_balance_args(disaggregation_mode="decode")
self.assertEqual(server_args.load_balance_method, "round_robin")
def test_pd_decode_radix_cache_rejects_hisparse(self):
server_args = ServerArgs(
model_path="dummy",
disaggregation_mode="decode",
disaggregation_decode_enable_radix_cache=True,
disaggregation_transfer_backend="nixl",
enable_hisparse=True,
)
with self.assertRaises(ValueError) as context:
server_args._handle_pd_disaggregation()
self.assertIn(
"--disaggregation-decode-enable-radix-cache is incompatible with "
"--enable-hisparse",
str(context.exception),
)
def test_pd_decode_radix_cache_rejects_fake_backend(self):
server_args = ServerArgs(
model_path="dummy",
disaggregation_mode="decode",
disaggregation_decode_enable_radix_cache=True,
disaggregation_transfer_backend="fake",
)
with self.assertRaises(ValueError) as context:
server_args._handle_pd_disaggregation()
self.assertIn(
"--disaggregation-decode-enable-radix-cache is incompatible "
"with --disaggregation-transfer-backend fake",
str(context.exception),
)
def test_pd_decode_radix_cache_allows_mooncake_tcp(self):
server_args = self._load_balance_args(
disaggregation_mode="decode",
disaggregation_decode_enable_radix_cache=True,
disaggregation_transfer_backend="mooncake_tcp",
)
self.assertFalse(server_args.disable_radix_cache)
self.assertEqual(server_args.disaggregation_transfer_backend, "mooncake")
class TestHiSparseDsaBackendPolicy(unittest.TestCase):
# The backend selection moved to the resolution pipeline; these policy
# tests drive the pass through its read-only view.
@staticmethod
def _resolve(kv_cache_dtype, **kw):
from types import SimpleNamespace
from sglang.srt.arg_groups.overrides import (
ResolvedView,
_dsa_split_backend_resolution,
)
hf = SimpleNamespace(architectures=["DeepseekV32ForCausalLM"])
defaults = dict(
kv_cache_dtype=kv_cache_dtype,
dsa_prefill_backend=None,
dsa_decode_backend=None,
enable_hisparse=True,
)
defaults.update(kw)
view = ResolvedView(
SimpleNamespace(
get_model_config=lambda: SimpleNamespace(hf_config=hf), **defaults
)
)
with (
patch("sglang.srt.configs.model_config.is_deepseek_dsa", return_value=True),
patch("sglang.srt.arg_groups.overrides.is_npu", return_value=False),
patch("sglang.srt.arg_groups.overrides.is_xpu", return_value=False),
patch("torch.cuda.get_device_capability", return_value=(9, 0)),
):
declared = _dsa_split_backend_resolution(view)
return {
"dsa_prefill_backend": declared.get(
"dsa_prefill_backend", defaults["dsa_prefill_backend"]
),
"dsa_decode_backend": declared.get(
"dsa_decode_backend", defaults["dsa_decode_backend"]
),
}
@patch("sglang.srt.server_args.is_hip", return_value=False)
def test_hisparse_defaults_to_flashmla_sparse_on_cuda_bfloat16(self, _mock_is_hip):
resolved = self._resolve("bfloat16")
self.assertEqual(resolved["dsa_prefill_backend"], "flashmla_sparse")
self.assertEqual(resolved["dsa_decode_backend"], "flashmla_sparse")
@patch("sglang.srt.server_args.is_hip", return_value=False)
def test_hisparse_defaults_to_flashmla_kv_on_cuda_fp8(self, _mock_is_hip):
resolved = self._resolve("fp8_e4m3")
self.assertEqual(resolved["dsa_prefill_backend"], "flashmla_kv")
self.assertEqual(resolved["dsa_decode_backend"], "flashmla_kv")
@patch("sglang.srt.server_args.is_hip", return_value=True)
def test_hisparse_defaults_to_tilelang_on_rocm(self, _mock_is_hip):
resolved = self._resolve("bfloat16")
self.assertEqual(resolved["dsa_prefill_backend"], "tilelang")
self.assertEqual(resolved["dsa_decode_backend"], "tilelang")
@patch("sglang.srt.server_args.is_hip", return_value=True)
def test_hisparse_preserves_rocm_user_backend_and_defaults_missing_side(
self, _mock_is_hip
):
resolved = self._resolve("bfloat16", dsa_prefill_backend="tilelang")
self.assertEqual(resolved["dsa_prefill_backend"], "tilelang")
self.assertEqual(resolved["dsa_decode_backend"], "tilelang")
@patch("sglang.srt.server_args.is_hip", return_value=True)
def test_hisparse_accepts_aiter_backend_on_rocm(self, _mock_is_hip):
server_args = ServerArgs(
model_path="dummy",
enable_hisparse=True,
kv_cache_dtype="bfloat16",
dsa_prefill_backend="aiter",
dsa_decode_backend="aiter",
)
server_args._validate_hisparse_dsa_backend("dsa_prefill_backend", "prefill")
server_args._validate_hisparse_dsa_backend("dsa_decode_backend", "decode")
@patch("sglang.srt.server_args.is_hip", return_value=True)
def test_hisparse_rejects_cuda_backend_on_rocm(self, _mock_is_hip):
server_args = ServerArgs(
model_path="dummy",
enable_hisparse=True,
kv_cache_dtype="bfloat16",
dsa_prefill_backend="flashmla_sparse",
)
with self.assertRaisesRegex(ValueError, "tilelang"):
server_args._validate_hisparse_dsa_backend("dsa_prefill_backend", "prefill")
@patch("sglang.srt.server_args.is_hip", return_value=False)
def test_hisparse_rejects_rocm_backend_on_cuda(self, _mock_is_hip):
server_args = ServerArgs(
model_path="dummy",
enable_hisparse=True,
kv_cache_dtype="bfloat16",
dsa_decode_backend="tilelang",
)
with self.assertRaisesRegex(ValueError, "flashmla_sparse"):
server_args._validate_hisparse_dsa_backend("dsa_decode_backend", "decode")
def test_hisparse_accepts_bfloat16_kv_cache_dtype(self):
server_args = ServerArgs(
model_path="dummy",
enable_hisparse=True,
kv_cache_dtype="bfloat16",
)
server_args._validate_hisparse_kv_cache_dtype()
def test_hisparse_accepts_fp8_e4m3_kv_cache_dtype(self):
server_args = ServerArgs(
model_path="dummy",
enable_hisparse=True,
kv_cache_dtype="fp8_e4m3",
)
server_args._validate_hisparse_kv_cache_dtype()
def test_hisparse_rejects_unsupported_kv_cache_dtype(self):
server_args = ServerArgs(
model_path="dummy",
enable_hisparse=True,
kv_cache_dtype="float16",
)
with self.assertRaisesRegex(ValueError, r"fp8_e4m3"):
server_args._validate_hisparse_kv_cache_dtype()
class TestFa4PageSizeAutoForce(CustomTestCase):
"""FA4 requires page_size 128 for non-MLA models on SM100. The auto-force
must trigger for `--attention-backend fa4` (combined) too, not only for the
explicit `--prefill-attention-backend fa4` path."""
def _make_args(self, attention_backend, prefill=None, decode=None, page_size=1):
args = ServerArgs(model_path="dummy")
args.attention_backend = attention_backend
args.prefill_attention_backend = prefill
args.decode_attention_backend = decode
args.page_size = page_size
# Short-circuit get_model_config(): the fa4 page_size branch only needs
# use_mla_backend() (mocked) and is_sm100_supported() (mocked), not a
# real model_config. Pre-set the attribute so get_model_config returns
# early without touching ModelConfig.from_server_args.
args.model_config = MagicMock()
args.model_config.hf_config.dual_chunk_attention_config = None
return args
@patch("sglang.srt.arg_groups.overrides.is_sm100_supported", return_value=True)
@patch("sglang.srt.server_args.ServerArgs.use_mla_backend", return_value=False)
def test_combined_attention_backend_fa4_forces_page_size_128(
self, _mock_mla, _mock_sm100
):
# `--attention-backend fa4` (combined): prefill/decode fields stay None.
args = self._make_args(attention_backend="fa4")
args._handle_attention_backend_compatibility()
from sglang.srt.arg_groups.overrides import resolved_view
self.assertEqual(args.page_size, 1) # dual-apply retired: pristine
self.assertEqual(resolved_view(args).page_size, 128)
@patch("sglang.srt.arg_groups.overrides.is_sm100_supported", return_value=True)
@patch("sglang.srt.server_args.ServerArgs.use_mla_backend", return_value=False)
def test_explicit_prefill_fa4_forces_page_size_128(self, _mock_mla, _mock_sm100):
# `--prefill-attention-backend fa4`: the previously-covered path.
args = self._make_args(attention_backend=None, prefill="fa4", page_size=1)
args._handle_attention_backend_compatibility()
from sglang.srt.arg_groups.overrides import resolved_view
self.assertEqual(args.page_size, 1) # dual-apply retired: pristine
self.assertEqual(resolved_view(args).page_size, 128)
class TestContextParallelServerArgs(CustomTestCase):
def setUp(self):
self.parser = server_args_module.argparse.ArgumentParser()
ServerArgs.add_cli_args(self.parser)
def _new_cp_args(self, **overrides):
server_args = object.__new__(ServerArgs)
defaults = dict(
enable_prefill_context_parallel=False,
enable_dsa_prefill_context_parallel=False,
enable_prefill_cp=False,
cp_strategy=None,
model_path="instance://127.0.0.1:8000/dummy",
dsa_prefill_cp_mode="round-robin-split",
prefill_cp_mode="in-seq-split",
attn_cp_size=1,
tp_size=1,
dp_size=1,
moe_dp_size=1,
ep_size=1,
pp_size=1,
enable_aiter_allreduce_fusion=False,
)
defaults.update(overrides)
for key, value in defaults.items():
setattr(server_args, key, value)
return server_args
def test_canonical_prefill_cp_requires_strategy(self):
args = self.parser.parse_args(["--model", "dummy", "--enable-prefill-cp"])
self.assertTrue(args.enable_prefill_cp)
self.assertIsNone(args.cp_strategy)
server_args = self._new_cp_args(
enable_prefill_cp=args.enable_prefill_cp,
cp_strategy=args.cp_strategy,
)
with self.assertRaisesRegex(ValueError, "--cp-strategy"):
server_args._handle_context_parallelism()
def test_deprecated_dsa_cp_mode_maps_to_unified_strategy(self):
args = self.parser.parse_args(
[
"--model",
"dummy",
"--enable-dsa-prefill-context-parallel",
"--dsa-prefill-cp-mode",
"round-robin-split",
]
)
server_args = self._new_cp_args(
enable_dsa_prefill_context_parallel=(
args.enable_dsa_prefill_context_parallel
),
dsa_prefill_cp_mode=args.dsa_prefill_cp_mode,
)
server_args._handle_legacy_cp_arguments()
self.assertTrue(server_args.enable_prefill_cp)
self.assertEqual(server_args.cp_strategy, "interleave")
self.assertEqual(server_args.dsa_prefill_cp_mode, "round-robin-split")
def test_canonical_interleave_cp_mirrors_to_dsa_runtime_aliases(self):
server_args = self._new_cp_args(
enable_prefill_cp=True,
cp_strategy="interleave",
attention_backend="dsa",
)
server_args._handle_legacy_cp_arguments()
server_args._handle_context_parallelism()
self.assertTrue(server_args.enable_dsa_prefill_context_parallel)
self.assertFalse(server_args.enable_prefill_context_parallel)
self.assertEqual(server_args.dsa_prefill_cp_mode, "round-robin-split")
self.assertEqual(server_args.prefill_cp_mode, "round-robin-split")
def test_context_parallel_handler_initializes_cp_strategy(self):
server_args = self._new_cp_args(
enable_prefill_cp=True,
cp_strategy="interleave",
attn_cp_size=2,
tp_size=2,
)
server_args._handle_context_parallelism()
self.assertTrue(is_cp_enabled())
self.assertTrue(is_interleave())
def test_registered_cp_legacy_args_map_to_unified_strategy(self):
cases = [
(
"deepseek_v3_mla_cp",
dict(enable_prefill_context_parallel=True),
"zigzag",
"in-seq-split",
False,
True,
),
(
"qwen3_gqa_cp",
dict(
enable_prefill_context_parallel=True,
tp_size=4,
attn_cp_size=2,
),
"zigzag",
"in-seq-split",
False,
True,
),
(
"deepseek_v32_dsa_in_seq_split",
dict(
enable_dsa_prefill_context_parallel=True,
dsa_prefill_cp_mode="in-seq-split",
tp_size=8,
dp_size=2,
attn_cp_size=4,
),
"zigzag",
"in-seq-split",
True,
False,
),
(
"deepseek_v32_dsa_round_robin_split",
dict(
enable_dsa_prefill_context_parallel=True,
tp_size=8,
attn_cp_size=8,
),
"interleave",
"round-robin-split",
True,
False,
),
(
"deepseek_v4_flash_fp4_b200_dsa_round_robin_split",
dict(
enable_dsa_prefill_context_parallel=True,
dsa_prefill_cp_mode="round-robin-split",
tp_size=4,
attn_cp_size=4,
),
"interleave",
"round-robin-split",
True,
False,
),
]
for name, overrides, strategy, mode, expect_dsa, expect_generic in cases:
with self.subTest(name=name):
server_args = self._new_cp_args(**overrides)
server_args._handle_legacy_cp_arguments()
server_args._handle_context_parallelism()
self.assertTrue(server_args.enable_prefill_cp)
self.assertEqual(server_args.cp_strategy, strategy)
self.assertEqual(server_args.dsa_prefill_cp_mode, mode)
self.assertEqual(server_args.prefill_cp_mode, mode)
self.assertEqual(
server_args.enable_dsa_prefill_context_parallel, expect_dsa
)
self.assertEqual(
server_args.enable_prefill_context_parallel, expect_generic
)
class TestPortArgs(unittest.TestCase):
@patch("sglang.srt.server_args.tempfile.NamedTemporaryFile")
def test_init_new_standard_case(self, mock_temp_file):
mock_temp_file.return_value.name = "temp_file"
server_args = ServerArgs(model_path="dummy")
server_args.port = 30000
server_args.nccl_port = None
server_args.enable_dp_attention = False
port_args = PortArgs.init_new(server_args)
self.assertTrue(port_args.tokenizer_ipc_name.startswith("ipc://"))
self.assertTrue(port_args.scheduler_input_ipc_name.startswith("ipc://"))
self.assertTrue(port_args.detokenizer_ipc_name.startswith("ipc://"))
self.assertIsInstance(port_args.nccl_port, int)
@patch("sglang.srt.server_args.tempfile.NamedTemporaryFile")
def test_init_new_builds_decoupled_spec_ipc_config(self, mock_temp_file):
mock_temp_file.return_value.name = "temp_file"
server_args = ServerArgs(model_path="dummy")
server_args.nccl_port = None
server_args.enable_dp_attention = False
server_args.decoupled_spec_role = "verifier"
server_args.decoupled_spec_bind_endpoint = "ipc:///tmp/v"
server_args.decoupled_spec_connect_endpoints = ["ipc:///tmp/d"]
server_args.decoupled_spec_rank = 0
port_args = PortArgs.init_new(server_args)
self.assertIsNotNone(port_args.decoupled_spec_ipc_config)
self.assertEqual(port_args.decoupled_spec_ipc_config.rank, 0)
self.assertEqual(
port_args.decoupled_spec_ipc_config.bind_endpoint, "ipc:///tmp/v"
)
self.assertEqual(
port_args.decoupled_spec_ipc_config.connect_endpoints, ("ipc:///tmp/d",)
)
@patch("sglang.srt.server_args.tempfile.NamedTemporaryFile")
def test_init_new_no_decoupled_config_when_role_null(self, mock_temp_file):
mock_temp_file.return_value.name = "temp_file"
server_args = ServerArgs(model_path="dummy")
server_args.nccl_port = None
server_args.enable_dp_attention = False
# decoupled_spec_role defaults to "null"
port_args = PortArgs.init_new(server_args)
self.assertIsNone(port_args.decoupled_spec_ipc_config)
def test_init_new_decoupled_role_requires_endpoints(self):
server_args = ServerArgs(model_path="dummy")
server_args.nccl_port = None
server_args.enable_dp_attention = False
server_args.decoupled_spec_role = "drafter"
# endpoints intentionally left as their None defaults
with self.assertRaises(ValueError):
PortArgs.init_new(server_args)
def test_init_new_with_single_node_dp_attention(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 = 1
server_args.dist_init_addr = None
port_args = PortArgs.init_new(server_args)
self.assertTrue(port_args.tokenizer_ipc_name.startswith("tcp://127.0.0.1:"))
self.assertTrue(
port_args.scheduler_input_ipc_name.startswith("tcp://127.0.0.1:")
)
self.assertTrue(port_args.detokenizer_ipc_name.startswith("tcp://127.0.0.1:"))
self.assertIsInstance(port_args.nccl_port, int)
def test_init_new_with_dp_rank(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 = 1
server_args.dist_init_addr = "192.168.1.1:25000"
worker_ports = [25006, 25007, 25008, 25009]
port_args = PortArgs.init_new(server_args, dp_rank=2, worker_ports=worker_ports)
self.assertTrue(port_args.scheduler_input_ipc_name.endswith(":25008"))
self.assertTrue(port_args.tokenizer_ipc_name.startswith("tcp://192.168.1.1:"))
self.assertTrue(port_args.detokenizer_ipc_name.startswith("tcp://192.168.1.1:"))
self.assertIsInstance(port_args.nccl_port, int)
def test_init_new_with_ipv4_address(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 = "192.168.1.1:25000"
port_args = PortArgs.init_new(server_args)
self.assertTrue(port_args.tokenizer_ipc_name.startswith("tcp://192.168.1.1:"))
self.assertTrue(
port_args.scheduler_input_ipc_name.startswith("tcp://192.168.1.1:")
)
self.assertTrue(port_args.detokenizer_ipc_name.startswith("tcp://192.168.1.1:"))
self.assertIsInstance(port_args.nccl_port, int)
def test_init_new_with_malformed_ipv4_address(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 = "192.168.1.1"
with self.assertRaises(ValueError) as context:
PortArgs.init_new(server_args)
self.assertIn("Missing port", str(context.exception))
def test_init_new_with_malformed_ipv4_address_invalid_port(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 = "192.168.1.1:abc"
with self.assertRaises(ValueError):
PortArgs.init_new(server_args)
class TestSSLArgs(unittest.TestCase):
def _validate_ssl(self, **kwargs):
server_args = ServerArgs(model_path="dummy", **kwargs)
server_args._handle_ssl_validation()
return server_args
def test_ssl_keyfile_without_certfile_raises(self):
with self.assertRaises(ValueError) as context:
self._validate_ssl(ssl_keyfile="key.pem")
self.assertIn("--ssl-certfile", str(context.exception))
def test_ssl_certfile_without_keyfile_raises(self):
with self.assertRaises(ValueError) as context:
self._validate_ssl(ssl_certfile="cert.pem")
self.assertIn("--ssl-keyfile", str(context.exception))
def test_url_returns_http_without_ssl(self):
server_args = ServerArgs(model_path="dummy")
self.assertTrue(server_args.url().startswith("http://"))
def test_url_rewrites_all_interfaces_to_loopback(self):
server_args = ServerArgs(model_path="dummy", host="0.0.0.0")
self.assertEqual(server_args.url(), "http://127.0.0.1:30000")
def test_url_rewrites_empty_host_to_loopback(self):
server_args = ServerArgs(model_path="dummy", host="")
self.assertEqual(server_args.url(), "http://127.0.0.1:30000")
@patch("os.path.isfile", return_value=True)
def test_url_returns_https_with_ssl(self, _mock_isfile):
server_args = self._validate_ssl(ssl_keyfile="key.pem", ssl_certfile="cert.pem")
self.assertTrue(server_args.url().startswith("https://"))
def test_ssl_verify_without_ssl(self):
server_args = ServerArgs(model_path="dummy")
self.assertIs(server_args.ssl_verify(), True)
@patch("os.path.isfile", return_value=True)
def test_ssl_verify_with_ssl_no_ca(self, _mock_isfile):
server_args = self._validate_ssl(ssl_keyfile="key.pem", ssl_certfile="cert.pem")
self.assertIs(server_args.ssl_verify(), False)
@patch("os.path.isfile", return_value=True)
def test_ssl_verify_with_ssl_and_ca(self, _mock_isfile):
server_args = self._validate_ssl(
ssl_keyfile="key.pem",
ssl_certfile="cert.pem",
ssl_ca_certs="ca.pem",
)
self.assertEqual(server_args.ssl_verify(), "ca.pem")
def test_ssl_ca_certs_without_certfile_raises(self):
with self.assertRaises(ValueError) as context:
self._validate_ssl(ssl_ca_certs="ca.pem")
self.assertIn("--ssl-ca-certs", str(context.exception))
def test_ssl_keyfile_password_without_certfile_raises(self):
with self.assertRaises(ValueError) as context:
self._validate_ssl(ssl_keyfile_password="secret")
self.assertIn("--ssl-keyfile-password", str(context.exception))
def test_ssl_keyfile_not_found_raises(self):
with self.assertRaises(ValueError) as context:
self._validate_ssl(
ssl_keyfile="/nonexistent/key.pem",
ssl_certfile="/nonexistent/cert.pem",
)
self.assertIn("not found", str(context.exception))
def test_ssl_certfile_not_found_raises(self):
with tempfile.NamedTemporaryFile(suffix=".pem") as keyfile:
with self.assertRaises(ValueError) as context:
self._validate_ssl(
ssl_keyfile=keyfile.name,
ssl_certfile="/nonexistent/cert.pem",
)
self.assertIn("SSL certificate file not found", str(context.exception))
def test_ssl_ca_certs_not_found_raises(self):
with tempfile.NamedTemporaryFile(suffix=".pem") as keyfile:
with tempfile.NamedTemporaryFile(suffix=".pem") as certfile:
with self.assertRaises(ValueError) as context:
self._validate_ssl(
ssl_keyfile=keyfile.name,
ssl_certfile=certfile.name,
ssl_ca_certs="/nonexistent/ca.pem",
)
self.assertIn(
"SSL CA certificates file not found", str(context.exception)
)
def test_enable_ssl_refresh_without_ssl_raises(self):
with self.assertRaises(ValueError) as context:
self._validate_ssl(enable_ssl_refresh=True)
self.assertIn("--enable-ssl-refresh", str(context.exception))
self.assertIn("--ssl-certfile", str(context.exception))
@patch("os.path.isfile", return_value=True)
def test_enable_ssl_refresh_with_ssl_accepted(self, _mock_isfile):
server_args = self._validate_ssl(
ssl_keyfile="key.pem",
ssl_certfile="cert.pem",
enable_ssl_refresh=True,
)
self.assertTrue(server_args.enable_ssl_refresh)
class TestHiCacheArgs(unittest.TestCase):
def _make_args(self, **overrides) -> ServerArgs:
args = ServerArgs(model_path="dummy")
for key, value in overrides.items():
setattr(args, key, value)
return args
def _assert_hicache_fields(
self,
args: ServerArgs,
*,
expected_io_backend: str,
expected_mem_layout: str,
expected_decode_backend: str | None = None,
):
self.assertEqual(args.hicache_io_backend, expected_io_backend)
self.assertEqual(args.hicache_mem_layout, expected_mem_layout)
if expected_decode_backend is not None:
self.assertEqual(args.decode_attention_backend, expected_decode_backend)
def test_hicache_io_backend_and_mem_layout_compatibility(self):
cases = [
{
"name": "default_kernel_page_first",
"overrides": {
"enable_hierarchical_cache": True,
},
"expected_io_backend": "kernel",
"expected_mem_layout": "page_first",
},
{
"name": "kernel_with_page_first_direct",
"overrides": {
"enable_hierarchical_cache": True,
"hicache_io_backend": "kernel",
"hicache_mem_layout": "page_first_direct",
},
"expected_io_backend": "direct",
"expected_mem_layout": "page_first_direct",
},
{
"name": "direct_with_page_first",
"overrides": {
"enable_hierarchical_cache": True,
"hicache_io_backend": "direct",
"hicache_mem_layout": "page_first",
},
"expected_io_backend": "direct",
"expected_mem_layout": "page_first_direct",
},
{
"name": "mooncake_with_layer_first",
"overrides": {
"enable_hierarchical_cache": True,
"hicache_storage_backend": "mooncake",
"hicache_io_backend": "direct",
"hicache_mem_layout": "layer_first",
},
"expected_io_backend": "direct",
"expected_mem_layout": "page_first_direct",
},
{
"name": "fa3_kernel_with_explicit_decode_backend",
"overrides": {
"enable_hierarchical_cache": True,
"hicache_io_backend": "kernel",
"hicache_mem_layout": "page_first",
"attention_backend": "triton",
"decode_attention_backend": "fa3",
},
"expected_io_backend": "kernel",
"expected_mem_layout": "page_first",
"expected_decode_backend": "fa3",
},
]
for case in cases:
with self.subTest(case=case["name"]):
args = self._make_args(**case["overrides"])
args._handle_hicache()
self._assert_hicache_fields(
args,
expected_io_backend=case["expected_io_backend"],
expected_mem_layout=case["expected_mem_layout"],
expected_decode_backend=case.get("expected_decode_backend"),
)
def test_hicache_kernel_keeps_implicit_fa3_decode_backend(self):
args = self._make_args(
enable_hierarchical_cache=True,
hicache_io_backend="kernel",
attention_backend="fa3",
decode_attention_backend=None,
)
args._handle_hicache()
self.assertEqual(args.hicache_io_backend, "kernel")
self.assertEqual(args.hicache_mem_layout, "page_first")
self.assertIsNone(args.decode_attention_backend)
class TestNgramExternalSamArgs(CustomTestCase):
def _make_dummy_ngram_args(self, **overrides):
args = ServerArgs(model_path="dummy")
args.speculative_algorithm = "NGRAM"
args.speculative_num_draft_tokens = 12
args.device = "cuda"
for key, value in overrides.items():
setattr(args, key, value)
return args
def test_external_sam_budget_must_fit_draft_budget(self):
args = self._make_dummy_ngram_args(
speculative_num_draft_tokens=4,
speculative_ngram_external_corpus_path="/tmp/ngram-corpus.jsonl",
speculative_ngram_external_sam_budget=4,
)
with self.assertRaises(ValueError) as context:
handle_speculative_decoding(args)
self.assertIn("speculative_num_draft_tokens - 1", str(context.exception))
def test_external_corpus_max_tokens_must_be_positive(self):
args = self._make_dummy_ngram_args(
speculative_ngram_external_corpus_path="/tmp/ngram-corpus.jsonl",
speculative_ngram_external_sam_budget=2,
speculative_ngram_external_corpus_max_tokens=0,
)
with self.assertRaises(ValueError) as context:
handle_speculative_decoding(args)
self.assertIn("external-corpus-max-tokens", str(context.exception))
class TestDecoupledSpecArgs(CustomTestCase):
"""Decoupled speculative-decoding CLI flags.
These flags are auto-derived from the ``A[...]`` field metadata on
``ServerArgs``; a bare annotation is silently skipped by
``add_cli_args_from_dataclass``. This guards against the regression where
the flags went missing (e.g. after rebasing onto the auto-gen
``add_cli_args``), which the direct-attribute ``PortArgs`` tests cannot
catch because they never exercise the CLI.
"""
def test_decoupled_spec_cli_flags_round_trip(self):
server_args = prepare_server_args(
[
"--model-path",
"dummy",
"--decoupled-spec-role",
"verifier",
"--decoupled-spec-bind-endpoint",
"ipc:///tmp/v",
"--decoupled-spec-connect-endpoints",
'["ipc:///tmp/d"]',
"--decoupled-spec-rank",
"0",
"--spec-trace-dir",
"/tmp/tr",
]
)
self.assertEqual(server_args.decoupled_spec_role, "verifier")
self.assertEqual(server_args.decoupled_spec_bind_endpoint, "ipc:///tmp/v")
self.assertEqual(server_args.decoupled_spec_connect_endpoints, ["ipc:///tmp/d"])
self.assertEqual(server_args.decoupled_spec_rank, 0)
self.assertEqual(server_args.spec_trace_dir, "/tmp/tr")
def test_decoupled_spec_role_rejects_invalid_choice(self):
with self.assertRaises(SystemExit):
prepare_server_args(
["--model-path", "dummy", "--decoupled-spec-role", "bogus"]
)
class TestAdaptiveSpecArgs(CustomTestCase):
def test_adaptive_defaults_to_config_step_when_spec_params_omitted(self):
with tempfile.NamedTemporaryFile("w", suffix=".json") as f:
json.dump(
{
"1": {"candidate_steps": [1, 3, 5]},
"8": {"candidate_steps": [1]},
},
f,
)
f.flush()
args = ServerArgs(model_path="dummy")
args.speculative_algorithm = "EAGLE"
args.speculative_adaptive = True
args.speculative_adaptive_config = f.name
args.device = "cuda"
args.get_model_config = lambda: SimpleNamespace(
hf_config=SimpleNamespace(
architectures=["LlamaForCausalLM"],
get_text_config=lambda: SimpleNamespace(),
)
)
handle_speculative_decoding(args)
self.assertTrue(args.speculative_adaptive)
self.assertEqual(args.speculative_eagle_topk, 1)
self.assertEqual(args.speculative_num_steps, 3)
self.assertEqual(args.speculative_num_draft_tokens, 4)
class TestWaterfillArgs(CustomTestCase):
def test_waterfill_enforces_shared_experts_fusion(self):
server_args = ServerArgs(
model_path="dummy",
moe_a2a_backend="deepep",
enable_waterfill=True,
disable_shared_experts_fusion=True,
)
# dummy-model path short-circuits __post_init__; invoke the handler directly.
server_args._handle_a2a_moe()
from sglang.srt.arg_groups.overrides import resolved_view
# dual-apply retired: the fields stay pristine, the declarations win
self.assertTrue(server_args.disable_shared_experts_fusion)
self.assertFalse(resolved_view(server_args).disable_shared_experts_fusion)
self.assertTrue(server_args.enforce_shared_experts_fusion)
def test_waterfill_overrides_moe_a2a_backend_to_deepep(self):
server_args = ServerArgs(
model_path="dummy",
moe_a2a_backend="none",
enable_waterfill=True,
)
# dummy-model path short-circuits __post_init__; invoke the handler directly.
server_args._handle_a2a_moe()
from sglang.srt.arg_groups.overrides import resolved_view
self.assertEqual(server_args.moe_a2a_backend, "none") # pristine
self.assertEqual(resolved_view(server_args).moe_a2a_backend, "deepep")
self.assertTrue(server_args.enforce_shared_experts_fusion)
def test_waterfill_keeps_megamoe_backend(self):
server_args = ServerArgs(
model_path="dummy",
moe_a2a_backend="megamoe",
enable_waterfill=True,
disable_shared_experts_fusion=True,
)
# dummy-model path short-circuits __post_init__; invoke the handler directly.
server_args._handle_a2a_moe()
from sglang.srt.arg_groups.overrides import resolved_view
self.assertEqual(resolved_view(server_args).moe_a2a_backend, "megamoe")
self.assertFalse(resolved_view(server_args).disable_shared_experts_fusion)
self.assertTrue(server_args.enforce_shared_experts_fusion)
def test_waterfill_supports_deepep_low_latency_mode(self):
server_args = ServerArgs(
model_path="dummy",
moe_a2a_backend="deepep",
enable_waterfill=True,
deepep_mode="low_latency",
)
# dummy-model path short-circuits __post_init__; invoke the handler directly.
server_args._handle_a2a_moe()
self.assertEqual(server_args.deepep_mode, "low_latency")
self.assertFalse(server_args.disable_cuda_graph)
self.assertTrue(server_args.enforce_shared_experts_fusion)
class TestPrefillOnlyDisableKvCache(unittest.TestCase):
"""Validation for --prefill-only-disable-kv-cache.
The flag wires NoOpMHATokenToKVPool, which is only safe when:
- the engine is in embedding mode (fa_skip_kv_cache active in FA backend),
- chunked_prefill_size == -1 (no inter-chunk K/V reuse),
- disable_radix_cache (radix cache otherwise indexes empty pool slots),
- no context-parallel attention (CP writes to the pool via set_kv_buffer),
- no HiSparse (uses a different pool family),
- kv_cache_dtype != fp4_e2m1 (FP4 pool is a separate allocation path).
All other configurations must be rejected before model load.
"""
def _base_kwargs(self, **overrides):
kwargs = dict(
model_path="dummy",
is_embedding=True,
chunked_prefill_size=-1,
disable_radix_cache=True,
prefill_only_disable_kv_cache=True,
)
kwargs.update(overrides)
return kwargs
def _validate_prefill_only_args(self, **overrides):
sa = ServerArgs(**self._base_kwargs(**overrides))
sa._handle_legacy_cp_arguments()
sa._validate_prefill_only_disable_kv_cache_args()
return sa
def test_valid_minimal_config_constructs(self):
sa = self._validate_prefill_only_args()
self.assertTrue(sa.prefill_only_disable_kv_cache)
def test_rejects_when_not_embedding(self):
with self.assertRaisesRegex(ValueError, "requires --is-embedding"):
self._validate_prefill_only_args(is_embedding=False)
def test_rejects_when_chunked_prefill_size_not_minus_one(self):
with self.assertRaisesRegex(ValueError, "--chunked-prefill-size=-1"):
self._validate_prefill_only_args(chunked_prefill_size=8192)
def test_rejects_when_radix_cache_enabled(self):
with self.assertRaisesRegex(ValueError, "--disable-radix-cache"):
self._validate_prefill_only_args(disable_radix_cache=False)
def test_rejects_attn_cp_size_greater_than_one(self):
with self.assertRaisesRegex(ValueError, "--attn-cp-size"):
self._validate_prefill_only_args(attn_cp_size=2, tp_size=2)
def test_rejects_prefill_context_parallel(self):
with self.assertRaisesRegex(ValueError, "--enable-prefill-cp"):
self._validate_prefill_only_args(enable_prefill_context_parallel=True)
def test_rejects_hisparse(self):
with self.assertRaisesRegex(ValueError, "--enable-hisparse"):
self._validate_prefill_only_args(enable_hisparse=True)
def test_rejects_fp4_kv_cache(self):
with self.assertRaisesRegex(ValueError, "fp4_e2m1"):
self._validate_prefill_only_args(kv_cache_dtype="fp4_e2m1")
class TestSessionRadixCacheServerArgs(unittest.TestCase):
def test_requires_priority_radix_eviction_policy(self):
server_args = ServerArgs(
model_path="dummy",
enable_session_radix_cache=True,
radix_eviction_policy="lru",
)
with self.assertRaisesRegex(ValueError, "--radix-eviction-policy priority"):
server_args._handle_cache_compatibility()
class TestCudaGraphConfigDataclassAccess(CustomTestCase):
@patch(
"sglang.srt.model_executor.runner_backend."
"tc_piecewise_cuda_graph_backend.get_moe_a2a_backend"
)
def test_tc_piecewise_build_config_reads_phase_config_dataclass(
self, mock_get_moe_a2a_backend
):
from sglang.srt.model_executor.runner_backend.tc_piecewise_cuda_graph_backend import (
TcPiecewiseCudaGraphBackend,
)
mock_backend = mock_get_moe_a2a_backend.return_value
mock_backend.is_deepep.return_value = False
mock_backend.is_mooncake.return_value = False
server_args = SimpleNamespace(
cuda_graph_config=CudaGraphConfig(
prefill=PhaseConfig(
backend=Backend.TC_PIECEWISE,
bs=[32, 64],
tc_compiler="eager",
)
),
enable_torch_compile_debug_mode=False,
)
config = TcPiecewiseCudaGraphBackend.build_compilation_config(server_args)
self.assertEqual(config.get_capture_sizes(), [32, 64])
self.assertEqual(config.compiler, "eager")
class TestCudaGraphDisaggregationRoles(CustomTestCase):
def _handled_args(self, **overrides):
args = ServerArgs(model_path="dummy", **overrides)
args.model_config = SimpleNamespace(
hf_config=SimpleNamespace(architectures=["LlamaForCausalLM"]),
is_piecewise_cuda_graph_disabled_model=False,
is_multimodal=False,
is_multimodal_piecewise_cuda_graph_supported=False,
)
with (
patch("sglang.srt.utils.is_cuda", return_value=True),
patch.object(ServerArgs, "use_mla_backend", return_value=False),
):
args._handle_cuda_graph_config()
return args
def test_cuda_graph_prefill_role_defaults_disable_decode_graph(self):
args = self._handled_args(disaggregation_mode="prefill")
self.assertFalse(args.disable_cuda_graph)
self.assertEqual(args.cuda_graph_config.decode.backend, Backend.DISABLED)
self.assertEqual(args.cuda_graph_config.prefill.backend, Backend.BREAKABLE)
def test_cuda_graph_decode_role_defaults_disable_prefill_graph(self):
args = self._handled_args(disaggregation_mode="decode")
self.assertEqual(args.cuda_graph_config.prefill.backend, Backend.DISABLED)
self.assertNotEqual(args.cuda_graph_config.decode.backend, Backend.DISABLED)
def test_cuda_graph_global_disable_still_disables_both_phases_for_all_roles(self):
for disaggregation_mode in ("prefill", "decode", "null"):
with self.subTest(disaggregation_mode=disaggregation_mode):
args = self._handled_args(
disaggregation_mode=disaggregation_mode,
disable_cuda_graph=True,
)
self.assertEqual(
args.cuda_graph_config.decode.backend, Backend.DISABLED
)
self.assertEqual(
args.cuda_graph_config.prefill.backend, Backend.DISABLED
)
def test_cuda_graph_explicit_decode_backend_survives_prefill_role(self):
args = self._handled_args(
disaggregation_mode="prefill",
cuda_graph_backend_decode=Backend.FULL,
)
self.assertEqual(args.cuda_graph_config.decode.backend, Backend.FULL)
self.assertIn((Phase.DECODE, "backend"), args._cuda_graph_config_locked)
class TestCutedslMoeMaxNumTokens(CustomTestCase):
"""The shared CuteDSL MoE per-forward token bound. Fields are set directly
to exercise the math independently of __post_init__ resolution.
cg-refactor: the legacy disable_piecewise_cuda_graph /
piecewise_cuda_graph_max_tokens / cuda_graph_max_bs fields were
consolidated into cuda_graph_config; the helper accepts the legacy
kwarg names for test readability and translates them to the per-phase
dataclasses.
"""
def _args(self, **overrides):
server_args = ServerArgs(model_path="dummy")
fields = dict(
speculative_algorithm=None,
speculative_num_draft_tokens=None,
max_prefill_tokens=16384,
disable_piecewise_cuda_graph=False,
piecewise_cuda_graph_max_tokens=2048,
cuda_graph_max_bs=512,
)
fields.update(overrides)
disable_piecewise = fields.pop("disable_piecewise_cuda_graph")
piecewise_max = fields.pop("piecewise_cuda_graph_max_tokens")
cg_max_bs = fields.pop("cuda_graph_max_bs")
for key, value in fields.items():
setattr(server_args, key, value)
server_args.cuda_graph_config = CudaGraphConfig(
decode=PhaseConfig(backend=Backend.FULL, max_bs=cg_max_bs),
prefill=PhaseConfig(
backend=(
Backend.DISABLED if disable_piecewise else Backend.TC_PIECEWISE
),
max_bs=piecewise_max,
tc_compiler="eager",
),
)
return server_args
def test_prefill_dominates_in_default_config(self):
self.assertEqual(self._args().cutedsl_moe_max_num_tokens(), 16384)
def test_speculative_decoding_scales_decode_bound(self):
# decode bound 512 * 8 dominates the small prefill/piecewise bounds
args = self._args(
max_prefill_tokens=512,
piecewise_cuda_graph_max_tokens=512,
speculative_algorithm="EAGLE",
speculative_num_draft_tokens=8,
)
self.assertEqual(args.cutedsl_moe_max_num_tokens(), 4096)
def test_piecewise_bound_excluded_when_disabled(self):
args = self._args(
max_prefill_tokens=512,
disable_piecewise_cuda_graph=True,
cuda_graph_max_bs=64,
)
self.assertEqual(args.cutedsl_moe_max_num_tokens(), 512)
class TestSamplingBackendTokenOracleEnvGate(CustomTestCase):
"""The 'token_oracle' choice is gated on SGLANG_KV_CANARY_ENABLE_TOKEN_ORACLE.
The choice set is built once at server_args.py import time, so each subtest
reloads the module with the env var set to the desired value.
"""
def _reload_server_args_with_env(self, *, enabled: bool):
previous = os.environ.get("SGLANG_KV_CANARY_ENABLE_TOKEN_ORACLE")
os.environ["SGLANG_KV_CANARY_ENABLE_TOKEN_ORACLE"] = "1" if enabled else "0"
try:
return importlib.reload(server_args_module)
finally:
if previous is None:
os.environ.pop("SGLANG_KV_CANARY_ENABLE_TOKEN_ORACLE", None)
else:
os.environ["SGLANG_KV_CANARY_ENABLE_TOKEN_ORACLE"] = previous
def test_token_oracle_rejected_when_env_disabled(self):
reloaded = self._reload_server_args_with_env(enabled=False)
self.assertNotIn("token_oracle", reloaded.SAMPLING_BACKEND_CHOICES)
with self.assertRaises(SystemExit):
reloaded.prepare_server_args(
[
"--model-path",
DEFAULT_SMALL_MODEL_NAME_FOR_TEST_QWEN,
"--sampling-backend",
"token_oracle",
]
)
def test_token_oracle_accepted_when_env_enabled(self):
reloaded = self._reload_server_args_with_env(enabled=True)
self.assertIn("token_oracle", reloaded.SAMPLING_BACKEND_CHOICES)
parsed = reloaded.prepare_server_args(
[
"--model-path",
DEFAULT_SMALL_MODEL_NAME_FOR_TEST_QWEN,
"--sampling-backend",
"token_oracle",
# Explicit device so ServerArgs.__post_init__ does not call
# get_device() (fails on CPU-only CI runners) and does not run
# _handle_cpu_backends (which would override sampling_backend
# to "pytorch", masking what we want to verify).
"--device",
"cuda",
]
)
self.assertEqual(parsed.sampling_backend, "token_oracle")
class TestHandleCrashDumpEnv(CustomTestCase):
_COREDUMP_ENV_KEYS = (
"CUDA_ENABLE_COREDUMP_ON_EXCEPTION",
"CUDA_ENABLE_USER_TRIGGERED_COREDUMP",
"CUDA_COREDUMP_SHOW_PROGRESS",
"CUDA_COREDUMP_GENERATION_FLAGS",
"CUDA_COREDUMP_FILE",
"CUDA_COREDUMP_PIPE",
)
def _run_handler(self, crash_dump_folder, preset_env=None):
server_args = ServerArgs.__new__(ServerArgs)
server_args.crash_dump_folder = crash_dump_folder
with patch.dict(os.environ, preset_env or {}):
for key in self._COREDUMP_ENV_KEYS:
if key not in (preset_env or {}):
os.environ.pop(key, None)
ServerArgs._handle_crash_dump_env(server_args)
def test_creates_coredump_dir_when_auto_set(self):
with tempfile.TemporaryDirectory() as tmp:
self._run_handler(tmp)
self.assertTrue(
os.path.isdir(os.path.join(tmp, socket.gethostname())),
"coredump dir not created for auto-set CUDA_COREDUMP_FILE",
)
def test_creates_coredump_dir_when_env_preset(self):
# Regression test: when CUDA_COREDUMP_FILE is preset, the coredump
# directory must still be created up front.
with tempfile.TemporaryDirectory() as tmp:
preset_dir = os.path.join(tmp, "preset-location")
self._run_handler(
tmp,
preset_env={"CUDA_COREDUMP_FILE": f"{preset_dir}/%h/core.cuda.%t.%p"},
)
self.assertTrue(
os.path.isdir(os.path.join(preset_dir, socket.gethostname())),
"coredump dir not created for preset CUDA_COREDUMP_FILE",
)
class TestGrpcServerArgs(CustomTestCase):
"""Native gRPC is enabled by --grpc-port (or SGLANG_GRPC_PORT) and runs
alongside HTTP; --smg-grpc-mode (and the deprecated --grpc-mode) select the
legacy SMG server. Worker-threads / max-prefill-tokens are env-only knobs.
The gRPC setup lives in ServerArgs._handle_deprecated_args, which
__post_init__ skips for dummy models, so these tests build a dummy
ServerArgs and invoke that handler directly (mirroring the real flow for a
concrete model path).
"""
@staticmethod
def _args(**kwargs):
return ServerArgs(model_path="dummy", **kwargs)
def test_http_only_high_port_does_not_derive_grpc_port(self):
sa = self._args(port=56000)
sa._handle_deprecated_args()
self.assertIsNone(sa.grpc_port)
def test_grpc_port_enables_native_and_env_knobs(self):
sa = self._args(grpc_port=50051)
with envs.SGLANG_GRPC_WORKER_THREADS.override(8):
sa._handle_deprecated_args()
self.assertEqual(sa.grpc_port, 50051)
self.assertEqual(sa.grpc_worker_threads, 8)
def test_env_grpc_port_enables_native(self):
sa = self._args(port=30000)
with envs.SGLANG_GRPC_PORT.override(45000):
sa._handle_deprecated_args()
self.assertEqual(sa.grpc_port, 45000)
def test_legacy_smg_derives_grpc_port_from_http_port(self):
sa = self._args(port=30000, smg_grpc_mode=True)
sa._handle_deprecated_args()
self.assertEqual(sa.grpc_port, 40000)
def test_grpc_mode_is_deprecated_alias_for_smg_grpc_mode(self):
sa = self._args(grpc_mode=True)
with self.assertLogs(server_args_module.logger, level="WARNING") as cm:
sa._handle_deprecated_args()
self.assertTrue(sa.smg_grpc_mode)
self.assertTrue(any("--grpc-mode is deprecated" in line for line in cm.output))
def test_legacy_smg_takes_precedence_over_grpc_port(self):
sa = self._args(grpc_port=50051, smg_grpc_mode=True)
sa._handle_deprecated_args()
self.assertTrue(sa.smg_grpc_mode)
self.assertEqual(sa.grpc_port, 50051)
def test_native_grpc_rejects_multi_tokenizer(self):
sa = self._args(grpc_port=40000, tokenizer_worker_num=2)
with self.assertRaises(ValueError):
sa._handle_deprecated_args()
def test_native_grpc_rejects_http_auth(self):
sa = self._args(grpc_port=40000, api_key="secret")
with self.assertRaises(ValueError):
sa._handle_deprecated_args()
def test_invalid_grpc_worker_threads_rejected(self):
sa = self._args(grpc_port=40000)
with envs.SGLANG_GRPC_WORKER_THREADS.override(0):
with self.assertRaises(ValueError):
sa._handle_deprecated_args()
def test_start_server_call_site_matches_native_signature(self):
"""Regression for the startup blocker: the native start_server binding
only accepts (host, port, runtime_handle, worker_threads, ...). The
arg-parsing tests above never call start_server, so a stray kwarg (e.g.
the removed max_prefill_tokens) would only surface as a TypeError at
launch. This mocks the native extension and locks the kwarg set."""
import sys
from sglang.srt.entrypoints import http_server
fake_core = SimpleNamespace(start_server=MagicMock(return_value="handle"))
fake_bridge = SimpleNamespace(RuntimeHandle=MagicMock(return_value="rt"))
server_args = SimpleNamespace(
host="127.0.0.1", grpc_port=50051, grpc_worker_threads=4
)
with patch.dict(
sys.modules,
{
"sglang.srt.grpc": SimpleNamespace(_core=fake_core),
"sglang.srt.grpc._core": fake_core,
"sglang.srt.entrypoints.grpc_bridge": fake_bridge,
},
):
handle = http_server._start_native_grpc_server_for_runtime(
server_args=server_args,
tokenizer_manager=MagicMock(),
template_manager=MagicMock(),
scheduler_info={},
)
self.assertEqual(handle, "handle")
_, kwargs = fake_core.start_server.call_args
self.assertEqual(
set(kwargs), {"host", "port", "runtime_handle", "worker_threads"}
)
self.assertNotIn("max_prefill_tokens", kwargs)
class TestTwoBatchOverlapBackend(CustomTestCase):
"""Non-EP DP two-batch-overlap backend requirement.
With no EP a2a backend (moe_a2a_backend='none'), --enable-two-batch-overlap
is only valid on the DeepSeek-V4 non-EP DP TP-MoE path (overlapping the DP
all_gatherv / reduce_scatterv with the other ubatch's compute), which
requires --enable-dp-attention. This replaced the removed opt-in
SGLANG_ENABLE_DP_TBO env: enabling DP TBO now needs no extra flag.
dummy-model short-circuits __post_init__, so the guard handler is invoked
directly (same pattern as TestWaterfillArgs)."""
def _args(self, **overrides):
args = ServerArgs(model_path="dummy")
args.enable_two_batch_overlap = True
args.moe_a2a_backend = "none"
args.enable_dp_attention = False
for key, value in overrides.items():
setattr(args, key, value)
return args
def test_no_a2a_without_dp_attention_raises(self):
args = self._args(enable_dp_attention=False)
with self.assertRaisesRegex(ValueError, "enable-dp-attention"):
args._check_two_batch_overlap()
def test_no_a2a_with_dp_attention_ok(self):
# DP TBO path is valid: --enable-dp-attention + --enable-two-batch-overlap
# with a2a backend 'none' must NOT raise (no SGLANG_ENABLE_DP_TBO needed).
args = self._args(enable_dp_attention=True)
args._check_two_batch_overlap()
def test_ep_a2a_backend_ok_without_dp_attention(self):
# EP a2a path (e.g. deepep) overlaps dispatch/combine; the guard does not
# require dp-attention there.
args = self._args(moe_a2a_backend="deepep", enable_dp_attention=False)
args._check_two_batch_overlap()
if __name__ == "__main__":
unittest.main()