config: resolution reads the declarations, not the fields (#36253)

This commit is contained in:
Cheng Wan
2026-08-26 05:05:28 -07:00
committed by GitHub
parent ae5feb4b9c
commit 5b7fc61306
34 changed files with 1865 additions and 1495 deletions
@@ -125,6 +125,13 @@ def _registry_collection_is_after_the_build():
def _server_args_names(tree, path):
"""Every local that names the record, including the read views over it.
A resolution-time reader reads through `resolving_view(server_args)` (the
declaration stash over the fields): declaration-only resolvers write no
field, so a field read there answers with the raw input. `cfg.dtype` after `cfg = resolving_view(sa)` is
the same read this scan is looking for, so the local it binds counts.
"""
names = {"self"} if path.name == "server_args.py" else {"server_args"}
for node in ast.walk(tree):
if not isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef)):
@@ -142,6 +149,32 @@ def _server_args_names(tree, path):
continue
if text == "ServerArgs":
names.add(arg.arg)
# `cfg = resolving_view(server_args)` / `resolved_view(server_args)`
for _ in range(2): # a view over a view-holding local is still one
for node in ast.walk(tree):
if not isinstance(node, ast.Assign):
continue
value = node.value
bare = (
isinstance(value, ast.Call)
and isinstance(value.func, ast.Name)
and value.func.id in ("resolving_view", "resolved_view")
and value.args
and isinstance(value.args[0], ast.Name)
and value.args[0].id in names
)
# `resolved = self._resolved()` is the same view, spelled as the
# record's own member.
member = (
isinstance(value, ast.Call)
and isinstance(value.func, ast.Attribute)
and value.func.attr == "_resolved"
and isinstance(value.func.value, ast.Name)
and value.func.value.id in names
)
if not (bare or member):
continue
names |= {t.id for t in node.targets if isinstance(t, ast.Name)}
return names
@@ -191,11 +224,12 @@ def _late_resolution_fields():
for name in (
"server_args.py",
"arg_groups/overrides.py",
"utils/template_detection.py",
"parser/template_detection.py",
):
path = _SRT / name
if not path.exists():
continue
# A named file that moved away has to be loud; skipping it silently
# leaves the scan believing it read a module it never opened.
assert path.exists(), f"{name} is not where this scan looks for it"
tree = _parsed(path)
for node in ast.walk(tree):
if not isinstance(node, ast.Call):
@@ -764,7 +764,7 @@ class TestResolutionDeclarations(CustomTestCase):
# Snapshot before publishing: the bag serves the very object the record
# holds, so comparing them after the fact compares an object with
# itself and passes however the projection behaves.
expected = copy.deepcopy(server_args.cuda_graph_config)
expected = copy.deepcopy(resolution_result(server_args, "cuda_graph_config"))
publish(server_args, role="scheduler")
published = get_exec().graph.cuda_graph_config
resolved = expected
@@ -36,6 +36,7 @@ import unittest.mock
import torch
import sglang
from sglang.srt.arg_groups.overrides import resolution_result
from sglang.srt.environ import EnvField, envs
from sglang.srt.server_args import ServerArgs
from sglang.srt.utils import is_cuda
@@ -270,7 +271,10 @@ class TestResolutionIsReproducible(_RestoresProcessState, CustomTestCase):
for field in dataclasses.fields(server_args):
if field.name in _NOT_COMPARABLE:
continue
value = getattr(server_args, field.name)
# The resolution result, not the field: a declaration-only resolver
# never writes the field, so comparing fields would miss exactly
# the decisions a leak would shift.
value = resolution_result(server_args, field.name)
# Nested dataclasses (cuda_graph_config) compare structurally, and
# everything else is deep-copied: a snapshot that stored the live
# list/dict would follow an in-place mutation, which is exactly the
@@ -382,10 +386,13 @@ class TestResolutionIsReproducible(_RestoresProcessState, CustomTestCase):
# differs from the cpu that `default_before` resolved to.
expected = (
"cuda_ipc"
if intermediate.mm_feature_transport == "cuda_ipc"
if resolution_result(intermediate, "mm_feature_transport")
== "cuda_ipc"
else "cpu"
)
self.assertEqual(after.mm_feature_transport, expected)
self.assertEqual(
resolution_result(after, "mm_feature_transport"), expected
)
def test_resolving_a_sibling_leaves_the_first_alone(self):
for label, config, kwargs in _SHAPES:
@@ -847,9 +854,9 @@ class TestACopyStaysResolved(_RestoresProcessState, CustomTestCase):
self.assertEqual(get_parallel().config.dist_init_addr, "1.2.3.4:5000")
self.assertEqual(
get_schedule().chunked_prefill_size,
parent.chunked_prefill_size,
"publishing the copy re-ran resolution; the bag disagrees with the "
"record the parent resolved",
resolution_result(parent, "chunked_prefill_size"),
"publishing the copy re-ran resolution; the bag disagrees with what "
"the parent's resolution decided",
)
def test_no_bare_replace_of_a_record_outside_the_helper(self):
@@ -10,6 +10,7 @@ from unittest.mock import MagicMock, patch
import sglang.srt.server_args as server_args_module
from sglang.srt.arg_groups import pd_disaggregation_hook
from sglang.srt.arg_groups.overrides import resolution_result
from sglang.srt.arg_groups.speculative_hook import handle_speculative_decoding
from sglang.srt.entrypoints.sidecar import (
SGLANG_GRPC_ENDPOINT_ENV,
@@ -61,7 +62,7 @@ class TestPrepareServerArgs(CustomTestCase):
args.resolve_once()
self.assertTrue(args.enable_w4a4_mxfp4_megamoe)
self.assertTrue(resolution_result(args, "enable_w4a4_mxfp4_megamoe"))
self.assertEqual(os.environ["DG_USE_FP4_ACTS"], "1")
self.assertEqual(os.environ["DG_USE_MXF4_KIND"], "1")
@@ -76,14 +77,14 @@ class TestPrepareServerArgs(CustomTestCase):
# nothing to be untouched by.
args.resolve_once()
self.assertFalse(args.enable_w4a4_mxfp4_megamoe)
self.assertFalse(resolution_result(args, "enable_w4a4_mxfp4_megamoe"))
self.assertEqual(os.environ["DG_USE_FP4_ACTS"], "0")
self.assertEqual(os.environ["DG_USE_MXF4_KIND"], "0")
def test_prefill_decode_interval(self):
args = ServerArgs(model_path="dummy", prefill_decode_interval=16)
args.resolve_once()
self.assertEqual(args.prefill_decode_interval, 16)
self.assertEqual(resolution_result(args, "prefill_decode_interval"), 16)
with self.assertRaisesRegex(
ValueError, "--prefill-decode-interval must be non-negative"
@@ -114,22 +115,24 @@ class TestPrepareServerArgs(CustomTestCase):
return server_args
disabled = _resolved(model_path="dummy")
self.assertFalse(disabled.enable_return_hidden_states)
self.assertIsNone(disabled.return_hidden_states_mode)
self.assertFalse(resolution_result(disabled, "enable_return_hidden_states"))
self.assertIsNone(resolution_result(disabled, "return_hidden_states_mode"))
last = _resolved(
model_path="dummy",
return_hidden_states_mode="last",
)
self.assertTrue(last.enable_return_hidden_states)
self.assertEqual(last.return_hidden_states_mode, "last")
self.assertTrue(resolution_result(last, "enable_return_hidden_states"))
self.assertEqual(resolution_result(last, "return_hidden_states_mode"), "last")
legacy_full = _resolved(
model_path="dummy",
enable_return_hidden_states=True,
)
self.assertTrue(legacy_full.enable_return_hidden_states)
self.assertEqual(legacy_full.return_hidden_states_mode, "full")
self.assertTrue(resolution_result(legacy_full, "enable_return_hidden_states"))
self.assertEqual(
resolution_result(legacy_full, "return_hidden_states_mode"), "full"
)
parsed_last = prepare_server_args(
[
@@ -140,8 +143,10 @@ class TestPrepareServerArgs(CustomTestCase):
]
)
parsed_last.resolve_once()
self.assertTrue(parsed_last.enable_return_hidden_states)
self.assertEqual(parsed_last.return_hidden_states_mode, "last")
self.assertTrue(resolution_result(parsed_last, "enable_return_hidden_states"))
self.assertEqual(
resolution_result(parsed_last, "return_hidden_states_mode"), "last"
)
# The rejection is resolution's, not the constructor's.
with self.assertRaisesRegex(
@@ -156,13 +161,24 @@ class TestPrepareServerArgs(CustomTestCase):
def test_draft_quantization_explicitness_survives_asdict_round_trip(self):
inherited = ServerArgs(model_path="dummy", quantization="modelopt_fp4")
inherited._handle_missing_default_values()
self.assertEqual(inherited.speculative_draft_model_quantization, "modelopt_fp4")
self.assertFalse(inherited._speculative_draft_quantization_explicitly_set)
self.assertEqual(
resolution_result(inherited, "speculative_draft_model_quantization"),
"modelopt_fp4",
)
self.assertFalse(
resolution_result(
inherited, "_speculative_draft_quantization_explicitly_set"
)
)
reconstructed = ServerArgs(**dataclasses.asdict(inherited))
reconstructed._handle_missing_default_values()
self.assertFalse(reconstructed._speculative_draft_quantization_explicitly_set)
self.assertFalse(
resolution_result(
reconstructed, "_speculative_draft_quantization_explicitly_set"
)
)
def test_config_nested_dict_args_are_json(self):
with tempfile.NamedTemporaryFile(mode="w", suffix=".yaml", delete=False) as f:
@@ -219,8 +235,10 @@ class TestImageProcessorBackend(CustomTestCase):
server_args._handle_deprecated_args()
self.assertEqual(server_args.image_processor_backend, "pil")
self.assertFalse(server_args.disable_fast_image_processor)
self.assertEqual(
resolution_result(server_args, "image_processor_backend"), "pil"
)
self.assertFalse(resolution_result(server_args, "disable_fast_image_processor"))
def test_legacy_flag_maps_to_pil_with_one_warning(self):
server_args = ServerArgs(model_path="dummy", disable_fast_image_processor=True)
@@ -228,8 +246,10 @@ class TestImageProcessorBackend(CustomTestCase):
with self.assertLogs(server_args_module.logger, level="WARNING") as logs:
server_args._handle_deprecated_args()
self.assertEqual(server_args.image_processor_backend, "pil")
self.assertTrue(server_args.disable_fast_image_processor)
self.assertEqual(
resolution_result(server_args, "image_processor_backend"), "pil"
)
self.assertTrue(resolution_result(server_args, "disable_fast_image_processor"))
self.assertEqual(
sum(
"--disable-fast-image-processor is deprecated" in x for x in logs.output
@@ -266,7 +286,9 @@ class TestMultimodalFeatureTransport(CustomTestCase):
with self.assertLogs(server_args_module.logger, level="INFO") as logs:
server_args._handle_multimodal_feature_transport()
self.assertEqual(server_args.mm_feature_transport, "cuda_ipc")
self.assertEqual(
resolution_result(server_args, "mm_feature_transport"), "cuda_ipc"
)
self.assertTrue(envs.SGLANG_USE_CUDA_IPC_TRANSPORT.get())
output = "\n".join(logs.output)
@@ -281,8 +303,12 @@ class TestMultimodalFeatureTransport(CustomTestCase):
with self.assertLogs(server_args_module.logger, level="WARNING") as logs:
server_args._handle_multimodal_feature_transport()
self.assertEqual(server_args.mm_feature_transport, "cuda_ipc")
self.assertFalse(server_args.keep_mm_feature_on_device)
self.assertEqual(
resolution_result(server_args, "mm_feature_transport"), "cuda_ipc"
)
self.assertFalse(
resolution_result(server_args, "keep_mm_feature_on_device")
)
self.assertTrue(envs.SGLANG_USE_CUDA_IPC_TRANSPORT.get())
self.assertIn("deprecated", logs.output[0])
@@ -305,7 +331,9 @@ class TestMultimodalFeatureTransport(CustomTestCase):
with self.assertLogs(server_args_module.logger, level="WARNING") as logs:
server_args._handle_multimodal_feature_transport()
self.assertEqual(server_args.mm_feature_transport, "cpu")
self.assertEqual(
resolution_result(server_args, "mm_feature_transport"), "cpu"
)
self.assertFalse(envs.SGLANG_USE_CUDA_IPC_TRANSPORT.get())
self.assertIn("overrides", logs.output[0])
@@ -316,7 +344,9 @@ class TestMultimodalFeatureTransport(CustomTestCase):
with patch.dict(os.environ, {"SGLANG_USE_CUDA_IPC_TRANSPORT": "0"}):
server_args._handle_multimodal_feature_transport()
self.assertEqual(server_args.mm_feature_transport, "cpu")
self.assertEqual(
resolution_result(server_args, "mm_feature_transport"), "cpu"
)
self.assertFalse(envs.SGLANG_USE_CUDA_IPC_TRANSPORT.get())
@patch("sglang.srt.server_args.is_cuda", return_value=True)
@@ -329,7 +359,9 @@ class TestMultimodalFeatureTransport(CustomTestCase):
with self.assertNoLogs(server_args_module.logger, level="INFO"):
server_args._handle_multimodal_feature_transport()
self.assertEqual(server_args.mm_feature_transport, "cpu")
self.assertEqual(
resolution_result(server_args, "mm_feature_transport"), "cpu"
)
self.assertFalse(envs.SGLANG_USE_CUDA_IPC_TRANSPORT.get())
@patch("sglang.srt.server_args.is_cuda", return_value=True)
@@ -342,7 +374,9 @@ class TestMultimodalFeatureTransport(CustomTestCase):
with self.assertNoLogs(server_args_module.logger, level="INFO"):
server_args._handle_multimodal_feature_transport()
self.assertEqual(server_args.mm_feature_transport, "cpu")
self.assertEqual(
resolution_result(server_args, "mm_feature_transport"), "cpu"
)
self.assertFalse(envs.SGLANG_USE_CUDA_IPC_TRANSPORT.get())
@patch("sglang.srt.server_args.os.path.exists", return_value=True)
@@ -367,7 +401,9 @@ class TestMultimodalFeatureTransport(CustomTestCase):
with self.assertLogs(server_args_module.logger, level="INFO") as logs:
server_args._handle_multimodal_feature_transport()
self.assertEqual(server_args.mm_feature_transport, "cuda_vmm")
self.assertEqual(
resolution_result(server_args, "mm_feature_transport"), "cuda_vmm"
)
self.assertFalse(envs.SGLANG_USE_CUDA_IPC_TRANSPORT.get())
output = "\n".join(logs.output)
@@ -394,7 +430,7 @@ class TestMultimodalFeatureTransport(CustomTestCase):
with self.assertLogs(server_args_module.logger, level="INFO") as logs:
server_args._handle_multimodal_feature_transport()
self.assertEqual(server_args.mm_feature_transport, "cpu")
self.assertEqual(resolution_result(server_args, "mm_feature_transport"), "cpu")
self.assertIn("has not opted into CUDA VMM", "\n".join(logs.output))
@patch("sglang.srt.server_args.os.path.exists", return_value=False)
@@ -411,7 +447,9 @@ class TestMultimodalFeatureTransport(CustomTestCase):
with self.assertLogs(server_args_module.logger, level="INFO") as logs:
server_args._handle_multimodal_feature_transport()
self.assertEqual(server_args.mm_feature_transport, "cpu")
self.assertEqual(
resolution_result(server_args, "mm_feature_transport"), "cpu"
)
self.assertIn("no IMEX channel", "\n".join(logs.output))
@@ -427,7 +465,9 @@ class TestMultimodalFeatureTransport(CustomTestCase):
envs.SGLANG_USE_CUDA_IPC_TRANSPORT.clear()
server_args._handle_multimodal_feature_transport()
self.assertEqual(server_args.mm_feature_transport, "cpu")
self.assertEqual(
resolution_result(server_args, "mm_feature_transport"), "cpu"
)
self.assertFalse(envs.SGLANG_USE_CUDA_IPC_TRANSPORT.get())
@patch("sglang.srt.server_args.is_cuda", return_value=True)
@@ -439,7 +479,9 @@ class TestMultimodalFeatureTransport(CustomTestCase):
envs.SGLANG_USE_CUDA_IPC_TRANSPORT.clear()
server_args._handle_multimodal_feature_transport()
self.assertEqual(server_args.mm_feature_transport, "cpu")
self.assertEqual(
resolution_result(server_args, "mm_feature_transport"), "cpu"
)
self.assertFalse(envs.SGLANG_USE_CUDA_IPC_TRANSPORT.get())
@patch("sglang.srt.server_args.is_cuda", return_value=False)
@@ -474,7 +516,9 @@ class TestMultimodalFeatureTransport(CustomTestCase):
with self.assertLogs(server_args_module.logger, level="INFO") as logs:
server_args._handle_multimodal_feature_transport()
self.assertEqual(server_args.mm_feature_transport, "cuda_vmm")
self.assertEqual(
resolution_result(server_args, "mm_feature_transport"), "cuda_vmm"
)
self.assertFalse(envs.SGLANG_USE_CUDA_IPC_TRANSPORT.get())
output = "\n".join(logs.output)
@@ -555,15 +599,22 @@ class TestLoadBalanceMethod(unittest.TestCase):
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")
self.assertEqual(
resolution_result(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")
self.assertEqual(
resolution_result(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")
self.assertEqual(
resolution_result(server_args, "load_balance_method"), "round_robin"
)
def test_pd_prefill_dcp_warns_about_performance(self):
server_args = ServerArgs(
@@ -581,7 +632,7 @@ class TestLoadBalanceMethod(unittest.TestCase):
disaggregation_transfer_backend="mooncake",
dcp_size=4,
)
self.assertTrue(server_args.disable_radix_cache)
self.assertTrue(resolution_result(server_args, "disable_radix_cache"))
def test_pd_decode_dcp_rejects_unsupported_transfer_backend(self):
server_args = ServerArgs(
@@ -601,7 +652,7 @@ class TestLoadBalanceMethod(unittest.TestCase):
disaggregation_transfer_backend="fake",
dcp_size=4,
)
self.assertTrue(server_args.disable_radix_cache)
self.assertTrue(resolution_result(server_args, "disable_radix_cache"))
def test_pd_decode_dcp_rejects_radix_cache(self):
server_args = ServerArgs(
@@ -665,8 +716,11 @@ class TestLoadBalanceMethod(unittest.TestCase):
disaggregation_transfer_backend="mooncake_tcp",
)
self.assertFalse(server_args.disable_radix_cache)
self.assertEqual(server_args.disaggregation_transfer_backend, "mooncake")
self.assertFalse(resolution_result(server_args, "disable_radix_cache"))
self.assertEqual(
resolution_result(server_args, "disaggregation_transfer_backend"),
"mooncake",
)
class TestSkipTokenizerInit(unittest.TestCase):
@@ -681,8 +735,8 @@ class TestSkipTokenizerInit(unittest.TestCase):
server_args._handle_tokenizer_batching()
# Tokenizer fanout preserved; detokenizer coerced to 1 (no decode work).
self.assertEqual(server_args.tokenizer_worker_num, 4)
self.assertEqual(server_args.detokenizer_worker_num, 1)
self.assertEqual(resolution_result(server_args, "tokenizer_worker_num"), 4)
self.assertEqual(resolution_result(server_args, "detokenizer_worker_num"), 1)
class TestHiSparseDsaBackendPolicy(unittest.TestCase):
@@ -870,7 +924,7 @@ class TestFa4PageSizeAutoForce(CustomTestCase):
from sglang.srt.arg_groups.overrides import resolved_view
self.assertEqual(args.page_size, 1) # dual-apply retired: pristine
self.assertEqual(args.page_size, 1) # the field stays pristine
self.assertEqual(resolved_view(args).page_size, 128)
@patch("sglang.srt.arg_groups.overrides.is_sm100_supported", return_value=True)
@@ -883,7 +937,7 @@ class TestFa4PageSizeAutoForce(CustomTestCase):
from sglang.srt.arg_groups.overrides import resolved_view
self.assertEqual(args.page_size, 1) # dual-apply retired: pristine
self.assertEqual(args.page_size, 1) # the field stays pristine
self.assertEqual(resolved_view(args).page_size, 128)
@@ -918,12 +972,12 @@ class TestContextParallelServerArgs(CustomTestCase):
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)
self.assertTrue(resolution_result(args, "enable_prefill_cp"))
self.assertIsNone(resolution_result(args, "cp_strategy"))
server_args = self._new_cp_args(
enable_prefill_cp=args.enable_prefill_cp,
cp_strategy=args.cp_strategy,
enable_prefill_cp=resolution_result(args, "enable_prefill_cp"),
cp_strategy=resolution_result(args, "cp_strategy"),
)
with self.assertRaisesRegex(ValueError, "--cp-strategy"):
server_args._handle_context_parallelism()
@@ -940,16 +994,18 @@ class TestContextParallelServerArgs(CustomTestCase):
)
server_args = self._new_cp_args(
enable_dsa_prefill_context_parallel=(
args.enable_dsa_prefill_context_parallel
resolution_result(args, "enable_dsa_prefill_context_parallel")
),
dsa_prefill_cp_mode=args.dsa_prefill_cp_mode,
dsa_prefill_cp_mode=resolution_result(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")
self.assertTrue(resolution_result(server_args, "enable_prefill_cp"))
self.assertEqual(resolution_result(server_args, "cp_strategy"), "interleave")
self.assertEqual(
resolution_result(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(
@@ -961,10 +1017,18 @@ class TestContextParallelServerArgs(CustomTestCase):
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")
self.assertTrue(
resolution_result(server_args, "enable_dsa_prefill_context_parallel")
)
self.assertFalse(
resolution_result(server_args, "enable_prefill_context_parallel")
)
self.assertEqual(
resolution_result(server_args, "dsa_prefill_cp_mode"), "round-robin-split"
)
self.assertEqual(
resolution_result(server_args, "prefill_cp_mode"), "round-robin-split"
)
def test_context_parallel_handler_initializes_cp_strategy(self):
server_args = self._new_cp_args(
@@ -1049,15 +1113,25 @@ class TestContextParallelServerArgs(CustomTestCase):
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.assertTrue(resolution_result(server_args, "enable_prefill_cp"))
self.assertEqual(
server_args.enable_dsa_prefill_context_parallel, expect_dsa
resolution_result(server_args, "cp_strategy"), strategy
)
self.assertEqual(
server_args.enable_prefill_context_parallel, expect_generic
resolution_result(server_args, "dsa_prefill_cp_mode"), mode
)
self.assertEqual(
resolution_result(server_args, "prefill_cp_mode"), mode
)
self.assertEqual(
resolution_result(
server_args, "enable_dsa_prefill_context_parallel"
),
expect_dsa,
)
self.assertEqual(
resolution_result(server_args, "enable_prefill_context_parallel"),
expect_generic,
)
@@ -1308,7 +1382,7 @@ class TestSSLArgs(unittest.TestCase):
ssl_certfile="cert.pem",
enable_ssl_refresh=True,
)
self.assertTrue(server_args.enable_ssl_refresh)
self.assertTrue(resolution_result(server_args, "enable_ssl_refresh"))
class TestHiCacheArgs(unittest.TestCase):
@@ -1328,10 +1402,17 @@ class TestHiCacheArgs(unittest.TestCase):
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)
self.assertEqual(
resolution_result(args, "hicache_io_backend"), expected_io_backend
)
self.assertEqual(
resolution_result(args, "hicache_mem_layout"), expected_mem_layout
)
if expected_decode_backend is not None:
self.assertEqual(args.decode_attention_backend, expected_decode_backend)
self.assertEqual(
resolution_result(args, "decode_attention_backend"),
expected_decode_backend,
)
def test_hicache_io_backend_and_mem_layout_compatibility(self):
cases = [
@@ -1409,9 +1490,9 @@ class TestHiCacheArgs(unittest.TestCase):
)
args._handle_hicache()
self.assertEqual(args.hicache_io_backend, "kernel")
self.assertEqual(args.hicache_mem_layout, "page_first")
self.assertIsNone(args.decode_attention_backend)
self.assertEqual(resolution_result(args, "hicache_io_backend"), "kernel")
self.assertEqual(resolution_result(args, "hicache_mem_layout"), "page_first")
self.assertIsNone(resolution_result(args, "decode_attention_backend"))
def test_decode_offload_rejects_host_pool_retraction(self):
args = self._make_args(
@@ -1494,11 +1575,19 @@ class TestDecoupledSpecArgs(CustomTestCase):
"/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")
self.assertEqual(
resolution_result(server_args, "decoupled_spec_role"), "verifier"
)
self.assertEqual(
resolution_result(server_args, "decoupled_spec_bind_endpoint"),
"ipc:///tmp/v",
)
self.assertEqual(
resolution_result(server_args, "decoupled_spec_connect_endpoints"),
["ipc:///tmp/d"],
)
self.assertEqual(resolution_result(server_args, "decoupled_spec_rank"), 0)
self.assertEqual(resolution_result(server_args, "spec_trace_dir"), "/tmp/tr")
def test_decoupled_spec_role_rejects_invalid_choice(self):
with self.assertRaises(SystemExit):
@@ -1533,10 +1622,10 @@ class TestAdaptiveSpecArgs(CustomTestCase):
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)
self.assertTrue(resolution_result(args, "speculative_adaptive"))
self.assertEqual(resolution_result(args, "speculative_eagle_topk"), 1)
self.assertEqual(resolution_result(args, "speculative_num_steps"), 3)
self.assertEqual(resolution_result(args, "speculative_num_draft_tokens"), 4)
class TestWaterfillArgs(CustomTestCase):
@@ -1552,10 +1641,9 @@ class TestWaterfillArgs(CustomTestCase):
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)
self.assertTrue(resolution_result(server_args, "enforce_shared_experts_fusion"))
def test_waterfill_overrides_moe_a2a_backend_to_deepep(self):
server_args = ServerArgs(
@@ -1570,7 +1658,7 @@ class TestWaterfillArgs(CustomTestCase):
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)
self.assertTrue(resolution_result(server_args, "enforce_shared_experts_fusion"))
def test_waterfill_keeps_megamoe_backend(self):
server_args = ServerArgs(
@@ -1586,7 +1674,7 @@ class TestWaterfillArgs(CustomTestCase):
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)
self.assertTrue(resolution_result(server_args, "enforce_shared_experts_fusion"))
def test_waterfill_supports_deepep_low_latency_mode(self):
server_args = ServerArgs(
@@ -1598,9 +1686,9 @@ class TestWaterfillArgs(CustomTestCase):
# 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)
self.assertEqual(resolution_result(server_args, "deepep_mode"), "low_latency")
self.assertFalse(resolution_result(server_args, "disable_cuda_graph"))
self.assertTrue(resolution_result(server_args, "enforce_shared_experts_fusion"))
class TestPrefillOnlyDisableKvCache(unittest.TestCase):
@@ -1635,7 +1723,7 @@ class TestPrefillOnlyDisableKvCache(unittest.TestCase):
def test_valid_minimal_config_constructs(self):
sa = self._validate_prefill_only_args()
self.assertTrue(sa.prefill_only_disable_kv_cache)
self.assertTrue(resolution_result(sa, "prefill_only_disable_kv_cache"))
def test_rejects_when_not_embedding(self):
with self.assertRaisesRegex(ValueError, "requires --is-embedding"):
@@ -1725,15 +1813,27 @@ class TestCudaGraphDisaggregationRoles(CustomTestCase):
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)
self.assertFalse(resolution_result(args, "disable_cuda_graph"))
self.assertEqual(
resolution_result(args, "cuda_graph_config").decode.backend,
Backend.DISABLED,
)
self.assertEqual(
resolution_result(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)
self.assertEqual(
resolution_result(args, "cuda_graph_config").prefill.backend,
Backend.DISABLED,
)
self.assertNotEqual(
resolution_result(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"):
@@ -1744,10 +1844,12 @@ class TestCudaGraphDisaggregationRoles(CustomTestCase):
)
self.assertEqual(
args.cuda_graph_config.decode.backend, Backend.DISABLED
resolution_result(args, "cuda_graph_config").decode.backend,
Backend.DISABLED,
)
self.assertEqual(
args.cuda_graph_config.prefill.backend, Backend.DISABLED
resolution_result(args, "cuda_graph_config").prefill.backend,
Backend.DISABLED,
)
def test_cuda_graph_explicit_decode_backend_survives_prefill_role(self):
@@ -1756,7 +1858,9 @@ class TestCudaGraphDisaggregationRoles(CustomTestCase):
cuda_graph_backend_decode=Backend.FULL,
)
self.assertEqual(args.cuda_graph_config.decode.backend, Backend.FULL)
self.assertEqual(
resolution_result(args, "cuda_graph_config").decode.backend, Backend.FULL
)
self.assertIn((Phase.DECODE, "backend"), args._cuda_graph_config_locked)
@@ -1782,12 +1886,18 @@ class TestPrefillCudaGraphLoRACompatibility(CustomTestCase):
def test_enable_lora_keeps_breakable_prefill_graph(self):
args = self._handled_args(enable_lora=True)
self.assertEqual(args.cuda_graph_config.prefill.backend, Backend.BREAKABLE)
self.assertEqual(
resolution_result(args, "cuda_graph_config").prefill.backend,
Backend.BREAKABLE,
)
def test_lora_paths_keep_breakable_prefill_graph(self):
args = self._handled_args(lora_paths=["dummy/lora-adapter"])
self.assertEqual(args.cuda_graph_config.prefill.backend, Backend.BREAKABLE)
self.assertEqual(
resolution_result(args, "cuda_graph_config").prefill.backend,
Backend.BREAKABLE,
)
def test_lora_still_disables_tc_piecewise_prefill_graph(self):
# Pin the tc_piecewise LoRA rule itself, with the hardware rule
@@ -1811,7 +1921,10 @@ class TestPrefillCudaGraphLoRACompatibility(CustomTestCase):
):
args._disable_tc_piecewise_cudagraph_if_incompatible()
self.assertEqual(args.cuda_graph_config.prefill.backend, Backend.DISABLED)
self.assertEqual(
resolution_result(args, "cuda_graph_config").prefill.backend,
Backend.DISABLED,
)
class TestBreakableCudaGraphMultimodalAllowlist(CustomTestCase):
@@ -1840,7 +1953,10 @@ class TestBreakableCudaGraphMultimodalAllowlist(CustomTestCase):
is_multimodal=True,
allowlisted=False,
)
self.assertEqual(args.cuda_graph_config.prefill.backend, Backend.DISABLED)
self.assertEqual(
resolution_result(args, "cuda_graph_config").prefill.backend,
Backend.DISABLED,
)
def test_allowlisted_multimodal_arch_keeps_prefill_breakable(self):
args = self._handled_args(
@@ -1848,7 +1964,10 @@ class TestBreakableCudaGraphMultimodalAllowlist(CustomTestCase):
is_multimodal=True,
allowlisted=True,
)
self.assertEqual(args.cuda_graph_config.prefill.backend, Backend.BREAKABLE)
self.assertEqual(
resolution_result(args, "cuda_graph_config").prefill.backend,
Backend.BREAKABLE,
)
def test_allowlist_membership(self):
from sglang.srt.configs.model_config import (
@@ -2046,20 +2165,20 @@ class TestGrpcServerArgs(CustomTestCase):
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)
self.assertIsNone(resolution_result(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(resolution_result(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)
self.assertEqual(resolution_result(sa, "grpc_port"), 45000)
@staticmethod
def _sidecar_parser():
@@ -2196,20 +2315,20 @@ class TestGrpcServerArgs(CustomTestCase):
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)
self.assertEqual(resolution_result(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(resolution_result(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)
self.assertTrue(resolution_result(sa, "smg_grpc_mode"))
self.assertEqual(resolution_result(sa, "grpc_port"), 50051)
def test_native_grpc_rejects_multi_tokenizer(self):
sa = self._args(grpc_port=40000, tokenizer_worker_num=2)
@@ -2254,7 +2373,7 @@ class TestGrpcServerArgs(CustomTestCase):
tokenizer_manager=MagicMock(),
template_manager=MagicMock(),
scheduler_info={},
grpc_port=server_args.grpc_port,
grpc_port=resolution_result(server_args, "grpc_port"),
)
self.assertEqual(handle, "handle")