[refactor] Add the resolved-flags tier + resolvable-field metadata (stack 4/15) (#30066)

Co-authored-by: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
Cheng Wan
2026-07-04 02:20:58 -07:00
committed by GitHub
co-authored by Claude Fable 5
parent ff57171f98
commit 10f258257f
4 changed files with 284 additions and 5 deletions
@@ -0,0 +1,42 @@
"""Unit tests for the model-override machinery: whitelist metadata (V3a)."""
from sglang.test.ci.ci_register import register_cpu_ci
register_cpu_ci(est_time=5, suite="base-a-test-cpu")
import dataclasses
import unittest
from typing import Optional
from sglang.srt.arg_groups.arg_utils import A, Arg, model_overridable_fields
from sglang.test.test_utils import CustomTestCase
@dataclasses.dataclass
class _FakeArgs:
plain: A[int, "help text only"] = 0
resolved_by_model: A[str, Arg(help="x", model_overridable=True)] = "auto"
also_resolved: A[Optional[int], Arg(help="y", model_overridable=True)] = None
metadata_but_not_overridable: A[bool, Arg(help="z")] = False
class TestModelOverridableWhitelist(CustomTestCase):
def test_arg_defaults_to_not_overridable(self):
self.assertFalse(Arg().model_overridable)
def test_whitelist_derivation_from_annotated_metadata(self):
self.assertEqual(
model_overridable_fields(_FakeArgs),
frozenset({"resolved_by_model", "also_resolved"}),
)
def test_server_args_whitelist_empty_at_skeleton(self):
# No ServerArgs field is tagged yet: the V3 sweeps whitelist fields
# one family at a time. This pin makes accidental tagging visible.
from sglang.srt.server_args import ServerArgs
self.assertEqual(model_overridable_fields(ServerArgs), frozenset())
if __name__ == "__main__":
unittest.main()
@@ -4,17 +4,23 @@ from sglang.test.ci.ci_register import register_cpu_ci
register_cpu_ci(est_time=5, suite="base-a-test-cpu")
import dataclasses
import unittest
from unittest.mock import patch
import sglang.srt.server_args as server_args_module
from sglang.srt.runtime_context import (
Flags,
ParallelContext,
RuntimeContext,
_FlagGroupBase,
_StaticFlags,
get_context,
get_flags,
get_parallel,
get_server_args,
reset_context,
resolve_flag_leaf,
)
from sglang.test.test_utils import CustomTestCase
@@ -208,5 +214,92 @@ class TestServerArgsOwnership(_IsolatedServerArgs):
self.assertFalse(hasattr(server_args_module, "_global_server_args"))
@dataclasses.dataclass
class _FakeStaticGroup(_StaticFlags):
alpha: int = 1
beta: str = "b"
@dataclasses.dataclass
class _FakeCaptureGroup(_FlagGroupBase):
gamma: int = 0
class TestFlagsTier(_IsolatedServerArgs):
"""V3a skeleton: typed dataclass groups, freeze guard, override primitive."""
def test_wiring_and_groups(self):
flags = get_flags()
self.assertIs(flags, get_context().flags)
self.assertIsInstance(flags, Flags)
for group in ("attn", "moe", "capture"):
self.assertTrue(hasattr(flags, group))
self.assertFalse(flags.frozen)
def test_typo_safety(self):
group = _FakeStaticGroup()
with self.assertRaises(AttributeError):
group.alpha_misspelled = 2 # undeclared leaf
with self.assertRaises(AttributeError):
get_flags().not_a_flag = 1
def test_static_group_writable_until_freeze(self):
group = _FakeStaticGroup()
group.alpha = 5
self.assertEqual(group.alpha, 5)
group.freeze()
with self.assertRaises(RuntimeError):
group.alpha = 6
self.assertEqual(group.alpha, 5)
def test_override_is_transactional_and_works_on_frozen(self):
group = _FakeStaticGroup()
group.freeze()
with group.override(alpha=99, beta="x"):
self.assertEqual(group.alpha, 99)
self.assertEqual(group.beta, "x")
self.assertEqual(group.alpha, 1)
self.assertEqual(group.beta, "b")
with self.assertRaises(ValueError):
with group.override(alpha=2, gamma=3): # gamma undeclared
pass
self.assertEqual(group.alpha, 1) # validated before any write
def test_non_static_group_has_no_freeze(self):
group = _FakeCaptureGroup()
group.gamma = 42
self.assertEqual(group.gamma, 42)
self.assertFalse(hasattr(group, "freeze"))
def test_container_freeze_cascades_except_capture(self):
flags = Flags() # fresh container, not the process singleton
flags.freeze()
self.assertTrue(flags.frozen)
self.assertTrue(flags.attn.frozen)
self.assertTrue(flags.moe.frozen)
with self.assertRaises(RuntimeError):
flags.attn = flags.attn # container leaves lock too
self.assertFalse(getattr(flags.capture, "_frozen", False))
def test_resolve_flag_leaf_flat_default_and_mapped(self):
flags = Flags()
owner, leaf = resolve_flag_leaf(flags, "some_field")
self.assertIs(owner, flags)
self.assertEqual(leaf, "some_field")
owner, leaf = resolve_flag_leaf(flags, "x", leaf_map={"x": "attn.x"})
self.assertIs(owner, flags.attn)
self.assertEqual(leaf, "x")
def test_reset_context_installs_fresh_unfrozen_flags(self):
try:
old = get_flags()
old.freeze()
reset_context()
self.assertIsNot(get_flags(), old)
self.assertFalse(get_flags().frozen)
finally:
reset_context() # never leave the singleton frozen for other tests
if __name__ == "__main__":
unittest.main()