diff --git a/python/sglang/benchmark/endpoint.py b/python/sglang/benchmark/endpoint.py index c8a93b988..44c0b63e6 100644 --- a/python/sglang/benchmark/endpoint.py +++ b/python/sglang/benchmark/endpoint.py @@ -10,6 +10,8 @@ import os import time from typing import Callable, Optional +import msgspec +import msgspec.structs import requests from sglang.srt.entrypoints.http_server import launch_server @@ -21,7 +23,7 @@ DEFAULT_TIMEOUT = 600 # Field defaults of ServerArgs, used to detect when --host/--port were set # explicitly (and would be silently ignored in connect mode). -_SERVER_ARGS_DEFAULTS = {f.name: f.default for f in dataclasses.fields(ServerArgs)} +_SERVER_ARGS_DEFAULTS = {f.name: f.default for f in msgspec.structs.fields(ServerArgs)} def server_is_up(base_url: str, timeout: float = DEFAULT_TIMEOUT) -> bool: diff --git a/python/sglang/lang/backend/runtime_endpoint.py b/python/sglang/lang/backend/runtime_endpoint.py index 84849d4ea..c57d7249a 100644 --- a/python/sglang/lang/backend/runtime_endpoint.py +++ b/python/sglang/lang/backend/runtime_endpoint.py @@ -6,6 +6,8 @@ import warnings from typing import Dict, List, Optional, Union import aiohttp +import msgspec +import msgspec.structs import requests from sglang.lang.backend.base_backend import BaseBackend @@ -383,7 +385,12 @@ class Runtime: # Pre-allocate a port before building the config, so the config is born # with the port this runtime will serve on. requested_port = kwargs.pop( - "port", ServerArgs.__dataclass_fields__["port"].default + "port", + next( + f.default + for f in msgspec.structs.fields(ServerArgs) + if f.name == "port" + ), ) for port in range(requested_port, 40000): if is_port_available(port): diff --git a/python/sglang/srt/arg_groups/arg_utils.py b/python/sglang/srt/arg_groups/arg_utils.py index 8b3a63f6c..0c54e930f 100644 --- a/python/sglang/srt/arg_groups/arg_utils.py +++ b/python/sglang/srt/arg_groups/arg_utils.py @@ -53,12 +53,14 @@ from typing import ( get_type_hints, ) +import msgspec +import msgspec.structs + A = Annotated -@dataclasses.dataclass(frozen=True) -class Arg: - """CLI argument metadata attached to a dataclass field via ``Annotated``.""" +class Arg(msgspec.Struct, frozen=True): + """CLI argument metadata attached to a field via ``Annotated``.""" help: str = "" choices: list | None = None @@ -90,8 +92,7 @@ class Arg: fallback: Any = None -@dataclasses.dataclass(frozen=True) -class Derived: +class Derived(msgspec.Struct, frozen=True): """Metadata for a field the configuration implies, not one anyone types. The other half of a namespace. An ``Arg`` field is the operator's input and @@ -118,8 +119,7 @@ class Derived: fn: str = "" -@dataclasses.dataclass(frozen=True) -class NS: +class NS(msgspec.Struct, frozen=True): """Namespace-path marker for a ServerArgs field, attached alongside the field's metadata in ``Annotated``: @@ -151,7 +151,7 @@ def namespace_of(cls) -> dict: the per-field ``NS`` marker. A field with neither is absent from the map (the coverage lint flags them). Non-dataclass types yield an empty map. """ - if not dataclasses.is_dataclass(cls): + if not is_record(cls): return {} # An assembled record: the collector recorded who declared each field, # because there are no base classes left to ask. @@ -164,10 +164,10 @@ def namespace_of(cls) -> dict: continue for name in getattr(base, "__annotations__", {}): out.setdefault(name, path) - if len(out) == len(dataclasses.fields(cls)): + if len(out) == len(record_fields(cls)): return out hints = get_type_hints(cls, include_extras=True) - for field in dataclasses.fields(cls): + for field in record_fields(cls): if field.name in out: continue tp = hints.get(field.name, field.type) @@ -182,9 +182,9 @@ def namespace_of(cls) -> dict: @functools.cache def field_names(cls) -> frozenset: """Names of ``cls`` dataclass fields — what a declaration may name.""" - if not dataclasses.is_dataclass(cls): + if not is_record(cls): return frozenset() - return frozenset(field.name for field in dataclasses.fields(cls)) + return frozenset(field.name for field in record_fields(cls)) @functools.cache @@ -194,11 +194,11 @@ def resolvable_fields(cls) -> frozenset: Non-dataclass types (e.g. mock config objects in tests) have no Arg metadata and yield an empty whitelist.""" - if not dataclasses.is_dataclass(cls): + if not is_record(cls): return frozenset() hints = get_type_hints(cls, include_extras=True) names = set() - for field in dataclasses.fields(cls): + for field in record_fields(cls): _, arg = _unwrap_annotated(hints.get(field.name, field.type)) if arg is not None and arg.resolvable: names.add(field.name) @@ -213,11 +213,11 @@ def fallbacks_of(cls) -> dict: beside the help text of the field it belongs to rather than in whatever hook used to fill it in. """ - if not dataclasses.is_dataclass(cls): + if not is_record(cls): return {} hints = get_type_hints(cls, include_extras=True) out = {} - for field in dataclasses.fields(cls): + for field in record_fields(cls): _, arg = _unwrap_annotated(hints.get(field.name, field.type)) if arg is not None and arg.fallback is not None: # Two things `with_fallback` relies on and cannot check itself, @@ -254,6 +254,26 @@ def with_fallback(cls, name: str, value: Any) -> Any: return fallbacks_of(cls).get(name, value) +def record_fields(cls): + """The declared fields of a record, Struct or dataclass. + + `ServerArgs` and the namespace classes are `msgspec.Struct`; the config-bag + tests build ad-hoc dataclasses spanning namespaces, and the helpers here are + driven with both. Anything else yields nothing. + """ + if isinstance(cls, type) and issubclass(cls, msgspec.Struct): + return msgspec.structs.fields(cls) + if dataclasses.is_dataclass(cls): + return dataclasses.fields(cls) + return () + + +def is_record(cls) -> bool: + """Whether ``cls`` declares fields the way a record does.""" + target = cls if isinstance(cls, type) else type(cls) + return issubclass(target, msgspec.Struct) or dataclasses.is_dataclass(cls) + + # --------------------------------------------------------------------------- # Internal helpers # --------------------------------------------------------------------------- @@ -314,10 +334,16 @@ def _infer_type_func(tp): def _field_default(field): - """Return the default value for a dataclass field, or _MISSING.""" - if field.default is not _MISSING: + """Return the default value for a field, or `_MISSING`. + + The two record shapes spell "no default" differently -- a Struct field says + `msgspec.NODEFAULT`, a dataclass field `dataclasses.MISSING` -- so both are + normalized here and every caller below tests against `_MISSING` alone. + """ + absent = (_MISSING, msgspec.NODEFAULT) + if field.default not in absent: return field.default - if field.default_factory is not _MISSING: + if field.default_factory not in absent: return field.default_factory() return _MISSING @@ -347,7 +373,7 @@ def add_cli_args_from_dataclass(parser, cls, *, fields: list[str] | None = None) """ hints = get_type_hints(cls, include_extras=True) - for field in dataclasses.fields(cls): + for field in record_fields(cls): if fields is not None and field.name not in fields: continue diff --git a/python/sglang/srt/arg_groups/fields/__init__.py b/python/sglang/srt/arg_groups/fields/__init__.py index 48ce80a05..b93347605 100644 --- a/python/sglang/srt/arg_groups/fields/__init__.py +++ b/python/sglang/srt/arg_groups/fields/__init__.py @@ -11,6 +11,9 @@ from __future__ import annotations import dataclasses from typing import Any, Dict, List, Tuple, get_type_hints +import msgspec +import msgspec.structs + from sglang.srt.arg_groups.field_order import POSITIONAL_FIELD_ORDER @@ -46,7 +49,7 @@ def collect_input_fields( defaults: Dict[str, Any] = {} for source in sources: hints = get_type_hints(source, include_extras=True) - for field in dataclasses.fields(source): + for field in msgspec.structs.fields(source): if field.name in annotations: raise ValueError( f"{field.name!r} is declared by both " @@ -54,10 +57,10 @@ def collect_input_fields( "a field belongs to exactly one namespace" ) annotations[field.name] = (source, hints[field.name]) - if field.default is not dataclasses.MISSING: + if field.default is not msgspec.NODEFAULT: defaults[field.name] = field.default - elif field.default_factory is not dataclasses.MISSING: - defaults[field.name] = dataclasses.field( + elif field.default_factory is not msgspec.NODEFAULT: + defaults[field.name] = msgspec.field( default_factory=field.default_factory ) known = [n for n in POSITIONAL_FIELD_ORDER if n in annotations] diff --git a/python/sglang/srt/arg_groups/fields/device.py b/python/sglang/srt/arg_groups/fields/device.py index 48ed843da..8b329fd13 100644 --- a/python/sglang/srt/arg_groups/fields/device.py +++ b/python/sglang/srt/arg_groups/fields/device.py @@ -9,18 +9,18 @@ how config is shaped at runtime. from __future__ import annotations -import dataclasses from typing import ( Callable, List, Optional, ) +import msgspec + from sglang.srt.arg_groups.arg_utils import A -@dataclasses.dataclass -class Device: +class Device(msgspec.Struct): """Namespace ``device``.""" _NS_PATH = "device" diff --git a/python/sglang/srt/arg_groups/fields/disagg.py b/python/sglang/srt/arg_groups/fields/disagg.py index a7251916e..bbd6c4da1 100644 --- a/python/sglang/srt/arg_groups/fields/disagg.py +++ b/python/sglang/srt/arg_groups/fields/disagg.py @@ -9,13 +9,14 @@ how config is shaped at runtime. from __future__ import annotations -import dataclasses from typing import ( List, Literal, Optional, ) +import msgspec + from sglang.srt.arg_groups.arg_utils import ( A, Arg, @@ -24,8 +25,7 @@ from sglang.srt.arg_groups.choices import DISAGG_TRANSFER_BACKEND_CHOICES from sglang.srt.utils.common import json_list_type -@dataclasses.dataclass -class Disagg: +class Disagg(msgspec.Struct): """Namespace ``disagg``.""" _NS_PATH = "disagg" @@ -137,7 +137,7 @@ class Disagg: choices=["auto", "zmq_to_scheduler", "zmq_to_tokenizer", "mooncake"], ), ] = "auto" - encoder_urls: A[List[str], "List of encoder server urls."] = dataclasses.field( + encoder_urls: A[List[str], "List of encoder server urls."] = msgspec.field( default_factory=list ) encoder_bootstrap_port: A[ @@ -147,7 +147,7 @@ class Disagg: encoder_register_urls: A[ List[str], "One or more EncoderBootstrapServer URLs to register this encoder with on startup, for dynamic encoder discovery. Example: --encoder-register-urls http://prefill0:8997 http://prefill1:8997. Used with --encoder-only servers.", - ] = dataclasses.field(default_factory=list) + ] = msgspec.field(default_factory=list) enable_adaptive_dispatch_to_encoder: A[ bool, "When enabled, adaptively dispatch: multi-image requests go to encoder in language_only epd mode, single-image requests are processed locally.", diff --git a/python/sglang/srt/arg_groups/fields/exec_.py b/python/sglang/srt/arg_groups/fields/exec_.py index 2c05b9957..f501d490e 100644 --- a/python/sglang/srt/arg_groups/fields/exec_.py +++ b/python/sglang/srt/arg_groups/fields/exec_.py @@ -10,13 +10,14 @@ how config is shaped at runtime. from __future__ import annotations import argparse -import dataclasses from typing import ( List, Literal, Optional, ) +import msgspec + from sglang.srt.arg_groups.arg_utils import ( A, Arg, @@ -38,8 +39,7 @@ from sglang.srt.model_executor.cuda_graph_config import ( ) -@dataclasses.dataclass -class ExecFeatures: +class ExecFeatures(msgspec.Struct): """Namespace ``exec.features``.""" _NS_PATH = "exec.features" @@ -105,8 +105,7 @@ class ExecFeatures: ] = False -@dataclasses.dataclass -class ExecKernel: +class ExecKernel(msgspec.Struct): """Namespace ``exec.kernel``.""" _NS_PATH = "exec.kernel" @@ -302,8 +301,7 @@ class ExecKernel: ] = False -@dataclasses.dataclass -class ExecMamba: +class ExecMamba(msgspec.Struct): """Namespace ``exec.mamba``.""" _NS_PATH = "exec.mamba" @@ -443,8 +441,7 @@ class ExecMamba: ] = False -@dataclasses.dataclass -class ExecGraph: +class ExecGraph(msgspec.Struct): """Namespace ``exec.graph``.""" _NS_PATH = "exec.graph" @@ -533,8 +530,7 @@ class ExecGraph: ] = 32 -@dataclasses.dataclass -class ExecComm: +class ExecComm(msgspec.Struct): """Namespace ``exec.comm``.""" _NS_PATH = "exec.comm" @@ -609,8 +605,7 @@ class ExecComm: ] = False -@dataclasses.dataclass -class ExecMoe: +class ExecMoe(msgspec.Struct): """Namespace ``exec.moe``.""" _NS_PATH = "exec.moe" @@ -812,8 +807,7 @@ class ExecMoe: ] = None -@dataclasses.dataclass -class ExecOverlap: +class ExecOverlap(msgspec.Struct): """Namespace ``exec.overlap``.""" _NS_PATH = "exec.overlap" @@ -834,8 +828,7 @@ class ExecOverlap: ] = 0.48 -@dataclasses.dataclass -class ExecOffload: +class ExecOffload(msgspec.Struct): """Namespace ``exec.offload``.""" _NS_PATH = "exec.offload" @@ -872,8 +865,7 @@ class ExecOffload: ] = None -@dataclasses.dataclass -class ExecDllm: +class ExecDllm(msgspec.Struct): """Namespace ``exec.dllm``.""" _NS_PATH = "exec.dllm" @@ -897,8 +889,7 @@ class ExecDllm: ] = True -@dataclasses.dataclass -class ExecDeterministic: +class ExecDeterministic(msgspec.Struct): """Namespace ``exec.deterministic``.""" _NS_PATH = "exec.deterministic" diff --git a/python/sglang/srt/arg_groups/fields/lora.py b/python/sglang/srt/arg_groups/fields/lora.py index 57262eb8b..674aa504f 100644 --- a/python/sglang/srt/arg_groups/fields/lora.py +++ b/python/sglang/srt/arg_groups/fields/lora.py @@ -10,13 +10,14 @@ how config is shaped at runtime. from __future__ import annotations import argparse -import dataclasses from typing import ( List, Optional, Union, ) +import msgspec + from sglang.srt.arg_groups.arg_utils import ( A, Arg, @@ -29,8 +30,7 @@ from sglang.srt.utils.common import ( ) -@dataclasses.dataclass -class Lora: +class Lora(msgspec.Struct): """Namespace ``lora``.""" _NS_PATH = "lora" diff --git a/python/sglang/srt/arg_groups/fields/memory.py b/python/sglang/srt/arg_groups/fields/memory.py index e184b0530..05d0c074d 100644 --- a/python/sglang/srt/arg_groups/fields/memory.py +++ b/python/sglang/srt/arg_groups/fields/memory.py @@ -9,7 +9,6 @@ how config is shaped at runtime. from __future__ import annotations -import dataclasses import json from typing import ( Any, @@ -17,6 +16,8 @@ from typing import ( Optional, ) +import msgspec + from sglang.srt.arg_groups.arg_utils import ( A, Arg, @@ -24,8 +25,7 @@ from sglang.srt.arg_groups.arg_utils import ( from sglang.srt.arg_groups.choices import RADIX_EVICTION_POLICY_CHOICES -@dataclasses.dataclass -class Memory: +class Memory(msgspec.Struct): """Namespace ``memory``.""" _NS_PATH = "memory" diff --git a/python/sglang/srt/arg_groups/fields/mm.py b/python/sglang/srt/arg_groups/fields/mm.py index c42832bef..7e6a98ba7 100644 --- a/python/sglang/srt/arg_groups/fields/mm.py +++ b/python/sglang/srt/arg_groups/fields/mm.py @@ -9,7 +9,6 @@ how config is shaped at runtime. from __future__ import annotations -import dataclasses import json from typing import ( Any, @@ -20,14 +19,15 @@ from typing import ( Union, ) +import msgspec + from sglang.srt.arg_groups.arg_utils import ( A, Arg, ) -@dataclasses.dataclass -class Mm: +class Mm(msgspec.Struct): """Namespace ``mm``.""" _NS_PATH = "mm" @@ -91,7 +91,7 @@ class Mm: "Restrict client-supplied HTTP(S) image, video, and audio URLs to these " "exact hostnames. Redirect destinations are checked against the same " "allowlist. When unset, remote media from any domain is allowed.", - ] = dataclasses.field(default_factory=list) + ] = msgspec.field(default_factory=list) media_url_max_file_size_mb: A[ int, "Maximum size in MiB for one client-supplied remote media download. " diff --git a/python/sglang/srt/arg_groups/fields/model.py b/python/sglang/srt/arg_groups/fields/model.py index 088139e95..606279c01 100644 --- a/python/sglang/srt/arg_groups/fields/model.py +++ b/python/sglang/srt/arg_groups/fields/model.py @@ -9,7 +9,6 @@ how config is shaped at runtime. from __future__ import annotations -import dataclasses from typing import ( Dict, List, @@ -18,6 +17,8 @@ from typing import ( Union, ) +import msgspec + from sglang.srt.arg_groups.arg_utils import ( A, Arg, @@ -34,8 +35,7 @@ from sglang.srt.utils.common import ( ) -@dataclasses.dataclass -class Model: +class Model(msgspec.Struct): """Namespace ``model``.""" _NS_PATH = "model" diff --git a/python/sglang/srt/arg_groups/fields/observability.py b/python/sglang/srt/arg_groups/fields/observability.py index 98f8f51e0..1e7289f1c 100644 --- a/python/sglang/srt/arg_groups/fields/observability.py +++ b/python/sglang/srt/arg_groups/fields/observability.py @@ -10,7 +10,6 @@ how config is shaped at runtime. from __future__ import annotations import argparse -import dataclasses import json from typing import ( Any, @@ -19,6 +18,8 @@ from typing import ( Optional, ) +import msgspec + from sglang.srt.arg_groups.arg_utils import ( A, Arg, @@ -26,8 +27,7 @@ from sglang.srt.arg_groups.arg_utils import ( from sglang.srt.utils.common import json_list_type -@dataclasses.dataclass -class Observability: +class Observability(msgspec.Struct): """Namespace ``observability``.""" _NS_PATH = "observability" @@ -68,7 +68,7 @@ class Observability: help="Exclude uvicorn access logs whose request path starts with any of these prefixes. Defaults to empty (disabled). Example: --uvicorn-access-log-exclude-prefixes /metrics /health", nargs="*", ), - ] = dataclasses.field(default_factory=list) + ] = msgspec.field(default_factory=list) crash_dump_folder: A[ Optional[str], "Folder path to dump requests from the last 5 min before a crash (if any). If not specified, crash dumping is disabled.", diff --git a/python/sglang/srt/arg_groups/fields/parallel.py b/python/sglang/srt/arg_groups/fields/parallel.py index c323d65ec..4e1c6eb52 100644 --- a/python/sglang/srt/arg_groups/fields/parallel.py +++ b/python/sglang/srt/arg_groups/fields/parallel.py @@ -10,9 +10,10 @@ how config is shaped at runtime. from __future__ import annotations import argparse -import dataclasses from typing import Optional +import msgspec + from sglang.srt.arg_groups.arg_utils import ( A, Arg, @@ -20,8 +21,7 @@ from sglang.srt.arg_groups.arg_utils import ( ) -@dataclasses.dataclass -class Parallel: +class Parallel(msgspec.Struct): """Namespace ``parallel``.""" _NS_PATH = "parallel" diff --git a/python/sglang/srt/arg_groups/fields/schedule.py b/python/sglang/srt/arg_groups/fields/schedule.py index bee8e21ee..d067290a0 100644 --- a/python/sglang/srt/arg_groups/fields/schedule.py +++ b/python/sglang/srt/arg_groups/fields/schedule.py @@ -9,12 +9,13 @@ how config is shaped at runtime. from __future__ import annotations -import dataclasses from typing import ( List, Optional, ) +import msgspec + from sglang.srt.arg_groups.arg_utils import ( A, Arg, @@ -22,8 +23,7 @@ from sglang.srt.arg_groups.arg_utils import ( from sglang.srt.utils.common import human_readable_int -@dataclasses.dataclass -class Schedule: +class Schedule(msgspec.Struct): """Namespace ``schedule``.""" _NS_PATH = "schedule" diff --git a/python/sglang/srt/arg_groups/fields/serving.py b/python/sglang/srt/arg_groups/fields/serving.py index 7877140a7..d7a2ec0c7 100644 --- a/python/sglang/srt/arg_groups/fields/serving.py +++ b/python/sglang/srt/arg_groups/fields/serving.py @@ -9,7 +9,6 @@ how config is shaped at runtime. from __future__ import annotations -import dataclasses import json from typing import ( Any, @@ -18,6 +17,8 @@ from typing import ( Optional, ) +import msgspec + from sglang.srt.arg_groups.arg_utils import ( A, Arg, @@ -25,8 +26,7 @@ from sglang.srt.arg_groups.arg_utils import ( from sglang.srt.utils.common import json_list_type -@dataclasses.dataclass -class Serving: +class Serving(msgspec.Struct): """Namespace ``serving``.""" _NS_PATH = "serving" diff --git a/python/sglang/srt/arg_groups/fields/spec.py b/python/sglang/srt/arg_groups/fields/spec.py index e59274bee..f9dde3e74 100644 --- a/python/sglang/srt/arg_groups/fields/spec.py +++ b/python/sglang/srt/arg_groups/fields/spec.py @@ -9,12 +9,13 @@ how config is shaped at runtime. from __future__ import annotations -import dataclasses from typing import ( Literal, Optional, ) +import msgspec + from sglang.srt.arg_groups.arg_utils import ( A, Arg, @@ -26,8 +27,7 @@ from sglang.srt.arg_groups.choices import ( ) -@dataclasses.dataclass -class Spec: +class Spec(msgspec.Struct): """Namespace ``spec``.""" _NS_PATH = "spec" diff --git a/python/sglang/srt/arg_groups/overrides.py b/python/sglang/srt/arg_groups/overrides.py index 796290ee2..4ece0f599 100644 --- a/python/sglang/srt/arg_groups/overrides.py +++ b/python/sglang/srt/arg_groups/overrides.py @@ -32,7 +32,6 @@ Two declaration forms, keyed on ``hf_config.architectures[0]``: from __future__ import annotations -import dataclasses import json import logging import math @@ -41,6 +40,7 @@ from typing import Any, Callable, Dict, List, Optional, Sequence, Tuple from sglang.srt.arg_groups import model_override_base from sglang.srt.arg_groups.arg_utils import ( field_names, + is_record, resolvable_fields, with_fallback, ) @@ -159,7 +159,7 @@ def declare_resolution(server_args: Any, source: str, **fields: Any) -> None: again, so a declaration afterwards is a silent no-op; post-publish changes go to the bags through ``get_context().override(...)``. """ - if dataclasses.is_dataclass(type(server_args)): + if is_record(server_args): unknown = sorted(set(fields) - field_names(type(server_args))) if unknown: raise AttributeError(f"{source}: {unknown} are not ServerArgs fields") @@ -229,7 +229,7 @@ def record_foreign_defaults( A stand-in record (tests drive the hooks with a plain namespace) has no view to read, so the resolver runs against it directly and uncaptured. """ - if not dataclasses.is_dataclass(server_args): + if not is_record(server_args): return resolve(server_args) recorder = _ForeignDefaults(server_args) result = resolve(recorder) @@ -1726,7 +1726,7 @@ def validate_declarations( """ # Non-dataclass fixtures carry no Arg metadata (mirrors the # resolvable_fields escape); only real ServerArgs is validated. - if not dataclasses.is_dataclass(type(server_args)): + if not is_record(server_args): return whitelist = resolvable_fields(type(server_args)) for source, decl in declarations: diff --git a/python/sglang/srt/arg_groups/pd_disaggregation_hook.py b/python/sglang/srt/arg_groups/pd_disaggregation_hook.py index 285d686a6..f66449f3f 100644 --- a/python/sglang/srt/arg_groups/pd_disaggregation_hook.py +++ b/python/sglang/srt/arg_groups/pd_disaggregation_hook.py @@ -1,10 +1,10 @@ from __future__ import annotations -import dataclasses import logging import os from typing import TYPE_CHECKING, Any +from sglang.srt.arg_groups.arg_utils import record_fields from sglang.srt.arg_groups.overrides import ( declare_resolution, model_config_of, @@ -158,7 +158,7 @@ def _alias_bootstrap_port_to_api_port(server_args: ServerArgs) -> None: cfg = resolving_view(server_args) default_port = next( f.default - for f in dataclasses.fields(server_args) + for f in record_fields(type(server_args)) if f.name == "disaggregation_bootstrap_port" ) if cfg.disaggregation_bootstrap_port not in ( diff --git a/python/sglang/srt/arg_groups/pipeline.py b/python/sglang/srt/arg_groups/pipeline.py index fd5cee710..ef35ba1da 100644 --- a/python/sglang/srt/arg_groups/pipeline.py +++ b/python/sglang/srt/arg_groups/pipeline.py @@ -8,9 +8,9 @@ it. from __future__ import annotations -import dataclasses from typing import Any +from sglang.srt.arg_groups.arg_utils import record_fields from sglang.srt.arg_groups.overrides import ( _page_size_default, _pipeline_parallel_overlap_disable, @@ -49,7 +49,7 @@ def run_resolution_pipeline(server_args: Any) -> None: # stash is the resolution result the projection reads. server_args._raw_input = { field.name: getattr(server_args, field.name) - for field in dataclasses.fields(server_args) + for field in record_fields(type(server_args)) } # Preserve launcher-stage declarations made before Engine starts. They are diff --git a/python/sglang/srt/entrypoints/engine.py b/python/sglang/srt/entrypoints/engine.py index 29575196c..df05675ea 100644 --- a/python/sglang/srt/entrypoints/engine.py +++ b/python/sglang/srt/entrypoints/engine.py @@ -44,6 +44,7 @@ from typing import ( cast, ) +import msgspec import torch import uvloop import zmq @@ -266,7 +267,7 @@ class Engine(EngineScoreMixin, EngineBase): # There was no command line, so the call is what the operator # asked for. `log_level` is filled in above when absent, so it # shows here even when the caller did not pass it. - object.__setattr__( + msgspec.Struct.__setattr__( server_args, "_launch_command", "Engine(" + ", ".join(f"{k}={v!r}" for k, v in kwargs.items()) + ")", diff --git a/python/sglang/srt/runtime_context.py b/python/sglang/srt/runtime_context.py index 4b0a9739b..567f7f654 100644 --- a/python/sglang/srt/runtime_context.py +++ b/python/sglang/srt/runtime_context.py @@ -47,7 +47,6 @@ test-only ``override(**kw)``. from __future__ import annotations -import dataclasses import functools import logging import math @@ -56,6 +55,8 @@ import sys from contextlib import contextmanager from typing import TYPE_CHECKING, Any, Dict, Optional +import msgspec + if TYPE_CHECKING: from sglang.srt.server_args import ServerArgs @@ -461,28 +462,29 @@ def _install_derived_widths() -> None: _install_derived_widths() -class _FlagGroupBase: +class _FlagGroupBase(msgspec.Struct): """Shared flag-group behavior: typo-safe writes + transactional ``override()``. - Groups are plain dataclasses; ``__dataclass_fields__`` is the single source - of truth for which leaves exist, so a mistyped name fails loudly instead of - creating a stray attribute. + ``__struct_fields__`` is the single source of truth for which leaves exist, + so a mistyped name fails loudly instead of creating a stray attribute. The + write goes through ``super().__setattr__``: a ``Struct`` keeps its fields in + its own layout, so ``object.__setattr__`` does not reach them. """ def __setattr__(self, name: str, value: Any) -> None: - if name not in type(self).__dataclass_fields__: + if name not in type(self).__struct_fields__: raise AttributeError( f"{type(self).__name__} has no flag '{name}' (leaves are " - "declared as dataclass fields; check for typos)" + "declared as struct fields; check for typos)" ) - object.__setattr__(self, name, value) + super().__setattr__(name, value) @contextmanager def override(self, **kwargs): """Temporarily force flag values, restoring on exit. Transactional (keys validated before any write) — the test-only injection primitive.""" - fields = type(self).__dataclass_fields__ + fields = type(self).__struct_fields__ unknown = set(kwargs) - set(fields) if unknown: raise ValueError( @@ -490,15 +492,14 @@ class _FlagGroupBase: ) saved = {name: getattr(self, name) for name in kwargs} for name, value in kwargs.items(): - object.__setattr__(self, name, value) + setattr(self, name, value) try: yield self finally: for name, value in saved.items(): - object.__setattr__(self, name, value) + setattr(self, name, value) -@dataclasses.dataclass class CaptureFlags(_FlagGroupBase): """Capture-time flags; never frozen (written during cuda-graph capture).""" @@ -512,7 +513,6 @@ class CaptureFlags(_FlagGroupBase): disable_dispose_tensor: bool = False -@dataclasses.dataclass class MoeFlags(_FlagGroupBase): """MoE runtime flags, materialized by ``initialize_moe_config`` (scheduler init, after distributed setup). ``a2a_backend`` / ``runner_backend`` / @@ -552,7 +552,6 @@ class MoeFlags(_FlagGroupBase): speculative_context: bool = False -@dataclasses.dataclass class DpFlags(_FlagGroupBase): """DP-attention runtime flags, materialized by ``initialize_dp_attention`` (after distributed setup; reads the model config). Topology values @@ -576,7 +575,6 @@ class DpFlags(_FlagGroupBase): buffer_device: Any = None -@dataclasses.dataclass class SpFlags(_FlagGroupBase): """LayerNorm sequence-parallelism flags, materialized by ``initialize_layernorm_sp`` (after distributed setup; reads the model @@ -585,7 +583,6 @@ class SpFlags(_FlagGroupBase): enabled: bool = False -@dataclasses.dataclass class Flags(_FlagGroupBase): """Root of the runtime-flags tier. @@ -595,13 +592,12 @@ class Flags(_FlagGroupBase): by lifecycle (``capture``) or subsystem (``moe`` / ``dp`` / ``sp``). """ - capture: CaptureFlags = dataclasses.field(default_factory=CaptureFlags) - moe: MoeFlags = dataclasses.field(default_factory=MoeFlags) - dp: DpFlags = dataclasses.field(default_factory=DpFlags) - sp: SpFlags = dataclasses.field(default_factory=SpFlags) + capture: CaptureFlags = msgspec.field(default_factory=CaptureFlags) + moe: MoeFlags = msgspec.field(default_factory=MoeFlags) + dp: DpFlags = msgspec.field(default_factory=DpFlags) + sp: SpFlags = msgspec.field(default_factory=SpFlags) -@dataclasses.dataclass class Resources(_FlagGroupBase): """Process-level resource handles: named slots with one reset lifecycle, scoped test injection via ``override()``, and the creation/publish @@ -616,16 +612,16 @@ class Resources(_FlagGroupBase): expert_distribution_recorder: Any = None expert_location_metadata: Any = None # LPLB: layer_id -> solver. - lplb_solvers: dict = dataclasses.field(default_factory=dict) + lplb_solvers: dict = msgspec.field(default_factory=dict) # Named side streams (see RuntimeContext.get_stream): name -> stream. - streams: dict = dataclasses.field(default_factory=dict) + streams: dict = msgspec.field(default_factory=dict) # Named persistent buffers (see RuntimeContext.get_buffer): name -> tensor. # Accessors with bespoke semantics (grow-only, per-device keys) manage # their entries directly. - buffers: dict = dataclasses.field(default_factory=dict) + buffers: dict = msgspec.field(default_factory=dict) # Persistent reusable CUDA events for non-EP DP TBO, keyed by # (kind, subbatch) — see dp_attention._tbo_event for why reuse matters. - tbo_event_pool: dict = dataclasses.field(default_factory=dict) + tbo_event_pool: dict = msgspec.field(default_factory=dict) # State capturers (installed by their subsystems when capture is on). indexer_capturer: Any = None experts_capturer: Any = None @@ -1290,7 +1286,7 @@ class _ServerArgsOverride: # Underscore names seed private property caches (the strict guard # exempts them); everything else must be a real config field. unknown = {name for name in self._fields if not name.startswith("_")} - set( - type(server_args).__dataclass_fields__ + type(server_args).__struct_fields__ ) if unknown: raise ValueError( @@ -1303,7 +1299,7 @@ class _ServerArgsOverride: # real field, and seeding it as a raw attribute would leave the earlier # declaration authoritative, so `resolution_result` and the bag would # both keep answering the pre-override value. - fields = set(type(server_args).__dataclass_fields__) + fields = set(type(server_args).__struct_fields__) declared = {n: v for n, v in self._fields.items() if n in fields} if declared: declare_resolution(server_args, "override_server_args", **declared) @@ -1311,7 +1307,7 @@ class _ServerArgsOverride: # and friends), which are not configuration and never were. seeds = {n: v for n, v in self._fields.items() if n not in fields} for name, value in seeds.items(): - object.__setattr__(server_args, name, value) + msgspec.Struct.__setattr__(server_args, name, value) ctx.set_server_args(server_args) self._installed = True return server_args @@ -1680,7 +1676,7 @@ def set_global_dwdp_manager(manager: Any) -> None: def _group_leaves(group: _FlagGroupBase) -> dict[str, Any]: """The leaf values of a flag group, recursively.""" leaves: dict[str, Any] = {} - for name in type(group).__dataclass_fields__: + for name in type(group).__struct_fields__: value = getattr(group, name) if isinstance(value, _FlagGroupBase): leaves[name] = _group_leaves(value) diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index 509129065..0daef9243 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -42,9 +42,13 @@ import tempfile import uuid from typing import Any, NoReturn +import msgspec + from sglang.kernels.ops.kv_canary.consts import RealKvHashMode from sglang.srt.arg_groups.arg_utils import ( add_cli_args_from_dataclass, + is_record, + record_fields, ) from sglang.srt.arg_groups.argparse_actions import ( DeprecatedStoreTrueAction, @@ -170,13 +174,14 @@ from sglang.srt.utils.common import ( # noqa: F401 def _plain(value: Any) -> Any: - """``dataclasses.asdict``'s conversion, applied to one value: dataclasses - become dicts, containers recurse, everything else is deep-copied (a caller - mutating the dump must not reach the live configuration).""" - if dataclasses.is_dataclass(value) and not isinstance(value, type): + """``asdict``'s conversion, applied to one value: a record -- Struct or + dataclass, since nested config values are both -- becomes a dict, containers + recurse, everything else is deep-copied (a caller mutating the dump must not + reach the live configuration).""" + if not isinstance(value, type) and is_record(value): return { field.name: _plain(getattr(value, field.name)) - for field in dataclasses.fields(value) + for field in record_fields(type(value)) } if isinstance(value, tuple) and hasattr(value, "_fields"): # namedtuple return type(value)(*(_plain(item) for item in value)) @@ -307,7 +312,7 @@ class ServerArgs: """This configuration as a plain dict of resolved field values. What the whole-object readbacks report (`/server_info` and its gRPC and - in-process twins). `dataclasses.asdict(self)` reads the fields, which + in-process twins). A plain `asdict` reads the fields, which carry the raw input; this reads the declarations, so it answers with what resolution decided. Nested dataclass fields are expanded the way `asdict` expands them; the private resolution bookkeeping and the @@ -316,7 +321,7 @@ class ServerArgs: return { field.name: _plain(resolution_result(self, field.name)) - for field in dataclasses.fields(self) + for field in record_fields(type(self)) } # ------------------------------------------------------------------ @@ -350,7 +355,7 @@ class ServerArgs: "--sampling-backend", type=str, choices=sampling_backend_choices, - default=ServerArgs.sampling_backend, + default=_declared_default("sampling_backend"), help="Choose the kernels for sampling layers.", ) @@ -359,7 +364,7 @@ class ServerArgs: "--reasoning-parser", type=str, choices=["auto"] + reasoning_parser_choices, - default=ServerArgs.reasoning_parser, + default=_declared_default("reasoning_parser"), help=f"Specify the parser for reasoning models. " f"Use 'auto' to detect from chat template. " f"Options include: {reasoning_parser_choices}.", @@ -369,7 +374,7 @@ class ServerArgs: "--tool-call-parser", type=str, choices=["auto"] + tool_call_parser_choices, - default=ServerArgs.tool_call_parser, + default=_declared_default("tool_call_parser"), help=f"Specify the parser for handling tool-call interactions. " f"Use 'auto' to detect from chat template. " f"Options include: {tool_call_parser_choices}.", @@ -377,7 +382,7 @@ class ServerArgs: parser.add_argument( "--kv-canary-real-data", type=str, - default=ServerArgs.kv_canary_real_data, + default=_declared_default("kv_canary_real_data"), choices=[m.name.lower() for m in RealKvHashMode], help=( "Check the real KV-cache in the canary. " @@ -415,9 +420,7 @@ class ServerArgs: # Some dataclass fields (e.g. stat_loggers) intentionally have no CLI # surface and won't appear on the argparse Namespace. Skip them so the # dataclass default applies. - attrs = [ - attr.name for attr in dataclasses.fields(cls) if hasattr(args, attr.name) - ] + attrs = [attr.name for attr in record_fields(cls) if hasattr(args, attr.name)] return cls(**{attr: getattr(args, attr) for attr in attrs}) def get_tokenizer_worker_class(self): @@ -469,7 +472,26 @@ class ServerArgs: "resolved config; a value one runner owns travels as a " "constructor argument." ) - object.__setattr__(self, name, value) + # The Struct's own setter, spelled explicitly: this method is copied + # into the class `defstruct` builds, so a zero-argument `super()` would + # still close over the class it was written in. `object.__setattr__` + # does not reach a Struct's fields at all. + msgspec.Struct.__setattr__(self, name, value) + + def __reduce__(self): + """Pickle the record *and* what resolution left on it. + + A Struct pickles its fields; everything else lives in the `dict=True` + namespace and would be dropped, which for this record means the input + snapshot, the declaration stash and the resolution flags -- the whole + reason a child can publish what its parent decided without resolving + again. Reconstruction restores the fields first and the bookkeeping + after, so the seal is re-armed only once the fields are in place. + """ + return ( + _rebuild_server_args, + (type(self), msgspec.structs.asdict(self), dict(self.__dict__)), + ) def check_server_args(self): from sglang.srt.arg_groups.validation_hook import check_server_args @@ -522,9 +544,23 @@ ServerArgs._NS_BY_FIELD = _namespaces # The classes themselves, so the bag projection can find the declarations # that are not fields -- the derived half of each namespace. ServerArgs._NAMESPACES = _INPUT_NAMESPACES -for _name, _value in _defaults.items(): - setattr(ServerArgs, _name, _value) -ServerArgs = dataclasses.dataclass(ServerArgs) +# `dict=True` so the record can carry what is not configuration -- the input +# snapshot, the declaration stash, the resolution flags, the memo slots. A +# Struct has no `__dict__` without it, and those are exactly the underscore +# names `_underscore_field_names()` is careful *not* to include. +ServerArgs = msgspec.defstruct( + "ServerArgs", + [ + (_name, _ann, _defaults[_name]) if _name in _defaults else (_name, _ann) + for _name, _ann in ServerArgs.__annotations__.items() + ], + namespace={ + _k: _v + for _k, _v in vars(ServerArgs).items() + if _k not in ("__dict__", "__weakref__", "__annotations__") + }, + dict=True, +) # -------------------------------------------------------------------------- @@ -586,9 +622,7 @@ def _underscore_field_names() -> frozenset: by spelling would leave exactly one leaf writable on a read-only record. """ return frozenset( - field.name - for field in dataclasses.fields(ServerArgs) - if field.name.startswith("_") + field.name for field in record_fields(ServerArgs) if field.name.startswith("_") ) @@ -631,6 +665,28 @@ def get_global_server_args() -> NoReturn: ) +def _rebuild_server_args(cls, fields, bookkeeping): + """Rebuild a pickled record: fields through the constructor, the rest after.""" + record = cls(**fields) + record.__dict__.update(bookkeeping) + return record + + +def _declared_default(name: str): + """The declared default of a field, for a manual `add_argument`. + + `ServerArgs.` used to answer with it. The record is a Struct now, so + that expression returns the slot descriptor instead -- which argparse + happily stores as the default, and the first reader gets a + `member_descriptor` where it expected a string. + """ + return next( + field.default + for field in msgspec.structs.fields(ServerArgs) + if field.name == name + ) + + def prepare_server_args(argv: list[str]) -> ServerArgs: """ Prepare the server arguments from the command line arguments. @@ -669,7 +725,7 @@ def prepare_server_args(argv: list[str]) -> ServerArgs: # Not a field: the record's fields are the configuration, and this is how # the configuration was asked for. It rides along on the record so a # subprocess copy can answer the same question the launcher can. - object.__setattr__(server_args, "_launch_command", " ".join(argv)) + server_args._launch_command = " ".join(argv) return server_args diff --git a/python/sglang/test/test_utils.py b/python/sglang/test/test_utils.py index 0920589d2..440a93b65 100644 --- a/python/sglang/test/test_utils.py +++ b/python/sglang/test/test_utils.py @@ -26,6 +26,7 @@ from types import ModuleType, SimpleNamespace from typing import Any, Awaitable, Callable, List, Optional, Tuple import aiohttp +import msgspec import numpy as np import requests import torch @@ -2118,7 +2119,7 @@ def server_args_variant(server_args, **fields): unknown = { name for name in fields - if name not in cls.__dataclass_fields__ + if name not in cls.__struct_fields__ and not hasattr(cls, name) and name not in _RUNNER_WRITTEN_NAMES } @@ -2130,16 +2131,14 @@ def server_args_variant(server_args, **fields): stash = getattr(variant, "_resolved_overrides", None) if stash is None: stash = [] - object.__setattr__(variant, "_resolved_overrides", stash) + msgspec.Struct.__setattr__(variant, "_resolved_overrides", stash) declared = { - name: value - for name, value in fields.items() - if name in cls.__dataclass_fields__ + name: value for name, value in fields.items() if name in cls.__struct_fields__ } if declared: stash.append(("server_args_variant", dict(declared))) for name, value in fields.items(): - object.__setattr__(variant, name, value) + msgspec.Struct.__setattr__(variant, name, value) return variant diff --git a/test/registered/cpu/test_server_args_backend.py b/test/registered/cpu/test_server_args_backend.py index 5427990f4..0bedb61ef 100644 --- a/test/registered/cpu/test_server_args_backend.py +++ b/test/registered/cpu/test_server_args_backend.py @@ -17,7 +17,7 @@ register_cpu_ci(est_time=10, suite="base-b-test-cpu-arm64") class TestServerArgsCPUBackend(CustomTestCase): def _make_server_args(self, attention_backend=None): - server_args = ServerArgs.__new__(ServerArgs) + server_args = ServerArgs(model_path="dummy") server_args.device = "cpu" server_args.attention_backend = attention_backend server_args.sampling_backend = None diff --git a/test/registered/prefill_only/test_multi_item_scoring.py b/test/registered/prefill_only/test_multi_item_scoring.py index c5c5ab855..e94346408 100644 --- a/test/registered/prefill_only/test_multi_item_scoring.py +++ b/test/registered/prefill_only/test_multi_item_scoring.py @@ -21,6 +21,7 @@ import torch from transformers import AutoConfig, AutoTokenizer from sglang.srt.entrypoints.engine import Engine +from sglang.srt.server_args import _declared_default from sglang.test.ci.ci_register import register_cuda_ci from sglang.test.test_utils import ( DEFAULT_SMALL_MODEL_NAME_FOR_TEST, @@ -56,9 +57,8 @@ class TestMISServerArgsValidation(unittest.TestCase): def test_enable_mis_default(self): """Test that enable_mis defaults to False.""" - from sglang.srt.server_args import ServerArgs - self.assertEqual(ServerArgs.enable_mis, False) + self.assertEqual(_declared_default("enable_mis"), False) class TestMultiItemScoringOptimization(CustomTestCase): diff --git a/test/registered/unit/entrypoints/test_server_info.py b/test/registered/unit/entrypoints/test_server_info.py index 0dee582f8..191c0bb22 100644 --- a/test/registered/unit/entrypoints/test_server_info.py +++ b/test/registered/unit/entrypoints/test_server_info.py @@ -21,11 +21,13 @@ Current coverage: """ import asyncio -import dataclasses import json import unittest from types import SimpleNamespace +import msgspec +import msgspec.structs + from sglang.srt.arg_groups.validation_hook import check_load_publish_args from sglang.srt.entrypoints import http_server from sglang.srt.lora.lora_registry import LoRARef @@ -457,7 +459,7 @@ class TestServerInfoExistingFieldsPreserved(CustomTestCase): info = _call_server_info_with(args) - for field in dataclasses.fields(ServerArgs): + for field in msgspec.structs.fields(ServerArgs): self.assertIn( field.name, info, diff --git a/test/registered/unit/managers/test_tokenizer_config_updates.py b/test/registered/unit/managers/test_tokenizer_config_updates.py index 1c20e2865..da73e9dce 100644 --- a/test/registered/unit/managers/test_tokenizer_config_updates.py +++ b/test/registered/unit/managers/test_tokenizer_config_updates.py @@ -10,6 +10,9 @@ import re import unittest from pathlib import Path +import msgspec +import msgspec.structs + import sglang from sglang.srt.managers.tokenizer_manager import TokenizerManager from sglang.srt.runtime_context import get_context, publish, reset_context @@ -85,14 +88,15 @@ class TestTokenizerConfigUpdates(CustomTestCase): ) def test_the_dump_snapshot_identifies_the_running_checkpoint(self): - import dataclasses manager = _manager(self, load_format="auto") manager.model_path = "at-startup" manager.served_model_name = "at-startup" manager._update_model_path_info("after-reload", "dummy") - snapshot = manager.resolved_config_dict(dataclasses.asdict(manager.server_args)) + snapshot = manager.resolved_config_dict( + msgspec.structs.asdict(manager.server_args) + ) self.assertEqual(snapshot["model_path"], "after-reload") self.assertEqual(snapshot["served_model_name"], "after-reload") self.assertEqual(snapshot["load_format"], "dummy") @@ -116,7 +120,6 @@ class TestTokenizerConfigUpdates(CustomTestCase): self.assertIsNone(manager._dump_config_snapshot()) def test_an_unpickleable_field_does_not_lose_the_dump(self): - import dataclasses import pickle # What --custom-sigquit-handler leaves on a real ServerArgs. @@ -128,7 +131,7 @@ class TestTokenizerConfigUpdates(CustomTestCase): "server_args": manager.server_args, "config_updates": get_context().overrides_log(), "resolved_config": manager.resolved_config_dict( - dataclasses.asdict(manager.server_args) + msgspec.structs.asdict(manager.server_args) ), "requests": [], } diff --git a/test/registered/unit/parser/test_template_manager.py b/test/registered/unit/parser/test_template_manager.py index 52edc6b83..33ebdab73 100644 --- a/test/registered/unit/parser/test_template_manager.py +++ b/test/registered/unit/parser/test_template_manager.py @@ -4,6 +4,8 @@ import unittest from types import ModuleType, SimpleNamespace from unittest.mock import Mock, patch +import msgspec + from sglang.srt.parser.template_detection import ( REASONING_PARSER_RULES, TOOL_CALL_PARSER_RULES, @@ -892,7 +894,9 @@ class TestResolveAutoParsers(unittest.TestCase): reasoning_parser="auto", tool_call_parser="auto", ) - object.__setattr__(args, "model_path", "nonexistent/model-does-not-exist-xyz") + msgspec.Struct.__setattr__( + 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")), diff --git a/test/registered/unit/server_args/test_page_major_backend_allowlist.py b/test/registered/unit/server_args/test_page_major_backend_allowlist.py index bacd79616..a8ac16cd6 100644 --- a/test/registered/unit/server_args/test_page_major_backend_allowlist.py +++ b/test/registered/unit/server_args/test_page_major_backend_allowlist.py @@ -31,6 +31,8 @@ the unified pool today. import unittest from types import SimpleNamespace +import msgspec + from sglang.srt.arg_groups.kv_cache_hook import handle_page_major_kv_layout from sglang.srt.configs.model_config import AttentionArch from sglang.srt.server_args import ServerArgs @@ -50,7 +52,7 @@ def _accepts( ) -> bool: """Run just `handle_page_major_kv_layout` against a minimal stand-in, since ServerArgs' real constructor pulls in a model config.""" - sa = ServerArgs.__new__(ServerArgs) + sa = ServerArgs(model_path="dummy") for name, value in { "enable_unified_memory": unified, # The unified pool sets this itself; without it the flag must be explicit @@ -64,8 +66,8 @@ def _accepts( "linear_attn_prefill_backend": linear_prefill, "mamba_backend": "triton", }.items(): - object.__setattr__(sa, name, value) - object.__setattr__( + msgspec.Struct.__setattr__(sa, name, value) + setattr( sa, "_model_config", SimpleNamespace( diff --git a/test/registered/unit/server_args/test_record_holds_the_raw_input.py b/test/registered/unit/server_args/test_record_holds_the_raw_input.py index 537592f89..f9b153511 100644 --- a/test/registered/unit/server_args/test_record_holds_the_raw_input.py +++ b/test/registered/unit/server_args/test_record_holds_the_raw_input.py @@ -8,13 +8,15 @@ value the caller still holds and the snapshot cannot see it. """ import copy -import dataclasses import json import os import shutil import tempfile import unittest +import msgspec +import msgspec.structs + from sglang.srt.server_args import ServerArgs from sglang.test.ci.ci_register import register_cpu_ci from sglang.test.test_utils import CustomTestCase @@ -108,7 +110,7 @@ class TestRecordHoldsTheRawInput(CustomTestCase): moved = { field.name: (raw[field.name], getattr(server_args, field.name)) - for field in dataclasses.fields(server_args) + for field in msgspec.structs.fields(server_args) if _moved(getattr(server_args, field.name), raw[field.name]) } self.assertEqual( diff --git a/test/registered/unit/server_args/test_resolution_declarations.py b/test/registered/unit/server_args/test_resolution_declarations.py index b7739665a..5daba2674 100644 --- a/test/registered/unit/server_args/test_resolution_declarations.py +++ b/test/registered/unit/server_args/test_resolution_declarations.py @@ -9,7 +9,6 @@ or is projected into the wrong namespace therefore fails on observed state. import ast import copy -import dataclasses import json import os import pathlib @@ -18,6 +17,9 @@ import tempfile import unittest import unittest.mock +import msgspec +import msgspec.structs + import sglang from sglang.srt.arg_groups.overrides import resolution_result from sglang.srt.server_args import ServerArgs @@ -30,7 +32,7 @@ _SRT = pathlib.Path(sglang.__file__).resolve().parent / "srt" # Every field of the record: resolution has no bare-assignment writer left, so # the scan states that as a whole rather than a converted-so-far list. -_RESOLVED_FIELDS = frozenset(field.name for field in dataclasses.fields(ServerArgs)) +_RESOLVED_FIELDS = frozenset(field.name for field in msgspec.structs.fields(ServerArgs)) # Shapes the agreement check runs on. Each needs a real config.json: # `model_path="dummy"` takes the pipeline's early return. @@ -201,14 +203,14 @@ class TestResolutionDeclarations(CustomTestCase): supplied = {"random_seed": 42, **shape} server_args = self._resolve(shape) overlay = _stash_overlay(server_args) - for field in dataclasses.fields(server_args): + for field in msgspec.structs.fields(server_args): if field.name in ("model_path", "device") or field.name in overlay: continue if field.name in supplied: before = supplied[field.name] - elif field.default is not dataclasses.MISSING: + elif field.default is not msgspec.NODEFAULT: before = field.default - elif field.default_factory is not dataclasses.MISSING: + elif field.default_factory is not msgspec.NODEFAULT: before = field.default_factory() else: continue @@ -278,7 +280,7 @@ class TestResolutionDeclarations(CustomTestCase): dump = server_args.resolved_dict() self.assertEqual( sorted(dump), - sorted(field.name for field in dataclasses.fields(server_args)), + sorted(field.name for field in msgspec.structs.fields(server_args)), "the readback dump is no longer exactly the fields", ) leaked = sorted( @@ -532,7 +534,7 @@ class TestResolutionDeclarations(CustomTestCase): overlay = _stash_overlay(server_args) raw_input = getattr(server_args, "_raw_input", None) self.assertTrue(raw_input, f"{shape}: the record kept no raw snapshot") - for field in dataclasses.fields(server_args): + for field in msgspec.structs.fields(server_args): name = field.name if name in overlay or name not in raw_input: continue diff --git a/test/registered/unit/server_args/test_resolution_is_reproducible.py b/test/registered/unit/server_args/test_resolution_is_reproducible.py index b6bc293ee..4ed946b1c 100644 --- a/test/registered/unit/server_args/test_resolution_is_reproducible.py +++ b/test/registered/unit/server_args/test_resolution_is_reproducible.py @@ -33,6 +33,8 @@ import tempfile import unittest import unittest.mock +import msgspec +import msgspec.structs import torch from sglang.srt.arg_groups.overrides import ( @@ -270,7 +272,7 @@ class TestResolutionIsReproducible(_RestoresProcessState, CustomTestCase): the one of those that a shared mutable could corrupt. """ out = {} - for field in dataclasses.fields(server_args): + for field in msgspec.structs.fields(server_args): if field.name in _NOT_COMPARABLE: continue # The resolution result, not the field: a declaration-only resolver @@ -517,7 +519,7 @@ class TestALateDeclarationKeepsTheResolution(_RestoresProcessState, CustomTestCa record arrives resolved, and a copy would throw that away. """ parent = self._resolved() - bare = dataclasses.replace(parent, dist_init_addr="1.2.3.4:5000") + bare = msgspec.structs.replace(parent, dist_init_addr="1.2.3.4:5000") self.assertFalse( getattr(bare, "_resolution_finished", False), "a bare replace carried the flag; then this test proves nothing", @@ -528,7 +530,7 @@ class TestALateDeclarationKeepsTheResolution(_RestoresProcessState, CustomTestCa resolution_result(parent, field.name), resolution_result(bare, field.name), ) - for field in dataclasses.fields(parent) + for field in msgspec.structs.fields(parent) if field.name not in ("dist_init_addr", "random_seed") and repr(resolution_result(parent, field.name)) != repr(resolution_result(bare, field.name)) diff --git a/test/registered/unit/server_args/test_resolution_reads_the_declarations.py b/test/registered/unit/server_args/test_resolution_reads_the_declarations.py index d04a1a0be..c608c43ef 100644 --- a/test/registered/unit/server_args/test_resolution_reads_the_declarations.py +++ b/test/registered/unit/server_args/test_resolution_reads_the_declarations.py @@ -22,10 +22,12 @@ the two scopes it can derive exactly. """ import ast -import dataclasses import pathlib import re +import msgspec +import msgspec.structs + import sglang from sglang.srt.server_args import ServerArgs from sglang.test.ci.ci_register import register_cpu_ci @@ -34,7 +36,7 @@ from sglang.test.test_utils import CustomTestCase register_cpu_ci(est_time=45, suite="base-a-test-cpu") _SRT = pathlib.Path(sglang.__file__).resolve().parent / "srt" -_FIELDS = frozenset(field.name for field in dataclasses.fields(ServerArgs)) +_FIELDS = frozenset(field.name for field in msgspec.structs.fields(ServerArgs)) # Names a config travels under. `args` is included because the platform hooks # use it; a false positive would be a function taking an argparse Namespace and diff --git a/test/registered/unit/server_args/test_server_args.py b/test/registered/unit/server_args/test_server_args.py index 46c4ec554..2fa31fe7f 100644 --- a/test/registered/unit/server_args/test_server_args.py +++ b/test/registered/unit/server_args/test_server_args.py @@ -1,5 +1,4 @@ import argparse -import dataclasses import json import os import pickle @@ -10,6 +9,9 @@ import unittest from types import SimpleNamespace from unittest.mock import MagicMock, patch +import msgspec +import msgspec.structs + import sglang.srt.server_args as server_args_module from sglang.srt.arg_groups import parallel_hook, pd_disaggregation_hook, serving_hook from sglang.srt.arg_groups.attention_hook import ( @@ -267,7 +269,7 @@ class TestPrepareServerArgs(CustomTestCase): # And across the hop that matters: the scheduler and the draft worker # rebuild the record from its fields and resolve again, so the bit has # to survive `asdict` and come back the same the second time. - reconstructed = ServerArgs(**dataclasses.asdict(inherited)) + reconstructed = ServerArgs(**msgspec.structs.asdict(inherited)) handle_missing_default_values(reconstructed) self.assertFalse( @@ -1046,7 +1048,9 @@ class TestContextParallelServerArgs(CustomTestCase): ServerArgs.add_cli_args(self.parser) def _new_cp_args(self, **overrides): - server_args = object.__new__(ServerArgs) + # Constructed, not conjured: a Struct has no uninitialized form, and + # every field this case does not name wants its declared default + # anyway. defaults = dict( enable_prefill_cp=False, cp_strategy=None, @@ -1061,9 +1065,7 @@ class TestContextParallelServerArgs(CustomTestCase): enable_aiter_allreduce_fusion=False, ) defaults.update(overrides) - for key, value in defaults.items(): - setattr(server_args, key, value) - return server_args + return ServerArgs(**defaults) def test_canonical_prefill_cp_requires_strategy(self): args = self.parser.parse_args(["--model", "dummy", "--enable-prefill-cp"]) @@ -2170,7 +2172,7 @@ class TestDeepEPv2Args(CustomTestCase): prefill=PhaseConfig(backend=Backend.FULL, max_bs=512), ) server_args._resolved_overrides = [] - valid = {f.name for f in dataclasses.fields(ServerArgs)} + valid = {f.name for f in msgspec.structs.fields(ServerArgs)} for key, value in overrides.items(): # Reject stale field names before setattr silently accepts them. assert key in valid, f"{key} is not a ServerArgs field" @@ -2489,7 +2491,7 @@ class TestHandleCrashDumpEnv(CustomTestCase): ) def _run_handler(self, crash_dump_folder, preset_env=None): - server_args = ServerArgs.__new__(ServerArgs) + server_args = ServerArgs(model_path="dummy") server_args.crash_dump_folder = crash_dump_folder with patch.dict(os.environ, preset_env or {}): for key in self._COREDUMP_ENV_KEYS: @@ -2868,7 +2870,7 @@ class TestTheInputIsSealedDuringResolution(CustomTestCase): server_args = ServerArgs(model_path="dummy", device="cuda") # The seal is what the pipeline runs under; drive it directly rather # than injecting a violation into a real handler. - object.__setattr__(server_args, "_input_frozen", True) + msgspec.Struct.__setattr__(server_args, "_input_frozen", True) with self.assertRaisesRegex(AttributeError, "during resolution"): server_args.tp_size = 4 # and the message says what to do instead @@ -2901,7 +2903,7 @@ class TestTheInputIsSealedDuringResolution(CustomTestCase): ) server_args = ServerArgs(model_path="dummy", device="cuda") - object.__setattr__(server_args, "_input_frozen", True) + msgspec.Struct.__setattr__(server_args, "_input_frozen", True) def foreign(config): # What a plugin does: read what is decided, assign a default. @@ -2961,7 +2963,7 @@ class TestLaunchCommand(CustomTestCase): """It describes how the configuration was asked for, so it is not part of the configuration: no CLI flag, no namespace, not in the bags.""" self.assertNotIn( - "launch_command", {f.name for f in dataclasses.fields(ServerArgs)} + "launch_command", {f.name for f in msgspec.structs.fields(ServerArgs)} ) self.assertNotIn( "launch_command", ServerArgs(model_path="/tmp/x").resolved_dict() diff --git a/test/registered/unit/server_args/test_unified_prefill_cuda_graph_gate.py b/test/registered/unit/server_args/test_unified_prefill_cuda_graph_gate.py index 41b5da611..a46f6069f 100644 --- a/test/registered/unit/server_args/test_unified_prefill_cuda_graph_gate.py +++ b/test/registered/unit/server_args/test_unified_prefill_cuda_graph_gate.py @@ -32,6 +32,8 @@ user's stated intent), and decode capture is untouched either way. import unittest from types import SimpleNamespace +import msgspec + from sglang.srt.arg_groups.kv_cache_hook import handle_unified_memory_pool from sglang.srt.model_executor.cuda_graph_config import Backend from sglang.srt.server_args import ServerArgs @@ -42,7 +44,7 @@ register_cpu_ci(est_time=8, suite="base-a-test-cpu") def _run_handler(*, prefill_backend, explicit): """Run just `handle_unified_memory_pool` over a minimal stand-in.""" - sa = ServerArgs.__new__(ServerArgs) + sa = ServerArgs(model_path="dummy") cg = SimpleNamespace( prefill=SimpleNamespace(backend=prefill_backend), decode=SimpleNamespace(backend=Backend.FULL), @@ -59,7 +61,7 @@ def _run_handler(*, prefill_backend, explicit): "cuda_graph_config": cg, "cuda_graph_backend_prefill": prefill_backend if explicit else None, }.items(): - object.__setattr__(sa, name, value) + msgspec.Struct.__setattr__(sa, name, value) handle_unified_memory_pool(sa) return cg diff --git a/test/registered/unit/server_args/test_unified_tbo_gate.py b/test/registered/unit/server_args/test_unified_tbo_gate.py index d3c290ad4..f0f62c7c5 100644 --- a/test/registered/unit/server_args/test_unified_tbo_gate.py +++ b/test/registered/unit/server_args/test_unified_tbo_gate.py @@ -23,6 +23,8 @@ pair, so without this gate a running server crashes mid-serving. import unittest from types import SimpleNamespace +import msgspec + from sglang.srt.arg_groups.kv_cache_hook import handle_unified_memory_pool from sglang.srt.model_executor.cuda_graph_config import Backend from sglang.srt.server_args import ServerArgs @@ -33,7 +35,7 @@ register_cpu_ci(est_time=5, suite="base-a-test-cpu") def _run_handler(*, unified, tbo): """Run just `handle_unified_memory_pool` over a minimal stand-in.""" - sa = ServerArgs.__new__(ServerArgs) + sa = ServerArgs(model_path="dummy") for name, value in { "enable_unified_memory": unified, "enable_two_batch_overlap": tbo, @@ -49,7 +51,7 @@ def _run_handler(*, unified, tbo): ), "cuda_graph_backend_prefill": Backend.DISABLED, }.items(): - object.__setattr__(sa, name, value) + msgspec.Struct.__setattr__(sa, name, value) handle_unified_memory_pool(sa) diff --git a/test/registered/unit/test_model_overrides.py b/test/registered/unit/test_model_overrides.py index 51c8f1e13..fff1cdc6b 100644 --- a/test/registered/unit/test_model_overrides.py +++ b/test/registered/unit/test_model_overrides.py @@ -36,6 +36,7 @@ from sglang.srt.runtime_context import ( override_platform, reset_context, ) +from sglang.srt.server_args import _declared_default from sglang.test.test_utils import CustomTestCase @@ -1613,14 +1614,13 @@ class TestGoldenModelOverrides(_IsolatedPublish): from sglang.srt.arg_groups.model_overrides.deepseek_v4 import ( _deepseek_v4_overrides, ) - from sglang.srt.server_args import ServerArgs hf = SimpleNamespace(architectures=["DeepseekV4ForCausalLM"]) def _args(**kw): defaults = dict( device="cuda", - swa_full_tokens_ratio=ServerArgs.swa_full_tokens_ratio, + swa_full_tokens_ratio=_declared_default("swa_full_tokens_ratio"), moe_a2a_backend="none", moe_runner_backend="auto", _model_config=SimpleNamespace(is_fp4_experts=True, nvfp4_moe_meta=None), diff --git a/test/registered/unit/test_runtime_context.py b/test/registered/unit/test_runtime_context.py index 851416ea5..717f9d838 100644 --- a/test/registered/unit/test_runtime_context.py +++ b/test/registered/unit/test_runtime_context.py @@ -14,6 +14,9 @@ import unittest import warnings from unittest.mock import patch +import msgspec +import msgspec.structs + import sglang as _sglang import sglang.srt.server_args as server_args_module from sglang.srt.arg_groups.arg_utils import NS, A, Arg @@ -438,7 +441,7 @@ class TestServerArgsScopedOverride(_IsolatedServerArgs): from sglang.srt.runtime_context import get_spec name = "_speculative_draft_quantization_explicitly_set" - self.assertIn(name, ServerArgs.__dataclass_fields__) + self.assertIn(name, ServerArgs.__struct_fields__) published = get_context().override_server_args(**{name: True}).install() # The record keeps the operator's input, as it does for every other @@ -473,7 +476,6 @@ class TestServerArgsScopedOverride(_IsolatedServerArgs): override.install() -@dataclasses.dataclass class _FakeCaptureGroup(_FlagGroupBase): gamma: int = 0 @@ -1639,13 +1641,12 @@ class TestTheDerivedHalfIsDeclared(CustomTestCase): def test_a_declared_quotient_is_not_a_record_field(self): """It has no operator input to preserve, and the record is what crosses a process boundary.""" - import dataclasses from sglang.srt.arg_groups.arg_utils import Derived from sglang.srt.arg_groups.fields.parallel import Parallel from sglang.srt.server_args import ServerArgs - fields = {f.name for f in dataclasses.fields(ServerArgs)} + fields = {f.name for f in msgspec.structs.fields(ServerArgs)} for name, value in vars(Parallel).items(): if isinstance(value, Derived): self.assertNotIn(name, fields) diff --git a/test/registered/unit/test_runtime_context_config_bags.py b/test/registered/unit/test_runtime_context_config_bags.py index ac0d39fd9..27a9ed635 100644 --- a/test/registered/unit/test_runtime_context_config_bags.py +++ b/test/registered/unit/test_runtime_context_config_bags.py @@ -9,6 +9,9 @@ import dataclasses import unittest from unittest import mock +import msgspec +import msgspec.structs + from sglang.srt import runtime_context as rc from sglang.srt.arg_groups.arg_utils import NS, A from sglang.srt.arg_groups.overrides import resolution_result @@ -83,7 +86,7 @@ class TestConfigBags(CustomTestCase): import dataclasses sa, reference = self._resolve_published_and_sibling() - defaults = {f.name: f.default for f in dataclasses.fields(ServerArgs)} + defaults = {f.name: f.default for f in msgspec.structs.fields(ServerArgs)} # Leaves resolution writes on this input on both CI device shapes # (CUDA host and CPU-only runner): each starts at a None default. sampled = ( diff --git a/test/registered/unit/test_runtime_context_override.py b/test/registered/unit/test_runtime_context_override.py index 44e0e0a70..25185adf9 100644 --- a/test/registered/unit/test_runtime_context_override.py +++ b/test/registered/unit/test_runtime_context_override.py @@ -8,6 +8,8 @@ field aborts before any write; provenance is recorded. import unittest +import msgspec + from sglang.srt import runtime_context as rc from sglang.srt.server_args import ServerArgs from sglang.test.ci.ci_register import register_cpu_ci @@ -110,7 +112,7 @@ class TestContextOverride(CustomTestCase): # server_args is read-only after resolution: resolved config changes go # to the bags, a per-runner config to a derived variant. sa = ServerArgs(model_path="dummy") - object.__setattr__(sa, "_resolution_finished", True) + msgspec.Struct.__setattr__(sa, "_resolution_finished", True) with self.assertRaises(AttributeError): sa.page_size = 999 diff --git a/test/registered/unit/test_server_args_cli_metadata.py b/test/registered/unit/test_server_args_cli_metadata.py index 8708c245a..f89b9e068 100644 --- a/test/registered/unit/test_server_args_cli_metadata.py +++ b/test/registered/unit/test_server_args_cli_metadata.py @@ -3,7 +3,7 @@ import argparse import unittest -from sglang.srt.server_args import ServerArgs +from sglang.srt.server_args import ServerArgs, _declared_default from sglang.srt.utils.common import human_readable_int from sglang.test.ci.ci_register import register_cpu_ci from sglang.test.test_utils import CustomTestCase @@ -23,7 +23,9 @@ class TestServerArgsMigratedCliMetadata(CustomTestCase): } def test_argparse_shape_is_preserved_for_representative_migrated_options(self): - self.assertEqual(self.actions_by_option["--dtype"].default, ServerArgs.dtype) + self.assertEqual( + self.actions_by_option["--dtype"].default, _declared_default("dtype") + ) self.assertEqual( self.actions_by_option["--dtype"].choices, ["auto", "half", "float16", "bfloat16", "float", "float32"], diff --git a/test/registered/unit/test_server_args_namespaces.py b/test/registered/unit/test_server_args_namespaces.py index dac4df82f..6767ec8ef 100644 --- a/test/registered/unit/test_server_args_namespaces.py +++ b/test/registered/unit/test_server_args_namespaces.py @@ -11,9 +11,11 @@ This is the guardrail that fails when an upstream PR adds a field to a namespace class that has no ``_NS_PATH``, or adds one outside the taxonomy below. """ -import dataclasses import unittest +import msgspec +import msgspec.structs + from sglang.srt.arg_groups.arg_utils import namespace_of from sglang.srt.server_args import ServerArgs from sglang.test.ci.ci_register import register_cpu_ci @@ -48,7 +50,7 @@ VALID_NAMESPACES = { def _field_names(): - return {f.name for f in dataclasses.fields(ServerArgs)} + return {f.name for f in msgspec.structs.fields(ServerArgs)} class TestServerArgsNamespaces(CustomTestCase): diff --git a/test/registered/unit/test_split_attention_backend_decisions.py b/test/registered/unit/test_split_attention_backend_decisions.py index 164910c2c..5d73b2dad 100644 --- a/test/registered/unit/test_split_attention_backend_decisions.py +++ b/test/registered/unit/test_split_attention_backend_decisions.py @@ -153,13 +153,13 @@ class TestSplitBackendsReachTheDecisions(CustomTestCase): def test_the_flashinfer_version_guard_sees_a_split_launch(self): # The launcher runs before any publish, so it asks the record; the # member and the accessor answer the same pair. - args = ServerArgs.__new__(ServerArgs) + args = ServerArgs(model_path="dummy") for name, value in ( ("attention_backend", None), ("prefill_attention_backend", None), ("decode_attention_backend", "flashinfer"), ): - object.__setattr__(args, name, value) + setattr(args, name, value) self.assertIn("flashinfer", attention_backends_of(resolved_view(args))) def test_support_triton_is_the_regression_being_guarded(self): diff --git a/test/registered/unit/test_supplied_instance_exposure_ratchet.py b/test/registered/unit/test_supplied_instance_exposure_ratchet.py index a2d9219ab..5ad79c344 100644 --- a/test/registered/unit/test_supplied_instance_exposure_ratchet.py +++ b/test/registered/unit/test_supplied_instance_exposure_ratchet.py @@ -43,7 +43,6 @@ process-wide at all, so neither the read nor the field is on this axis. """ import ast -import dataclasses import json import os import shutil @@ -51,6 +50,9 @@ import tempfile import unittest from pathlib import Path +import msgspec +import msgspec.structs + import sglang from sglang.srt.server_args import ServerArgs from sglang.test.ci.ci_register import register_cpu_ci, register_cuda_ci @@ -367,10 +369,10 @@ class TestSuppliedInstanceExposure(CustomTestCase): "would drift" ) defaults = {} - for field in dataclasses.fields(resolved): - if field.default is not dataclasses.MISSING: + for field in msgspec.structs.fields(resolved): + if field.default is not msgspec.NODEFAULT: defaults[field.name] = field.default - elif field.default_factory is not dataclasses.MISSING: + elif field.default_factory is not msgspec.NODEFAULT: defaults[field.name] = field.default_factory() for field_name, default in defaults.items(): if field_name in _PASSED or field_name in extra: