diff --git a/tests/compass/test_clock_lp_identity.py b/tests/compass/test_clock_lp_identity.py index 28e843adc1..fad037a5d1 100644 --- a/tests/compass/test_clock_lp_identity.py +++ b/tests/compass/test_clock_lp_identity.py @@ -402,6 +402,11 @@ def test_the_package_imports_only_the_standard_library_it_names(module): roots += [alias.name.split(".")[0] for alias in node.names] elif isinstance(node, ast.ImportFrom) and not node.level: roots.append((node.module or "").split(".")[0]) + elif isinstance(node, ast.ImportFrom) and not module.parents[ + node.level - 1 + ].is_relative_to(CLOCK_PACKAGE): + # A relative import that climbs out of this package, named as written. + roots.append("." * node.level + (node.module or "")) strays = sorted({root for root in roots if root not in allowed}) assert not strays, f"{module.name} imports {strays}; allowed: {sorted(allowed)}" @@ -415,16 +420,16 @@ def test_the_package_builds_no_set_at_all(module): # see -- and the package has no use for one -- so the rule it enforces is the # one it can prove: none is constructed, so none can be iterated. tree = ast.parse(module.read_text()) + # Every mention of either type, called or not: `frozenset.union` or an + # alias builds one without the call a narrower match would look for. + called = {id(n.func) for n in ast.walk(tree) if isinstance(n, ast.Call)} offenders = [] for node in ast.walk(tree): if isinstance(node, (ast.Set, ast.SetComp)): offenders.append(f"{type(node).__name__} at line {node.lineno}") - elif ( - isinstance(node, ast.Call) - and isinstance(node.func, ast.Name) - and node.func.id in ("set", "frozenset") - ): - offenders.append(f"{node.func.id}() at line {node.lineno}") + elif isinstance(node, ast.Name) and node.id in ("set", "frozenset"): + how = "()" if id(node) in called else " named" + offenders.append(f"{node.id}{how} at line {node.lineno}") assert not offenders, ( f"{module.name} builds {offenders}; use a dict with None values as an " "ordered set, or sort at the point of iteration" diff --git a/tests/compass/test_ir_data_model.py b/tests/compass/test_ir_data_model.py index d86f4107bf..7e136d118b 100644 --- a/tests/compass/test_ir_data_model.py +++ b/tests/compass/test_ir_data_model.py @@ -740,6 +740,11 @@ def test_the_package_imports_only_the_standard_library_it_names(module): roots += [alias.name.split(".")[0] for alias in node.names] elif isinstance(node, ast.ImportFrom) and not node.level: roots.append((node.module or "").split(".")[0]) + elif isinstance(node, ast.ImportFrom) and not module.parents[ + node.level - 1 + ].is_relative_to(IR_PACKAGE): + # A relative import that climbs out of this package, named as written. + roots.append("." * node.level + (node.module or "")) strays = sorted({root for root in roots if root not in allowed}) assert not strays, f"{module.name} imports {strays}; allowed: {sorted(allowed)}" diff --git a/tests/compass/test_kv_blob_site.py b/tests/compass/test_kv_blob_site.py index 3257766b82..34632ac65b 100644 --- a/tests/compass/test_kv_blob_site.py +++ b/tests/compass/test_kv_blob_site.py @@ -28,9 +28,11 @@ def _blob_site(path: Path) -> list[ast.Assign]: """Every `.kv_transfer_params_output = ...` in one module. - Deliberately narrow: plain `ast.Assign` to an `ast.Attribute`, which is the - form both connectors use. A `setattr`, an `AnnAssign` or a walrus would be - invisible to it, so it is narrower than the one-site claim it pins. + Narrow on purpose: plain `ast.Assign` to an `ast.Attribute`, which is the + form both connectors use. Every other write of the name is what + `_writes_not_read` returns, and the one-site test refuses those, so a + second site written another way fails there by name instead of passing + unseen. """ tree = ast.parse(path.read_text(encoding="utf-8")) return [ @@ -43,9 +45,39 @@ def _blob_site(path: Path) -> list[ast.Assign]: ] +def _writes_not_read(path: Path) -> list[str]: + """Every write of the name `_blob_site` does not read, by line and source. + + An annotated or augmented assignment, the name inside a tuple target, the + name as a string anywhere in the module -- and any mention of `setattr`, + `__setattr__`, `__dict__` or `vars` at all, since a name built at run time + cannot be read and neither connector writes attributes that way. + """ + source = path.read_text(encoding="utf-8") + tree = ast.parse(source) + read = { + id(t) for n in ast.walk(tree) if isinstance(n, ast.Assign) for t in n.targets + } + return [ + f"{path.name}:{n.lineno}: {source.splitlines()[n.lineno - 1].strip()}" + for n in ast.walk(tree) + if ( + isinstance(n, ast.Attribute) + and n.attr == BLOB_ATTR + and not isinstance(n.ctx, ast.Load) + and id(n) not in read + ) + or (isinstance(n, ast.Constant) and n.value == BLOB_ATTR) + or getattr(n, "id", getattr(n, "attr", None)) + in ("setattr", "__setattr__", "__dict__", "vars") + ] + + @pytest.mark.parametrize("backend", sorted(CONNECTORS)) def test_each_backend_builds_the_blob_at_exactly_one_site(backend): """A second site would mean the backend emits one of two shapes.""" + unread = _writes_not_read(CONNECTORS[backend]) + assert not unread, f"{BLOB_ATTR} is written in a form not read here: {unread}" sites = _blob_site(CONNECTORS[backend]) assert len(sites) == 1, ( f"{CONNECTORS[backend].name} assigns {BLOB_ATTR} at " diff --git a/tests/compass/test_runner_non_allocating.py b/tests/compass/test_runner_non_allocating.py index cb8660b6ab..00d5031c4f 100644 --- a/tests/compass/test_runner_non_allocating.py +++ b/tests/compass/test_runner_non_allocating.py @@ -35,6 +35,7 @@ import pytest import torch +from test_runner_rpc_surface import _methods from torch.utils._python_dispatch import TorchDispatchMode from atom.compass.runner import COMPASS_RUNNER_QUALNAME @@ -64,13 +65,11 @@ def _classes(path): tree = ast.parse(path.read_text()) + for n in ast.walk(tree): + n.file = path.name return {n.name: n for n in ast.walk(tree) if isinstance(n, ast.ClassDef)} -def _methods(node): - return {n.name for n in node.body if isinstance(n, ast.FunctionDef)} - - def _self_calls(node): """Names of `self.x(...)` calls anywhere inside a function definition.""" return { @@ -207,8 +206,16 @@ def _method_def(node, name): ) -def _self_assigned(node): - """Every `self.x = ...` in *node*, as (name, the source of its value). +# Every name through which an attribute can be written without an assignment. +_WRITERS = ("setattr", "__setattr__", "__dict__", "vars") + + +def _on_self(node): + return isinstance(node, ast.Attribute) and ast.unparse(node.value) == "self" + + +def _self_assigned(node, unreadable=()): + """Every binding of `self.x` in *node*, as (name, the source of its value). Pairs, not a mapping keyed by name. A name can be assigned more than once -- `forward_vars` is bound to the dict of buffers and later rebound to a slot @@ -216,16 +223,50 @@ def _self_assigned(node): walked, which here is the rebind. The rebind names no buffer, so keying by name dropped `forward_vars` out of the holder set entirely. Keeping the pairs is what lets a name count as a holder when *any* of its bindings is. + + Read in every spelling that states the name: a target of `=`, including + one inside a tuple or list, of an annotated or augmented `=`, and + `setattr(self, "x", ...)`. Any other write -- a `for` or `with` target, and + every other mention of `setattr`, `__setattr__`, `__dict__` or `vars`, + however it is reached -- binds something this cannot read. Each of those + must be listed in *unreadable* by its source text, or this refuses, so a + spelling it cannot read fails here rather than leaving the set smaller + than the class. """ - return { - (t.attr, ast.unparse(n.value)) + bound, read = set(), set() + for n in ast.walk(node): + if isinstance(n, (ast.Assign, ast.AnnAssign, ast.AugAssign)): + for target in n.targets if isinstance(n, ast.Assign) else [n.target]: + for t in ast.walk(target): + if _on_self(t) and isinstance(t.ctx, ast.Store) and n.value: + bound.add((t.attr, ast.unparse(n.value))) + read.add(id(t)) + elif ( + isinstance(n, ast.Call) + and ast.unparse(n.func) == "setattr" + and [ast.unparse(a) for a in n.args[:1]] == ["self"] + and len(n.args) == 3 + and isinstance(n.args[1], ast.Constant) + and isinstance(n.args[1].value, str) + ): + bound.add((n.args[1].value, ast.unparse(n.args[2]))) + read.add(id(n.func)) + calls = {id(n.func): n for n in ast.walk(node) if isinstance(n, ast.Call)} + unread = [ + (n.lineno, ast.unparse(calls.get(id(n), n))) for n in ast.walk(node) - if isinstance(n, ast.Assign) - for t in n.targets - if isinstance(t, ast.Attribute) - and isinstance(t.value, ast.Name) - and t.value.id == "self" - } + if id(n) not in read + and ( + (_on_self(n) and isinstance(n.ctx, ast.Store)) + or getattr(n, "id", getattr(n, "attr", None)) in _WRITERS + ) + ] + assert sorted(text for _, text in unread) == sorted(unreadable), ( + f"{node.name} binds attributes on self in a form not read here: " + + "; ".join(f"{node.file}:{line}: {text}" for line, text in sorted(unread)) + + f" -- listed as expected: {sorted(unreadable)}" + ) + return bound def test_exactly_two_attributes_of_the_runner_hold_the_ring(): @@ -244,7 +285,12 @@ def test_exactly_two_attributes_of_the_runner_hold_the_ring(): buffer, one attribute deeper -- true, and outside a claim about attributes on the runner. """ - assigned = _self_assigned(_classes(ATOM_RUNNER)["ModelRunner"]) + # The three `setattr`s bind whatever the attention builders return -- + # the KV cache and the per-request state, by names the source never states. + assigned = _self_assigned( + _classes(ATOM_RUNNER)["ModelRunner"], + unreadable=["setattr(self, name, value)"] * 3, + ) holders = {n for n, v in assigned if any(t in v for t in BUFFER_TERMS)} assert holders == {"forward_vars", "_fv_ring"} @@ -488,7 +534,17 @@ def test_only_the_binding_module_reaches_the_engine(path): The exemption is exactly the two packages whose closure something asserts; widening it to `atom.compass` would exempt packages nothing has checked. """ - imported = _import_time_imports(path.read_text()) + # A relative import is resolved against this module's package first, so + # one that climbs out of the package is read as the module it names. + package = path.relative_to(REPO).parent.parts + imported = { + ( + ".".join([*package[: len(package) + 1 - level], name[level:]]).strip(".") + if (level := len(name) - len(name.lstrip("."))) + else name + ) + for name in _import_time_imports(path.read_text()) + } engine = {m for m in imported if m.split(".")[0] == "atom"} - { m for m in imported diff --git a/tests/compass/test_runner_rpc_surface.py b/tests/compass/test_runner_rpc_surface.py index ef73bb45bd..14d661fde6 100644 --- a/tests/compass/test_runner_rpc_surface.py +++ b/tests/compass/test_runner_rpc_surface.py @@ -68,16 +68,51 @@ class Site(NamedTuple): line: int waits: bool aggregated: bool - arity: int # values unpacked from the reply; 0 when it is discarded + arity: int | None # values unpacked; 0 when discarded, None when unread def _classes(path): tree = ast.parse(path.read_text()) + for n in ast.walk(tree): + n.file = path.name return {n.name: n for n in ast.walk(tree) if isinstance(n, ast.ClassDef)} +UNREAD_CLASS_BODY: list[str] = [] + + +def _is_docstring(n): + return ( + isinstance(n, ast.Expr) + and isinstance(n.value, ast.Constant) + and isinstance(n.value.value, str) + ) + + def _methods(node): - return {n.name for n in node.body if isinstance(n, ast.FunctionDef)} + """Every name the body of class *node* binds, and so answers `getattr` with. + + A `def` and an `async def` bind a name the same way an assignment in the + class body does, and the worker's `getattr` finds all three. A statement in + the body that could bind a name some other way -- an `if`, a `try`, a + loop, an expression such as `vars().update(...)` -- is recorded in + `UNREAD_CLASS_BODY` rather than skipped, and + `test_every_class_body_the_method_sets_are_read_from_is_read` names it. + Recorded, not raised: this runs at import, where a raise would stop every + module importing this one from collecting. + """ + names = set() + for n in node.body: + if isinstance(n, (ast.FunctionDef, ast.AsyncFunctionDef, ast.ClassDef)): + names.add(n.name) + elif isinstance(n, (ast.Assign, ast.AnnAssign, ast.AugAssign)): + for target in n.targets if isinstance(n, ast.Assign) else [n.target]: + names |= {t.id for t in ast.walk(target) if isinstance(t, ast.Name)} + elif not isinstance(n, ast.Pass) and not _is_docstring(n): + UNREAD_CLASS_BODY.append( + f"{node.file}:{n.lineno}: {ast.unparse(n).splitlines()[0]}" + ) + return names def _busy_loop(): @@ -99,20 +134,30 @@ def _arity(parent): `0` means the reply is discarded -- the broadcast is a bare statement and nothing can read what came back. `1` means it is used whole: bound to a - name, handed straight back to this function's own caller, or passed on as - an argument. Anything above `1` is a tuple unpack, which is the only shape - that fixes a length rather than just a type. + name, or handed straight back to this function's own caller. Anything + above `1` is a tuple unpack, which is the only shape that fixes a length + rather than just a type. `ast.Return` is the case worth naming, because reading only `ast.Assign` scores `return self.runner_mgr.call_func(...)` as a discard when the value is in fact the function's result -- `engine_core.py:749`, `dummy_execution`. + + Any other shape -- a list or starred target, an annotated one, a chained + target, an argument to another call -- is scored `None` rather than `1`, + which would read an unpack as a use of the whole reply, and + `test_every_reply_is_taken_in_a_shape_the_arity_reads` names its site. """ if isinstance(parent, ast.Expr): return 0 - if isinstance(parent, ast.Assign): - target = parent.targets[0] - return len(target.elts) if isinstance(target, ast.Tuple) else 1 - return 1 + if isinstance(parent, ast.Return): + return 1 + target = parent.targets[0] if isinstance(parent, ast.Assign) else None + if isinstance(target, ast.Name) and len(parent.targets) == 1: + return 1 + starred = any(isinstance(e, ast.Starred) for e in getattr(target, "elts", ())) + if isinstance(target, ast.Tuple) and len(parent.targets) == 1 and not starred: + return len(target.elts) + return None def _mentions(roots, needle, base=REPO): @@ -150,13 +195,13 @@ def _call_sites(): isinstance(node, ast.Call) and isinstance(node.func, ast.Attribute) ): continue - if node.func.attr not in BROADCAST or not node.args: + if node.func.attr not in BROADCAST: continue where = f"{path.relative_to(REPO)}:{node.lineno}" if "compass" in path.parts: from_compass.append(where) continue - if not isinstance(node.args[0], ast.Constant): + if not node.args or not isinstance(node.args[0], ast.Constant): non_literal.append(where) continue aggregated = node.func.attr == "call_func_with_aggregation" @@ -239,6 +284,19 @@ def test_both_filters_in_the_derivation_drop_nothing(): assert FROM_COMPASS == [] +def test_every_reply_is_taken_in_a_shape_the_arity_reads(): + """A site `_arity` could not score would otherwise vanish from every arity set.""" + unscored = [ + f"{s.file}:{s.line}" for v in SITES.values() for s in v if s.arity is None + ] + assert unscored == [], f"replies taken in a shape not scored: {unscored}" + + +def test_every_class_body_the_method_sets_are_read_from_is_read(): + """A statement `_methods` could not read may bind a name it would miss.""" + assert UNREAD_CLASS_BODY == [], f"class bodies not read: {UNREAD_CLASS_BODY}" + + def test_the_dispatched_names_outside_the_surface_belong_to_other_runners(): """Without this the intersection above could shrink and look like a pass. @@ -511,13 +569,18 @@ def test_the_base_capture_reaches_a_device_before_it_reaches_the_model(): def test_get_num_blocks_refuses_and_the_keys_its_caller_reads_are_named(): site = SITES["get_num_blocks"][0] tree = ast.parse((REPO / site.file).read_text()) - required = { - n.slice.value + subscripts = [ + n for n in ast.walk(tree) - if isinstance(n, ast.Subscript) - and getattr(n.value, "id", None) == "block_info" - and isinstance(n.slice, ast.Constant) - } + if isinstance(n, ast.Subscript) and getattr(n.value, "id", None) == "block_info" + ] + unread = [ + f"{site.file}:{n.lineno}: {ast.unparse(n)}" + for n in subscripts + if not isinstance(n.slice, ast.Constant) + ] + assert not unread, f"a key the caller reads is not written as a literal: {unread}" + required = {n.slice.value for n in subscripts} optional = { n.args[0].value for n in ast.walk(tree) @@ -927,8 +990,21 @@ def test_what_a_hole_at_exit_loses_is_what_exit_does(): for n in ast.walk(_classes(ATOM_RUNNER)["ModelRunner"]) if isinstance(n, ast.FunctionDef) and n.name == "exit" ) - calls = {ast.unparse(n.func) for n in ast.walk(body) if isinstance(n, ast.Call)} + # Statements of the body itself: a call inside a lambda, a branch or a + # nested def is one `exit` may never make. + calls = { + ast.unparse(n.value.func) + for n in body.body + if isinstance(n, ast.Expr) and isinstance(n.value, ast.Call) + } assert {"destroy_dist_env", "torch.cuda.empty_cache"} <= calls + # And nothing leaves `exit` between them: its only ways out are the guard + # that opens it and the `return True` that closes it. + exits = sorted( + ast.unparse(n) for n in ast.walk(body) if isinstance(n, (ast.Return, ast.Raise)) + ) + assert exits == ["return", "return True"], f"exit leaves early: {exits}" + assert ast.unparse(body.body[-1]) == "return True" deleted = { ast.unparse(t) for n in ast.walk(body) @@ -943,7 +1019,11 @@ def test_what_a_hole_at_exit_loses_is_what_exit_does(): ] assert len(literal_loops) == 1 kv = literal_loops[0] - assert {e.value for e in kv.iter.elts if isinstance(e, ast.Constant)} == { + unread = [ast.unparse(e) for e in kv.iter.elts if not isinstance(e, ast.Constant)] + assert ( + not unread + ), f"line {kv.lineno} deletes names not written as literals: {unread}" + assert {e.value for e in kv.iter.elts} == { "kv_cache", "kv_scale", "index_cache", @@ -961,16 +1041,29 @@ def test_the_unanswered_helper_returns_one_list_its_one_caller_partitions(): """`unanswered_rpc_names` returns its list unpartitioned, and its only caller in the package splits it on `RPC_SURFACE` before reporting it. """ - callers = [ - str(f.relative_to(REPO)) - for f in sorted((REPO / "atom").rglob("*.py")) - if any( - isinstance(n, ast.Call) - and getattr(n.func, "id", None) == "unanswered_rpc_names" - for n in ast.walk(ast.parse(f.read_text())) - ) - ] - assert callers == ["atom/compass/runner/model_runner.py"] + # Called by bare name or through a module, both are a call; any other + # mention -- an `import ... as`, a `partial`, a callback -- is a caller + # this cannot follow, so it is refused rather than left out of the count. + callers, unread = set(), [] + for f in sorted((REPO / "atom").rglob("*.py")): + tree = ast.parse(f.read_text()) + called = {id(n.func) for n in ast.walk(tree) if isinstance(n, ast.Call)} + for n in ast.walk(tree): + if isinstance(n, ast.alias) and n.name == "unanswered_rpc_names": + if n.asname: + unread.append(f"{f.relative_to(REPO)}:{n.lineno}: as {n.asname}") + continue + if "unanswered_rpc_names" not in ( + getattr(n, "id", 0), + getattr(n, "attr", 0), + ): + continue + if id(n) in called: + callers.add(str(f.relative_to(REPO))) + else: + unread.append(f"{f.relative_to(REPO)}:{n.lineno}: {ast.unparse(n)}") + assert not unread, f"the helper is reached other than by calling it: {unread}" + assert callers == {"atom/compass/runner/model_runner.py"} tree = ast.parse(COMPOSED) bound = next( n.targets[0].id diff --git a/tests/compass/test_runner_step_semantics.py b/tests/compass/test_runner_step_semantics.py index bf9628e38a..0edc1697b3 100644 --- a/tests/compass/test_runner_step_semantics.py +++ b/tests/compass/test_runner_step_semantics.py @@ -453,17 +453,48 @@ def _decorators(path, class_name, name): return [ast.unparse(d) for d in _function(path, class_name, name).decorator_list] +# The methods the reply is handed to whole. `postprocess` reads it under +# `fwd_output`, which the walk below sees; `send_tokens` pickles it and reads +# nothing. A hand-off to any other method is refused: it would read the reply +# under a name this walk does not follow. +HANDED_TO = ("postprocess", "send_tokens") + + def _reply_attribute_reads(): - """Every attribute ATOM reads off a forward reply, from ATOM's own source.""" - names = set() + """Every attribute ATOM reads off a forward reply, from ATOM's own source. + + A read is `fwd_out.x` or `fwd_output.x`. The reply may also be handed on + whole -- passed positionally to one of `HANDED_TO`, returned, or tested + against `None`. Any other use of it, a `getattr`, an alias, a hand-off to + any other function or method, is refused by file and line: it + could read an attribute this set would never contain. + """ + names, unread = set(), [] for path in (REPO / "atom" / "model_engine").glob("*.py"): - for node in ast.walk(ast.parse(path.read_text())): - if ( - isinstance(node, ast.Attribute) - and isinstance(node.value, ast.Name) - and node.value.id in {"fwd_out", "fwd_output"} + tree = ast.parse(path.read_text()) + parent = {c: n for n in ast.walk(tree) for c in ast.iter_child_nodes(n)} + for node in ast.walk(tree): + if not ( + isinstance(node, ast.Name) + and node.id in {"fwd_out", "fwd_output"} + and isinstance(node.ctx, ast.Load) ): - names.add(node.attr) + continue + up = parent[node] + handed_on = ( + isinstance(up, ast.Return) + or ast.unparse(up) == f"{node.id} is None" + or ( + getattr(up, "func", None) is not None + and getattr(up.func, "attr", None) in HANDED_TO + and node in up.args + ) + ) + if isinstance(up, ast.Attribute): + names.add(up.attr) + elif not handed_on: + unread.append(f"{path.name}:{node.lineno}: {ast.unparse(up)}") + assert not unread, f"the forward reply is used in a form not read here: {unread}" return names