[CI] Prune redundant CPU test overhead (#34309)
Co-authored-by: hnyls2002 <lsyincs@gmail.com>
This commit is contained in:
co-authored by
hnyls2002
parent
6ad3f2d8fd
commit
85cdf1178d
@@ -1,39 +0,0 @@
|
||||
#!/usr/bin/env python3
|
||||
"""Check that the rust-ext cache_key_prefix defaults stay in sync.
|
||||
|
||||
The build workflow saves cache entries under its default; the download action
|
||||
restores with its own. Neither file can reference the other, and a mismatch
|
||||
makes every pool silently fall back to source builds at install time.
|
||||
"""
|
||||
|
||||
import sys
|
||||
|
||||
import yaml
|
||||
|
||||
BUILD_WORKFLOW = ".github/workflows/_pr-test-rust-ext-build.yml"
|
||||
DOWNLOAD_ACTION = ".github/actions/download-rust-ext/action.yml"
|
||||
|
||||
|
||||
def main() -> int:
|
||||
with open(BUILD_WORKFLOW, encoding="utf-8") as f:
|
||||
workflow = yaml.safe_load(f)
|
||||
with open(DOWNLOAD_ACTION, encoding="utf-8") as f:
|
||||
action = yaml.safe_load(f)
|
||||
|
||||
# yaml 1.1 parses the `on:` key as boolean True
|
||||
triggers = workflow.get("on", workflow.get(True))
|
||||
save_prefix = triggers["workflow_call"]["inputs"]["cache_key_prefix"]["default"]
|
||||
restore_prefix = action["inputs"]["cache_key_prefix"]["default"]
|
||||
|
||||
if save_prefix != restore_prefix:
|
||||
print("ERROR: rust-ext cache_key_prefix defaults do not match.")
|
||||
print(f" {BUILD_WORKFLOW} saves under: {save_prefix}")
|
||||
print(f" {DOWNLOAD_ACTION} restores with: {restore_prefix}")
|
||||
print("Bump both together, or every pool falls back to source builds.")
|
||||
return 1
|
||||
|
||||
return 0
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
sys.exit(main())
|
||||
@@ -104,12 +104,6 @@ _FILE_EXTENSIONS = (
|
||||
".onnx",
|
||||
)
|
||||
|
||||
# Non-test helper files under test/registered/ (skipped by basename, matching
|
||||
# scripts/ci/check_registered_tests.py). run_suite.py skips `cpu/utils.py` by
|
||||
# path; excluding every `utils.py` by basename is a superset that drops no
|
||||
# CUDA-registered test (the other `utils.py` registers CPU only).
|
||||
_NON_TEST_BASENAMES = frozenset({"conftest.py", "__init__.py", "utils.py"})
|
||||
|
||||
|
||||
def looks_like_model_id(value: str, deny: Optional[Set[str]] = None) -> bool:
|
||||
"""Heuristic: does ``value`` look like a HuggingFace repo id?
|
||||
@@ -273,11 +267,13 @@ def collect_suite_files(
|
||||
ci_register = _load_ci_register(repo_root)
|
||||
backend = getattr(ci_register.HWBackend, backend_name.upper())
|
||||
|
||||
pattern = os.path.join(repo_root, "test", "registered", "**", "*.py")
|
||||
# Same exclusion as run_suite.py: pytest+package structure files.
|
||||
files = sorted(
|
||||
f
|
||||
for f in glob.glob(pattern, recursive=True)
|
||||
if os.path.basename(f) not in _NON_TEST_BASENAMES
|
||||
for f in glob.glob(
|
||||
os.path.join(repo_root, "test", "registered", "**", "*.py"), recursive=True
|
||||
)
|
||||
if os.path.basename(f) not in ("conftest.py", "__init__.py")
|
||||
)
|
||||
|
||||
suite_files: Dict[str, List[str]] = {}
|
||||
|
||||
@@ -88,7 +88,8 @@ def discover_files(repo_root: str) -> list[str]:
|
||||
for f in glob.glob(
|
||||
os.path.join(test_dir, "registered", "**", "*.py"), recursive=True
|
||||
)
|
||||
if not f.endswith("/conftest.py") and not f.endswith("/__init__.py")
|
||||
# Same exclusion as run_suite.py: pytest+package structure files.
|
||||
if os.path.basename(f) not in ("conftest.py", "__init__.py")
|
||||
]
|
||||
jit_kernel_dir = os.path.join(repo_root, "python", "sglang", "jit_kernel")
|
||||
files += glob.glob(
|
||||
|
||||
@@ -0,0 +1,235 @@
|
||||
"""Ownership contract for per-request bookkeeping clocks.
|
||||
|
||||
Per-request accounting state (`decode_batch_idx` / `extend_batch_idx` iter
|
||||
clocks, `kv_committed_len` / `kv_allocated_len` KV watermarks,
|
||||
`spec_verify_ct`, and the `maybe_evict_swa()` call) must only be advanced by
|
||||
the reviewed owner sites in _OWNER_SITES; spec-v2 draft workers must not
|
||||
repeat any of them (the scheduler-driven free function / resolve path already
|
||||
does).
|
||||
A clock that runs fast fires SWA eviction in the overlap race window and
|
||||
releases the SWA prefix lock early; neither shows up in e2e CI or the idle
|
||||
leak checker, hence this AST-level guard.
|
||||
"""
|
||||
|
||||
import ast
|
||||
import warnings
|
||||
from collections import Counter
|
||||
from pathlib import Path
|
||||
|
||||
_REPO_ROOT = Path(__file__).resolve().parents[2]
|
||||
_SRT_DIR = _REPO_ROOT / "python" / "sglang" / "srt"
|
||||
_SPECULATIVE_DIR = _SRT_DIR / "speculative"
|
||||
assert _SRT_DIR.is_dir(), f"srt dir not found: {_SRT_DIR}"
|
||||
|
||||
_TRACKED_ATTRS = (
|
||||
"decode_batch_idx",
|
||||
"extend_batch_idx",
|
||||
"kv_committed_len",
|
||||
"kv_allocated_len",
|
||||
"spec_verify_ct",
|
||||
)
|
||||
_EVICT_METHOD = "maybe_evict_swa"
|
||||
|
||||
# {(path relative to srt/, scope, kind): mutation count}. Kind is the mutated
|
||||
# attribute (`= 0` resets exempt) or "evict" for a `maybe_evict_swa()` call.
|
||||
# Any added/removed/recounted site fails until reviewed here.
|
||||
_SB = "managers/schedule_batch.py"
|
||||
_EAGLE_DECODE = ("speculative/eagle_utils.py", "eagle_prepare_for_decode")
|
||||
_RESOLVE = (
|
||||
"managers/scheduler_components/batch_result_processor.py",
|
||||
"SchedulerBatchResultProcessor._resolve_spec_v2_tokens",
|
||||
)
|
||||
_SS = "session/streaming_session.py"
|
||||
_OWNER_SITES = {
|
||||
# non-spec scheduler
|
||||
(_SB, "ScheduleBatch.prepare_for_decode", "decode_batch_idx"): 1,
|
||||
(_SB, "ScheduleBatch.prepare_for_decode", "kv_committed_len"): 1,
|
||||
(_SB, "ScheduleBatch.prepare_for_extend", "extend_batch_idx"): 1,
|
||||
(_SB, "ScheduleBatch.prepare_for_extend", "kv_committed_len"): 1,
|
||||
# kv_allocated_len is settled inside the owned-kv alloc functions (op28).
|
||||
("mem_cache/allocation.py", "alloc_for_extend", "evict"): 1,
|
||||
("mem_cache/allocation.py", "alloc_for_extend", "kv_allocated_len"): 1,
|
||||
("mem_cache/allocation.py", "alloc_for_decode", "evict"): 1,
|
||||
("mem_cache/allocation.py", "alloc_for_decode", "kv_allocated_len"): 1,
|
||||
# spec v2: no pre-claim; resolve commits the full accepted run uniformly.
|
||||
# kv_allocated_len for spec v2 draft decode (eagle + dflash) is settled
|
||||
# inside the owned-kv alloc_for_spec_decode function (op42).
|
||||
(*_EAGLE_DECODE, "decode_batch_idx"): 1,
|
||||
(*_EAGLE_DECODE, "evict"): 1,
|
||||
(
|
||||
"mem_cache/allocation.py",
|
||||
"alloc_for_spec_decode",
|
||||
"kv_allocated_len",
|
||||
): 1,
|
||||
(*_RESOLVE, "kv_committed_len"): 1,
|
||||
(*_RESOLVE, "spec_verify_ct"): 1,
|
||||
# disaggregation decode prealloc: kv_allocated_len is settled inside the
|
||||
# owned-kv alloc_for_decode_prealloc(_hisparse) functions (op42).
|
||||
(
|
||||
"disaggregation/decode.py",
|
||||
"DecodePreallocQueue._pre_alloc",
|
||||
"kv_committed_len",
|
||||
): 1,
|
||||
("disaggregation/decode.py", "alloc_for_decode_prealloc", "kv_allocated_len"): 1,
|
||||
(
|
||||
"disaggregation/decode.py",
|
||||
"alloc_for_decode_prealloc_hisparse",
|
||||
"kv_allocated_len",
|
||||
): 1,
|
||||
# streaming session slot save/restore and tail trimming
|
||||
(_SS, "SessionSlot.save_from_req", "kv_committed_len"): 1,
|
||||
(_SS, "SessionSlot.restore_to_req", "kv_committed_len"): 1,
|
||||
(_SS, "StreamingSession._free_tail", "kv_committed_len"): 2,
|
||||
(_SS, "StreamingSession._free_tail", "kv_allocated_len"): 2,
|
||||
(_SS, "StreamingSession._trim_overshoot", "kv_committed_len"): 1,
|
||||
(_SS, "StreamingSession._trim_overshoot", "kv_allocated_len"): 1,
|
||||
(_SS, "StreamingSession.try_cache_finished_req", "kv_allocated_len"): 1,
|
||||
# Inherit the authoritative finished length (not the lagging req clock).
|
||||
(_SS, "StreamingSession.try_cache_finished_req", "kv_committed_len"): 1,
|
||||
}
|
||||
|
||||
|
||||
def _iter_scoped_nodes(tree):
|
||||
"""Yield (node, dotted Class.method scope) for every node."""
|
||||
scope_of = {}
|
||||
|
||||
def visit(node, scope):
|
||||
if isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef, ast.ClassDef)):
|
||||
scope = f"{scope}.{node.name}" if scope else node.name
|
||||
scope_of[node] = scope
|
||||
for child in ast.iter_child_nodes(node):
|
||||
visit(child, scope)
|
||||
|
||||
visit(tree, "")
|
||||
return scope_of.items()
|
||||
|
||||
|
||||
def _is_zero_reset(node):
|
||||
return isinstance(node, ast.Assign) and (
|
||||
isinstance(node.value, ast.Constant) and node.value.value == 0
|
||||
)
|
||||
|
||||
|
||||
def _scan_tree(tree):
|
||||
"""Count bookkeeping sites in an AST as Counter[(scope, kind)]."""
|
||||
sites = Counter()
|
||||
for node, scope in _iter_scoped_nodes(tree):
|
||||
if isinstance(node, (ast.AugAssign, ast.Assign)):
|
||||
targets = node.targets if isinstance(node, ast.Assign) else [node.target]
|
||||
for target in targets:
|
||||
if (
|
||||
isinstance(target, ast.Attribute)
|
||||
and target.attr in _TRACKED_ATTRS
|
||||
and not _is_zero_reset(node)
|
||||
):
|
||||
sites[(scope, target.attr)] += 1
|
||||
if (
|
||||
isinstance(node, ast.Call)
|
||||
and isinstance(node.func, ast.Attribute)
|
||||
and node.func.attr == _EVICT_METHOD
|
||||
):
|
||||
sites[(scope, "evict")] += 1
|
||||
return sites
|
||||
|
||||
|
||||
def _parse(path: Path):
|
||||
# utf-8-sig: some srt files carry a BOM that breaks plain-utf-8 ast.parse.
|
||||
with warnings.catch_warnings():
|
||||
warnings.simplefilter("ignore", SyntaxWarning)
|
||||
return ast.parse(path.read_text(encoding="utf-8-sig"))
|
||||
|
||||
|
||||
def _scan_srt():
|
||||
"""Count all bookkeeping sites in srt/ as Counter[(rel, scope, kind)]."""
|
||||
found = Counter()
|
||||
for path in sorted(_SRT_DIR.rglob("*.py")):
|
||||
rel = path.relative_to(_SRT_DIR).as_posix()
|
||||
for (scope, kind), count in _scan_tree(_parse(path)).items():
|
||||
found[(rel, scope, kind)] += count
|
||||
return found
|
||||
|
||||
|
||||
def _draft_worker_classes():
|
||||
"""All transitive EagleDraftWorkerBase subclasses under speculative/."""
|
||||
by_name = {}
|
||||
for path in sorted(_SPECULATIVE_DIR.glob("*.py")):
|
||||
rel = path.relative_to(_SRT_DIR).as_posix()
|
||||
for node in ast.walk(_parse(path)):
|
||||
if isinstance(node, ast.ClassDef):
|
||||
bases = {
|
||||
b.id if isinstance(b, ast.Name) else getattr(b, "attr", None)
|
||||
for b in node.bases
|
||||
}
|
||||
by_name[node.name] = (rel, node, bases)
|
||||
|
||||
workers = {"EagleDraftWorkerBase"}
|
||||
changed = True
|
||||
while changed:
|
||||
changed = False
|
||||
for name, (_, _, bases) in by_name.items():
|
||||
if name not in workers and bases & workers:
|
||||
workers.add(name)
|
||||
changed = True
|
||||
return [
|
||||
(rel, node)
|
||||
for name, (rel, node, _) in sorted(by_name.items())
|
||||
if name in workers and name != "EagleDraftWorkerBase"
|
||||
]
|
||||
|
||||
|
||||
def _scan_class_subtree(class_node):
|
||||
"""Scan one ClassDef subtree; returns (method_scope, kind) sites."""
|
||||
module = ast.Module(body=[class_node], type_ignores=[])
|
||||
sites = set()
|
||||
for scope, kind in _scan_tree(module):
|
||||
# Strip the leading class name; keep method-level scope.
|
||||
sites.add((scope.split(".", 1)[1] if "." in scope else scope, kind))
|
||||
return sites
|
||||
|
||||
|
||||
def check_bookkeeping_sites_match_owner_allowlist():
|
||||
found = _scan_srt()
|
||||
allow = Counter(_OWNER_SITES)
|
||||
unexpected = found - allow
|
||||
missing = allow - found
|
||||
messages = []
|
||||
if unexpected:
|
||||
messages.append(
|
||||
"New bookkeeping mutation(s) beyond the recorded counts:\n "
|
||||
+ "\n ".join(f"{site} x{n}" for site, n in sorted(unexpected.items()))
|
||||
+ "\nThese are owned by the sites in _OWNER_SITES -- do not "
|
||||
"repeat them; a genuinely new owner must be recorded there."
|
||||
)
|
||||
if missing:
|
||||
messages.append(
|
||||
"Recorded site(s) no longer exist (update _OWNER_SITES):\n "
|
||||
+ "\n ".join(f"{site} x{n}" for site, n in sorted(missing.items()))
|
||||
)
|
||||
if messages:
|
||||
raise AssertionError("\n\n".join(messages))
|
||||
|
||||
|
||||
def check_spec_v2_draft_workers_do_no_scheduler_bookkeeping():
|
||||
classes = _draft_worker_classes()
|
||||
names = {node.name for _, node in classes}
|
||||
for expected in ("EagleDraftWorker", "FrozenKVMTPDraftWorker"):
|
||||
if expected not in names:
|
||||
raise AssertionError(f"draft worker discovery missed {expected}")
|
||||
|
||||
violations = []
|
||||
for rel, node in classes:
|
||||
for scope, kind in _scan_class_subtree(node):
|
||||
violations.append((rel, f"{node.name}.{scope}", kind))
|
||||
if violations:
|
||||
raise AssertionError(
|
||||
"Spec-v2 draft worker(s) repeat scheduler-owned bookkeeping:\n "
|
||||
+ "\n ".join(map(str, sorted(violations)))
|
||||
+ "\nUnder spec v2 the iter-clock ticks, `maybe_evict_swa`, and "
|
||||
"KV watermark settlement are owned by the scheduler-driven "
|
||||
"free function / resolve path. Remove these from the worker.",
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
check_bookkeeping_sites_match_owner_allowlist()
|
||||
check_spec_v2_draft_workers_do_no_scheduler_bookkeeping()
|
||||
@@ -0,0 +1,471 @@
|
||||
"""Guard: business code never reads a config field off the process-global record.
|
||||
|
||||
``get_server_args()`` returns the published ``ServerArgs`` -- one process's
|
||||
startup record. Config decisions read the namespace accessors instead
|
||||
(``get_exec()`` / ``get_memory()`` / ...); per-runner values come from the
|
||||
runner that owns them. Both baselines are zero, over the whole package minus
|
||||
the modules that own the slot.
|
||||
|
||||
The scanners match ``get_server_args`` and ``configured_*_size`` by their
|
||||
literal names, which is why import-renaming them is banned below. A name
|
||||
computed at runtime, or indirection deeper than a local name copy, is invisible
|
||||
here -- the census tool in the context repo audits that shape.
|
||||
"""
|
||||
|
||||
import ast
|
||||
from functools import cache
|
||||
from pathlib import Path
|
||||
|
||||
# srt is the migrated surface; the rest of the package has no reads today and is
|
||||
# scanned so a new one cannot appear there unnoticed.
|
||||
_PACKAGE_ROOT = Path(__file__).resolve().parents[2] / "python" / "sglang"
|
||||
|
||||
# The modules that own the slot: runtime_context publishes it and exposes the
|
||||
# named accessors for the derived members, server_args/arg_groups ARE the
|
||||
# resolution pipeline.
|
||||
_SLOT_OWNERS = ("srt/runtime_context.py", "srt/server_args.py", "srt/arg_groups/")
|
||||
|
||||
# Every call site of a ``configured_*_size()`` accessor, with the reason the
|
||||
# live topology cannot answer there. The checker below asserts this map is exactly
|
||||
# the set of call sites, so the reasons cannot drift away from the code.
|
||||
_CONFIGURED_SIZE_CALL_SITES = {
|
||||
("srt/layers/attention/dsa/dsa_indexer.py", "configured_pp_size"): (
|
||||
"gates `pp_size > 1 and not get_pp_group()...`; the short circuit is the "
|
||||
"point, since with PP off the group is never touched, which is what lets "
|
||||
"the Indexer be constructed before distributed init"
|
||||
),
|
||||
("srt/mem_cache/kv_cache_configurator.py", "configured_pp_size"): (
|
||||
"decides whether the token capacity needs a cross-PP all-reduce at all; "
|
||||
"asking the configured size keeps that decision independent of whether a "
|
||||
"PP group is installed in this process"
|
||||
),
|
||||
("srt/layers/dp_attention.py", "configured_attn_cp_size"): (
|
||||
"compared against the configured moe_dp_size below"
|
||||
),
|
||||
("srt/layers/dp_attention.py", "configured_moe_dp_size"): (
|
||||
"the configuration this predicate detects (attn_cp_size > moe_dp_size) is "
|
||||
"the one where initialize_model_parallel aliases _MOE_DP to _ATTN_CP, so "
|
||||
"the live sizes are equal there and a live comparison is always false"
|
||||
),
|
||||
("srt/model_loader/loader.py", "configured_moe_dp_size"): (
|
||||
"the same dict already carries the live moe_dp_size under 'dp'; this entry "
|
||||
"is the configured intent"
|
||||
),
|
||||
("srt/models/kimi_k25.py", "configured_tp_size"): (
|
||||
"the IPC refcount must match the configured TP consumer count captured "
|
||||
"when the tokenizer creates MmItemMemoryPool; a live attention subgroup "
|
||||
"size could strand leases in the bounded pool"
|
||||
),
|
||||
("srt/models/kimi_k3.py", "configured_tp_size"): (
|
||||
"same as kimi_k25: the IPC refcount must agree with the recycler's waiter"
|
||||
),
|
||||
}
|
||||
|
||||
# A dynamic read whose name is set nowhere in the tree, so the predicate it
|
||||
# feeds is inert (the ``getattr`` default decides it). Converting it would mean
|
||||
# choosing what it should have named, which is the CP path's call, not this
|
||||
# sweep's -- so it is listed here rather than silently counted or "fixed".
|
||||
_INERT_DYNAMIC_READS = frozenset({("srt/layers/cp/base.py", "_is_dsa_model_arch")})
|
||||
|
||||
_DIRECT_BASELINE = 0
|
||||
_ALIAS_BASELINE = 0
|
||||
|
||||
|
||||
def _is_global_call(node) -> bool:
|
||||
"""``get_server_args()`` however it is spelled: bare, or module-qualified
|
||||
(``ctx.get_server_args()``), which an ast.Name check alone would miss."""
|
||||
if not isinstance(node, ast.Call):
|
||||
return False
|
||||
func = node.func
|
||||
if isinstance(func, ast.Name):
|
||||
return func.id == "get_server_args"
|
||||
return isinstance(func, ast.Attribute) and func.attr == "get_server_args"
|
||||
|
||||
|
||||
def _collect(rel: str, tree: ast.AST, inert: frozenset = frozenset()):
|
||||
"""The (direct, alias) field reads in one module.
|
||||
|
||||
``inert`` names the fields listed in ``_INERT_DYNAMIC_READS`` for this file;
|
||||
they are dropped here, at the point the read is recognized, so the filter
|
||||
matches on the field name rather than on the rendered message.
|
||||
"""
|
||||
direct, alias = [], []
|
||||
|
||||
def counted(attr: str) -> bool:
|
||||
return attr not in inert
|
||||
|
||||
def _getattr_name(node):
|
||||
"""``getattr(<record>, "field")`` names a field just as ``.field`` does;
|
||||
matching only ast.Attribute would let a dynamic read walk past."""
|
||||
if not (
|
||||
isinstance(node, ast.Call)
|
||||
and isinstance(node.func, ast.Name)
|
||||
and node.func.id == "getattr"
|
||||
and len(node.args) >= 2
|
||||
and isinstance(node.args[1], ast.Constant)
|
||||
and isinstance(node.args[1].value, str)
|
||||
):
|
||||
return None
|
||||
return node.args[1].value
|
||||
|
||||
for node in ast.walk(tree):
|
||||
if (
|
||||
isinstance(node, ast.Attribute)
|
||||
and _is_global_call(node.value)
|
||||
and counted(node.attr)
|
||||
):
|
||||
direct.append(f"{rel}:{node.lineno}: get_server_args().{node.attr}")
|
||||
|
||||
name = _getattr_name(node)
|
||||
if name is not None and _is_global_call(node.args[0]) and counted(name):
|
||||
direct.append(f"{rel}:{node.lineno}: getattr(get_server_args(), {name!r})")
|
||||
|
||||
if not isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef)):
|
||||
continue
|
||||
params = {a.arg for a in list(node.args.args) + list(node.args.kwonlyargs)}
|
||||
bound = {}
|
||||
for inner in ast.walk(node):
|
||||
# ``sa = get_server_args()`` and its annotated form
|
||||
# ``sa: ServerArgs = get_server_args()``.
|
||||
if isinstance(inner, (ast.Assign, ast.AnnAssign)) and _is_global_call(
|
||||
getattr(inner, "value", None)
|
||||
):
|
||||
targets = (
|
||||
inner.targets if isinstance(inner, ast.Assign) else [inner.target]
|
||||
)
|
||||
for target in targets:
|
||||
if not isinstance(target, ast.Name):
|
||||
continue
|
||||
# A parameter reassigned from the global is the
|
||||
# optional-injection shape (``f(server_args=None)`` then
|
||||
# ``server_args = get_server_args()``): the reads that
|
||||
# follow are global reads wearing a parameter's name, so
|
||||
# they count from the bind on.
|
||||
bound.setdefault(target.id, inner.lineno)
|
||||
if not bound:
|
||||
continue
|
||||
# A copy of an alias reaches the same record (``cfg = sa`` after
|
||||
# ``sa = get_server_args()``), so follow Name-to-Name assignments to a
|
||||
# fixpoint. Deeper indirection (through containers, attributes of
|
||||
# other objects, cross-scope copies) stays census-tool territory.
|
||||
changed = True
|
||||
while changed:
|
||||
changed = False
|
||||
for inner in ast.walk(node):
|
||||
if not isinstance(inner, (ast.Assign, ast.AnnAssign)):
|
||||
continue
|
||||
value = getattr(inner, "value", None)
|
||||
if not (isinstance(value, ast.Name) and value.id in bound):
|
||||
continue
|
||||
targets = (
|
||||
inner.targets if isinstance(inner, ast.Assign) else [inner.target]
|
||||
)
|
||||
for target in targets:
|
||||
if isinstance(target, ast.Name) and target.id not in bound:
|
||||
bound[target.id] = inner.lineno
|
||||
changed = True
|
||||
for inner in ast.walk(node):
|
||||
if (
|
||||
isinstance(inner, ast.Attribute)
|
||||
and isinstance(inner.value, ast.Name)
|
||||
and inner.value.id in bound
|
||||
and inner.lineno >= bound[inner.value.id]
|
||||
and counted(inner.attr)
|
||||
):
|
||||
alias.append(
|
||||
f"{rel}:{inner.lineno}: {inner.value.id}.{inner.attr} "
|
||||
f"(bound from get_server_args() at line {bound[inner.value.id]})"
|
||||
)
|
||||
name = _getattr_name(inner)
|
||||
if (
|
||||
name is not None
|
||||
and isinstance(inner.args[0], ast.Name)
|
||||
and inner.args[0].id in bound
|
||||
and inner.lineno >= bound[inner.args[0].id]
|
||||
and counted(name)
|
||||
):
|
||||
alias.append(
|
||||
f"{rel}:{inner.lineno}: getattr({inner.args[0].id}, {name!r}) "
|
||||
f"(bound from get_server_args() at line {bound[inner.args[0].id]})"
|
||||
)
|
||||
# A module-level alias is visible to every function in the file, so it needs
|
||||
# its own pass -- the per-function scan above deliberately does not reach
|
||||
# across scopes.
|
||||
module_bound = {}
|
||||
module_stack = list(tree.body)
|
||||
while module_stack:
|
||||
stmt = module_stack.pop()
|
||||
# A module-level bind can sit inside an `if` / `try` / `with`, so the
|
||||
# walk descends into those bodies -- but not into a nested function or
|
||||
# class, whose binds are that scope's own.
|
||||
if isinstance(
|
||||
stmt, (ast.FunctionDef, ast.AsyncFunctionDef, ast.ClassDef, ast.Lambda)
|
||||
):
|
||||
continue
|
||||
module_stack.extend(ast.iter_child_nodes(stmt))
|
||||
if isinstance(stmt, (ast.Assign, ast.AnnAssign)) and _is_global_call(
|
||||
getattr(stmt, "value", None)
|
||||
):
|
||||
targets = stmt.targets if isinstance(stmt, ast.Assign) else [stmt.target]
|
||||
for target in targets:
|
||||
if isinstance(target, ast.Name):
|
||||
module_bound.setdefault(target.id, stmt.lineno)
|
||||
if module_bound:
|
||||
# Shadowing is per lexical scope: a function with its own `sa` hides the
|
||||
# module alias *inside that function only*. Aggregating the names
|
||||
# file-wide would suppress every read in the module, including the
|
||||
# top-level ones and the ones in functions that do resolve to the alias.
|
||||
parents = {}
|
||||
scope_binds = {}
|
||||
stack = [tree]
|
||||
while stack:
|
||||
node = stack.pop()
|
||||
enclosing = parents.get(id(node))
|
||||
for child in ast.iter_child_nodes(node):
|
||||
parents[id(child)] = (
|
||||
node
|
||||
if isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef))
|
||||
else enclosing
|
||||
)
|
||||
stack.append(child)
|
||||
if isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef)):
|
||||
names = {
|
||||
a.arg for a in list(node.args.args) + list(node.args.kwonlyargs)
|
||||
}
|
||||
# Only this scope's own stores: a nested function's local `sa`
|
||||
# shadows the alias inside *that* function, not in its parent.
|
||||
pending = list(node.body)
|
||||
while pending:
|
||||
inner = pending.pop()
|
||||
if isinstance(
|
||||
inner,
|
||||
(
|
||||
ast.FunctionDef,
|
||||
ast.AsyncFunctionDef,
|
||||
ast.Lambda,
|
||||
ast.ClassDef,
|
||||
),
|
||||
):
|
||||
continue
|
||||
if isinstance(inner, ast.Name) and isinstance(inner.ctx, ast.Store):
|
||||
names.add(inner.id)
|
||||
pending.extend(ast.iter_child_nodes(inner))
|
||||
scope_binds[id(node)] = names
|
||||
|
||||
def _shadowed(node, name):
|
||||
scope = parents.get(id(node))
|
||||
while scope is not None:
|
||||
if name in scope_binds.get(id(scope), ()):
|
||||
return True
|
||||
scope = parents.get(id(scope))
|
||||
return False
|
||||
|
||||
for node in ast.walk(tree):
|
||||
base = attr = None
|
||||
if (
|
||||
isinstance(node, ast.Attribute)
|
||||
and isinstance(node.value, ast.Name)
|
||||
and node.value.id in module_bound
|
||||
):
|
||||
base, attr = node.value.id, node.attr
|
||||
shown = f"{base}.{attr}"
|
||||
else:
|
||||
attr_name = _getattr_name(node)
|
||||
if (
|
||||
attr_name is not None
|
||||
and isinstance(node.args[0], ast.Name)
|
||||
and node.args[0].id in module_bound
|
||||
):
|
||||
base, attr = node.args[0].id, attr_name
|
||||
shown = f"getattr({base}, {attr!r})"
|
||||
if base and not _shadowed(node, base) and counted(attr):
|
||||
alias.append(
|
||||
f"{rel}:{node.lineno}: {shown} "
|
||||
f"(module-level bind from get_server_args() at line "
|
||||
f"{module_bound[base]})"
|
||||
)
|
||||
|
||||
# An alias parked on an instance attribute (``self._sa = get_server_args()``
|
||||
# in one method, ``self._sa.field`` in another) reaches the same slot and
|
||||
# crosses function scopes, so it is collected per class rather than per
|
||||
# function.
|
||||
for node in ast.walk(tree):
|
||||
if not isinstance(node, ast.ClassDef):
|
||||
continue
|
||||
attr_bound = {}
|
||||
for inner in ast.walk(node):
|
||||
if isinstance(inner, (ast.Assign, ast.AnnAssign)) and _is_global_call(
|
||||
getattr(inner, "value", None)
|
||||
):
|
||||
targets = (
|
||||
inner.targets if isinstance(inner, ast.Assign) else [inner.target]
|
||||
)
|
||||
for target in targets:
|
||||
if (
|
||||
isinstance(target, ast.Attribute)
|
||||
and isinstance(target.value, ast.Name)
|
||||
and target.value.id in ("self", "cls")
|
||||
):
|
||||
attr_bound.setdefault(
|
||||
(target.value.id, target.attr), inner.lineno
|
||||
)
|
||||
if not attr_bound:
|
||||
continue
|
||||
|
||||
def _bound_attr(value):
|
||||
"""``self._sa`` when that attribute was bound from the global."""
|
||||
if (
|
||||
isinstance(value, ast.Attribute)
|
||||
and isinstance(value.value, ast.Name)
|
||||
and (value.value.id, value.attr) in attr_bound
|
||||
):
|
||||
return (value.value.id, value.attr)
|
||||
return None
|
||||
|
||||
for inner in ast.walk(node):
|
||||
key = shown = None
|
||||
if isinstance(inner, ast.Attribute):
|
||||
key = _bound_attr(inner.value)
|
||||
if key is not None and counted(inner.attr):
|
||||
shown = f"{key[0]}.{key[1]}.{inner.attr}"
|
||||
else:
|
||||
name = _getattr_name(inner)
|
||||
if name is not None:
|
||||
key = _bound_attr(inner.args[0])
|
||||
if key is not None and counted(name):
|
||||
shown = f"getattr({key[0]}.{key[1]}, {name!r})"
|
||||
if shown is not None:
|
||||
alias.append(
|
||||
f"{rel}:{inner.lineno}: {shown} "
|
||||
f"(attribute bind from get_server_args() at line "
|
||||
f"{attr_bound[key]})"
|
||||
)
|
||||
return direct, alias
|
||||
|
||||
|
||||
@cache
|
||||
def _parsed_modules():
|
||||
"""(rel, tree) per parseable module; the three scanners below share it."""
|
||||
modules = []
|
||||
for path in sorted(_PACKAGE_ROOT.rglob("*.py")):
|
||||
try:
|
||||
tree = ast.parse(path.read_text())
|
||||
except SyntaxError:
|
||||
continue
|
||||
modules.append((path.relative_to(_PACKAGE_ROOT).as_posix(), tree))
|
||||
return modules
|
||||
|
||||
|
||||
def _field_reads():
|
||||
direct, alias = [], []
|
||||
for rel, tree in _parsed_modules():
|
||||
if rel.startswith(_SLOT_OWNERS):
|
||||
continue
|
||||
inert = frozenset(name for path_, name in _INERT_DYNAMIC_READS if path_ == rel)
|
||||
module_direct, module_alias = _collect(rel, tree, inert)
|
||||
direct += module_direct
|
||||
alias += module_alias
|
||||
return direct, alias
|
||||
|
||||
|
||||
def _configured_size_call_sites():
|
||||
found = set()
|
||||
for rel, tree in _parsed_modules():
|
||||
if rel.startswith(_SLOT_OWNERS):
|
||||
continue
|
||||
for node in ast.walk(tree):
|
||||
if not isinstance(node, ast.Call):
|
||||
continue
|
||||
func = node.func
|
||||
name = (
|
||||
func.id
|
||||
if isinstance(func, ast.Name)
|
||||
else (func.attr if isinstance(func, ast.Attribute) else None)
|
||||
)
|
||||
if name and name.startswith("configured_") and name.endswith("_size"):
|
||||
found.add((rel, name))
|
||||
return found
|
||||
|
||||
|
||||
def _renamed_accessor_imports():
|
||||
offenders = []
|
||||
for rel, tree in _parsed_modules():
|
||||
for node in ast.walk(tree):
|
||||
if not isinstance(node, (ast.ImportFrom, ast.Import)):
|
||||
continue
|
||||
for imported in node.names:
|
||||
if imported.asname is None or imported.asname == imported.name:
|
||||
continue
|
||||
base = imported.name.rsplit(".", 1)[-1]
|
||||
if base == "get_server_args" or (
|
||||
base.startswith("configured_") and base.endswith("_size")
|
||||
):
|
||||
offenders.append(
|
||||
f"{rel}:{node.lineno}: {imported.name} as {imported.asname}"
|
||||
)
|
||||
return offenders
|
||||
|
||||
|
||||
def _check_count(kind, reads, baseline):
|
||||
if len(reads) > baseline:
|
||||
raise AssertionError(
|
||||
f"{kind} process-global config field reads grew: {len(reads)} > "
|
||||
f"baseline {baseline}. Read the namespace accessor for the "
|
||||
"field's namespace, or the owning runner for a per-runner "
|
||||
"field:\n" + "\n".join(reads)
|
||||
)
|
||||
|
||||
|
||||
def check_global_config_read_ratchet():
|
||||
direct, alias = _field_reads()
|
||||
_check_count("direct", direct, _DIRECT_BASELINE)
|
||||
_check_count("alias-form", alias, _ALIAS_BASELINE)
|
||||
|
||||
|
||||
def check_configured_size_call_sites():
|
||||
"""The configured-vs-live exceptions are enumerated, with reasons.
|
||||
|
||||
``configured_*_size()`` answers what the user asked for where
|
||||
``get_parallel()`` would answer what the process ended up with. Each such
|
||||
exception is listed above with why the live property cannot serve it, and
|
||||
this case fails if the code and that list disagree.
|
||||
|
||||
The unit is **(file, accessor)**, not the individual call: a second
|
||||
`configured_pp_size()` in a file already registered for it collapses into
|
||||
the same entry, so the reason has to cover the file's use of that accessor
|
||||
rather than one line. A new file, or a new accessor in a listed file, is
|
||||
what this catches -- in either call form (bare or module-qualified).
|
||||
"""
|
||||
|
||||
found = _configured_size_call_sites()
|
||||
documented = set(_CONFIGURED_SIZE_CALL_SITES)
|
||||
if documented != found:
|
||||
raise AssertionError(
|
||||
"configured-size call sites drifted from their documented reasons.\n"
|
||||
f" undocumented: {sorted(found - documented)}\n"
|
||||
f" stale entries: {sorted(documented - found)}",
|
||||
)
|
||||
|
||||
|
||||
def check_no_renamed_accessor_imports():
|
||||
"""The scanners above match ``get_server_args`` and ``configured_*_size``
|
||||
by their literal names, so an ``import ... as`` rename would walk a read
|
||||
straight past both the zero baseline and the call-site registry. Renaming
|
||||
these accessors buys nothing (the names are already short and unambiguous),
|
||||
so it is banned outright — which is exactly what makes literal-name
|
||||
matching sound."""
|
||||
|
||||
offenders = _renamed_accessor_imports()
|
||||
if offenders:
|
||||
raise AssertionError(
|
||||
"get_server_args / configured_*_size imported under another name; "
|
||||
"the read ratchet and the configured-size registry match these "
|
||||
"accessors by their literal names, so a rename silently escapes "
|
||||
"both:\n" + "\n".join(offenders),
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
check_global_config_read_ratchet()
|
||||
check_configured_size_call_sites()
|
||||
check_no_renamed_accessor_imports()
|
||||
@@ -0,0 +1,52 @@
|
||||
"""Ratchet guard: legacy global-accessor call-sites may only decrease.
|
||||
|
||||
The process-wide ``ServerArgs`` is owned by the runtime context; the legacy
|
||||
``get_global_server_args`` / ``set_global_server_args_for_*`` names survive as
|
||||
thin shims for the existing call-sites. New code should use the
|
||||
``sglang.srt.runtime_context`` accessors (``get_server_args()`` /
|
||||
``get_context().set_server_args()``), so the shim call-site counts below must
|
||||
never grow. When your change removes call-sites, lower the matching baseline
|
||||
to the new count.
|
||||
"""
|
||||
|
||||
import re
|
||||
from pathlib import Path
|
||||
|
||||
_SRT_ROOT = Path(__file__).resolve().parents[2] / "python" / "sglang" / "srt"
|
||||
|
||||
# Baselines counted over python/sglang/srt/**/*.py, including each function's
|
||||
# own def line. Ratchet: decrease-only.
|
||||
_RATCHETS = [
|
||||
# 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*\(",
|
||||
4,
|
||||
),
|
||||
]
|
||||
|
||||
|
||||
def check_legacy_global_ratchet():
|
||||
sources = [
|
||||
path.read_text(encoding="utf-8", errors="replace")
|
||||
for path in sorted(_SRT_ROOT.rglob("*.py"))
|
||||
]
|
||||
for name, pattern, baseline in _RATCHETS:
|
||||
count = sum(len(re.findall(pattern, source)) for source in sources)
|
||||
if count > baseline:
|
||||
raise AssertionError(
|
||||
f"{name} call-sites grew: {count} > baseline {baseline}. "
|
||||
"New code must use the sglang.srt.runtime_context accessors "
|
||||
"(get_server_args() / get_context().set_server_args())."
|
||||
)
|
||||
if count < baseline:
|
||||
raise AssertionError(
|
||||
f"{name} call-sites shrank: {count} < baseline {baseline}. "
|
||||
"Lower the baseline in this file to lock in the progress."
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
check_legacy_global_ratchet()
|
||||
@@ -0,0 +1,52 @@
|
||||
"""Ratchet guard: module-level runtime state in the flag-owning layers may
|
||||
only shrink.
|
||||
|
||||
Runtime flags belong on ``get_flags()`` groups, which have lifecycle reset and
|
||||
a scoped test-override primitive; a module-level ``global`` has neither and
|
||||
leaks across test teardowns. The pin below names the survivors -- migrating one
|
||||
must shrink it.
|
||||
"""
|
||||
|
||||
import ast
|
||||
from pathlib import Path
|
||||
|
||||
_SRT_ROOT = Path(__file__).resolve().parents[2] / "python" / "sglang" / "srt"
|
||||
|
||||
_PINNED_GLOBALS = {
|
||||
"layers/moe/utils.py": frozenset(),
|
||||
"layers/dp_attention.py": frozenset(
|
||||
{
|
||||
# DP-attention topology (parallel vertical scope).
|
||||
"_ATTN_DP_RANK",
|
||||
"_ATTN_DP_SIZE",
|
||||
}
|
||||
),
|
||||
}
|
||||
|
||||
|
||||
def check_module_state_ratchet():
|
||||
for rel, pinned in _PINNED_GLOBALS.items():
|
||||
tree = ast.parse((_SRT_ROOT / rel).read_text())
|
||||
declared = {
|
||||
name
|
||||
for node in ast.walk(tree)
|
||||
if isinstance(node, ast.Global)
|
||||
for name in node.names
|
||||
}
|
||||
grown = declared - pinned
|
||||
if grown:
|
||||
raise AssertionError(
|
||||
f"{rel} declares new module-level runtime state {sorted(grown)}; "
|
||||
"put runtime flags on a get_flags() group instead "
|
||||
"(see runtime_context.MoeFlags / DpFlags).",
|
||||
)
|
||||
shrunk = pinned - declared
|
||||
if shrunk:
|
||||
raise AssertionError(
|
||||
f"{rel} no longer declares {sorted(shrunk)}; "
|
||||
"shrink the pin in this file to lock in the progress.",
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
check_module_state_ratchet()
|
||||
Executable
+154
@@ -0,0 +1,154 @@
|
||||
#!/usr/bin/env python3
|
||||
|
||||
import ast
|
||||
import pathlib
|
||||
import re
|
||||
import sys
|
||||
|
||||
_PYTEST_MAIN = re.compile(r"pytest\s*\.\s*main")
|
||||
|
||||
|
||||
def is_main_guard(node: ast.expr) -> bool:
|
||||
if not isinstance(node, ast.Compare) or len(node.ops) != 1:
|
||||
return False
|
||||
if not isinstance(node.ops[0], ast.Eq):
|
||||
return False
|
||||
sides = [node.left, *node.comparators]
|
||||
has_name = any(
|
||||
isinstance(side, ast.Name) and side.id == "__name__" for side in sides
|
||||
)
|
||||
has_main = any(
|
||||
isinstance(side, ast.Constant) and side.value == "__main__" for side in sides
|
||||
)
|
||||
return has_name and has_main
|
||||
|
||||
|
||||
def is_pytest_main_call(node: ast.AST) -> bool:
|
||||
if not isinstance(node, ast.Call):
|
||||
return False
|
||||
func = node.func
|
||||
return (
|
||||
isinstance(func, ast.Attribute)
|
||||
and func.attr == "main"
|
||||
and isinstance(func.value, ast.Name)
|
||||
and func.value.id == "pytest"
|
||||
)
|
||||
|
||||
|
||||
def is_exit_call(node: ast.AST, parents: dict[int, ast.AST]) -> bool:
|
||||
"""``sys.exit(...)``, or a ``SystemExit(...)`` that is actually raised."""
|
||||
if not isinstance(node, ast.Call):
|
||||
return False
|
||||
func = node.func
|
||||
if (
|
||||
isinstance(func, ast.Attribute)
|
||||
and func.attr == "exit"
|
||||
and isinstance(func.value, ast.Name)
|
||||
and func.value.id == "sys"
|
||||
):
|
||||
return True
|
||||
parent = parents.get(id(node))
|
||||
return (
|
||||
isinstance(func, ast.Name)
|
||||
and func.id == "SystemExit"
|
||||
and isinstance(parent, ast.Raise)
|
||||
and parent.exc is node
|
||||
)
|
||||
|
||||
|
||||
def assigned_names(node: ast.AST) -> list[str]:
|
||||
if isinstance(node, ast.Assign):
|
||||
return [t.id for t in node.targets if isinstance(t, ast.Name)]
|
||||
if isinstance(node, (ast.AnnAssign, ast.NamedExpr)):
|
||||
return [node.target.id] if isinstance(node.target, ast.Name) else []
|
||||
return []
|
||||
|
||||
|
||||
def exited_names(nodes: list[ast.AST], parents: dict[int, ast.AST]) -> set[str]:
|
||||
"""Names handed to an exit call, so the two-step form still propagates."""
|
||||
return {
|
||||
arg.id
|
||||
for node in nodes
|
||||
if is_exit_call(node, parents)
|
||||
for arg in node.args
|
||||
if isinstance(arg, ast.Name)
|
||||
}
|
||||
|
||||
|
||||
def propagates_exit_code(
|
||||
node: ast.Call, parents: dict[int, ast.AST], exited: set[str]
|
||||
) -> bool:
|
||||
parent = parents.get(id(node))
|
||||
if (
|
||||
isinstance(parent, ast.Call)
|
||||
and node in parent.args
|
||||
and is_exit_call(parent, parents)
|
||||
):
|
||||
return True
|
||||
return any(name in exited for name in assigned_names(parent))
|
||||
|
||||
|
||||
def runtime_nodes(node: ast.AST):
|
||||
yield node
|
||||
if isinstance(
|
||||
node, (ast.FunctionDef, ast.AsyncFunctionDef, ast.ClassDef, ast.Lambda)
|
||||
):
|
||||
return
|
||||
for child in ast.iter_child_nodes(node):
|
||||
yield from runtime_nodes(child)
|
||||
|
||||
|
||||
def find_bare_pytest_main(path: pathlib.Path) -> int | None:
|
||||
try:
|
||||
source = path.read_text(encoding="utf-8")
|
||||
except (OSError, UnicodeDecodeError):
|
||||
return None
|
||||
if "__main__" not in source or _PYTEST_MAIN.search(source) is None:
|
||||
return None
|
||||
try:
|
||||
tree = ast.parse(source, filename=str(path))
|
||||
except SyntaxError:
|
||||
return None
|
||||
|
||||
for node in ast.walk(tree):
|
||||
if not isinstance(node, ast.If) or not is_main_guard(node.test):
|
||||
continue
|
||||
# Whole body at once: the call and the sys.exit() that propagates it
|
||||
# are separate statements.
|
||||
nodes = [n for statement in node.body for n in runtime_nodes(statement)]
|
||||
parents = {
|
||||
id(child): parent
|
||||
for parent in nodes
|
||||
for child in ast.iter_child_nodes(parent)
|
||||
}
|
||||
exited = exited_names(nodes, parents)
|
||||
for candidate in nodes:
|
||||
if is_pytest_main_call(candidate) and not propagates_exit_code(
|
||||
candidate, parents, exited
|
||||
):
|
||||
return candidate.lineno
|
||||
return None
|
||||
|
||||
|
||||
def main(paths: list[str]) -> int:
|
||||
offenders = []
|
||||
for path_string in paths:
|
||||
path = pathlib.Path(path_string)
|
||||
line = find_bare_pytest_main(path)
|
||||
if line is not None:
|
||||
offenders.append(f"{path}:{line}")
|
||||
|
||||
if not offenders:
|
||||
return 0
|
||||
|
||||
print(
|
||||
"ERROR: pytest.main(...) in an __main__ block must propagate its exit "
|
||||
"code with sys.exit(...) or raise SystemExit(...):"
|
||||
)
|
||||
for offender in offenders:
|
||||
print(f" {offender}")
|
||||
return 1
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
sys.exit(main(sys.argv[1:]))
|
||||
@@ -0,0 +1,63 @@
|
||||
"""Guard: no legacy parallel-getter calls in the swept directories.
|
||||
|
||||
``models/`` and ``layers/`` read parallel topology through
|
||||
``get_parallel().<dim>`` (the read-through wrapper in ``runtime_context``),
|
||||
which gives one import, one naming scheme, and the scoped ``override()`` test
|
||||
primitive. Exemptions are pinned in ``_EXEMPT``, each with its reason; sweeping
|
||||
one must remove it from there.
|
||||
"""
|
||||
|
||||
import re
|
||||
from pathlib import Path
|
||||
|
||||
_SRT_ROOT = Path(__file__).resolve().parents[2] / "python" / "sglang" / "srt"
|
||||
|
||||
_BANNED_CALLS = re.compile(
|
||||
r"\b(?:dcp_enabled|get_(?:"
|
||||
r"tensor_model_parallel_(?:world_size|rank)"
|
||||
r"|pipeline_model_parallel_(?:world_size|rank)"
|
||||
r"|moe_expert_parallel_(?:world_size|rank)"
|
||||
r"|moe_tensor_parallel_(?:world_size|rank)"
|
||||
r"|moe_data_parallel_(?:world_size|rank)"
|
||||
r"|attn_tensor_model_parallel_(?:world_size|rank)"
|
||||
r"|attn_context_model_parallel_(?:world_size|rank)"
|
||||
r"|dcp_(?:world_size|rank)"
|
||||
r"|dcp_group(?:_no_assert)?"
|
||||
r"|attention_dcp_(?:world_size|rank)"
|
||||
r"|attention_(?:tp|cp)_(?:group|rank|size)"
|
||||
r"))\(\)"
|
||||
)
|
||||
|
||||
# The whole package is swept; the exemptions are the substrate itself.
|
||||
_SWEPT_DIRS = ("",)
|
||||
|
||||
_EXEMPT = (
|
||||
"distributed/", # parallel_state: defines the canonical getters
|
||||
"runtime_context.py", # delegates DCP reads to canonical getters
|
||||
"layers/dp_attention.py", # delegation substrate for the attn-DP dims
|
||||
"layers/dcp/comm.py", # deprecated out-of-tree DCP compatibility shims
|
||||
# The dumper's megatron plugin calls third-party getters that share the
|
||||
# parallel_state names (self._mpu.get_tensor_model_parallel_rank()).
|
||||
"debug_utils/dumper.py",
|
||||
)
|
||||
|
||||
|
||||
def check_parallel_adoption_ratchet():
|
||||
offenders = []
|
||||
for top in _SWEPT_DIRS:
|
||||
for path in sorted((_SRT_ROOT / top).rglob("*.py")):
|
||||
rel = path.relative_to(_SRT_ROOT).as_posix()
|
||||
if rel.startswith(_EXEMPT):
|
||||
continue
|
||||
for line_number, line in enumerate(path.read_text().split("\n"), 1):
|
||||
if _BANNED_CALLS.search(line):
|
||||
offenders.append(f"{rel}:{line_number}")
|
||||
if offenders:
|
||||
raise AssertionError(
|
||||
"legacy parallel-getter calls in swept directories (use "
|
||||
f"get_parallel().<dim> instead): {offenders}",
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
check_parallel_adoption_ratchet()
|
||||
@@ -80,11 +80,11 @@ def main() -> int:
|
||||
spec.loader.exec_module(ci_register)
|
||||
cuda = ci_register.HWBackend.CUDA
|
||||
|
||||
# Same filter as run_suite.py: skip conftest.py, __init__.py, and utils.py
|
||||
# Same exclusion as run_suite.py: pytest+package structure files.
|
||||
files = sorted(
|
||||
f
|
||||
for f in glob.glob("test/registered/**/*.py", recursive=True)
|
||||
if os.path.basename(f) not in ("conftest.py", "__init__.py", "utils.py")
|
||||
if os.path.basename(f) not in ("conftest.py", "__init__.py")
|
||||
)
|
||||
if not files:
|
||||
return 0
|
||||
Executable
+66
@@ -0,0 +1,66 @@
|
||||
#!/usr/bin/env python3
|
||||
"""Check that the rust-ext cache key stays in sync across its sites.
|
||||
|
||||
The build workflow looks up and saves cache entries under its own prefix and
|
||||
hashed inputs; the download action restores with its own. Neither file can
|
||||
reference the other, and a mismatch in EITHER half makes every pool silently
|
||||
fall back to source builds at install time.
|
||||
"""
|
||||
|
||||
import re
|
||||
import sys
|
||||
|
||||
import yaml
|
||||
|
||||
BUILD_WORKFLOW = ".github/workflows/_pr-test-rust-ext-build.yml"
|
||||
DOWNLOAD_ACTION = ".github/actions/download-rust-ext/action.yml"
|
||||
|
||||
_HASH_FILES = re.compile(r"hashFiles\(([^)]*)\)")
|
||||
_QUOTED = re.compile(r"'([^']*)'")
|
||||
|
||||
|
||||
def hashed_inputs(path: str) -> list[tuple[str, ...]]:
|
||||
"""The argument tuple of every ``hashFiles(...)`` cache key in a file."""
|
||||
with open(path, encoding="utf-8") as f:
|
||||
text = f.read()
|
||||
return [tuple(_QUOTED.findall(args)) for args in _HASH_FILES.findall(text)]
|
||||
|
||||
|
||||
def main() -> int:
|
||||
with open(BUILD_WORKFLOW, encoding="utf-8") as f:
|
||||
workflow = yaml.safe_load(f)
|
||||
with open(DOWNLOAD_ACTION, encoding="utf-8") as f:
|
||||
action = yaml.safe_load(f)
|
||||
|
||||
# yaml 1.1 parses the `on:` key as boolean True
|
||||
triggers = workflow.get("on", workflow.get(True))
|
||||
save_prefix = triggers["workflow_call"]["inputs"]["cache_key_prefix"]["default"]
|
||||
restore_prefix = action["inputs"]["cache_key_prefix"]["default"]
|
||||
|
||||
if save_prefix != restore_prefix:
|
||||
print("ERROR: rust-ext cache_key_prefix defaults do not match.")
|
||||
print(f" {BUILD_WORKFLOW} saves under: {save_prefix}")
|
||||
print(f" {DOWNLOAD_ACTION} restores with: {restore_prefix}")
|
||||
print("Bump both together, or every pool falls back to source builds.")
|
||||
return 1
|
||||
|
||||
# Adding a file to one key alone permanently misses the other's entries.
|
||||
sites = [(BUILD_WORKFLOW, inputs) for inputs in hashed_inputs(BUILD_WORKFLOW)]
|
||||
sites += [(DOWNLOAD_ACTION, inputs) for inputs in hashed_inputs(DOWNLOAD_ACTION)]
|
||||
|
||||
if not sites:
|
||||
print("ERROR: no hashFiles(...) cache key found; this check is dead.")
|
||||
return 1
|
||||
|
||||
if len({inputs for _, inputs in sites}) > 1:
|
||||
print("ERROR: rust-ext cache key inputs do not match.")
|
||||
for path, inputs in sites:
|
||||
print(f" {path}: {list(inputs)}")
|
||||
print("Every lookup/save/restore site must hash the same inputs.")
|
||||
return 1
|
||||
|
||||
return 0
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
sys.exit(main())
|
||||
@@ -0,0 +1,57 @@
|
||||
"""Guard: no server_args mutation outside the resolution pipeline, pinned at 0.
|
||||
|
||||
``ServerArgs.__setattr__`` already raises on a bare assignment after
|
||||
resolution; this static scan is what reaches the sites tests never execute.
|
||||
"""
|
||||
|
||||
import re
|
||||
from pathlib import Path
|
||||
|
||||
_SGLANG_ROOT = Path(__file__).resolve().parents[2] / "python" / "sglang"
|
||||
|
||||
# Assignments to a server_args attribute (``server_args.x = ...``,
|
||||
# ``self.server_args.x = ...``, and the ``sa`` alias used by a few helpers).
|
||||
# ``==`` comparisons are excluded by the negative lookahead.
|
||||
_MUTATION_PATTERNS = [
|
||||
# (?![=}]) skips ``==`` comparisons and f-string ``{x=}`` debug specs.
|
||||
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) and multimodal_gen,
|
||||
# whose ServerArgs is a different class outside this contract.
|
||||
_EXCLUDED = (
|
||||
"srt/server_args.py",
|
||||
"srt/arg_groups",
|
||||
"multimodal_gen",
|
||||
)
|
||||
|
||||
_BASELINE = 0
|
||||
|
||||
|
||||
def check_server_args_mutation_ratchet():
|
||||
count = 0
|
||||
for path in sorted(_SGLANG_ROOT.rglob("*.py")):
|
||||
rel = path.relative_to(_SGLANG_ROOT).as_posix()
|
||||
if rel.startswith(_EXCLUDED):
|
||||
continue
|
||||
source = path.read_text()
|
||||
count += sum(len(pattern.findall(source)) for pattern in _MUTATION_PATTERNS)
|
||||
if count > _BASELINE:
|
||||
raise AssertionError(
|
||||
f"server_args mutations outside the resolution pipeline grew: "
|
||||
f"{count} > baseline {_BASELINE}. Configuration is resolved in "
|
||||
"ServerArgs.__post_init__; declare through the pipeline "
|
||||
"(passes / declare_late_resolution), change resolved config "
|
||||
"with get_context().override(source, ...), or hand the value "
|
||||
"to its runner as a constructor argument — do not assign fields."
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
check_server_args_mutation_ratchet()
|
||||
Executable
+48
@@ -0,0 +1,48 @@
|
||||
#!/usr/bin/env python3
|
||||
"""Runs every static ratchet in one process, so the package is parsed once."""
|
||||
|
||||
import sys
|
||||
|
||||
from check_decode_bookkeeping_ownership import (
|
||||
check_bookkeeping_sites_match_owner_allowlist,
|
||||
check_spec_v2_draft_workers_do_no_scheduler_bookkeeping,
|
||||
)
|
||||
from check_global_config_read_ratchet import (
|
||||
check_configured_size_call_sites,
|
||||
check_global_config_read_ratchet,
|
||||
check_no_renamed_accessor_imports,
|
||||
)
|
||||
from check_legacy_global_ratchet import check_legacy_global_ratchet
|
||||
from check_module_state_ratchet import check_module_state_ratchet
|
||||
from check_parallel_adoption_ratchet import check_parallel_adoption_ratchet
|
||||
from check_server_args_mutation_ratchet import check_server_args_mutation_ratchet
|
||||
|
||||
|
||||
def main():
|
||||
checks = (
|
||||
check_bookkeeping_sites_match_owner_allowlist,
|
||||
check_spec_v2_draft_workers_do_no_scheduler_bookkeeping,
|
||||
check_global_config_read_ratchet,
|
||||
check_configured_size_call_sites,
|
||||
check_no_renamed_accessor_imports,
|
||||
check_legacy_global_ratchet,
|
||||
check_module_state_ratchet,
|
||||
check_parallel_adoption_ratchet,
|
||||
check_server_args_mutation_ratchet,
|
||||
)
|
||||
# They guard independent invariants, so one failure must not hide the rest.
|
||||
failures = []
|
||||
for check in checks:
|
||||
try:
|
||||
check()
|
||||
except AssertionError as exc:
|
||||
failures.append(f"[{check.__name__}] {exc}")
|
||||
|
||||
for failure in failures:
|
||||
print(failure, file=sys.stderr)
|
||||
print(file=sys.stderr)
|
||||
return 1 if failures else 0
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
sys.exit(main())
|
||||
@@ -0,0 +1,86 @@
|
||||
import pathlib
|
||||
import tempfile
|
||||
import unittest
|
||||
|
||||
from check_no_bare_pytest_main import find_bare_pytest_main
|
||||
|
||||
|
||||
class TestFindBarePytestMain(unittest.TestCase):
|
||||
def check_source(self, source: str) -> int | None:
|
||||
with tempfile.TemporaryDirectory() as directory:
|
||||
path = pathlib.Path(directory) / "example.py"
|
||||
path.write_text(source, encoding="utf-8")
|
||||
return find_bare_pytest_main(path)
|
||||
|
||||
def test_rejects_discarded_result(self):
|
||||
source = """
|
||||
if __name__ == "__main__":
|
||||
pytest.main([__file__])
|
||||
"""
|
||||
self.assertEqual(self.check_source(source), 3)
|
||||
|
||||
def test_rejects_discarded_result_with_whitespace(self):
|
||||
source = """
|
||||
if __name__ == "__main__":
|
||||
pytest . main([__file__])
|
||||
"""
|
||||
self.assertEqual(self.check_source(source), 3)
|
||||
|
||||
def test_accepts_propagated_result(self):
|
||||
source = """
|
||||
if __name__ == "__main__":
|
||||
sys.exit(pytest.main([__file__]))
|
||||
"""
|
||||
self.assertIsNone(self.check_source(source))
|
||||
|
||||
def test_rejects_assigned_result(self):
|
||||
source = """
|
||||
if "__main__" == __name__:
|
||||
exit_code = pytest.main([__file__])
|
||||
"""
|
||||
self.assertEqual(self.check_source(source), 3)
|
||||
|
||||
def test_accepts_assigned_result_that_is_later_exited(self):
|
||||
source = """
|
||||
if __name__ == "__main__":
|
||||
exit_code = pytest.main([__file__])
|
||||
sys.exit(exit_code)
|
||||
"""
|
||||
self.assertIsNone(self.check_source(source))
|
||||
|
||||
def test_accepts_assigned_result_that_is_raised(self):
|
||||
source = """
|
||||
if __name__ == "__main__":
|
||||
exit_code = pytest.main([__file__])
|
||||
raise SystemExit(exit_code)
|
||||
"""
|
||||
self.assertIsNone(self.check_source(source))
|
||||
|
||||
def test_rejects_nested_discarded_result(self):
|
||||
source = """
|
||||
if __name__ == "__main__":
|
||||
if enabled:
|
||||
pytest.main([__file__])
|
||||
"""
|
||||
self.assertEqual(self.check_source(source), 4)
|
||||
|
||||
def test_accepts_raised_system_exit(self):
|
||||
source = """
|
||||
if __name__ == "__main__":
|
||||
raise SystemExit(pytest.main([__file__]))
|
||||
"""
|
||||
self.assertIsNone(self.check_source(source))
|
||||
|
||||
def test_rejects_unraised_system_exit(self):
|
||||
source = """
|
||||
if __name__ == "__main__":
|
||||
error = SystemExit(pytest.main([__file__]))
|
||||
"""
|
||||
self.assertEqual(self.check_source(source), 3)
|
||||
|
||||
def test_ignores_call_outside_main_guard(self):
|
||||
self.assertIsNone(self.check_source("pytest.main([__file__])\n"))
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
Reference in New Issue
Block a user