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
@@ -1,6 +1,7 @@
import unittest
from types import SimpleNamespace
from sglang.srt.arg_groups.overrides import resolution_result
from sglang.srt.arg_groups.speculative_hook import (
_handle_dspark,
_target_checkpoint_bundles_dspark_draft,
@@ -63,8 +64,13 @@ class TestDsparkDraftPathDefaulting(CustomTestCase):
model_path=_BUNDLED_MODEL_PATH, hf_config=_bundled_hf_config()
)
_handle_dspark(server_args)
self.assertEqual(server_args.speculative_draft_model_path, _BUNDLED_MODEL_PATH)
self.assertEqual(server_args.speculative_num_draft_tokens, 6)
self.assertEqual(
resolution_result(server_args, "speculative_draft_model_path"),
_BUNDLED_MODEL_PATH,
)
self.assertEqual(
resolution_result(server_args, "speculative_num_draft_tokens"), 6
)
def test_plain_target_without_draft_path_raises(self):
server_args = _make_dspark_server_args(
@@ -80,7 +86,7 @@ class TestDsparkDraftPathDefaulting(CustomTestCase):
server_args.speculative_draft_model_path = "deepseek-ai/some-other-dspark-draft"
_handle_dspark(server_args)
self.assertEqual(
server_args.speculative_draft_model_path,
resolution_result(server_args, "speculative_draft_model_path"),
"deepseek-ai/some-other-dspark-draft",
)
@@ -4,6 +4,7 @@ import unittest
from types import SimpleNamespace
from unittest.mock import patch
from sglang.srt.arg_groups.overrides import resolution_result
from sglang.srt.configs.embedding_model_spec import resolve_embedding_model_spec
from sglang.srt.configs.model_config import (
is_multimodal_piecewise_cuda_graph_supported,
@@ -90,7 +91,10 @@ class TestMultimodalPiecewiseCudaGraph(CustomTestCase):
):
args._apply_cuda_graph_compatibility()
self.assertEqual(args.cuda_graph_config.prefill.backend, Backend.TC_PIECEWISE)
self.assertEqual(
resolution_result(args, "cuda_graph_config").prefill.backend,
Backend.TC_PIECEWISE,
)
disable_if_incompatible.assert_called_once()
def test_trtllm_mla_stays_on_breakable(self):
@@ -118,7 +122,10 @@ class TestMultimodalPiecewiseCudaGraph(CustomTestCase):
):
args._apply_cuda_graph_compatibility()
self.assertEqual(args.cuda_graph_config.prefill.backend, Backend.BREAKABLE)
self.assertEqual(
resolution_result(args, "cuda_graph_config").prefill.backend,
Backend.BREAKABLE,
)
def test_explicit_tc_piecewise_overrides_trtllm_mla_default(self):
args = ServerArgs(model_path="dummy")
@@ -134,7 +141,10 @@ class TestMultimodalPiecewiseCudaGraph(CustomTestCase):
):
args._apply_cuda_graph_compatibility()
self.assertEqual(args.cuda_graph_config.prefill.backend, Backend.TC_PIECEWISE)
self.assertEqual(
resolution_result(args, "cuda_graph_config").prefill.backend,
Backend.TC_PIECEWISE,
)
def test_multimodal_inputs_keep_tc_piecewise_prefill_enabled(self):
runner = self._make_prefill_runner(Backend.TC_PIECEWISE)
@@ -178,10 +188,16 @@ class TestMultimodalPiecewiseCudaGraph(CustomTestCase):
):
args._handle_model_capability_adjustments()
self.assertTrue(args.disable_radix_cache)
self.assertEqual(args.chunked_prefill_size, -1)
self.assertEqual(args.cuda_graph_config.decode.backend, Backend.DISABLED)
self.assertEqual(args.cuda_graph_config.prefill.backend, Backend.BREAKABLE)
self.assertTrue(resolution_result(args, "disable_radix_cache"))
self.assertEqual(resolution_result(args, "chunked_prefill_size"), -1)
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_encoder_embedding_model_enables_embedding_mode_without_flag(self):
args = ServerArgs(model_path="dummy")
@@ -199,7 +215,7 @@ class TestMultimodalPiecewiseCudaGraph(CustomTestCase):
with patch.object(args, "get_model_config", return_value=args.model_config):
args._handle_model_capability_adjustments()
self.assertTrue(args.is_embedding)
self.assertTrue(resolution_result(args, "is_embedding"))
if __name__ == "__main__":
@@ -15,6 +15,7 @@ import zmq.asyncio
from fastapi import HTTPException
from PIL import Image
from sglang.srt.arg_groups.overrides import resolution_result
from sglang.srt.disaggregation.encoder.preprocessor import (
EncoderPreprocessor,
EncoderPreprocessResult,
@@ -195,7 +196,7 @@ def test_epd_rejection_reads_the_resolved_transfer_backend():
finally:
shutil.rmtree(config_dir, ignore_errors=True)
assert resolved.encoder_transfer_backend == "zmq_to_tokenizer"
assert resolution_result(resolved, "encoder_transfer_backend") == "zmq_to_tokenizer"
# Publish that record: the guard reads the resolved value out of the bags,
# so a raw record does not silently disable the rejection.
publish(resolved, role="tokenizer")
@@ -6,6 +6,7 @@ from unittest.mock import MagicMock, patch
import torch
from sglang.srt.arg_groups.overrides import resolution_result
from sglang.srt.environ import envs
from sglang.srt.runtime_context import get_context
from sglang.srt.server_args import ServerArgs
@@ -25,17 +26,20 @@ class TestMmProcessConfigValidation(CustomTestCase):
def test_valid_config_accepted(self):
args = self._validate_config({"image": {"max_pixels": 5000000}})
self.assertEqual(args.mm_process_config, {"image": {"max_pixels": 5000000}})
self.assertEqual(
resolution_result(args, "mm_process_config"),
{"image": {"max_pixels": 5000000}},
)
def test_empty_config_accepted(self):
args = self._validate_config({})
self.assertEqual(args.mm_process_config, {})
self.assertEqual(resolution_result(args, "mm_process_config"), {})
def test_none_config_defaults_to_empty_dict(self):
args = self._validate_config(None)
# None is kept as-is for dummy models (default happens after early return)
# but for real models it would be set to {}
self.assertIsNone(args.mm_process_config)
self.assertIsNone(resolution_result(args, "mm_process_config"))
def test_top_level_non_dict_rejected(self):
with self.assertRaises(TypeError) as ctx:
@@ -64,7 +68,7 @@ class TestMmProcessConfigValidation(CustomTestCase):
"audio": {"sample_rate": 16000},
}
args = self._validate_config(config)
self.assertEqual(args.mm_process_config, config)
self.assertEqual(resolution_result(args, "mm_process_config"), config)
class TestBaseProcessorConfigExtraction(CustomTestCase):
@@ -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")
@@ -1,6 +1,7 @@
import unittest
from types import SimpleNamespace
from sglang.srt.arg_groups.overrides import resolution_result
from sglang.srt.arg_groups.speculative_hook import handle_speculative_decoding
from sglang.srt.server_args import ServerArgs
from sglang.test.ci.ci_register import register_cpu_ci
@@ -33,18 +34,18 @@ def _make_spec_args(device: str, algorithm: str = "EAGLE", **overrides) -> Serve
class TestSpecCPUOverlapConstraint(CustomTestCase):
def test_cpu_eagle_forces_disable_overlap_schedule(self):
args = _make_spec_args(device="cpu")
self.assertFalse(args.disable_overlap_schedule)
self.assertFalse(resolution_result(args, "disable_overlap_schedule"))
handle_speculative_decoding(args)
self.assertTrue(args.disable_overlap_schedule)
self.assertTrue(resolution_result(args, "disable_overlap_schedule"))
def test_cpu_eagle3_forces_disable_overlap_schedule(self):
args = _make_spec_args(device="cpu", algorithm="EAGLE3")
handle_speculative_decoding(args)
self.assertTrue(args.disable_overlap_schedule)
self.assertTrue(resolution_result(args, "disable_overlap_schedule"))
def test_cpu_explicit_disable_overlap_is_preserved(self):
args = _make_spec_args(device="cpu", disable_overlap_schedule=True)
@@ -56,7 +57,7 @@ class TestSpecCPUOverlapConstraint(CustomTestCase):
) as logs:
handle_speculative_decoding(args)
self.assertTrue(args.disable_overlap_schedule)
self.assertTrue(resolution_result(args, "disable_overlap_schedule"))
self.assertFalse(
any("Overlap schedule" in message for message in logs.output),
f"hook warned about overriding an already-disabled overlap: {logs.output}",
@@ -68,7 +69,7 @@ class TestSpecCPUOverlapConstraint(CustomTestCase):
handle_speculative_decoding(args)
self.assertFalse(args.disable_overlap_schedule)
self.assertFalse(resolution_result(args, "disable_overlap_schedule"))
if __name__ == "__main__":
@@ -4,6 +4,7 @@ import unittest
from types import SimpleNamespace
from unittest.mock import MagicMock
from sglang.srt.arg_groups.overrides import resolution_result
from sglang.srt.arg_groups.speculative_hook import handle_speculative_decoding
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
from sglang.srt.speculative.spec_registry import (
@@ -237,7 +238,9 @@ class TestServerArgsHook(_RegistryIsolated):
handle_speculative_decoding(server_args)
self.assertEqual(server_args.speculative_algorithm, "MY_HANDLE_ARGS")
self.assertEqual(
resolution_result(server_args, "speculative_algorithm"), "MY_HANDLE_ARGS"
)
self.assertEqual(server_args.custom_spec_handle_seen, "MY_HANDLE_ARGS")
self.assertEqual(server_args.speculative_num_draft_tokens, 7)
+8 -2
View File
@@ -280,10 +280,16 @@ class TestPublishInstallsSlot(_IsolatedPublish):
set_global_server_args_for_scheduler(sa)
self.assertIs(get_server_args(), sa)
# Publishing is what resolved it; the handlers ahead of the dummy
# short-circuit still declare.
# short-circuit still declare. What they decided is the projection --
# the fields keep what the caller passed.
from sglang.srt.arg_groups.overrides import resolution_result
self.assertTrue(sa._resolved_overrides, "publishing declared nothing")
for source, declared in sa._resolved_overrides:
for field, value in declared.items():
self.assertEqual(getattr(sa, field), value, f"{source}: {field}")
self.assertEqual(
resolution_result(sa, field), value, f"{source}: {field}"
)
class TestGoldenModelOverrides(_IsolatedPublish):
@@ -31,11 +31,15 @@ class TestContextOverride(CustomTestCase):
def test_override_writes_bag_not_server_args(self):
sa = self._publish()
before = sa.hicache_ratio
# The published leaf, not the field: `hicache_ratio` is resolved by
# declaration, so the field still holds what the caller passed.
before = rc.get_memory().hicache_ratio
pristine = sa.hicache_ratio
rc.get_context().override("test", hicache_ratio=before + 1.0)
self.assertEqual(rc.get_memory().hicache_ratio, before + 1.0)
# server_args stays the pristine startup record.
self.assertEqual(sa.hicache_ratio, before)
# server_args stays the pristine startup record: the override does not
# touch it, and neither did resolution.
self.assertEqual(sa.hicache_ratio, pristine)
def test_override_routes_across_namespaces(self):
self._publish()
@@ -7,6 +7,7 @@ translates field annotations into argparse arguments.
import argparse
import unittest
from sglang.srt.arg_groups.overrides import resolution_result
from sglang.srt.server_args import ServerArgs
from sglang.srt.utils.common import configure_media_url_security
from sglang.test.ci.ci_register import register_cpu_ci
@@ -95,8 +96,10 @@ class TestServerArgsAnnotatedCli(CustomTestCase):
"32",
]
)
# The normalization is a declaration.
self.assertEqual(
sa.allowed_media_domains, ["127.0.0.1", "media.example.com"]
resolution_result(sa, "allowed_media_domains"),
["127.0.0.1", "media.example.com"],
)
self.assertEqual(sa.media_url_max_file_size_mb, 32)
finally:
@@ -134,10 +134,10 @@ _ENV_MATRIX = (({}, {"SGLANG_IS_IN_CI": "true"}),)
# are step-12 exposure like any other pair.
_PASSED = frozenset({"model_path", "device", "random_seed"})
# The reads that still take a value off the supplied instance. `initialize_moe_config`
# is handed the record until the replay goes away; the rest are pre-publish launcher
# reads.
_EXPOSED = {
("dllm/config.py", "max_running_requests"),
("dllm/config.py", "model_path"),
("speculative/spec_registry.py", "disable_overlap_schedule"),
("disaggregation/encoder/server.py", "model_loader_extra_config"),
("layers/moe/utils.py", "deepep_mode"),
("layers/moe/utils.py", "disable_shared_experts_fusion"),
@@ -150,40 +150,16 @@ _EXPOSED = {
("configs/embedding_model_spec.py", "disable_radix_cache"),
("configs/embedding_model_spec.py", "is_embedding"),
("configs/embedding_model_spec.py", "prefill_only_disable_kv_cache"),
("configs/model_config.py", "_speculative_draft_quantization_explicitly_set"),
("configs/model_config.py", "disable_hybrid_swa_memory"),
("configs/model_config.py", "dtype"),
("configs/model_config.py", "enable_multi_layer_eagle"),
("configs/model_config.py", "is_embedding"),
("configs/model_config.py", "model_path"),
("configs/model_config.py", "quantization"),
("configs/model_config.py", "speculative_algorithm"),
("configs/model_config.py", "speculative_draft_model_quantization"),
("dllm/config.py", "max_running_requests"),
("dllm/config.py", "model_path"),
("entrypoints/engine.py", "enable_symm_mem"),
("entrypoints/engine.py", "reasoning_parser"),
("entrypoints/engine.py", "tool_call_parser"),
("layers/cp/base.py", "attn_cp_size"),
("layers/cp/base.py", "cp_strategy"),
("layers/cp/base.py", "enable_prefill_cp"),
("layers/cp/bcg.py", "cp_strategy"),
("layers/cp/bcg.py", "enable_prefill_cp"),
("layers/flashinfer_comm_fusion.py", "flashinfer_allreduce_fusion_backend"),
("layers/moe/utils.py", "deepep_mode"),
("layers/moe/utils.py", "moe_a2a_backend"),
("layers/moe/utils.py", "moe_runner_backend"),
("layers/moe/utils.py", "quantization"),
("layers/moe/utils.py", "speculative_moe_runner_backend"),
("lora/marlin_lora_temp/policy.py", "lora_paths"),
("model_loader/expert_pack_runtime.py", "model_path"),
("model_loader/expert_pack_runtime.py", "tokenizer_path"),
("parser/template_detection.py", "model_path"),
("speculative/adaptive_spec_params.py", "speculative_algorithm"),
("speculative/adaptive_spec_params.py", "speculative_eagle_topk"),
("speculative/draft_worker_common.py", "speculative_draft_attention_backend"),
("speculative/spec_info.py", "enable_multi_layer_eagle"),
("speculative/spec_registry.py", "disable_overlap_schedule"),
("utils/common.py", "speculative_num_draft_tokens"),
("utils/common.py", "speculative_num_steps"),
("utils/hf_transformers/processor.py", "image_processor_backend"),
@@ -216,26 +192,19 @@ _EXPOSED_CUDA_ONLY: frozenset = frozenset()
# some code overrides post-publish. Each needs an ordering judgment, not a blanket
# conversion; the list exists so a new one is a decision made when it is written.
_OVERRIDDEN_AND_READ = {
("configs/model_config.py", "dtype"),
("configs/model_config.py", "model_path"),
("dllm/config.py", "model_path"),
("entrypoints/engine.py", "reasoning_parser"),
("entrypoints/engine.py", "tool_call_parser"),
("model_loader/expert_pack_runtime.py", "model_path"),
("weight_cache/daemon.py", "dp_size"),
("weight_cache/daemon.py", "dtype"),
("weight_cache/daemon.py", "ep_size"),
("weight_cache/daemon.py", "load_format"),
("weight_cache/daemon.py", "model_path"),
("configs/model_config.py", "dtype"),
("configs/model_config.py", "model_path"),
("mem_cache/pool_host/common.py", "hicache_storage_backend"),
("mem_cache/pool_host/common.py", "hicache_storage_backend_extra_config"),
("mem_cache/unified_radix_cache.py", "hicache_storage_backend"),
("mem_cache/unified_radix_cache.py", "hicache_storage_backend_extra_config"),
("mem_cache/unified_radix_cache.py", "hicache_storage_prefetch_policy"),
("mem_cache/unified_radix_cache.py", "hicache_write_policy"),
("parser/template_detection.py", "model_path"),
("utils/common.py", "speculative_num_draft_tokens"),
("utils/common.py", "speculative_num_steps"),
("weight_cache/daemon.py", "dp_size"),