[refactor] Retire the legacy config accessor and the remaining process singletons (#30493)
This commit is contained in:
@@ -5,7 +5,7 @@ from types import SimpleNamespace
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
from sglang.srt.runtime_context import get_parallel
|
||||
from sglang.srt.runtime_context import get_parallel, get_server_args
|
||||
|
||||
_parallel_override = get_parallel().override(attn_tp_size=1)
|
||||
_parallel_override.__enter__()
|
||||
@@ -26,7 +26,6 @@ from sglang.srt.model_executor.forward_context import (
|
||||
)
|
||||
from sglang.srt.server_args import (
|
||||
ServerArgs,
|
||||
get_global_server_args,
|
||||
set_global_server_args_for_scheduler,
|
||||
)
|
||||
from sglang.srt.utils import is_flashinfer_available
|
||||
@@ -222,7 +221,7 @@ class MockModelRunner:
|
||||
self.page_size = config["page_size"]
|
||||
|
||||
# Server args stub - needed by attention backends
|
||||
self.server_args = get_global_server_args()
|
||||
self.server_args = get_server_args()
|
||||
|
||||
# Model-config stub with MLA attributes
|
||||
self.model_config = type(
|
||||
|
||||
+7
-4
@@ -22,7 +22,6 @@ from sglang.srt.layers.attention.hybrid_linear_attn_backend import (
|
||||
HybridLinearAttnBackend,
|
||||
)
|
||||
from sglang.srt.model_executor.forward_batch_info import ForwardMode
|
||||
from sglang.srt.server_args import set_global_server_args_for_scheduler
|
||||
from sglang.test.ci.ci_register import register_cuda_ci
|
||||
from sglang.test.kits.attention_unittest.attention_methods.mla_attention import (
|
||||
DEFAULT_KV_LORA_RANK,
|
||||
@@ -46,9 +45,13 @@ class _ChunkKVMLARunner(MockMLAModelRunner):
|
||||
|
||||
def __init__(self, **kwargs):
|
||||
super().__init__(**kwargs)
|
||||
self.server_args.disable_chunked_prefix_cache = False
|
||||
self.server_args.flashinfer_mla_disable_ragged = False
|
||||
set_global_server_args_for_scheduler(self.server_args)
|
||||
# The fixture's config is already published; adjust it through the
|
||||
# audited entry point (bare writes raise under the strict guard).
|
||||
self.server_args.override(
|
||||
source="attention-unittest",
|
||||
disable_chunked_prefix_cache=False,
|
||||
flashinfer_mla_disable_ragged=False,
|
||||
)
|
||||
|
||||
|
||||
def _make_case() -> MLAAttentionCase:
|
||||
|
||||
@@ -418,7 +418,7 @@ class TestAiterAllreduceFusionGate(CustomTestCase):
|
||||
)
|
||||
)
|
||||
stack.enter_context(
|
||||
mock.patch.object(comm, "get_global_server_args", lambda: server_args)
|
||||
mock.patch.object(comm, "get_server_args", lambda: server_args)
|
||||
)
|
||||
from sglang.srt.runtime_context import get_flags
|
||||
|
||||
|
||||
@@ -7,9 +7,9 @@ import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
from sglang.srt.layers.logits_processor import LogitsProcessor
|
||||
from sglang.srt.runtime_context import get_server_args
|
||||
from sglang.srt.server_args import (
|
||||
ServerArgs,
|
||||
get_global_server_args,
|
||||
set_global_server_args_for_scheduler,
|
||||
)
|
||||
from sglang.srt.utils import get_device
|
||||
@@ -44,7 +44,7 @@ class TestLMHeadFP32(unittest.TestCase):
|
||||
|
||||
def _make_logprocessor(self, vocab_size, enable_fp32):
|
||||
set_global_server_args_for_scheduler(ServerArgs(model_path="dummy"))
|
||||
get_global_server_args().enable_fp32_lm_head = enable_fp32
|
||||
get_server_args().enable_fp32_lm_head = enable_fp32
|
||||
cfg = SimpleNamespace(vocab_size=vocab_size, final_logit_softcapping=None)
|
||||
return LogitsProcessor(cfg, skip_all_gather=True, logit_scale=None)
|
||||
|
||||
|
||||
@@ -4,7 +4,7 @@
|
||||
``ForwardBatch.num_token_non_padded`` is a scalar tensor on the model device
|
||||
(see ``ForwardBatch.compute``, which does ``.to(device, ...)``). The eager TBO
|
||||
split path already honors this -- ``compute_tbo_children_num_token_non_padded_raw``
|
||||
moves the tensor to ``get_global_server_args().device`` -- but
|
||||
moves the tensor to ``get_server_args().device`` -- but
|
||||
``TboCudaGraphRunnerPlugin`` preallocated its persistent buffer with a bare
|
||||
``torch.zeros((2,), dtype=torch.int32)``, leaving it on CPU.
|
||||
|
||||
@@ -39,7 +39,7 @@ class TestTboCudaGraphNumTokenDevice(CustomTestCase):
|
||||
# Use 'meta' so the configured device differs from the implicit CPU
|
||||
# default; a bare torch.zeros() would leave the buffer on CPU and fail.
|
||||
fake_args = SimpleNamespace(device="meta")
|
||||
with patch.object(tbo, "get_global_server_args", lambda: fake_args):
|
||||
with patch.object(tbo, "get_server_args", lambda: fake_args):
|
||||
plugin = TboCudaGraphRunnerPlugin()
|
||||
|
||||
buf = plugin._tbo_children_num_token_non_padded
|
||||
@@ -51,7 +51,7 @@ class TestTboCudaGraphNumTokenDevice(CustomTestCase):
|
||||
# Both the preallocated cuda-graph buffer and the eager split tensor must
|
||||
# land on the same (model) device, matching ForwardBatch's contract.
|
||||
fake_args = SimpleNamespace(device="meta")
|
||||
with patch.object(tbo, "get_global_server_args", lambda: fake_args):
|
||||
with patch.object(tbo, "get_server_args", lambda: fake_args):
|
||||
eager = (
|
||||
TboForwardBatchPreparer.compute_tbo_children_num_token_non_padded_raw(
|
||||
tbo_split_token_index=3, num_token_non_padded=8
|
||||
@@ -68,7 +68,7 @@ class TestTboCudaGraphNumTokenDevice(CustomTestCase):
|
||||
# value_a = min(split, n); value_b = max(0, n - split). Computed on CPU
|
||||
# so the values are materializable.
|
||||
fake_args = SimpleNamespace(device="cpu")
|
||||
with patch.object(tbo, "get_global_server_args", lambda: fake_args):
|
||||
with patch.object(tbo, "get_server_args", lambda: fake_args):
|
||||
eager = (
|
||||
TboForwardBatchPreparer.compute_tbo_children_num_token_non_padded_raw(
|
||||
tbo_split_token_index=3, num_token_non_padded=8
|
||||
|
||||
@@ -39,7 +39,7 @@ def _make_target_verify_batch(bs: int) -> ForwardBatch:
|
||||
def _filter(batch: ForwardBatch, *, lo: int, hi: int) -> ForwardBatch:
|
||||
fake_args = SimpleNamespace(moe_dense_tp_size=None, attention_backend="fa3")
|
||||
with get_parallel().override(attn_tp_size=1), patch.object(
|
||||
tbo, "get_global_server_args", lambda: fake_args
|
||||
tbo, "get_server_args", lambda: fake_args
|
||||
):
|
||||
return TboForwardBatchPreparer.filter_batch(
|
||||
batch,
|
||||
|
||||
@@ -1257,7 +1257,7 @@ class TestMlxOverlapScheduler(unittest.TestCase):
|
||||
logits_output = SimpleNamespace(customized_info=None)
|
||||
original_release = batch_result_processor_module.release_kv_cache
|
||||
original_get_indexer = batch_result_processor_module.get_global_indexer_capturer
|
||||
original_get_server_args = batch_result_processor_module.get_global_server_args
|
||||
original_get_server_args = batch_result_processor_module.get_server_args
|
||||
|
||||
def fake_release_kv_cache(release_req, tree_cache, is_insert=False):
|
||||
events.append(("release", release_req.rid))
|
||||
@@ -1265,7 +1265,7 @@ class TestMlxOverlapScheduler(unittest.TestCase):
|
||||
|
||||
batch_result_processor_module.release_kv_cache = fake_release_kv_cache
|
||||
batch_result_processor_module.get_global_indexer_capturer = lambda: None
|
||||
batch_result_processor_module.get_global_server_args = lambda: SimpleNamespace(
|
||||
batch_result_processor_module.get_server_args = lambda: SimpleNamespace(
|
||||
enable_mamba_extra_buffer_lazy=lambda: False
|
||||
)
|
||||
try:
|
||||
@@ -1279,9 +1279,7 @@ class TestMlxOverlapScheduler(unittest.TestCase):
|
||||
batch_result_processor_module.get_global_indexer_capturer = (
|
||||
original_get_indexer
|
||||
)
|
||||
batch_result_processor_module.get_global_server_args = (
|
||||
original_get_server_args
|
||||
)
|
||||
batch_result_processor_module.get_server_args = original_get_server_args
|
||||
|
||||
self.assertEqual(
|
||||
events,
|
||||
|
||||
@@ -56,10 +56,10 @@ from sglang.srt.mem_cache.unified_radix_cache import (
|
||||
UnifiedRadixCache,
|
||||
UnifiedTreeNode,
|
||||
)
|
||||
from sglang.srt.runtime_context import get_server_args
|
||||
from sglang.srt.sampling.sampling_params import SamplingParams
|
||||
from sglang.srt.server_args import (
|
||||
ServerArgs,
|
||||
get_global_server_args,
|
||||
set_global_server_args_for_scheduler,
|
||||
)
|
||||
from sglang.srt.utils import get_device
|
||||
@@ -942,13 +942,13 @@ class UnifiedRadixCacheSuite:
|
||||
req.mamba_last_track_seqlen = kv_len
|
||||
req.reasoning_tokens = 1
|
||||
|
||||
get_global_server_args().strip_thinking_cache = True
|
||||
get_server_args().strip_thinking_cache = True
|
||||
try:
|
||||
avail_before = allocator.available_size()
|
||||
cache.cache_finished_req(req, is_insert=True)
|
||||
start_p, end_p = req.pop_overallocated_kv_cache()
|
||||
finally:
|
||||
get_global_server_args().strip_thinking_cache = False
|
||||
get_server_args().strip_thinking_cache = False
|
||||
if ps > 1:
|
||||
start_p = ((start_p + ps - 1) // ps) * ps
|
||||
if start_p < end_p:
|
||||
@@ -3527,7 +3527,7 @@ class UnifiedRadixCacheSuite:
|
||||
if not self.cfg.has_mamba or self.cfg.has_swa or self.cfg.page_size != 1:
|
||||
self.skipTest("requires page_size=1 Full+Mamba")
|
||||
cache, allocator, req_to_token_pool = build_fixture(self.cfg)
|
||||
chunk_size = get_global_server_args().mamba_cache_chunk_size
|
||||
chunk_size = get_server_args().mamba_cache_chunk_size
|
||||
tokens = self._make_seq(1, chunk_size + 1)
|
||||
self._insert(cache, allocator, req_to_token_pool, tokens)
|
||||
leaf = cache.match_prefix(
|
||||
|
||||
@@ -21,6 +21,7 @@ from sglang.srt.observability.trace import (
|
||||
TraceThreadContext,
|
||||
TraceThreadInfo,
|
||||
extract_trace_headers,
|
||||
get_global_trace_level,
|
||||
get_global_tracing_enabled,
|
||||
process_tracing_init,
|
||||
set_global_trace_level,
|
||||
@@ -51,19 +52,29 @@ class TestTraceFunctions(unittest.TestCase):
|
||||
self.assertEqual(extract_trace_headers({}), {})
|
||||
|
||||
def test_set_global_trace_level(self):
|
||||
orig = mod.global_trace_level
|
||||
set_global_trace_level(5)
|
||||
self.assertEqual(mod.global_trace_level, 5)
|
||||
mod.global_trace_level = orig
|
||||
from sglang.srt.runtime_context import get_resources
|
||||
|
||||
orig = get_resources().trace_level
|
||||
try:
|
||||
set_global_trace_level(5)
|
||||
self.assertEqual(get_global_trace_level(), 5)
|
||||
finally:
|
||||
get_resources().trace_level = orig
|
||||
|
||||
def test_global_trace_level_env_var(self):
|
||||
import importlib
|
||||
# The level lives on ctx.resources and is seeded lazily from the env
|
||||
# on first read after a reset (no module reload involved).
|
||||
from sglang.srt.runtime_context import get_resources
|
||||
|
||||
with patch.dict(os.environ, {"SGLANG_TRACE_LEVEL": "2"}):
|
||||
importlib.reload(mod)
|
||||
self.assertEqual(mod.global_trace_level, 2)
|
||||
importlib.reload(mod) # restore default (SGLANG_TRACE_LEVEL unset → 3)
|
||||
self.assertEqual(mod.global_trace_level, 3)
|
||||
orig = get_resources().trace_level
|
||||
try:
|
||||
with patch.dict(os.environ, {"SGLANG_TRACE_LEVEL": "2"}):
|
||||
get_resources().trace_level = None
|
||||
self.assertEqual(get_global_trace_level(), 2)
|
||||
get_resources().trace_level = None # SGLANG_TRACE_LEVEL unset → 3
|
||||
self.assertEqual(get_global_trace_level(), 3)
|
||||
finally:
|
||||
get_resources().trace_level = orig
|
||||
|
||||
def test_get_global_tracing_enabled(self):
|
||||
self.assertEqual(get_global_tracing_enabled(), mod.opentelemetry_initialized)
|
||||
@@ -244,7 +255,9 @@ class TestTraceReqContextEnabled(unittest.TestCase):
|
||||
self.orig_initialized = mod.opentelemetry_initialized
|
||||
self.orig_tracer = mod.tracer
|
||||
self.orig_threads = mod.threads_info.copy()
|
||||
self.orig_level = mod.global_trace_level
|
||||
from sglang.srt.runtime_context import get_resources
|
||||
|
||||
self.orig_level = get_resources().trace_level
|
||||
|
||||
# Reset OTel global TracerProvider so set_tracer_provider works each test
|
||||
otel_trace._TRACER_PROVIDER_SET_ONCE._done = False
|
||||
@@ -254,14 +267,16 @@ class TestTraceReqContextEnabled(unittest.TestCase):
|
||||
otel_trace.set_tracer_provider(self.provider)
|
||||
mod.opentelemetry_initialized = True
|
||||
mod.tracer = otel_trace.get_tracer("test")
|
||||
mod.global_trace_level = 3
|
||||
set_global_trace_level(3)
|
||||
|
||||
def tearDown(self):
|
||||
mod.opentelemetry_initialized = self.orig_initialized
|
||||
mod.tracer = self.orig_tracer
|
||||
mod.threads_info.clear()
|
||||
mod.threads_info.update(self.orig_threads)
|
||||
mod.global_trace_level = self.orig_level
|
||||
from sglang.srt.runtime_context import get_resources
|
||||
|
||||
get_resources().trace_level = self.orig_level
|
||||
|
||||
def test_trace_set_thread_info(self):
|
||||
trace_set_thread_info("scheduler", tp_rank=0, dp_rank=0)
|
||||
|
||||
@@ -602,10 +602,18 @@ class TestResolveAutoParsers(unittest.TestCase):
|
||||
|
||||
qwen3_template = "{% set enable_thinking = enable_thinking if enable_thinking is defined else true %}"
|
||||
|
||||
class _Args(SimpleNamespace):
|
||||
# Write-through override, per the runtime-context testing idiom:
|
||||
# production adjusts parsers through override(source, ...), so the
|
||||
# stand-in needs the method (a bare SimpleNamespace would raise).
|
||||
def override(self, source, **fields):
|
||||
for key, value in fields.items():
|
||||
setattr(self, key, value)
|
||||
|
||||
def _make_server_args(
|
||||
self, reasoning_parser=None, tool_call_parser=None, chat_template=None
|
||||
):
|
||||
return SimpleNamespace(
|
||||
return self._Args(
|
||||
reasoning_parser=reasoning_parser,
|
||||
tool_call_parser=tool_call_parser,
|
||||
model_path="Qwen/Qwen3-0.6B",
|
||||
@@ -650,12 +658,8 @@ class TestResolveAutoParsers(unittest.TestCase):
|
||||
self.assertEqual(args.tool_call_parser, "qwen")
|
||||
|
||||
def test_nonexistent_model_disables_both_parsers(self):
|
||||
args = SimpleNamespace(
|
||||
reasoning_parser="auto",
|
||||
tool_call_parser="auto",
|
||||
model_path="nonexistent/model-does-not-exist-xyz",
|
||||
trust_remote_code=False,
|
||||
)
|
||||
args = self._make_server_args(reasoning_parser="auto", tool_call_parser="auto")
|
||||
args.model_path = "nonexistent/model-does-not-exist-xyz"
|
||||
with _patch_hf_transformers_utils(
|
||||
Mock(side_effect=RuntimeError("tokenizer unavailable")),
|
||||
Mock(side_effect=RuntimeError("config unavailable")),
|
||||
|
||||
@@ -452,7 +452,7 @@ class TestFromScheduleBatch(CustomTestCase):
|
||||
req.tokenizer.eos_token_id = eos_id
|
||||
return req
|
||||
|
||||
@patch("sglang.srt.sampling.sampling_batch_info.get_global_server_args")
|
||||
@patch("sglang.srt.sampling.sampling_batch_info.get_server_args")
|
||||
def test_basic_construction(self, mock_server_args):
|
||||
"""Test that from_schedule_batch correctly extracts sampling params from requests."""
|
||||
mock_server_args.return_value.enable_deterministic_inference = False
|
||||
@@ -469,7 +469,7 @@ class TestFromScheduleBatch(CustomTestCase):
|
||||
self.assertAlmostEqual(info.top_ps[0].item(), 0.9, places=5)
|
||||
self.assertEqual(info.top_ks[0].item(), 50)
|
||||
|
||||
@patch("sglang.srt.sampling.sampling_batch_info.get_global_server_args")
|
||||
@patch("sglang.srt.sampling.sampling_batch_info.get_server_args")
|
||||
def test_greedy_detection(self, mock_server_args):
|
||||
"""Test that top_k=1 sets is_all_greedy=True."""
|
||||
mock_server_args.return_value.enable_deterministic_inference = False
|
||||
@@ -482,7 +482,7 @@ class TestFromScheduleBatch(CustomTestCase):
|
||||
info = SamplingBatchInfo.from_schedule_batch(batch, VOCAB_SIZE)
|
||||
self.assertTrue(info.is_all_greedy)
|
||||
|
||||
@patch("sglang.srt.sampling.sampling_batch_info.get_global_server_args")
|
||||
@patch("sglang.srt.sampling.sampling_batch_info.get_server_args")
|
||||
def test_logit_bias_construction(self, mock_server_args):
|
||||
"""Test that logit_bias dict is converted to a tensor with correct values."""
|
||||
mock_server_args.return_value.enable_deterministic_inference = False
|
||||
@@ -498,7 +498,7 @@ class TestFromScheduleBatch(CustomTestCase):
|
||||
self.assertAlmostEqual(info.logit_bias[0, 10].item(), -1.0)
|
||||
self.assertAlmostEqual(info.logit_bias[0, 0].item(), 0.0)
|
||||
|
||||
@patch("sglang.srt.sampling.sampling_batch_info.get_global_server_args")
|
||||
@patch("sglang.srt.sampling.sampling_batch_info.get_server_args")
|
||||
def test_deterministic_seed(self, mock_server_args):
|
||||
"""Test that explicit seed=123 is kept and missing seed defaults to 42."""
|
||||
mock_server_args.return_value.enable_deterministic_inference = True
|
||||
@@ -513,7 +513,7 @@ class TestFromScheduleBatch(CustomTestCase):
|
||||
self.assertEqual(info.sampling_seed[0].item(), 123)
|
||||
self.assertEqual(info.sampling_seed[1].item(), 42) # default
|
||||
|
||||
@patch("sglang.srt.sampling.sampling_batch_info.get_global_server_args")
|
||||
@patch("sglang.srt.sampling.sampling_batch_info.get_server_args")
|
||||
def test_from_schedule_batch_sampling_flags(self, mock_server_args):
|
||||
"""Test that sampling flags (need_top_p/top_k/min_p) are set correctly."""
|
||||
mock_server_args.return_value.enable_deterministic_inference = False
|
||||
@@ -529,7 +529,7 @@ class TestFromScheduleBatch(CustomTestCase):
|
||||
self.assertTrue(info.need_min_p_sampling) # 0.1 > 0
|
||||
self.assertFalse(info.is_all_greedy) # top_k=50 > 1
|
||||
|
||||
@patch("sglang.srt.sampling.sampling_batch_info.get_global_server_args")
|
||||
@patch("sglang.srt.sampling.sampling_batch_info.get_server_args")
|
||||
def test_no_logit_bias_when_all_none(self, mock_server_args):
|
||||
"""Test that logit_bias stays None when no request has logit_bias set."""
|
||||
mock_server_args.return_value.enable_deterministic_inference = False
|
||||
@@ -542,7 +542,7 @@ class TestFromScheduleBatch(CustomTestCase):
|
||||
info = SamplingBatchInfo.from_schedule_batch(batch, VOCAB_SIZE)
|
||||
self.assertIsNone(info.logit_bias)
|
||||
|
||||
@patch("sglang.srt.sampling.sampling_batch_info.get_global_server_args")
|
||||
@patch("sglang.srt.sampling.sampling_batch_info.get_server_args")
|
||||
def test_custom_logit_processor_merging(self, mock_server_args):
|
||||
"""Test deserialization and merging of custom logit processors."""
|
||||
from sglang.srt.sampling.custom_logit_processor import (
|
||||
|
||||
@@ -25,7 +25,9 @@ _SRT_ROOT = Path(next(iter(sglang.srt.__path__)))
|
||||
# Baselines counted over python/sglang/srt/**/*.py, including each function's
|
||||
# own def line. Ratchet: decrease-only.
|
||||
_RATCHETS = [
|
||||
("get_global_server_args", r"\bget_global_server_args\s*\(", 279),
|
||||
# Down to the shim definition itself; every call-site now goes through
|
||||
# runtime_context.get_server_args().
|
||||
("get_global_server_args", r"\bget_global_server_args\s*\(", 1),
|
||||
(
|
||||
"set_global_server_args_for_*",
|
||||
r"\bset_global_server_args_for_(?:scheduler|tokenizer)\s*\(",
|
||||
|
||||
@@ -312,12 +312,11 @@ class TestGoldenModelOverrides(_IsolatedPublish):
|
||||
|
||||
def _publish(self, server_args):
|
||||
from sglang.srt.server_args import (
|
||||
get_global_server_args,
|
||||
set_global_server_args_for_scheduler,
|
||||
)
|
||||
|
||||
set_global_server_args_for_scheduler(server_args)
|
||||
return get_global_server_args()
|
||||
return get_server_args()
|
||||
|
||||
def test_mistral_large3_forces_bfloat16(self):
|
||||
sa = self._construct("MistralLarge3ForCausalLM", "mistral")
|
||||
|
||||
@@ -5,6 +5,7 @@ from sglang.test.ci.ci_register import register_cpu_ci
|
||||
register_cpu_ci(est_time=5, suite="base-a-test-cpu")
|
||||
|
||||
import dataclasses
|
||||
import os
|
||||
import unittest
|
||||
from unittest.mock import patch
|
||||
|
||||
@@ -207,6 +208,86 @@ class TestServerArgsOwnership(_IsolatedServerArgs):
|
||||
with self.assertRaises(ValueError):
|
||||
get_server_args()
|
||||
|
||||
|
||||
class TestServerArgsScopedOverride(_IsolatedServerArgs):
|
||||
"""ctx.override_server_args: the config tier's scoped test override —
|
||||
tests force execution paths by overriding the context, not by
|
||||
hand-building and publishing config objects."""
|
||||
|
||||
def test_install_publishes_fresh_config_with_fields(self):
|
||||
reset_context()
|
||||
override = get_context().override_server_args(
|
||||
attention_backend="triton", chunked_prefill_size=-1
|
||||
)
|
||||
published = override.install()
|
||||
self.assertIs(get_server_args(), published)
|
||||
self.assertEqual(published.attention_backend, "triton")
|
||||
self.assertEqual(published.chunked_prefill_size, -1)
|
||||
# unnamed fields keep their dataclass defaults
|
||||
self.assertEqual(published.tp_size, 1)
|
||||
|
||||
def test_fields_carry_provenance(self):
|
||||
published = get_context().override_server_args(tp_size=4).install()
|
||||
self.assertIn(("test-override", {"tp_size": 4}), published._runtime_mutations)
|
||||
|
||||
def test_restore_reinstates_previous_publish(self):
|
||||
previous = object()
|
||||
get_context().set_server_args(previous)
|
||||
override = get_context().override_server_args(tp_size=8)
|
||||
override.install()
|
||||
self.assertEqual(get_server_args().tp_size, 8)
|
||||
override.restore()
|
||||
self.assertIs(get_server_args(), previous)
|
||||
|
||||
def test_restore_reinstates_the_empty_slot(self):
|
||||
reset_context()
|
||||
with get_context().override_server_args():
|
||||
get_server_args() # published inside the scope
|
||||
with self.assertRaises(ValueError):
|
||||
get_server_args()
|
||||
|
||||
def test_nesting_restores_in_order(self):
|
||||
reset_context()
|
||||
with get_context().override_server_args(tp_size=2) as outer:
|
||||
with get_context().override_server_args(tp_size=4):
|
||||
self.assertEqual(get_server_args().tp_size, 4)
|
||||
self.assertIs(get_server_args(), outer)
|
||||
self.assertEqual(get_server_args().tp_size, 2)
|
||||
|
||||
def test_private_attribute_seeding(self):
|
||||
# Property caches (e.g. _mamba_cache_chunk_size) are seeded through
|
||||
# the same call; the strict guard exempts underscore names.
|
||||
published = (
|
||||
get_context().override_server_args(_mamba_cache_chunk_size=64).install()
|
||||
)
|
||||
self.assertEqual(published.mamba_cache_chunk_size, 64)
|
||||
|
||||
def test_installed_config_arms_the_strict_guard(self):
|
||||
# The published dummy must behave like a resolved config: bare writes
|
||||
# raise under the strict harness; override() stays the entry point.
|
||||
published = get_context().override_server_args(tp_size=2).install()
|
||||
with self.assertRaises(AttributeError):
|
||||
published.tp_size = 4
|
||||
published.override(source="test", tp_size=4)
|
||||
self.assertEqual(published.tp_size, 4)
|
||||
|
||||
def test_restore_resets_the_capture_seed(self):
|
||||
# install() seeds flags.capture from the published dummy; restore()
|
||||
# must put back the pre-install runtime state on both restore paths.
|
||||
reset_context()
|
||||
self.assertFalse(get_flags().capture.enable_torch_compile)
|
||||
override = get_context().override_server_args(enable_torch_compile=True)
|
||||
override.install()
|
||||
self.assertTrue(get_flags().capture.enable_torch_compile)
|
||||
override.restore()
|
||||
self.assertFalse(get_flags().capture.enable_torch_compile)
|
||||
|
||||
def test_double_install_rejected(self):
|
||||
override = get_context().override_server_args()
|
||||
override.install()
|
||||
with self.assertRaises(AssertionError):
|
||||
override.install()
|
||||
|
||||
def test_module_global_removed(self):
|
||||
# The legacy storage must not survive: a stale _global_server_args would
|
||||
# silently fork the config into two objects.
|
||||
@@ -462,6 +543,58 @@ class TestNamedStreams(_IsolatedServerArgs):
|
||||
reset_context()
|
||||
self.assertEqual(get_context().resources.streams, {})
|
||||
|
||||
def test_capturer_slots_roundtrip_and_reset(self):
|
||||
from sglang.srt.state_capturer.indexer_topk import (
|
||||
get_global_indexer_capturer,
|
||||
set_global_indexer_capturer,
|
||||
)
|
||||
from sglang.srt.state_capturer.routed_experts import (
|
||||
get_global_experts_capturer,
|
||||
set_global_experts_capturer,
|
||||
)
|
||||
|
||||
reset_context()
|
||||
self.assertIsNone(get_global_indexer_capturer())
|
||||
self.assertIsNone(get_global_experts_capturer())
|
||||
indexer, experts = object(), object()
|
||||
set_global_indexer_capturer(indexer)
|
||||
set_global_experts_capturer(experts)
|
||||
self.assertIs(get_global_indexer_capturer(), indexer)
|
||||
self.assertIs(get_global_experts_capturer(), experts)
|
||||
reset_context()
|
||||
self.assertIsNone(get_global_indexer_capturer())
|
||||
self.assertIsNone(get_global_experts_capturer())
|
||||
|
||||
def test_tcp_store_slot_roundtrip_and_reset(self):
|
||||
from sglang.srt.distributed.utils import (
|
||||
get_global_tcp_store,
|
||||
set_global_tcp_store,
|
||||
)
|
||||
|
||||
reset_context()
|
||||
self.assertIsNone(get_global_tcp_store())
|
||||
store = object()
|
||||
set_global_tcp_store(store)
|
||||
self.assertIs(get_global_tcp_store(), store)
|
||||
reset_context()
|
||||
self.assertIsNone(get_global_tcp_store())
|
||||
|
||||
def test_trace_level_env_seeded_lazy_default(self):
|
||||
from sglang.srt.observability.trace import (
|
||||
get_global_trace_level,
|
||||
set_global_trace_level,
|
||||
)
|
||||
|
||||
reset_context()
|
||||
with patch.dict(os.environ, {}, clear=False):
|
||||
os.environ.pop("SGLANG_TRACE_LEVEL", None)
|
||||
self.assertEqual(get_global_trace_level(), 3)
|
||||
set_global_trace_level(5)
|
||||
self.assertEqual(get_global_trace_level(), 5)
|
||||
reset_context()
|
||||
with patch.dict(os.environ, {"SGLANG_TRACE_LEVEL": "1"}):
|
||||
self.assertEqual(get_global_trace_level(), 1)
|
||||
|
||||
|
||||
class TestEpBufferState(_IsolatedServerArgs):
|
||||
"""EP dispatcher buffer managers: state lives on ctx.resources; the
|
||||
|
||||
@@ -38,17 +38,18 @@ _MUTATION_PATTERNS = [
|
||||
re.compile(r"\bserver_args\.[a-z0-9_]+\s*=(?![=}])"),
|
||||
re.compile(r"\bsa\.[a-z0-9_]+\s*=(?![=}])"),
|
||||
re.compile(r"get_(?:global_)?server_args\(\)\.[a-z0-9_]+\s*=(?![=}])"),
|
||||
# setattr is the same write with the attribute name behind a variable.
|
||||
re.compile(
|
||||
r"setattr\(\s*(?:[\w.]+\.)?(?:server_args|sa|get_(?:global_)?server_args\(\))\s*,"
|
||||
),
|
||||
]
|
||||
|
||||
# The resolution pipeline itself (mutation is its job); multimodal_gen, whose
|
||||
# ServerArgs is a different class outside this contract; and the sanctioned
|
||||
# mock-fixture factory (bare object.__new__ instances never materialize, so
|
||||
# the strict guard does not apply to their construction).
|
||||
# The resolution pipeline itself (mutation is its job) and multimodal_gen,
|
||||
# whose ServerArgs is a different class outside this contract.
|
||||
_EXCLUDED = (
|
||||
"srt/server_args.py",
|
||||
"srt/arg_groups",
|
||||
"multimodal_gen",
|
||||
"test/kits/attention_unittest/mock_server_args.py",
|
||||
)
|
||||
|
||||
_BASELINE = 0
|
||||
|
||||
Reference in New Issue
Block a user