diff --git a/python/sglang/srt/layers/attention/dsa_backend.py b/python/sglang/srt/layers/attention/dsa_backend.py index 3ac2ae68e..8f9929a95 100644 --- a/python/sglang/srt/layers/attention/dsa_backend.py +++ b/python/sglang/srt/layers/attention/dsa_backend.py @@ -424,7 +424,7 @@ class DeepseekSparseAttnBackend( # `flashmla_sparse_q8` is prefill-only (FP8 decode goes through # `flashmla_kv`); reject it as a decode backend, since argparse accepts it - # via the shared DSA_CHOICES list. + # via the shared CLI choices. if self.dsa_decode_impl == "flashmla_sparse_q8": raise ValueError( "--dsa-decode-backend flashmla_sparse_q8 is not supported: " diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index c30fcdd6a..be0455e03 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -114,8 +114,6 @@ logger = logging.getLogger(__name__) DEFAULT_UVICORN_ACCESS_LOG_EXCLUDE_PREFIXES = () SAMPLING_BACKEND_CHOICES = {"flashinfer", "pytorch", "ascend"} -if envs.SGLANG_KV_CANARY_ENABLE_TOKEN_ORACLE.get(): - SAMPLING_BACKEND_CHOICES.add("token_oracle") LOAD_FORMAT_CHOICES = [ "auto", @@ -142,6 +140,7 @@ LOAD_FORMAT_CHOICES = [ "private", "runai_streamer", ] +add_load_format_choices = LOAD_FORMAT_CHOICES.extend # NOTE: LoadFormat.IPC_CACHE intentionally has no public --load-format choice. # It is an internal dispatch format set automatically by ModelRunner when the # weight cache is enabled (weight_cache_mode != "off"). Exposing it as a CLI @@ -186,6 +185,7 @@ QUANTIZATION_CHOICES = [ "unquant", "humming", ] +add_quantization_method_choices = QUANTIZATION_CHOICES.extend ATTENTION_BACKEND_CHOICES = [ # Common @@ -218,6 +218,7 @@ ATTENTION_BACKEND_CHOICES = [ "ascend", "intel_xpu", ] +add_attention_backend_choices = ATTENTION_BACKEND_CHOICES.extend # trtllm_mha is valid for decode-only dense-MQA drafts. DFLASH rejects it # earlier when its per-layer attention requirements are not met. @@ -229,6 +230,7 @@ DRAFT_ATTENTION_BACKEND_CHOICES = [ "ascend", "trtllm_mha", ] +add_draft_attention_backend_choices = DRAFT_ATTENTION_BACKEND_CHOICES.extend # Attention backends whose kernels read the chunked prefix-cache layout. # Out-of-tree platforms may extend this list (via @@ -244,6 +246,9 @@ CHUNKED_PREFIX_CACHE_SUPPORTED_ATTENTION_BACKENDS = [ "trtllm_mla", "tokenspeed_mla", ] +add_chunked_prefix_cache_attention_backend = ( + CHUNKED_PREFIX_CACHE_SUPPORTED_ATTENTION_BACKENDS.append +) DETERMINISTIC_ATTENTION_BACKEND_CHOICES = [ "ascend", @@ -253,8 +258,14 @@ DETERMINISTIC_ATTENTION_BACKEND_CHOICES = [ "intel_xpu", "triton", ] +add_deterministic_attention_backend_choices = ( + DETERMINISTIC_ATTENTION_BACKEND_CHOICES.extend +) RADIX_SUPPORTED_DETERMINISTIC_ATTENTION_BACKEND = ["ascend", "fa3", "fa4", "triton"] +add_radix_supported_deterministic_attention_backend_choices = ( + RADIX_SUPPORTED_DETERMINISTIC_ATTENTION_BACKEND.extend +) DISAGG_TRANSFER_BACKEND_CHOICES = [ "mooncake", @@ -264,8 +275,10 @@ DISAGG_TRANSFER_BACKEND_CHOICES = [ "mori", "mooncake_tcp", ] +add_disagg_transfer_backend_choices = DISAGG_TRANSFER_BACKEND_CHOICES.extend GRAMMAR_BACKEND_CHOICES = ["xgrammar", "outlines", "llguidance", "none"] +add_grammar_backend_choices = GRAMMAR_BACKEND_CHOICES.extend # Placeholder token inserted between items in Multi-Item Scoring sequences: # queryitem1item2... Positions are pre-computed from item @@ -293,20 +306,7 @@ MOE_RUNNER_BACKEND_CHOICES = [ "hpc_ops", # HPC-Ops (https://github.com/Tencent/hpc-ops), FP8 MoE on Hopper (SM90) only "megamoe", ] - -MOE_A2A_BACKEND_CHOICES = [ - "none", - "deepep", - "mooncake", - "nixl", - "mori", - "ascend_fuseep", - "flashinfer", - "megamoe", - "deepep_v2", - "pplx", - "ascend_tp", -] +add_moe_runner_backend_choices = MOE_RUNNER_BACKEND_CHOICES.extend # These architectures take the A2A MoE path and skip post-expert all-reduce. _DEEPEP_V2_VALIDATED_ARCHITECTURES = frozenset( @@ -323,6 +323,7 @@ MXFP8_MOE_RUNNER_BACKEND_CHOICES = [ "flashinfer_trtllm", "flashinfer_trtllm_routed", ] +add_mxfp8_moe_runner_backend_choices = MXFP8_MOE_RUNNER_BACKEND_CHOICES.extend FP8_GEMM_RUNNER_BACKEND_CHOICES = [ "auto", @@ -335,6 +336,7 @@ FP8_GEMM_RUNNER_BACKEND_CHOICES = [ "triton", "aiter", ] +add_fp8_gemm_runner_backend_choices = FP8_GEMM_RUNNER_BACKEND_CHOICES.extend FP4_GEMM_RUNNER_BACKEND_CHOICES = [ "auto", @@ -344,77 +346,20 @@ FP4_GEMM_RUNNER_BACKEND_CHOICES = [ "flashinfer_trtllm", "marlin", ] - -BF16_GEMM_BACKEND_CHOICES = ["auto", "cutedsl", "gemv", "torch"] +add_fp4_gemm_runner_backend_choices = FP4_GEMM_RUNNER_BACKEND_CHOICES.extend RADIX_EVICTION_POLICY_CHOICES = ["lru", "lfu", "slru", "priority"] -RETRACTION_POLICY_CHOICES = ["length", "priority"] +add_radix_eviction_policy_choices = RADIX_EVICTION_POLICY_CHOICES.extend RL_ON_POLICY_TARGET_CHOICES = ["fsdp"] +add_rl_on_policy_target_choices = RL_ON_POLICY_TARGET_CHOICES.extend # Speculative algorithms whose verify forward presents a uniform per-request # token width, which is what the LoRA segment layout assumes. _LORA_SPEC_ALGORITHMS = ("EAGLE", "EAGLE3", "DFLASH", "DSPARK") -LORA_BACKEND_CHOICES = ["triton", "csgmv", "ascend", "torch_native"] - -ENCODER_TRANSFER_BACKEND_CHOICES = [ - "auto", - "zmq_to_scheduler", - "zmq_to_tokenizer", - "mooncake", -] - - -def resolve_encoder_transfer_backend( - backend: str, model_arch: str, tp_size: int -) -> str: - if backend != "auto": - return backend - if model_arch == "KimiK3ForConditionalGeneration" and tp_size > 1: - return "zmq_to_tokenizer" - return "zmq_to_scheduler" - - -DSA_PREFILL_CP_SPLIT_CHOICES = ["in-seq-split", "round-robin-split"] -NSA_PREFILL_CP_SPLIT_CHOICES = DSA_PREFILL_CP_SPLIT_CHOICES # deprecated alias - -PREFILL_CP_SPLIT_CHOICES = ["in-seq-split"] - DEFAULT_LORA_EVICTION_POLICY = "lru" -DSA_CHOICES = [ - "flashmla_sparse", - "flashmla_sparse_q8", - "flashmla_kv", - "flashmla_auto", - "flashinfer_sparse_mla", - "fa3", - "tilelang", - "aiter", - "trtllm", -] -NSA_CHOICES = DSA_CHOICES # deprecated alias - -DSV4_PREFILL_BACKEND_CHOICES = [ - "auto", - "flashmla_sparse", - "flashmla_sparse_q8", -] - -DSA_TOPK_BACKEND_CHOICES = ["sgl-kernel", "torch", "flashinfer"] - -DSA_PAGED_MQA_LOGITS_BACKEND_CHOICES = ["auto", "deepgemm", "cutedsl", "aiter"] - -MAMBA_RADIX_CACHE_STRATEGY_CHOICES = [ - "auto", - "no_buffer", - "extra_buffer", - "extra_buffer_lazy", -] - -MAMBA_BACKEND_CHOICES = ["triton", "flashinfer"] - LINEAR_ATTN_KERNEL_BACKEND_CHOICES = [ "triton", "cutedsl", @@ -425,71 +370,7 @@ LINEAR_ATTN_KERNEL_BACKEND_CHOICES = [ "helion", "intel_xpu", ] - - -# Allow external code to add more choices -def add_load_format_choices(choices): - LOAD_FORMAT_CHOICES.extend(choices) - - -def add_quantization_method_choices(choices): - QUANTIZATION_CHOICES.extend(choices) - - -def add_attention_backend_choices(choices): - ATTENTION_BACKEND_CHOICES.extend(choices) - - -def add_draft_attention_backend_choices(choices): - DRAFT_ATTENTION_BACKEND_CHOICES.extend(choices) - - -def add_chunked_prefix_cache_attention_backend(backend_name): - CHUNKED_PREFIX_CACHE_SUPPORTED_ATTENTION_BACKENDS.append(backend_name) - - -def add_deterministic_attention_backend_choices(choices): - DETERMINISTIC_ATTENTION_BACKEND_CHOICES.extend(choices) - - -def add_radix_supported_deterministic_attention_backend_choices(choices): - RADIX_SUPPORTED_DETERMINISTIC_ATTENTION_BACKEND.extend(choices) - - -def add_disagg_transfer_backend_choices(choices): - DISAGG_TRANSFER_BACKEND_CHOICES.extend(choices) - - -def add_grammar_backend_choices(choices): - GRAMMAR_BACKEND_CHOICES.extend(choices) - - -def add_moe_runner_backend_choices(choices): - MOE_RUNNER_BACKEND_CHOICES.extend(choices) - - -def add_mxfp8_moe_runner_backend_choices(choices): - MXFP8_MOE_RUNNER_BACKEND_CHOICES.extend(choices) - - -def add_fp8_gemm_runner_backend_choices(choices): - FP8_GEMM_RUNNER_BACKEND_CHOICES.extend(choices) - - -def add_fp4_gemm_runner_backend_choices(choices): - FP4_GEMM_RUNNER_BACKEND_CHOICES.extend(choices) - - -def add_radix_eviction_policy_choices(choices): - RADIX_EVICTION_POLICY_CHOICES.extend(choices) - - -def add_rl_on_policy_target_choices(choices): - RL_ON_POLICY_TARGET_CHOICES.extend(choices) - - -def add_linear_attn_kernel_backend_choices(choices): - LINEAR_ATTN_KERNEL_BACKEND_CHOICES.extend(choices) +add_linear_attn_kernel_backend_choices = LINEAR_ATTN_KERNEL_BACKEND_CHOICES.extend @dataclasses.dataclass @@ -935,7 +816,7 @@ class ServerArgs: "requests first, using the same priority direction as priority " "scheduling." ), - choices=RETRACTION_POLICY_CHOICES, + choices=["length", "priority"], ), NS("schedule"), ] = "length" @@ -1802,7 +1683,7 @@ class ServerArgs: Optional[str], Arg( help="Choose the kernels for sampling layers.", - choices=SAMPLING_BACKEND_CHOICES, + no_cli=True, resolvable=True, ), NS("exec.kernel"), @@ -1863,7 +1744,7 @@ class ServerArgs: Arg( help="Choose the backend for unquantized BF16 GEMM operations. Options: 'auto' (default; selects 'cutedsl' on SM10x GPUs, except deterministic inference selects 'torch'; otherwise uses cuBLAS via torch.nn.functional.linear), 'cutedsl' (SGLang JIT CuTe DSL TGV BF16 GEMM on SM10x; dispatches between the CuTe DSL kernel and cuBLAS), 'torch' (always uses cuBLAS via torch.nn.functional.linear).", cli_name="--bf16-gemm-backend", - choices=BF16_GEMM_BACKEND_CHOICES, + choices=["auto", "cutedsl", "gemv", "torch"], ), NS("exec.kernel"), ] = "auto" @@ -1871,7 +1752,17 @@ class ServerArgs: Optional[str], Arg( help="DSA (DeepSeek Sparse Attention) prefill backend. If not specified, auto-detects based on hardware and kv_cache_dtype.", - choices=DSA_CHOICES, + choices=[ + "flashmla_sparse", + "flashmla_sparse_q8", + "flashmla_kv", + "flashmla_auto", + "flashinfer_sparse_mla", + "fa3", + "tilelang", + "aiter", + "trtllm", + ], resolvable=True, ), NS("exec.kernel"), @@ -1884,7 +1775,7 @@ class ServerArgs: "'flashmla_sparse' use the existing BF16 sparse prefill path; " "'flashmla_sparse_q8' enables the Q8KV8 sparse prefill path." ), - choices=DSV4_PREFILL_BACKEND_CHOICES, + choices=["auto", "flashmla_sparse", "flashmla_sparse_q8"], ), NS("exec.kernel"), ] = "auto" @@ -1892,7 +1783,17 @@ class ServerArgs: Optional[str], Arg( help="DSA (DeepSeek Sparse Attention) decode backend. If not specified, auto-detects based on hardware and kv_cache_dtype.", - choices=DSA_CHOICES, + choices=[ + "flashmla_sparse", + "flashmla_sparse_q8", + "flashmla_kv", + "flashmla_auto", + "flashinfer_sparse_mla", + "fa3", + "tilelang", + "aiter", + "trtllm", + ], resolvable=True, ), NS("exec.kernel"), @@ -1901,7 +1802,7 @@ class ServerArgs: str, Arg( help="DSA indexer paged MQA logits kernel backend. Options: 'auto' (default; DeepGEMM on CUDA, aiter on ROCm), 'deepgemm', 'cutedsl' (CuTe DSL kernel, SM 100 (Blackwell) only; wins at low batch size and long context), 'aiter' (ROCm only).", - choices=DSA_PAGED_MQA_LOGITS_BACKEND_CHOICES, + choices=["auto", "deepgemm", "cutedsl", "aiter"], ), NS("exec.kernel"), ] = "auto" @@ -1909,7 +1810,7 @@ class ServerArgs: str, Arg( help="DSA indexer top-k backend for the target model. Options: 'sgl-kernel', 'torch', 'flashinfer'. The 'torch' backend currently requires SGLANG_DSA_FUSE_TOPK=false.", - choices=DSA_TOPK_BACKEND_CHOICES, + choices=["sgl-kernel", "torch", "flashinfer"], ), NS("exec.kernel"), ] = "sgl-kernel" @@ -1932,7 +1833,7 @@ class ServerArgs: str, Arg( help="Choose the kernel backend for Mamba SSM operations. Default is 'triton'. Options: 'triton' (default), 'flashinfer' (requires FlashInfer with Mamba support).", - choices=MAMBA_BACKEND_CHOICES, + choices=["triton", "flashinfer"], ), NS("exec.mamba"), ] = "triton" @@ -2279,7 +2180,7 @@ class ServerArgs: str, Arg( help="DSA indexer top-k backend for speculative draft workers. Options: 'sgl-kernel', 'torch', 'flashinfer'. The 'torch' backend currently requires SGLANG_DSA_FUSE_TOPK=false.", - choices=DSA_TOPK_BACKEND_CHOICES, + choices=["sgl-kernel", "torch", "flashinfer"], ), NS("spec"), ] = "sgl-kernel" @@ -2315,7 +2216,19 @@ class ServerArgs: Optional[str], Arg( help="Choose the backend for MoE A2A in speculative decoding", - choices=MOE_A2A_BACKEND_CHOICES, + choices=[ + "none", + "deepep", + "mooncake", + "nixl", + "mori", + "ascend_fuseep", + "flashinfer", + "megamoe", + "deepep_v2", + "pplx", + "ascend_tp", + ], resolvable=True, ), NS("spec"), @@ -2459,7 +2372,19 @@ class ServerArgs: ], Arg( help="Choose the backend for MoE A2A.", - choices=MOE_A2A_BACKEND_CHOICES, + choices=[ + "none", + "deepep", + "mooncake", + "nixl", + "mori", + "ascend_fuseep", + "flashinfer", + "megamoe", + "deepep_v2", + "pplx", + "ascend_tp", + ], resolvable=True, ), NS("exec.moe"), @@ -2688,7 +2613,7 @@ class ServerArgs: str, Arg( help="The strategy to use for mamba radix cache.", - choices=MAMBA_RADIX_CACHE_STRATEGY_CHOICES, + choices=["auto", "no_buffer", "extra_buffer", "extra_buffer_lazy"], resolvable=True, ), NS("exec.mamba"), @@ -3057,7 +2982,7 @@ class ServerArgs: str, Arg( help="Choose the kernel backend for multi-LoRA serving.", - choices=LORA_BACKEND_CHOICES, + choices=["triton", "csgmv", "ascend", "torch_native"], ), NS("lora"), ] = "csgmv" @@ -3312,10 +3237,10 @@ class ServerArgs: str, Arg( help="The backend for encoder disaggregation transfer. Auto selects a model- and TP-aware backend.", - choices=ENCODER_TRANSFER_BACKEND_CHOICES, + choices=["auto", "zmq_to_scheduler", "zmq_to_tokenizer", "mooncake"], ), NS("disagg"), - ] = ENCODER_TRANSFER_BACKEND_CHOICES[0] + ] = "auto" encoder_urls: A[List[str], "List of encoder server urls.", NS("disagg")] = ( dataclasses.field(default_factory=list) ) @@ -9579,6 +9504,17 @@ class ServerArgs: add_cli_args_from_dataclass(parser, ServerArgs) # --- Fields with dynamic choices (computed at add_cli_args time) --- + sampling_backend_choices = set(SAMPLING_BACKEND_CHOICES) + if envs.SGLANG_KV_CANARY_ENABLE_TOKEN_ORACLE.get(): + sampling_backend_choices.add("token_oracle") + parser.add_argument( + "--sampling-backend", + type=str, + choices=sampling_backend_choices, + default=ServerArgs.sampling_backend, + help="Choose the kernels for sampling layers.", + ) + reasoning_parser_choices = list(ReasoningParser.DetectorMap.keys()) parser.add_argument( "--reasoning-parser", @@ -9657,7 +9593,17 @@ class ServerArgs: new_flag="--dsa-prefill-backend", default=argparse.SUPPRESS, type=str, - choices=DSA_CHOICES, + choices=[ + "flashmla_sparse", + "flashmla_sparse_q8", + "flashmla_kv", + "flashmla_auto", + "flashinfer_sparse_mla", + "fa3", + "tilelang", + "aiter", + "trtllm", + ], help="[Deprecated] Use --dsa-prefill-backend instead.", ) parser.add_argument( @@ -9667,7 +9613,17 @@ class ServerArgs: new_flag="--dsa-decode-backend", default=argparse.SUPPRESS, type=str, - choices=DSA_CHOICES, + choices=[ + "flashmla_sparse", + "flashmla_sparse_q8", + "flashmla_kv", + "flashmla_auto", + "flashinfer_sparse_mla", + "fa3", + "tilelang", + "aiter", + "trtllm", + ], help="[Deprecated] Use --dsa-decode-backend instead.", ) parser.add_argument( @@ -9797,7 +9753,7 @@ class ServerArgs: new_flag="--cp-strategy", type=str, default=ServerArgs.dsa_prefill_cp_mode, - choices=DSA_PREFILL_CP_SPLIT_CHOICES, + choices=["in-seq-split", "round-robin-split"], help=( "[Deprecated] Use --cp-strategy {zigzag,interleave} instead. " "'in-seq-split' maps to 'zigzag'; 'round-robin-split' maps to " @@ -9811,7 +9767,7 @@ class ServerArgs: new_flag="--cp-strategy", type=str, default=argparse.SUPPRESS, - choices=DSA_PREFILL_CP_SPLIT_CHOICES, + choices=["in-seq-split", "round-robin-split"], help="[Deprecated] Use --cp-strategy instead.", ) parser.add_argument( @@ -9821,7 +9777,7 @@ class ServerArgs: new_flag="--cp-strategy", type=str, default=ServerArgs.prefill_cp_mode, - choices=PREFILL_CP_SPLIT_CHOICES, + choices=["in-seq-split"], help=( "[Deprecated] Use --cp-strategy {zigzag,interleave} instead. " "'in-seq-split' maps to 'zigzag'." @@ -10752,6 +10708,16 @@ class ServerArgs: return cfg.expert_balancedness_report_mode in ("prometheus", "both") +def resolve_encoder_transfer_backend( + backend: str, model_arch: str, tp_size: int +) -> str: + if backend != "auto": + return backend + if model_arch == "KimiK3ForConditionalGeneration" and tp_size > 1: + return "zmq_to_tokenizer" + return "zmq_to_scheduler" + + def compute_world_size(config) -> int: """Return the total GPU count across all data-parallel replicas. diff --git a/test/manual/test_dsa_alias_cli_registry_env.py b/test/manual/test_dsa_alias_cli_registry_env.py index 105b95193..066b871c2 100644 --- a/test/manual/test_dsa_alias_cli_registry_env.py +++ b/test/manual/test_dsa_alias_cli_registry_env.py @@ -21,29 +21,30 @@ sys.path.insert(0, os.path.join(os.path.dirname(__file__), "../../python")) class TestDSAChoicesAndFields(unittest.TestCase): - """Verify DSA_CHOICES constant and ServerArgs field renaming.""" + """Verify DSA CLI choices and ServerArgs field renaming.""" def setUp(self): - from sglang.srt.server_args import ( - DSA_CHOICES, - NSA_CHOICES, - ServerArgs, - ) + from sglang.srt.server_args import ServerArgs self.ServerArgs = ServerArgs - self.DSA_CHOICES = DSA_CHOICES - self.NSA_CHOICES = NSA_CHOICES + parser = argparse.ArgumentParser() + ServerArgs.add_cli_args(parser) + self.actions = { + option: action + for action in parser._actions + for option in action.option_strings + } def test_dsa_choices_is_canonical(self): - self.assertIn("fa3", self.DSA_CHOICES) - self.assertIn("tilelang", self.DSA_CHOICES) - self.assertIn("flashinfer_sparse_mla", self.DSA_CHOICES) + choices = self.actions["--dsa-prefill-backend"].choices + self.assertIn("fa3", choices) + self.assertIn("tilelang", choices) + self.assertIn("flashinfer_sparse_mla", choices) - def test_nsa_choices_is_alias(self): - self.assertIs( - self.NSA_CHOICES, - self.DSA_CHOICES, - "NSA_CHOICES must be the same object as DSA_CHOICES", + def test_nsa_choices_match_dsa_choices(self): + self.assertEqual( + self.actions["--nsa-prefill-backend"].choices, + self.actions["--dsa-prefill-backend"].choices, ) def test_serverargs_has_dsa_fields(self): diff --git a/test/registered/unit/disaggregation/test_kimi_k3_encoder_mode.py b/test/registered/unit/disaggregation/test_kimi_k3_encoder_mode.py index dd1a19ebb..523ca0278 100644 --- a/test/registered/unit/disaggregation/test_kimi_k3_encoder_mode.py +++ b/test/registered/unit/disaggregation/test_kimi_k3_encoder_mode.py @@ -104,8 +104,8 @@ def test_epd_rejection_reads_the_resolved_transfer_backend(): The record is produced by actual resolution -- a language-only Kimi-K3 launch at TP2, whose `encoder_transfer_backend` starts at the argument - default `"auto"` (`ENCODER_TRANSFER_BACKEND_CHOICES[0]`) and is filled in - by `resolve_encoder_transfer_backend` to `"zmq_to_tokenizer"`. The guard + default `"auto"` and is filled in by `resolve_encoder_transfer_backend` to + `"zmq_to_tokenizer"`. The guard reads that resolved value out of the published bags, so the rejection survives the record going raw: what a reader must never do is go back to the record for this field. diff --git a/test/registered/unit/server_args/test_server_args.py b/test/registered/unit/server_args/test_server_args.py index a4b937185..ef4729154 100644 --- a/test/registered/unit/server_args/test_server_args.py +++ b/test/registered/unit/server_args/test_server_args.py @@ -1,5 +1,5 @@ +import argparse import dataclasses -import importlib import json import os import socket @@ -2067,55 +2067,62 @@ class TestCutedslMoeMaxNumTokens(CustomTestCase): 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. + The choice set is finalized when CLI arguments are registered, so each + parser must reflect the environment at construction time. """ - 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 patch.dict(os.environ, {"SGLANG_KV_CANARY_ENABLE_TOKEN_ORACLE": "0"}): + with self.assertRaises(SystemExit): + server_args_module.prepare_server_args( + [ + "--model-path", + DEFAULT_SMALL_MODEL_NAME_FOR_TEST_QWEN, + "--sampling-backend", + "token_oracle", + ] + ) - with self.assertRaises(SystemExit): - reloaded.prepare_server_args( + def test_token_oracle_accepted_when_env_enabled(self): + with patch.dict(os.environ, {"SGLANG_KV_CANARY_ENABLE_TOKEN_ORACLE": "1"}): + parsed = server_args_module.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", ] ) - - 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") + def test_gate_is_recomputed_for_each_parser(self): + with patch.dict(os.environ, {"SGLANG_KV_CANARY_ENABLE_TOKEN_ORACLE": "1"}): + enabled_parser = argparse.ArgumentParser() + ServerArgs.add_cli_args(enabled_parser) + + with patch.dict(os.environ, {"SGLANG_KV_CANARY_ENABLE_TOKEN_ORACLE": "0"}): + disabled_parser = argparse.ArgumentParser() + ServerArgs.add_cli_args(disabled_parser) + + enabled_action = next( + action + for action in enabled_parser._actions + if action.dest == "sampling_backend" + ) + disabled_action = next( + action + for action in disabled_parser._actions + if action.dest == "sampling_backend" + ) + self.assertIn("token_oracle", enabled_action.choices) + self.assertNotIn("token_oracle", disabled_action.choices) + class TestDeepEPv2Args(CustomTestCase): """DeepEP v2 server-argument resolution and validation."""