Add UT guarding per-request bookkeeping clock ownership (#27710)
This commit is contained in:
@@ -0,0 +1,252 @@
|
||||
"""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 mixin / 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 unittest
|
||||
import warnings
|
||||
from collections import Counter
|
||||
from pathlib import Path
|
||||
|
||||
from sglang.test.ci.ci_register import register_cpu_ci
|
||||
from sglang.test.test_utils import CustomTestCase
|
||||
|
||||
register_cpu_ci(est_time=8, suite="base-a-test-cpu")
|
||||
|
||||
_REPO_ROOT = Path(__file__).resolve().parents[4]
|
||||
_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"
|
||||
_MIXIN = ("speculative/eagle_info_v2.py", "EagleDraftInputV2Mixin.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_decode", "kv_allocated_len"): 1,
|
||||
(_SB, "ScheduleBatch.prepare_for_extend", "extend_batch_idx"): 1,
|
||||
(_SB, "ScheduleBatch.prepare_for_extend", "kv_committed_len"): 1,
|
||||
(_SB, "ScheduleBatch.prepare_for_extend", "kv_allocated_len"): 1,
|
||||
("mem_cache/common.py", "alloc_for_extend", "evict"): 1,
|
||||
("mem_cache/common.py", "alloc_for_decode", "evict"): 1,
|
||||
# spec v2: pre-claim in the scheduler-driven mixin, settle in resolve
|
||||
(*_MIXIN, "decode_batch_idx"): 1,
|
||||
(*_MIXIN, "evict"): 1,
|
||||
(*_MIXIN, "kv_committed_len"): 1,
|
||||
(*_MIXIN, "kv_allocated_len"): 1,
|
||||
(*_RESOLVE, "kv_committed_len"): 2,
|
||||
(*_RESOLVE, "spec_verify_ct"): 1,
|
||||
# spec v1: each verify path owns its own settlement
|
||||
("speculative/eagle_info.py", "EagleVerifyInput.verify", "kv_committed_len"): 1,
|
||||
("speculative/eagle_info.py", "EagleVerifyInput.verify", "kv_allocated_len"): 1,
|
||||
("speculative/eagle_info.py", "EagleVerifyInput.verify", "spec_verify_ct"): 1,
|
||||
(
|
||||
"speculative/ngram_info.py",
|
||||
"NgramVerifyInput._fill_requests",
|
||||
"spec_verify_ct",
|
||||
): 1,
|
||||
(
|
||||
"speculative/ngram_info.py",
|
||||
"NgramVerifyInput._free_cache",
|
||||
"kv_committed_len",
|
||||
): 1,
|
||||
(
|
||||
"speculative/ngram_info.py",
|
||||
"NgramVerifyInput._free_cache",
|
||||
"kv_allocated_len",
|
||||
): 1,
|
||||
("speculative/dflash_info.py", "DFlashVerifyInput.verify", "kv_committed_len"): 1,
|
||||
("speculative/dflash_info.py", "DFlashVerifyInput.verify", "kv_allocated_len"): 1,
|
||||
("speculative/dflash_info.py", "DFlashVerifyInput.verify", "spec_verify_ct"): 1,
|
||||
# disaggregation decode prealloc
|
||||
(
|
||||
"disaggregation/decode.py",
|
||||
"DecodePreallocQueue._pre_alloc",
|
||||
"kv_committed_len",
|
||||
): 1,
|
||||
(
|
||||
"disaggregation/decode.py",
|
||||
"DecodePreallocQueue._pre_alloc",
|
||||
"kv_allocated_len",
|
||||
): 1,
|
||||
# streaming session slot save/restore and tail trimming
|
||||
(_SS, "SessionSlot.save_from_req", "kv_committed_len"): 1,
|
||||
(_SS, "SessionSlot.save_from_req", "kv_allocated_len"): 1,
|
||||
(_SS, "SessionSlot.restore_to_req", "kv_committed_len"): 1,
|
||||
(_SS, "SessionSlot.restore_to_req", "kv_allocated_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,
|
||||
}
|
||||
|
||||
|
||||
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 BaseDraftWorker 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 = {"BaseDraftWorker"}
|
||||
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 != "BaseDraftWorker"
|
||||
]
|
||||
|
||||
|
||||
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
|
||||
|
||||
|
||||
class TestDecodeBookkeepingOwnership(CustomTestCase):
|
||||
def test_bookkeeping_sites_match_owner_allowlist(self):
|
||||
found = _scan_srt()
|
||||
allow = Counter(_OWNER_SITES)
|
||||
unexpected = found - allow
|
||||
missing = allow - found
|
||||
msg = []
|
||||
if unexpected:
|
||||
msg.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:
|
||||
msg.append(
|
||||
"Recorded site(s) no longer exist (update _OWNER_SITES):\n "
|
||||
+ "\n ".join(f"{site} x{n}" for site, n in sorted(missing.items()))
|
||||
)
|
||||
self.assertFalse(msg, "\n\n".join(msg))
|
||||
|
||||
def test_spec_v2_draft_workers_do_no_scheduler_bookkeeping(self):
|
||||
classes = _draft_worker_classes()
|
||||
names = {node.name for _, node in classes}
|
||||
# Discovery sanity: fail loudly instead of silently guarding nothing.
|
||||
self.assertIn("EagleDraftWorker", names)
|
||||
self.assertIn("FrozenKVMTPDraftWorker", names)
|
||||
|
||||
violations = []
|
||||
for rel, node in classes:
|
||||
for scope, kind in _scan_class_subtree(node):
|
||||
violations.append((rel, f"{node.name}.{scope}", kind))
|
||||
self.assertFalse(
|
||||
violations,
|
||||
"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 "
|
||||
"mixin / resolve path. Remove these from the worker.",
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main(verbosity=3)
|
||||
Reference in New Issue
Block a user