Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
81 changes: 76 additions & 5 deletions src/archunitpython/common/extraction/extract_graph.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,7 +5,10 @@
import ast
import os
import re
from collections import deque
from collections.abc import Iterator
from dataclasses import dataclass, field
from itertools import islice

from archunitpython.common.extraction.graph import Edge, Graph, ImportKind
from archunitpython.common.fluentapi.checkable import CheckOptions
Expand Down Expand Up @@ -39,6 +42,11 @@
r"(?P<modules>(?:\s+[\w.]+)*)\s*$"
)

_IMPORT_ANALYSIS_NODE_TYPES = (ast.Import, ast.ImportFrom, ast.Call, ast.If, ast.Try)
_IMPORT_LEAF_NODE_TYPES = frozenset((ast.Import, ast.ImportFrom, ast.Name, ast.Constant))
_PARALLEL_EXTRACTION_MIN_FILES = 64
_PARALLEL_EXTRACTION_WORKERS = 8


@dataclass(frozen=True)
class _LocatedImport:
Expand Down Expand Up @@ -225,12 +233,56 @@ def _load_archignore_patterns(project_path: str) -> list[str]:
def _extract_graph_from_files(session: _ExtractionSession, files: list[str]) -> Graph:
"""Assemble a graph from cached or newly parsed file edges."""
edges: list[Edge] = []
for file_path in files:
edges.extend(_extract_file_edges(session, file_path))
for file_edges in _iter_file_edges(session, files):
edges.extend(file_edges)
return _merge_edges(edges)


def _extract_file_edges(session: _ExtractionSession, file_path: str) -> list[Edge]:
def _iter_file_edges(session: _ExtractionSession, files: list[str]) -> Iterator[list[Edge]]:
"""Overlap independent source reads while assembling edges in original order.

Small/cached selections remain serial. The pending window is bounded, and
only the calling thread resolves targets or mutates the extraction session.
"""
uncached_count = sum(path not in session.edges_by_file for path in files)
if uncached_count < _PARALLEL_EXTRACTION_MIN_FILES:
for path in files:
yield _extract_file_edges(session, path)
return

# Avoid importing executor machinery for small or already cached checks.
from concurrent.futures import Future, ThreadPoolExecutor

paths = iter(files)
with ThreadPoolExecutor(max_workers=_PARALLEL_EXTRACTION_WORKERS) as executor:
pending: deque[tuple[str, Future[list[_LocatedImport]] | None]] = deque()

def schedule(path: str) -> tuple[str, Future[list[_LocatedImport]] | None]:
future = (
None
if path in session.edges_by_file
else executor.submit(_extract_located_imports, path)
)
return path, future

for path in islice(paths, _PARALLEL_EXTRACTION_WORKERS):
pending.append(schedule(path))

while pending:
path, future = pending.popleft()
imports = future.result() if future is not None else None
next_path = next(paths, None)
if next_path is not None:
pending.append(schedule(next_path))
yield _extract_file_edges(session, path, located_imports=imports)


def _extract_file_edges(
session: _ExtractionSession,
file_path: str,
*,
located_imports: list[_LocatedImport] | None = None,
) -> list[Edge]:
cached = session.edges_by_file.get(file_path)
if cached is not None:
return cached
Expand All @@ -244,7 +296,9 @@ def _extract_file_edges(session: _ExtractionSession, file_path: str) -> list[Edg
)
]

