refactor(workflow): contract shared core (#33)
This commit is contained in:
@@ -0,0 +1,72 @@
|
||||
"""Architecture contracts for the internal workflow modules."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import ast
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
THT_ROOT = Path(__file__).resolve().parents[1] / "tht"
|
||||
DOMAIN_PACKAGES = ("evidence", "memory")
|
||||
|
||||
|
||||
def _package(path: Path) -> tuple[str, ...]:
|
||||
relative = path.relative_to(THT_ROOT).with_suffix("").parts
|
||||
return ("tht", *relative[:-1])
|
||||
|
||||
|
||||
def _resolve_import_from(package: tuple[str, ...], node: ast.ImportFrom) -> set[str]:
|
||||
if node.level:
|
||||
keep = len(package) - (node.level - 1)
|
||||
base = package[: max(keep, 0)]
|
||||
else:
|
||||
base = ()
|
||||
module = (*base, *node.module.split(".")) if node.module else base
|
||||
resolved = {".".join(module)} if module else set()
|
||||
if not node.module or node.module == "tht":
|
||||
resolved.update(
|
||||
".".join((*module, alias.name))
|
||||
for alias in node.names
|
||||
if alias.name != "*"
|
||||
)
|
||||
return resolved
|
||||
|
||||
|
||||
def _imports(path: Path) -> set[str]:
|
||||
tree = ast.parse(path.read_text(), filename=str(path))
|
||||
package = _package(path)
|
||||
imported: set[str] = set()
|
||||
for node in ast.walk(tree):
|
||||
if isinstance(node, ast.Import):
|
||||
imported.update(alias.name for alias in node.names)
|
||||
elif isinstance(node, ast.ImportFrom):
|
||||
imported.update(_resolve_import_from(package, node))
|
||||
return imported
|
||||
|
||||
|
||||
def test_relative_import_resolution_reaches_a_sibling_domain() -> None:
|
||||
imports = (
|
||||
"from ..memory import core",
|
||||
"from .. import memory",
|
||||
"from tht import memory",
|
||||
)
|
||||
|
||||
for statement in imports:
|
||||
node = ast.parse(statement).body[0]
|
||||
assert isinstance(node, ast.ImportFrom)
|
||||
assert "tht.memory" in _resolve_import_from(("tht", "evidence"), node)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("domain", DOMAIN_PACKAGES)
|
||||
def test_python_domain_does_not_import_another_domain(domain: str) -> None:
|
||||
forbidden = {f"tht.{other}" for other in DOMAIN_PACKAGES if other != domain}
|
||||
violations: list[str] = []
|
||||
|
||||
for path in sorted((THT_ROOT / domain).rglob("*.py")):
|
||||
for imported in sorted(_imports(path)):
|
||||
if any(imported == root or imported.startswith(f"{root}.") for root in forbidden):
|
||||
violations.append(f"{path.relative_to(THT_ROOT)} -> {imported}")
|
||||
|
||||
assert violations == []
|
||||
Reference in New Issue
Block a user