diff --git a/src/archunitpython/common/extraction/extract_graph.py b/src/archunitpython/common/extraction/extract_graph.py index 2ca4972..bbb09b3 100644 --- a/src/archunitpython/common/extraction/extract_graph.py +++ b/src/archunitpython/common/extraction/extract_graph.py @@ -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 @@ -39,6 +42,11 @@ r"(?P(?:\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: @@ -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 @@ -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 @@ -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) @@ -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. diff --git a/tests/common/test_extract_graph.py b/tests/common/test_extract_graph.py index bf926d6..08491b3 100644 --- a/tests/common/test_extract_graph.py +++ b/tests/common/test_extract_graph.py @@ -1,9 +1,11 @@ """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 @@ -11,6 +13,7 @@ from archunitpython.common.extraction.extract_graph import ( _extract_imports, _find_python_files, + _import_analysis_nodes, _normalize, _resolve_exclude_patterns, clear_graph_cache, @@ -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__"])