Skip to content
17 changes: 11 additions & 6 deletions tests/compass/test_clock_lp_identity.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)}"

Expand All @@ -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"
Expand Down
5 changes: 5 additions & 0 deletions tests/compass/test_ir_data_model.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)}"

Expand Down
38 changes: 35 additions & 3 deletions tests/compass/test_kv_blob_site.py
Original file line number Diff line number Diff line change
Expand Up @@ -28,9 +28,11 @@
def _blob_site(path: Path) -> list[ast.Assign]:
"""Every `<something>.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 [
Expand All @@ -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 "
Expand Down
88 changes: 72 additions & 16 deletions tests/compass/test_runner_non_allocating.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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 {
Expand Down Expand Up @@ -207,25 +206,67 @@ 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")

Copy link
Copy Markdown
Owner Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

G5 (cycle 2). Not blocking, principle 6 and principle 8 (on the docstring's claim). F1, F2 and F7 now refuse every identifier spelling of the writers and of the helper; N10 __import__("builtins").setattr is red. A name looked up as data still passes:

mutation head 176cac3d4
; __builtins__["setattr"](self, "kv_cache", None) (overrides.py:246) 1153 passed
}; __builtins__["setattr"](seq, "kv_transfer_params" + "_output", …) (moriio_connector.py:997) 1153 passed
def _all(): from atom.compass.runner import overrides as o; return getattr(o, "unanswered_rpc_names")(object) (step_output.py:60) 1153 passed

Two cheap changes:

  • Refuse a literal-string mention of the watched names ("setattr", "__setattr__", "__dict__", "vars", "unanswered_rpc_names"), as _writes_not_read already does for BLOB_ATTR. That closes all three rows. None of the three scopes contains such a string today; measure that first.
  • Record the residue. A name built at run time ("set" + "attr") is the same class as the dynamic imports this PR already records as out of reach. Record it there, and drop "however it is reached" from _self_assigned's docstring, because that overclaims.



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
of the ring it already holds -- and a mapping keeps only the last binding
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))

Copy link
Copy Markdown
Owner Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

G1 (cycle 2). Blocking, principle 6. The ctx=Store half reads a write only when its base is exactly self, and skips a write on any other base without refusing it. A write through the type binds an attribute that every instance reads:

mutation, appended to overrides.py:246 tip ff9617f30 head 176cac3d4
; type(self).kv_cache = None 1164 passed 1153 passed, rc=0

self.__class__.kv_cache = None takes the same route: its base is self.__class__, not self, and __class__ is not in _WRITERS.

A generic "refuse every Store not on self" rule would hit 36 legitimate seq., config. and batch. stores in ModelRunner, which I measured at head. In NonAllocatingRunner the same count is 0, and it has 0 chains through __class__. So the refusal can be scoped to the overrides call: refuse any attribute store not rooted at self, and any chain through __class__. That fires on nothing today. This spelling was silent at the tip too, and neither inventory listed it.

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():
Expand All @@ -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"}

Expand Down Expand Up @@ -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
Expand Down
Loading