for located_import in _extract_located_imports(file_path):
if located_imports is None:
located_imports = _extract_located_imports(file_path)
for located_import in located_imports:
if (
session.ignore_type_checking_imports
and located_import.import_kind == ImportKind.TYPE_IMPORT
Expand Down Expand Up @@ -369,7 +423,7 @@ def _extract_located_imports(file_path: str) -> list[_LocatedImport]:

imports: list[_LocatedImport] = []
ignore_directives = _find_ignore_directives(source)
nodes = list(ast.walk(tree))
nodes = _import_analysis_nodes(tree)
type_checking_ranges = _find_type_checking_ranges(nodes)
conditional_import_ranges = _find_conditional_import_ranges(nodes)

Expand Down Expand Up @@ -425,6 +479,23 @@ def _extract_located_imports(file_path: str) -> list[_LocatedImport]:
return [import_ for import_ in imports if not _is_ignored_import(import_, ignore_directives)]


def _import_analysis_nodes(tree: ast.AST) -> list[ast.AST]:
"""Collect import/context nodes in ast.walk order without visiting inert leaves.

Unknown node types still use the standard child iterator, so new syntax is
not silently skipped. Only parser-produced terminal types are pruned.
"""
pending = deque([tree])
nodes: list[ast.AST] = []
while pending:
node = pending.popleft()
if isinstance(node, _IMPORT_ANALYSIS_NODE_TYPES):
nodes.append(node)
if node._fields and type(node) not in _IMPORT_LEAF_NODE_TYPES:
pending.extend(ast.iter_child_nodes(node))
return nodes


def _find_ignore_directives(source: str) -> dict[int, _IgnoreDirective]:
"""Find architecture-ignore directives.

Expand Down
184 changes: 184 additions & 0 deletions tests/common/test_extract_graph.py
Original file line number Diff line number Diff line change
@@ -1,16 +1,19 @@
"""Tests for graph extraction."""

import ast
import importlib
import os
import shutil
from pathlib import Path
from threading import Event, Lock, get_ident
from uuid import uuid4

import pytest

from archunitpython.common.extraction.extract_graph import (
_extract_imports,
_find_python_files,
_import_analysis_nodes,
_normalize,
_resolve_exclude_patterns,
clear_graph_cache,
Expand All @@ -26,6 +29,187 @@
SAMPLE_PROJECT = os.path.join(FIXTURES_DIR, "sample_project")


@pytest.mark.parametrize(
"source",
[
"import a, b\nfrom .c import d\n",
'@decorate(__import__("decorator"))\n'
'def f(x: __import__("annotation") = __import__("default")):\n'
' return [__import__("item") for x in values if __import__("guard")]\n',
'class C(__import__("base"), metaclass=__import__("meta")):\n'
' value = lambda: __import__("lambda")\n',
"value = f\"{__import__('embedded')}\"\n",
"try:\n import optional\nexcept ImportError:\n import fallback\n"
"if TYPE_CHECKING:\n from typing import Any\n",
'match value:\n case {"x": x} if __import__("guard"):\n import matched\n',
'with __import__("context"):\n import body\n',
],
)
def test_import_analysis_nodes_preserve_standard_walk_order(source):
tree = ast.parse(source)
expected = [
node
for node in ast.walk(tree)
if isinstance(node, (ast.Import, ast.ImportFrom, ast.Call, ast.If, ast.Try))
]
assert _import_analysis_nodes(tree) == expected


def test_import_analysis_nodes_traverse_unknown_node_types():
class FutureNode(ast.AST):
_fields = ("body",)

tree = FutureNode()
tree.body = ast.parse('import a\n__import__("b")\n').body
assert [type(node) for node in _import_analysis_nodes(tree)] == [ast.Import, ast.Call]


def test_parallel_extraction_preserves_graph_and_calling_thread_resolution(monkeypatch):
extraction = importlib.import_module("archunitpython.common.extraction.extract_graph")
clear_graph_cache()
expected = extract_graph(SAMPLE_PROJECT)
clear_graph_cache()
monkeypatch.setattr(extraction, "_PARALLEL_EXTRACTION_MIN_FILES", 1)
main_thread = get_ident()
worker_threads = set()
lock = Lock()
original_extract = extraction._extract_located_imports
original_resolve = extraction._resolve_import_targets

def extract(path):
with lock:
worker_threads.add(get_ident())
return original_extract(path)

def resolve(*args):
assert get_ident() == main_thread
return original_resolve(*args)

monkeypatch.setattr(extraction, "_extract_located_imports", extract)
monkeypatch.setattr(extraction, "_resolve_import_targets", resolve)
assert extract_graph(SAMPLE_PROJECT) == expected
assert worker_threads
assert main_thread not in worker_threads
clear_graph_cache()


def test_parallel_partial_then_full_graph_parses_each_file_once(monkeypatch):
import concurrent.futures

extraction = importlib.import_module("archunitpython.common.extraction.extract_graph")
clear_graph_cache()
monkeypatch.setattr(extraction, "_PARALLEL_EXTRACTION_MIN_FILES", 1)
original_extract = extraction._extract_located_imports
parsed = []
lock = Lock()

def extract(path):
with lock:
parsed.append(path)
return original_extract(path)

monkeypatch.setattr(extraction, "_extract_located_imports", extract)
selected = [RegexFactory.folder_matcher("**/services*")]
partial = extract_graph_for_sources(SAMPLE_PROJECT, selected)
assert all(matches_all_patterns(_normalize(path), selected) for path in parsed)
full = extract_graph(SAMPLE_PROJECT)
assert partial == [edge for edge in full if matches_all_patterns(edge.source, selected)]
assert len(parsed) == len(set(parsed))
assert set(parsed) == set(_find_python_files(SAMPLE_PROJECT, ["__pycache__"]))

def unexpected_extract(_path):
raise AssertionError("Cached files must not be read again")

monkeypatch.setattr(extraction, "_extract_located_imports", unexpected_extract)

def unexpected_executor(**_kwargs):
raise AssertionError("Cached selections must not create a worker pool")

monkeypatch.setattr(concurrent.futures, "ThreadPoolExecutor", unexpected_executor)
assert extract_graph_for_sources(SAMPLE_PROJECT, selected) == partial
clear_graph_cache()


def test_parallel_extraction_keeps_file_order_when_later_file_finishes_first(monkeypatch):
extraction = importlib.import_module("archunitpython.common.extraction.extract_graph")
monkeypatch.setattr(extraction, "_PARALLEL_EXTRACTION_MIN_FILES", 1)
monkeypatch.setattr(extraction, "_PARALLEL_EXTRACTION_WORKERS", 2)
second_finished = Event()
files = ["first.py", "second.py", "third.py"]
session = extraction._ExtractionSession("project", files, set(files), False)

def extract(path):
if path == "first.py":
assert second_finished.wait(timeout=10)
elif path == "second.py":
second_finished.set()
return [extraction._LocatedImport(f"dependency-{path}", ImportKind.IMPORT, 1)]

monkeypatch.setattr(extraction, "_extract_located_imports", extract)
monkeypatch.setattr(
extraction, "_resolve_import_targets", lambda import_, *_: [(import_.module_name, True)]
)
graph = extraction._extract_graph_from_files(session, files)
assert [(edge.source, edge.target) for edge in graph] == [
pair
for path in files
for pair in [(path, path), (path, f"dependency-{path}")]
]


@pytest.mark.parametrize("fail_first", [False, True])
def test_parallel_extraction_bounds_pending_tasks(monkeypatch, fail_first):
import concurrent.futures

extraction = importlib.import_module("archunitpython.common.extraction.extract_graph")
monkeypatch.setattr(extraction, "_PARALLEL_EXTRACTION_MIN_FILES", 1)
monkeypatch.setattr(extraction, "_PARALLEL_EXTRACTION_WORKERS", 3)
files = [f"file-{index}.py" for index in range(20)]
session = extraction._ExtractionSession("project", files, set(files), False)
outstanding = 0
maximum = 0
closed = False

class Result:
def __init__(self, path):
self.path = path

def result(self):
nonlocal outstanding
outstanding -= 1
if fail_first and self.path == files[0]:
raise RuntimeError("Unexpected extraction failure")
return []

class Executor:
def __init__(self, *, max_workers):
assert max_workers == 3

def __enter__(self):
return self

def __exit__(self, *_args):
nonlocal closed
closed = True

def submit(self, _function, path):
nonlocal outstanding, maximum
outstanding += 1
maximum = max(maximum, outstanding)
return Result(path)

monkeypatch.setattr(concurrent.futures, "ThreadPoolExecutor", Executor)
if fail_first:
with pytest.raises(RuntimeError, match="Unexpected extraction failure"):
extraction._extract_graph_from_files(session, files)
assert not session.edges_by_file
else:
assert len(extraction._extract_graph_from_files(session, files)) == len(files)
assert outstanding == 0
assert maximum == 3
assert closed


class TestFindPythonFiles:
def test_finds_all_py_files(self):
files = _find_python_files(SAMPLE_PROJECT, ["__pycache__"])
Expand Down
Loading