diff --git a/src/archunitpython/common/extraction/extract_graph.py b/src/archunitpython/common/extraction/extract_graph.py index e0ddc14..2ca4972 100644 --- a/src/archunitpython/common/extraction/extract_graph.py +++ b/src/archunitpython/common/extraction/extract_graph.py @@ -5,10 +5,12 @@ import ast import os import re -from dataclasses import dataclass +from dataclasses import dataclass, field from archunitpython.common.extraction.graph import Edge, Graph, ImportKind from archunitpython.common.fluentapi.checkable import CheckOptions +from archunitpython.common.pattern_matching import matches_all_patterns +from archunitpython.common.types import Filter GraphCacheKey = tuple[str, tuple[str, ...], bool] @@ -57,15 +59,28 @@ def matches(self, import_: _LocatedImport) -> bool: if not self.modules: return True return any( - import_.module_name == module - or import_.module_name.startswith(f"{module}.") + import_.module_name == module or import_.module_name.startswith(f"{module}.") for module in self.modules ) +@dataclass +class _ExtractionSession: + project_path: str + py_files: list[str] + normalized_py_file_set: set[str] + ignore_type_checking_imports: bool + edges_by_file: dict[str, list[Edge]] = field(default_factory=dict) + selected_files: dict[tuple[Filter, ...], list[str]] = field(default_factory=dict) + + +_extraction_sessions: dict[GraphCacheKey, _ExtractionSession] = {} + + def clear_graph_cache(options: CheckOptions | None = None) -> None: """Clear the cached dependency graphs.""" _graph_cache.clear() + _extraction_sessions.clear() def extract_graph( @@ -99,19 +114,73 @@ def extract_graph( if options and options.clear_cache: _graph_cache.pop(cache_key, None) + _extraction_sessions.pop(cache_key, None) if cache_key in _graph_cache: return _graph_cache[cache_key] - result = _extract_graph_uncached( - project_path, - excludes, - ignore_type_checking_imports=ignore_type_checking_imports, + session = _get_extraction_session( + cache_key, project_path, excludes, ignore_type_checking_imports ) + result = _extract_graph_from_files(session, session.py_files) _graph_cache[cache_key] = result return result +def extract_graph_for_sources( + project_path: str | None, + source_filters: list[Filter], + *, + options: CheckOptions | None = None, +) -> Graph: + """Extract edges from files selected by a direct file-dependency rule. + + All project files are inventoried to classify imported targets correctly. + Parsed edges are reused by later rules and full-graph checks in this process. + """ + if not source_filters: + return extract_graph(project_path, options=options) + + if project_path is None: + project_path = os.getcwd() + project_path = os.path.abspath(project_path) + excludes = _resolve_exclude_patterns(project_path, None) + ignore_type_checking_imports = bool(options and options.ignore_type_checking_imports) + cache_key = _build_cache_key(project_path, excludes, ignore_type_checking_imports) + if options and options.clear_cache: + _graph_cache.pop(cache_key, None) + _extraction_sessions.pop(cache_key, None) + + session = _get_extraction_session( + cache_key, project_path, excludes, ignore_type_checking_imports + ) + filter_key = tuple(source_filters) + if filter_key not in session.selected_files: + session.selected_files[filter_key] = [ + path for path in session.py_files if matches_all_patterns(path, source_filters) + ] + return _extract_graph_from_files(session, session.selected_files[filter_key]) + + +def _get_extraction_session( + cache_key: GraphCacheKey, + project_path: str, + excludes: list[str], + ignore_type_checking_imports: bool, +) -> _ExtractionSession: + session = _extraction_sessions.get(cache_key) + if session is None: + py_files = _find_python_files(project_path, excludes) + session = _ExtractionSession( + project_path=project_path, + py_files=py_files, + normalized_py_file_set={_normalize(path) for path in py_files}, + ignore_type_checking_imports=ignore_type_checking_imports, + ) + _extraction_sessions[cache_key] = session + return session + + def _build_cache_key( project_path: str, exclude_patterns: list[str], @@ -153,55 +222,51 @@ def _load_archignore_patterns(project_path: str) -> list[str]: return patterns -def _extract_graph_uncached( - project_path: str, - exclude_patterns: list[str], - *, - ignore_type_checking_imports: bool = False, -) -> Graph: - """Extract graph without caching.""" - py_files = _find_python_files(project_path, exclude_patterns) - +def _extract_graph_from_files(session: _ExtractionSession, files: list[str]) -> Graph: + """Assemble a graph from cached or newly parsed file edges.""" edges: list[Edge] = [] - py_files_set = set(py_files) - normalized_py_file_set = {_normalize(f) for f in py_files_set} - - for file_path in py_files: - # Add self-referencing edge (ensures the file appears as a node) - edges.append( - Edge( - source=_normalize(file_path), - target=_normalize(file_path), - external=False, - ) + for file_path in files: + edges.extend(_extract_file_edges(session, file_path)) + return _merge_edges(edges) + + +def _extract_file_edges(session: _ExtractionSession, file_path: str) -> list[Edge]: + cached = session.edges_by_file.get(file_path) + if cached is not None: + return cached + + source_label = _normalize(file_path) + edges = [ + Edge( + source=source_label, + target=source_label, + external=False, ) + ] - imports = _extract_located_imports(file_path) - for located_import in imports: - import_kind = located_import.import_kind - if ( - ignore_type_checking_imports - and import_kind == ImportKind.TYPE_IMPORT - ): - continue - for resolved, is_external in _resolve_import_targets( - located_import, file_path, project_path - ): - if resolved and resolved != _normalize(file_path): - # Check if the resolved path is in our project - if not is_external and resolved not in normalized_py_file_set: - continue - - edges.append( - Edge( - source=_normalize(file_path), - target=resolved, - external=is_external, - import_kinds=_edge_import_kinds(located_import), - ) + for located_import in _extract_located_imports(file_path): + if ( + session.ignore_type_checking_imports + and located_import.import_kind == ImportKind.TYPE_IMPORT + ): + continue + for resolved, is_external in _resolve_import_targets( + located_import, file_path, session.project_path + ): + if resolved and resolved != source_label: + if not is_external and resolved not in session.normalized_py_file_set: + continue + edges.append( + Edge( + source=source_label, + target=resolved, + external=is_external, + import_kinds=_edge_import_kinds(located_import), ) + ) - return _merge_edges(edges) + session.edges_by_file[file_path] = edges + return edges def _normalize(path: str) -> str: @@ -213,6 +278,9 @@ def _find_python_files(root: str, exclude: list[str]) -> list[str]: """Recursively find all .py files, excluding specified patterns.""" py_files: list[str] = [] root = os.path.abspath(root) + # Defaults can only match directory names, never a .py filename. Keep + # checking user patterns against files, including .archignore entries. + file_excludes = [pattern for pattern in exclude if pattern not in _DEFAULT_EXCLUDE] for dirpath, dirnames, filenames in os.walk(root): # Filter out excluded directories in-place dirnames[:] = [ @@ -223,8 +291,9 @@ def _find_python_files(root: str, exclude: list[str]) -> list[str]: for filename in filenames: full_path = os.path.join(dirpath, filename) - if filename.endswith(".py") and not _should_exclude_path( - full_path, root, exclude, is_dir=False + if filename.endswith(".py") and not ( + file_excludes + and _should_exclude_path(full_path, root, file_excludes, is_dir=False) ): py_files.append(os.path.abspath(full_path)) @@ -300,10 +369,11 @@ def _extract_located_imports(file_path: str) -> list[_LocatedImport]: imports: list[_LocatedImport] = [] ignore_directives = _find_ignore_directives(source) - type_checking_ranges = _find_type_checking_ranges(tree) - conditional_import_ranges = _find_conditional_import_ranges(tree) + nodes = list(ast.walk(tree)) + type_checking_ranges = _find_type_checking_ranges(nodes) + conditional_import_ranges = _find_conditional_import_ranges(nodes) - for node in ast.walk(tree): + for node in nodes: if isinstance(node, ast.Import): syntax_kind = ImportKind.IMPORT kind = _classify_import( @@ -313,9 +383,7 @@ def _extract_located_imports(file_path: str) -> list[_LocatedImport]: conditional_import_ranges, ) for alias in node.names: - imports.append( - _LocatedImport(alias.name, kind, node.lineno, syntax_kind) - ) + imports.append(_LocatedImport(alias.name, kind, node.lineno, syntax_kind)) elif isinstance(node, ast.ImportFrom): syntax_kind = ( @@ -329,9 +397,7 @@ def _extract_located_imports(file_path: str) -> list[_LocatedImport]: type_checking_ranges, conditional_import_ranges, ) - fallback_module_name = ( - "." * node.level if node.level and node.module is None else None - ) + fallback_module_name = "." * node.level if node.level and node.module is None else None aliases = _module_aliases(node) if node.module else () for module_name in _import_from_module_names(node): imports.append( @@ -354,15 +420,9 @@ def _extract_located_imports(file_path: str) -> list[_LocatedImport]: conditional_import_ranges, ) for module_name in _extract_dynamic_import_names(node): - imports.append( - _LocatedImport(module_name, kind, node.lineno, syntax_kind) - ) + imports.append(_LocatedImport(module_name, kind, node.lineno, syntax_kind)) - return [ - import_ - for import_ in imports - if not _is_ignored_import(import_, ignore_directives) - ] + return [import_ for import_ in imports if not _is_ignored_import(import_, ignore_directives)] def _find_ignore_directives(source: str) -> dict[int, _IgnoreDirective]: @@ -463,11 +523,11 @@ def _classify_import( return default_kind -def _find_type_checking_ranges(tree: ast.Module) -> list[tuple[int, int]]: +def _find_type_checking_ranges(nodes: list[ast.AST]) -> list[tuple[int, int]]: """Find line ranges of TYPE_CHECKING blocks.""" ranges: list[tuple[int, int]] = [] - for node in ast.walk(tree): + for node in nodes: if isinstance(node, ast.If): # Check for `if TYPE_CHECKING:` pattern test = node.test @@ -488,11 +548,11 @@ def _find_type_checking_ranges(tree: ast.Module) -> list[tuple[int, int]]: return sorted(ranges, key=lambda ele: ele[0]) -def _find_conditional_import_ranges(tree: ast.Module) -> list[tuple[int, int]]: +def _find_conditional_import_ranges(nodes: list[ast.AST]) -> list[tuple[int, int]]: """Find try/except ImportError ranges that contain optional imports.""" ranges: list[tuple[int, int]] = [] - for node in ast.walk(tree): + for node in nodes: if not isinstance(node, ast.Try): continue if not any(_handles_import_error(handler.type) for handler in node.handlers): @@ -555,14 +615,10 @@ def _resolve_import( Returns (resolved_path, is_external). The path is normalized with forward slashes. """ - if ( - kind - in ( - ImportKind.RELATIVE_IMPORT, - ImportKind.TYPE_IMPORT, - ) - and import_name.startswith(".") - ): + if kind in ( + ImportKind.RELATIVE_IMPORT, + ImportKind.TYPE_IMPORT, + ) and import_name.startswith("."): # Relative import return _resolve_relative_import(import_name, source_file, project_root) diff --git a/src/archunitpython/common/pattern_matching.py b/src/archunitpython/common/pattern_matching.py index 1956602..bb6cfb5 100644 --- a/src/archunitpython/common/pattern_matching.py +++ b/src/archunitpython/common/pattern_matching.py @@ -47,7 +47,8 @@ def matches_pattern(file_path: str, filter_: Filter) -> bool: else: target_string = normalize_path(file_path) - return bool(filter_.regexp.search(target_string)) + regexp = filter_.search_regexp or filter_.regexp + return bool(regexp.search(target_string)) def matches_pattern_classname(class_name: str, file_path: str, filter_: Filter) -> bool: @@ -65,7 +66,8 @@ def matches_pattern_classname(class_name: str, file_path: str, filter_: Filter) else: target_string = normalize_path(file_path) - return bool(filter_.regexp.search(target_string)) + regexp = filter_.search_regexp or filter_.regexp + return bool(regexp.search(target_string)) def matches_all_patterns(file_path: str, filters: list[Filter]) -> bool: diff --git a/src/archunitpython/common/regex_factory.py b/src/archunitpython/common/regex_factory.py index df470a1..ea792ea 100644 --- a/src/archunitpython/common/regex_factory.py +++ b/src/archunitpython/common/regex_factory.py @@ -29,6 +29,16 @@ def _pattern_to_regex(pattern: Pattern) -> re.Pattern[str]: return _glob_to_regex(pattern) +def _optimized_search_regex(pattern: Pattern) -> re.Pattern[str] | None: + """Avoid redundant leading glob stars only for the search path. + + Keep Filter.regexp unchanged for callers that use match/fullmatch directly. + """ + if isinstance(pattern, str) and pattern.startswith("*"): + return _glob_to_regex(pattern.lstrip("*")) + return None + + class RegexFactory: """Factory for creating Filter objects from patterns.""" @@ -38,6 +48,7 @@ def filename_matcher(name: Pattern) -> Filter: return Filter( regexp=_pattern_to_regex(name), options=PatternMatchingOptions(target="filename"), + search_regexp=_optimized_search_regex(name), ) @staticmethod @@ -46,6 +57,7 @@ def classname_matcher(name: Pattern) -> Filter: return Filter( regexp=_pattern_to_regex(name), options=PatternMatchingOptions(target="classname"), + search_regexp=_optimized_search_regex(name), ) @staticmethod @@ -54,6 +66,7 @@ def folder_matcher(folder: Pattern) -> Filter: return Filter( regexp=_pattern_to_regex(folder), options=PatternMatchingOptions(target="path-no-filename"), + search_regexp=_optimized_search_regex(folder), ) @staticmethod @@ -62,6 +75,7 @@ def path_matcher(path: Pattern) -> Filter: return Filter( regexp=_pattern_to_regex(path), options=PatternMatchingOptions(target="path"), + search_regexp=_optimized_search_regex(path), ) @staticmethod diff --git a/src/archunitpython/common/types.py b/src/archunitpython/common/types.py index 4e1839c..0e0ed12 100644 --- a/src/archunitpython/common/types.py +++ b/src/archunitpython/common/types.py @@ -3,7 +3,7 @@ from __future__ import annotations import re -from dataclasses import dataclass +from dataclasses import dataclass, field from typing import Literal, Union Pattern = Union[str, re.Pattern[str]] @@ -27,3 +27,4 @@ class Filter: regexp: re.Pattern[str] options: PatternMatchingOptions + search_regexp: re.Pattern[str] | None = field(default=None, compare=False, repr=False) diff --git a/src/archunitpython/files/assertion/depend_on_files.py b/src/archunitpython/files/assertion/depend_on_files.py index 69abd90..46826d8 100644 --- a/src/archunitpython/files/assertion/depend_on_files.py +++ b/src/archunitpython/files/assertion/depend_on_files.py @@ -39,13 +39,24 @@ def gather_depend_on_file_violations( List of violations. """ violations: list[Violation] = [] + subject_matches_by_label: dict[str, bool] = {} + target_matches_by_label: dict[str, bool] = {} for edge in edges: - source_matches = all(matches_pattern(edge.source_label, f) for f in subject_filters) - if not source_matches: + source_label = edge.source_label + if source_label not in subject_matches_by_label: + subject_matches_by_label[source_label] = all( + matches_pattern(source_label, filter_) for filter_ in subject_filters + ) + if not subject_matches_by_label[source_label]: continue - target_matches = all(matches_pattern(edge.target_label, f) for f in object_filters) + target_label = edge.target_label + if target_label not in target_matches_by_label: + target_matches_by_label[target_label] = all( + matches_pattern(target_label, filter_) for filter_ in object_filters + ) + target_matches = target_matches_by_label[target_label] if is_negated: # should_not(): violation if dependency EXISTS diff --git a/src/archunitpython/files/fluentapi/files.py b/src/archunitpython/files/fluentapi/files.py index 422ce32..7222ad3 100644 --- a/src/archunitpython/files/fluentapi/files.py +++ b/src/archunitpython/files/fluentapi/files.py @@ -13,7 +13,7 @@ from collections.abc import Sequence from archunitpython.common.assertion.violation import EmptyTestViolation, Violation -from archunitpython.common.extraction.extract_graph import extract_graph +from archunitpython.common.extraction.extract_graph import extract_graph, extract_graph_for_sources from archunitpython.common.fluentapi.checkable import CheckOptions, RuleRationaleMixin from archunitpython.common.pattern_matching import matches_all_patterns from archunitpython.common.projection.edge_projections import ( @@ -366,7 +366,9 @@ def __init__( self._is_negated = is_negated def check(self, options: CheckOptions | None = None) -> list[Violation]: - graph = extract_graph(self._project_path, options=options) + graph = extract_graph_for_sources( + self._project_path, self._subject_filters, options=options + ) edges = project_edges(graph, per_internal_edge()) return gather_depend_on_file_violations( diff --git a/tests/common/test_core_types.py b/tests/common/test_core_types.py index 32fb653..b15b95f 100644 --- a/tests/common/test_core_types.py +++ b/tests/common/test_core_types.py @@ -83,6 +83,19 @@ def test_filter_creation(self): assert f.regexp.pattern == r".*\.py$" assert f.options.target == "filename" + def test_search_regex_does_not_change_filter_identity(self): + regexp = re.compile(r".*\.py$") + options = PatternMatchingOptions(target="filename") + ordinary = Filter(regexp=regexp, options=options) + optimized = Filter( + regexp=regexp, + options=options, + search_regexp=re.compile(r"\.py$"), + ) + assert optimized == ordinary + assert hash(optimized) == hash(ordinary) + assert repr(optimized) == repr(ordinary) + def test_filter_frozen(self): f = Filter( regexp=re.compile(r"test"), diff --git a/tests/common/test_extract_graph.py b/tests/common/test_extract_graph.py index f0b25cb..bf926d6 100644 --- a/tests/common/test_extract_graph.py +++ b/tests/common/test_extract_graph.py @@ -1,5 +1,6 @@ """Tests for graph extraction.""" +import importlib import os import shutil from pathlib import Path @@ -14,9 +15,12 @@ _resolve_exclude_patterns, clear_graph_cache, extract_graph, + extract_graph_for_sources, ) from archunitpython.common.extraction.graph import Edge, ImportKind from archunitpython.common.fluentapi.checkable import CheckOptions +from archunitpython.common.pattern_matching import matches_all_patterns +from archunitpython.common.regex_factory import RegexFactory FIXTURES_DIR = os.path.join(os.path.dirname(os.path.dirname(__file__)), "fixtures") SAMPLE_PROJECT = os.path.join(FIXTURES_DIR, "sample_project") @@ -185,6 +189,146 @@ def test_cache_clear(self): graph2 = extract_graph(SAMPLE_PROJECT, options=CheckOptions(clear_cache=True)) assert graph1 is not graph2 # Different objects after cache clear + def test_partial_then_full_graph_preserves_all_edges(self): + service_filter = [RegexFactory.folder_matcher("**/services*")] + partial = extract_graph_for_sources(SAMPLE_PROJECT, service_filter) + full = extract_graph(SAMPLE_PROJECT) + + assert partial == [ + edge for edge in full if matches_all_patterns(edge.source, service_filter) + ] + assert any("/models/" in edge.source for edge in full) + + def test_full_then_partial_graph_preserves_selected_edges(self): + full = extract_graph(SAMPLE_PROJECT) + service_filter = [RegexFactory.folder_matcher("**/services*")] + partial = extract_graph_for_sources(SAMPLE_PROJECT, service_filter) + + assert partial == [ + edge for edge in full if matches_all_patterns(edge.source, service_filter) + ] + + def test_multiple_source_filters_select_only_matching_files(self): + filters = [ + RegexFactory.folder_matcher("**/services*"), + RegexFactory.filename_matcher("service_b.py"), + ] + partial = extract_graph_for_sources(SAMPLE_PROJECT, filters) + full = extract_graph(SAMPLE_PROJECT) + + assert partial == [edge for edge in full if matches_all_patterns(edge.source, filters)] + assert partial + assert {Path(edge.source).name for edge in partial} == {"service_b.py"} + + def test_no_source_filters_reuses_full_graph(self): + full = extract_graph(SAMPLE_PROJECT) + assert extract_graph_for_sources(SAMPLE_PROJECT, []) is full + + def test_partial_graph_refreshes_after_cache_clear(self, tmp_path): + api = tmp_path / "api" + api.mkdir() + source = api / "endpoint.py" + source.write_text("import infrastructure\n", encoding="utf-8") + selected = [RegexFactory.folder_matcher("**/api*")] + + before = extract_graph_for_sources(str(tmp_path), selected) + assert any(edge.external and edge.target == "infrastructure" for edge in before) + + infrastructure = tmp_path / "infrastructure.py" + infrastructure.write_text("VALUE = 1\n", encoding="utf-8") + after = extract_graph_for_sources( + str(tmp_path), selected, options=CheckOptions(clear_cache=True) + ) + assert any( + not edge.external and edge.target == _normalize(str(infrastructure)) for edge in after + ) + + def test_selective_graph_preserves_namespace_and_import_kinds(self, tmp_path): + api = tmp_path / "api" + domain = tmp_path / "domain" + api.mkdir() + domain.mkdir() + (domain / "model.py").write_text("VALUE = 1\n", encoding="utf-8") + (domain / "types.py").write_text("VALUE = 2\n", encoding="utf-8") + (domain / "optional.py").write_text("VALUE = 3\n", encoding="utf-8") + (api / "consumer.py").write_text( + "\n".join( + [ + "from typing import TYPE_CHECKING", + "from domain import model", + "if TYPE_CHECKING:", + " from domain import types", + "try:", + " from domain import optional", + "except ImportError:", + " pass", + "import importlib", + 'importlib.import_module("domain.model")', + "", + ] + ), + encoding="utf-8", + ) + selected = [RegexFactory.folder_matcher("**/api*")] + + partial = extract_graph_for_sources(str(tmp_path), selected) + full = extract_graph(str(tmp_path)) + assert partial == [edge for edge in full if matches_all_patterns(edge.source, selected)] + assert any( + edge.target.endswith("/domain/model.py") and not edge.external for edge in partial + ) + assert any(edge.target.endswith("/domain/types.py") for edge in partial) + assert any(edge.target.endswith("/domain/optional.py") for edge in partial) + model_edges = [edge for edge in partial if edge.target.endswith("/domain/model.py")] + assert len(model_edges) == 1 + assert set(model_edges[0].import_kinds) == { + ImportKind.FROM_IMPORT, + ImportKind.DYNAMIC_IMPORT, + } + optional_edges = [ + edge for edge in partial if edge.target.endswith("/domain/optional.py") + ] + assert ImportKind.CONDITIONAL_IMPORT in optional_edges[0].import_kinds + + ignore_types = CheckOptions(ignore_type_checking_imports=True) + partial_without_types = extract_graph_for_sources( + str(tmp_path), selected, options=ignore_types + ) + full_without_types = extract_graph(str(tmp_path), options=ignore_types) + assert partial_without_types == [ + edge for edge in full_without_types if matches_all_patterns(edge.source, selected) + ] + assert not any(edge.target.endswith("/domain/types.py") for edge in partial_without_types) + + def test_no_selected_sources_do_not_parse_imports(self, monkeypatch): + def unexpected_parse(_path): + raise AssertionError("Unselected files must not be parsed") + + extraction = importlib.import_module("archunitpython.common.extraction.extract_graph") + monkeypatch.setattr(extraction, "_extract_located_imports", unexpected_parse) + graph = extract_graph_for_sources( + SAMPLE_PROJECT, [RegexFactory.folder_matcher("**/not_present*")] + ) + assert graph == [] + + def test_selective_graph_respects_archignore_for_targets(self, tmp_path): + api = tmp_path / "api" + domain = tmp_path / "domain" + api.mkdir() + domain.mkdir() + (api / "consumer.py").write_text("from domain import model\n", encoding="utf-8") + (domain / "model.py").write_text("VALUE = 1\n", encoding="utf-8") + (tmp_path / ".archignore").write_text("domain/model.py\n", encoding="utf-8") + selected = [RegexFactory.folder_matcher("**/api*")] + + partial = extract_graph_for_sources(str(tmp_path), selected) + full = extract_graph(str(tmp_path)) + assert partial == [edge for edge in full if matches_all_patterns(edge.source, selected)] + assert not any( + edge.target == _normalize(str(domain / "model.py")) and not edge.external + for edge in partial + ) + def test_edge_has_import_kinds(self): graph = extract_graph(SAMPLE_PROJECT) edges_with_kinds = [e for e in graph if len(e.import_kinds) > 0] @@ -227,12 +371,24 @@ def test_archignore_excludes_files_and_directories(self): excludes = _resolve_exclude_patterns(str(self._temp_dir), ["__pycache__"]) files = _find_python_files(str(self._temp_dir), excludes) relative_files = { - Path(file_path).relative_to(self._temp_dir).as_posix() - for file_path in files + Path(file_path).relative_to(self._temp_dir).as_posix() for file_path in files } assert relative_files == {"keep.py"} + def test_explicit_file_excludes_survive_default_directory_shortcut(self): + self._write("keep.py") + self._write("manual.py") + self._write("__pycache__/cached.py") + + files = _find_python_files( + str(self._temp_dir), ["__pycache__", "manual.py"] + ) + relative_files = { + Path(file_path).relative_to(self._temp_dir).as_posix() for file_path in files + } + assert relative_files == {"keep.py"} + def test_archignore_ignored_files_are_not_dependency_targets(self): self._write(".archignore", "ignored.py\n") self._write("keep.py", "import ignored\n") @@ -396,9 +552,7 @@ def test_dynamic_import_resolves_to_internal_edge(self): ).replace("\\", "/") edges = [ - edge - for edge in graph - if edge.source == loader_path and edge.target == models_path + edge for edge in graph if edge.source == loader_path and edge.target == models_path ] assert len(edges) == 1 assert ImportKind.DYNAMIC_IMPORT in edges[0].import_kinds @@ -446,9 +600,7 @@ def _service_to_model_edges(self, project_root: str) -> list[Edge]: service_path = os.path.abspath( os.path.join(project_root, "namespace_pkg", "services", "service.py") ).replace("\\", "/") - return [ - edge for edge in graph if edge.source == service_path and edge.target == model_path - ] + return [edge for edge in graph if edge.source == service_path and edge.target == model_path] def _service_edges(self, project_root: str) -> list[Edge]: graph = extract_graph(project_root) @@ -456,15 +608,11 @@ def _service_edges(self, project_root: str) -> list[Edge]: os.path.join(project_root, "namespace_pkg", "services", "service.py") ).replace("\\", "/") return [ - edge - for edge in graph - if edge.source == service_path and edge.target != service_path + edge for edge in graph if edge.source == service_path and edge.target != service_path ] def test_absolute_from_import_resolves_namespace_package_submodule(self): - project_root = self._build_namespace_project( - "from namespace_pkg.domain import model\n" - ) + project_root = self._build_namespace_project("from namespace_pkg.domain import model\n") edges = self._service_to_model_edges(project_root) @@ -493,8 +641,7 @@ def test_mixed_aliases_preserve_internal_and_external_edges(self): assert any(edge.target == model_path and not edge.external for edge in edges) assert any( - edge.target == "namespace_pkg.domain.remote_model" and edge.external - for edge in edges + edge.target == "namespace_pkg.domain.remote_model" and edge.external for edge in edges ) def test_multiple_internal_aliases_resolve_to_each_submodule(self): @@ -516,9 +663,7 @@ def test_multiple_internal_aliases_resolve_to_each_submodule(self): assert internal_targets == expected_targets def test_all_external_aliases_keep_original_base_edge(self): - project_root = self._build_namespace_project( - "from vendor_sdk import Client, Config\n" - ) + project_root = self._build_namespace_project("from vendor_sdk import Client, Config\n") external_targets = { edge.target for edge in self._service_edges(project_root) if edge.external @@ -534,9 +679,7 @@ def test_as_alias_uses_original_submodule_name(self): assert len(self._service_to_model_edges(project_root)) == 1 def test_relative_mixed_aliases_preserve_unresolved_edge(self): - project_root = self._build_namespace_project( - "from ..domain import model, remote_model\n" - ) + project_root = self._build_namespace_project("from ..domain import model, remote_model\n") edges = self._service_edges(project_root) unresolved_path = os.path.abspath( @@ -544,14 +687,10 @@ def test_relative_mixed_aliases_preserve_unresolved_edge(self): ).replace("\\", "/") assert len(self._service_to_model_edges(project_root)) == 1 - assert any( - edge.target == unresolved_path and edge.external for edge in edges - ) + assert any(edge.target == unresolved_path and edge.external for edge in edges) def test_archignore_suppresses_resolved_namespace_target(self): - project_root = self._build_namespace_project( - "from namespace_pkg.domain import model\n" - ) + project_root = self._build_namespace_project("from namespace_pkg.domain import model\n") Path(project_root, ".archignore").write_text( "namespace_pkg/domain/model.py\n", encoding="utf-8", @@ -625,18 +764,16 @@ def test_import_error_fallback_imports_are_marked_conditional(self): os.path.join(project_root, "sample_project", "service.py") ).replace("\\", "/") target_paths = { - os.path.abspath( - os.path.join(project_root, "sample_project", "fast_model.py") - ).replace("\\", "/"), + os.path.abspath(os.path.join(project_root, "sample_project", "fast_model.py")).replace( + "\\", "/" + ), os.path.abspath( os.path.join(project_root, "sample_project", "fallback_model.py") ).replace("\\", "/"), } edges = [ - edge - for edge in graph - if edge.source == service_path and edge.target in target_paths + edge for edge in graph if edge.source == service_path and edge.target in target_paths ] assert len(edges) == 2 @@ -672,9 +809,7 @@ def test_relative_fallback_imports_resolve_sibling_modules(self): edges = self._conditional_edges(project_root) target_paths = { - os.path.abspath( - os.path.join(project_root, "sample_project", module) - ).replace("\\", "/") + os.path.abspath(os.path.join(project_root, "sample_project", module)).replace("\\", "/") for module in ("fast_model.py", "fallback_model.py") } @@ -698,9 +833,7 @@ def test_conditional_relative_import_resolves_multiple_aliased_modules(self): edges = self._conditional_edges(project_root) target_paths = { - os.path.abspath( - os.path.join(project_root, "sample_project", module) - ).replace("\\", "/") + os.path.abspath(os.path.join(project_root, "sample_project", module)).replace("\\", "/") for module in ("fast_model.py", "fallback_model.py") } @@ -751,9 +884,7 @@ def test_conditional_relative_import_resolves_module_and_package_attribute(self) edges = self._conditional_edges(project_root) target_paths = { - os.path.abspath( - os.path.join(project_root, "sample_project", target) - ).replace("\\", "/") + os.path.abspath(os.path.join(project_root, "sample_project", target)).replace("\\", "/") for target in ("fast_model.py", "__init__.py") } @@ -779,9 +910,7 @@ def test_parent_relative_fallback_imports_resolve_modules(self): edges = self._conditional_edges(project_root) target_paths = { - os.path.abspath( - os.path.join(project_root, "sample_project", module) - ).replace("\\", "/") + os.path.abspath(os.path.join(project_root, "sample_project", module)).replace("\\", "/") for module in ("fast_model.py", "fallback_model.py") } @@ -842,9 +971,7 @@ def test_non_import_error_handler_does_not_mark_imports_conditional(self): assert len(edges) == 2 assert all(ImportKind.FROM_IMPORT in edge.import_kinds for edge in edges) - assert all( - ImportKind.CONDITIONAL_IMPORT not in edge.import_kinds for edge in edges - ) + assert all(ImportKind.CONDITIONAL_IMPORT not in edge.import_kinds for edge in edges) def test_conditional_dynamic_import_retains_dynamic_kind(self): project_root = self._build_conditional_project( @@ -866,10 +993,7 @@ def test_conditional_dynamic_import_retains_dynamic_kind(self): ] assert len(internal_edges) == 2 - assert all( - ImportKind.CONDITIONAL_IMPORT in edge.import_kinds - for edge in internal_edges - ) + assert all(ImportKind.CONDITIONAL_IMPORT in edge.import_kinds for edge in internal_edges) assert all(ImportKind.DYNAMIC_IMPORT in edge.import_kinds for edge in internal_edges) def test_type_checking_import_takes_precedence_over_conditional_context(self): @@ -903,8 +1027,7 @@ def test_type_checking_import_takes_precedence_over_conditional_context(self): ), ) assert not any( - edge.source.endswith("/service.py") - and edge.target.endswith("/fast_model.py") + edge.source.endswith("/service.py") and edge.target.endswith("/fast_model.py") for edge in ignored_graph ) @@ -945,9 +1068,7 @@ def _service_to_model_edges(self, project_root: str) -> list[Edge]: os.path.join(project_root, "sample_project", "service.py") ).replace("\\", "/") return [ - edge - for edge in graph - if edge.source == service_path and edge.target == models_path + edge for edge in graph if edge.source == service_path and edge.target == models_path ] def test_inline_ignore_directive_removes_import_edge(self): diff --git a/tests/common/test_pattern_matching.py b/tests/common/test_pattern_matching.py index e111ab3..e4bb637 100644 --- a/tests/common/test_pattern_matching.py +++ b/tests/common/test_pattern_matching.py @@ -1,5 +1,6 @@ """Tests for pattern matching and regex factory.""" +import fnmatch import re from archunitpython.common.pattern_matching import ( @@ -57,6 +58,36 @@ def test_deep_path(self): class TestRegexFactory: + def test_leading_star_globs_keep_fnmatch_search_semantics(self): + patterns = ("*", "**", "**/api*", "***.py", "*?foo", "**/a*b", "*/x?") + paths = ( + "", + "api", + "api/file.py", + "/api/file.py", + "src/api/file.py", + "src/application/file.py", + "src/a/x/b", + "src/a\nx/b", + "foo", + "xfoo", + "a/b", + ) + for pattern in patterns: + original = re.compile(fnmatch.translate(pattern)) + filter_ = RegexFactory.path_matcher(pattern) + optimized = filter_.search_regexp + assert optimized is not None + for path in paths: + assert bool(filter_.regexp.match(path)) == bool(original.match(path)), ( + pattern, + path, + ) + assert bool(optimized.search(path)) == bool(original.search(path)), ( + pattern, + path, + ) + def test_filename_matcher_glob(self): f = RegexFactory.filename_matcher("*.py") assert f.options.target == "filename" diff --git a/tests/files/test_file_assertions.py b/tests/files/test_file_assertions.py index 8c1235d..c741d8e 100644 --- a/tests/files/test_file_assertions.py +++ b/tests/files/test_file_assertions.py @@ -87,6 +87,34 @@ def test_non_matching_subject_skipped(self): violations = gather_depend_on_file_violations(edges, subject, obj, is_negated=True) assert len(violations) == 0 + def test_repeated_labels_preserve_rule_semantics_and_order(self): + edges = [ + _edge("src/ui/view.py", "src/db/database.py"), + _edge("src/ui/view.py", "src/services/service.py"), + _edge("src/ui/other.py", "src/db/database.py"), + _edge("src/other/file.py", "src/db/database.py"), + ] + subject = [ + RegexFactory.folder_matcher("src/ui*"), + RegexFactory.filename_matcher("*.py"), + ] + obj = [RegexFactory.folder_matcher("src/db*")] + + forbidden = gather_depend_on_file_violations(edges, subject, obj, is_negated=True) + required = gather_depend_on_file_violations(edges, subject, obj, is_negated=False) + + assert [violation.dependency for violation in forbidden] == [edges[0], edges[2]] + assert [violation.dependency for violation in required] == [edges[1]] + + def test_match_results_do_not_leak_to_later_checks(self): + edges = [_edge("src/ui/view.py", "src/db/database.py")] + subject = [RegexFactory.folder_matcher("src/ui*")] + db_target = [RegexFactory.folder_matcher("src/db*")] + service_target = [RegexFactory.folder_matcher("src/services*")] + + assert len(gather_depend_on_file_violations(edges, subject, db_target, True)) == 1 + assert gather_depend_on_file_violations(edges, subject, service_target, True) == [] + class TestCycleFree: def test_no_cycles_no_violations(self):