72 lines
2.3 KiB
Python
72 lines
2.3 KiB
Python
"""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 == []
|