[refactor] Move ServerArgs ownership into the runtime context (stack 2/15) (#30064)
Co-authored-by: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
co-authored by
Claude Fable 5
parent
6d662c9245
commit
def20782cf
@@ -33,6 +33,7 @@ from sglang.multimodal_gen.runtime.pipelines_core import Req
|
|||||||
from sglang.srt import server_args as srt_server_args_module
|
from sglang.srt import server_args as srt_server_args_module
|
||||||
from sglang.srt.observability import trace as srt_trace
|
from sglang.srt.observability import trace as srt_trace
|
||||||
from sglang.srt.observability.trace import TraceNullContext, TraceReqContext
|
from sglang.srt.observability.trace import TraceNullContext, TraceReqContext
|
||||||
|
from sglang.srt.runtime_context import reset_context
|
||||||
from sglang.srt.server_args import set_global_server_args_for_scheduler
|
from sglang.srt.server_args import set_global_server_args_for_scheduler
|
||||||
|
|
||||||
try:
|
try:
|
||||||
@@ -62,12 +63,18 @@ def _enable_minimal_otel() -> None:
|
|||||||
|
|
||||||
@contextmanager
|
@contextmanager
|
||||||
def _srt_trace_server_args():
|
def _srt_trace_server_args():
|
||||||
prev_server_args = srt_server_args_module._global_server_args
|
try:
|
||||||
|
prev_server_args = srt_server_args_module.get_global_server_args()
|
||||||
|
except ValueError: # nothing published yet
|
||||||
|
prev_server_args = None
|
||||||
set_global_server_args_for_scheduler(SimpleNamespace(trace_modules="request"))
|
set_global_server_args_for_scheduler(SimpleNamespace(trace_modules="request"))
|
||||||
try:
|
try:
|
||||||
yield
|
yield
|
||||||
finally:
|
finally:
|
||||||
srt_server_args_module._global_server_args = prev_server_args
|
if prev_server_args is None:
|
||||||
|
reset_context()
|
||||||
|
else:
|
||||||
|
set_global_server_args_for_scheduler(prev_server_args)
|
||||||
|
|
||||||
|
|
||||||
def _traceparent_from(ctx) -> str | None:
|
def _traceparent_from(ctx) -> str | None:
|
||||||
|
|||||||
@@ -21,10 +21,12 @@ wrapper, not a cache. It gives call-sites one import and one naming scheme in
|
|||||||
place of a dozen free functions, plus a test-only ``override()`` hook to force a
|
place of a dozen free functions, plus a test-only ``override()`` hook to force a
|
||||||
topology without monkeypatching the underlying getters.
|
topology without monkeypatching the underlying getters.
|
||||||
|
|
||||||
``get_server_args()`` returns the process-wide ``ServerArgs`` (the config tier).
|
``get_server_args()`` returns the process-wide ``ServerArgs`` (the config
|
||||||
It is a read-through to ``server_args.get_global_server_args()`` — same object,
|
tier). The context owns the storage: publishing goes through
|
||||||
same pre-publish error — so new code can adopt the context accessor while the
|
``RuntimeContext.set_server_args`` (the legacy
|
||||||
legacy getter remains canonical.
|
``set_global_server_args_for_scheduler`` / ``get_global_server_args`` in
|
||||||
|
``server_args.py`` are thin shims over this slot), and the object is returned
|
||||||
|
by reference — the same live instance everywhere, never a copy.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
@@ -50,12 +52,6 @@ def _dp():
|
|||||||
return dp_attention
|
return dp_attention
|
||||||
|
|
||||||
|
|
||||||
def _sa():
|
|
||||||
from sglang.srt import server_args
|
|
||||||
|
|
||||||
return server_args
|
|
||||||
|
|
||||||
|
|
||||||
_PARALLEL_FIELDS = frozenset(
|
_PARALLEL_FIELDS = frozenset(
|
||||||
{
|
{
|
||||||
"world_size",
|
"world_size",
|
||||||
@@ -223,15 +219,29 @@ class RuntimeContext:
|
|||||||
"""Container for the structured runtime accessors; exposes ``parallel`` and
|
"""Container for the structured runtime accessors; exposes ``parallel`` and
|
||||||
``server_args``."""
|
``server_args``."""
|
||||||
|
|
||||||
__slots__ = ("parallel",)
|
__slots__ = ("parallel", "_server_args")
|
||||||
|
|
||||||
def __init__(self, parallel: ParallelContext):
|
def __init__(self, parallel: ParallelContext):
|
||||||
self.parallel = parallel
|
self.parallel = parallel
|
||||||
|
self._server_args: ServerArgs | None = None
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def server_args(self) -> ServerArgs:
|
def server_args(self) -> ServerArgs:
|
||||||
"""The process-wide ``ServerArgs``, read through the global getter."""
|
"""The process-wide ``ServerArgs`` (context-owned slot)."""
|
||||||
return _sa().get_global_server_args()
|
server_args = self._server_args
|
||||||
|
if server_args is None:
|
||||||
|
# Verbatim legacy message: tests and user scripts may match on it.
|
||||||
|
raise ValueError("Global server args is not set yet!")
|
||||||
|
return server_args
|
||||||
|
|
||||||
|
def set_server_args(self, server_args: ServerArgs) -> None:
|
||||||
|
"""Publish the process-wide ``ServerArgs`` into the context-owned slot.
|
||||||
|
|
||||||
|
Overwrite-allowed: a re-publish replaces the slot (test kits re-publish
|
||||||
|
per test; production ordering discipline lives at the call-sites, e.g.
|
||||||
|
the draft-worker guard in ``ModelRunner.__init__``).
|
||||||
|
"""
|
||||||
|
self._server_args = server_args
|
||||||
|
|
||||||
|
|
||||||
_PARALLEL = ParallelContext()
|
_PARALLEL = ParallelContext()
|
||||||
@@ -248,3 +258,11 @@ def get_parallel() -> ParallelContext:
|
|||||||
|
|
||||||
def get_server_args() -> ServerArgs:
|
def get_server_args() -> ServerArgs:
|
||||||
return _CONTEXT.server_args
|
return _CONTEXT.server_args
|
||||||
|
|
||||||
|
|
||||||
|
def reset_context() -> None:
|
||||||
|
"""Clear the context-owned store (unit-test teardown).
|
||||||
|
|
||||||
|
Wrapper subsystems (``parallel``) hold no state and are unaffected.
|
||||||
|
"""
|
||||||
|
_CONTEXT._server_args = None
|
||||||
|
|||||||
@@ -7594,23 +7594,23 @@ class ServerArgs:
|
|||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
# NOTE: This is a global variable to hold the server args for scheduler.
|
# NOTE: The process-wide ServerArgs is owned by the runtime context
|
||||||
_global_server_args: Optional[ServerArgs] = None
|
# (sglang.srt.runtime_context). The two functions below are thin shims kept for
|
||||||
|
# the existing call-sites; they publish/read the same live object by reference.
|
||||||
|
# Imports are in-function so the two modules stay cycle-free at import time.
|
||||||
def set_global_server_args_for_scheduler(server_args: ServerArgs):
|
def set_global_server_args_for_scheduler(server_args: ServerArgs):
|
||||||
global _global_server_args
|
from sglang.srt.runtime_context import get_context
|
||||||
_global_server_args = server_args
|
|
||||||
|
get_context().set_server_args(server_args)
|
||||||
|
|
||||||
|
|
||||||
set_global_server_args_for_tokenizer = set_global_server_args_for_scheduler
|
set_global_server_args_for_tokenizer = set_global_server_args_for_scheduler
|
||||||
|
|
||||||
|
|
||||||
def get_global_server_args() -> ServerArgs:
|
def get_global_server_args() -> ServerArgs:
|
||||||
if _global_server_args is None:
|
from sglang.srt.runtime_context import get_context
|
||||||
raise ValueError("Global server args is not set yet!")
|
|
||||||
|
|
||||||
return _global_server_args
|
return get_context().server_args
|
||||||
|
|
||||||
|
|
||||||
def prepare_server_args(argv: List[str]) -> ServerArgs:
|
def prepare_server_args(argv: List[str]) -> ServerArgs:
|
||||||
|
|||||||
@@ -7,18 +7,19 @@ register_cpu_ci(est_time=5, suite="base-a-test-cpu")
|
|||||||
import unittest
|
import unittest
|
||||||
from unittest.mock import patch
|
from unittest.mock import patch
|
||||||
|
|
||||||
|
import sglang.srt.server_args as server_args_module
|
||||||
from sglang.srt.runtime_context import (
|
from sglang.srt.runtime_context import (
|
||||||
ParallelContext,
|
ParallelContext,
|
||||||
RuntimeContext,
|
RuntimeContext,
|
||||||
get_context,
|
get_context,
|
||||||
get_parallel,
|
get_parallel,
|
||||||
get_server_args,
|
get_server_args,
|
||||||
|
reset_context,
|
||||||
)
|
)
|
||||||
from sglang.test.test_utils import CustomTestCase
|
from sglang.test.test_utils import CustomTestCase
|
||||||
|
|
||||||
_PS = "sglang.srt.distributed.parallel_state"
|
_PS = "sglang.srt.distributed.parallel_state"
|
||||||
_DP = "sglang.srt.layers.dp_attention"
|
_DP = "sglang.srt.layers.dp_attention"
|
||||||
_SA = "sglang.srt.server_args"
|
|
||||||
|
|
||||||
SIZE_RANK_DELEGATIONS = [
|
SIZE_RANK_DELEGATIONS = [
|
||||||
("world_size", f"{_PS}.get_world_size"),
|
("world_size", f"{_PS}.get_world_size"),
|
||||||
@@ -144,41 +145,67 @@ class TestParallelOverride(_IsolatedOverrides):
|
|||||||
self.assertEqual(p._overrides, {})
|
self.assertEqual(p._overrides, {})
|
||||||
|
|
||||||
|
|
||||||
class TestServerArgsReadThrough(CustomTestCase):
|
class _IsolatedServerArgs(CustomTestCase):
|
||||||
"""``server_args`` delegates live to the global getter (read-through, V2a)."""
|
"""Save/restore the published ServerArgs around each test (the slot is
|
||||||
|
process-global; another test file sharing the process may have published)."""
|
||||||
|
|
||||||
def test_delegates_to_global_getter(self):
|
def setUp(self):
|
||||||
sentinel = object()
|
super().setUp()
|
||||||
with patch(f"{_SA}.get_global_server_args", return_value=sentinel):
|
self._saved_server_args = get_context()._server_args
|
||||||
self.assertIs(get_server_args(), sentinel)
|
|
||||||
self.assertIs(get_context().server_args, sentinel)
|
|
||||||
|
|
||||||
def test_identity_with_global_getter(self):
|
def tearDown(self):
|
||||||
import sglang.srt.server_args as server_args_module
|
if self._saved_server_args is None:
|
||||||
|
reset_context()
|
||||||
|
else:
|
||||||
|
get_context().set_server_args(self._saved_server_args)
|
||||||
|
super().tearDown()
|
||||||
|
|
||||||
|
|
||||||
|
class TestServerArgsOwnership(_IsolatedServerArgs):
|
||||||
|
"""V2b: the context owns the slot; the legacy getters are identity shims."""
|
||||||
|
|
||||||
|
def test_legacy_setter_publishes_into_context(self):
|
||||||
# Identity (not equality) is the contract; publish accepts any object.
|
# Identity (not equality) is the contract; publish accepts any object.
|
||||||
sentinel = object()
|
sentinel = object()
|
||||||
saved = server_args_module._global_server_args
|
server_args_module.set_global_server_args_for_scheduler(sentinel)
|
||||||
try:
|
self.assertIs(server_args_module.get_global_server_args(), sentinel)
|
||||||
server_args_module.set_global_server_args_for_scheduler(sentinel)
|
self.assertIs(get_server_args(), sentinel)
|
||||||
self.assertIs(
|
self.assertIs(get_context().server_args, sentinel)
|
||||||
get_server_args(), server_args_module.get_global_server_args()
|
|
||||||
)
|
|
||||||
self.assertIs(get_server_args(), sentinel)
|
|
||||||
finally:
|
|
||||||
server_args_module._global_server_args = saved
|
|
||||||
|
|
||||||
def test_pre_publish_error_passes_through(self):
|
def test_context_publish_visible_through_legacy_getter(self):
|
||||||
import sglang.srt.server_args as server_args_module
|
sentinel = object()
|
||||||
|
get_context().set_server_args(sentinel)
|
||||||
|
self.assertIs(server_args_module.get_global_server_args(), sentinel)
|
||||||
|
|
||||||
saved = server_args_module._global_server_args
|
def test_tokenizer_alias_is_same_function(self):
|
||||||
server_args_module._global_server_args = None
|
self.assertIs(
|
||||||
try:
|
server_args_module.set_global_server_args_for_tokenizer,
|
||||||
|
server_args_module.set_global_server_args_for_scheduler,
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_pre_publish_error_verbatim(self):
|
||||||
|
reset_context()
|
||||||
|
for accessor in (get_server_args, server_args_module.get_global_server_args):
|
||||||
with self.assertRaises(ValueError) as cm:
|
with self.assertRaises(ValueError) as cm:
|
||||||
get_server_args()
|
accessor()
|
||||||
self.assertEqual(str(cm.exception), "Global server args is not set yet!")
|
self.assertEqual(str(cm.exception), "Global server args is not set yet!")
|
||||||
finally:
|
|
||||||
server_args_module._global_server_args = saved
|
def test_republish_overwrite_allowed(self):
|
||||||
|
first, second = object(), object()
|
||||||
|
server_args_module.set_global_server_args_for_scheduler(first)
|
||||||
|
server_args_module.set_global_server_args_for_scheduler(second)
|
||||||
|
self.assertIs(get_server_args(), second)
|
||||||
|
|
||||||
|
def test_reset_context_clears_owned_store(self):
|
||||||
|
server_args_module.set_global_server_args_for_scheduler(object())
|
||||||
|
reset_context()
|
||||||
|
with self.assertRaises(ValueError):
|
||||||
|
get_server_args()
|
||||||
|
|
||||||
|
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.
|
||||||
|
self.assertFalse(hasattr(server_args_module, "_global_server_args"))
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
|
|||||||
Reference in New Issue
Block a user