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
31 changes: 25 additions & 6 deletions src/archunitpython/common/extraction/extract_graph.py
Original file line number Diff line number Diff line change
Expand Up @@ -80,6 +80,7 @@ class _ExtractionSession:
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)
absolute_resolutions: dict[str, tuple[str, bool]] = field(default_factory=dict)


_extraction_sessions: dict[GraphCacheKey, _ExtractionSession] = {}
Expand Down Expand Up @@ -305,7 +306,10 @@ def _extract_file_edges(
):
continue
for resolved, is_external in _resolve_import_targets(
located_import, file_path, session.project_path
located_import,
file_path,
session.project_path,
absolute_resolutions=session.absolute_resolutions,
):
if resolved and resolved != source_label:
if not is_external and resolved not in session.normalized_py_file_set:
Expand Down Expand Up @@ -680,6 +684,7 @@ def _resolve_import(
source_file: str,
project_root: str,
kind: ImportKind,
absolute_resolutions: dict[str, tuple[str, bool]] | None = None,
) -> tuple[str, bool]:
"""Resolve an import name to an absolute file path.

Expand All @@ -693,14 +698,23 @@ def _resolve_import(
# Relative import
return _resolve_relative_import(import_name, source_file, project_root)

# Absolute import: try to resolve within the project
return _resolve_absolute_import(import_name, project_root)
# Absolute targets are independent of the source file within this session.
if absolute_resolutions is not None:
cached = absolute_resolutions.get(import_name)
if cached is not None:
return cached
resolved = _resolve_absolute_import(import_name, project_root)
if absolute_resolutions is not None:
absolute_resolutions[import_name] = resolved
return resolved


def _resolve_import_targets(
import_: _LocatedImport,
source_file: str,
project_root: str,
*,
absolute_resolutions: dict[str, tuple[str, bool]] | None = None,
) -> list[tuple[str, bool]]:
"""Resolve an import, including namespace-package submodule aliases."""
resolution_kind = import_.resolution_kind or import_.import_kind
Expand All @@ -709,13 +723,15 @@ def _resolve_import_targets(
source_file,
project_root,
resolution_kind,
absolute_resolutions,
)
if is_external and import_.fallback_module_name is not None:
fallback, fallback_is_external = _resolve_import(
import_.fallback_module_name,
source_file,
project_root,
resolution_kind,
absolute_resolutions,
)
if fallback and not fallback_is_external:
resolved, is_external = fallback, False
Expand All @@ -731,6 +747,7 @@ def _resolve_import_targets(
source_file,
project_root,
resolution_kind,
absolute_resolutions,
)
alias_targets.append((alias_resolved, alias_is_external))
found_internal_alias = found_internal_alias or not alias_is_external
Expand Down Expand Up @@ -811,15 +828,17 @@ def _resolve_absolute_import(
candidate_file = candidate_base + ".py"
if os.path.isfile(candidate_file):
resolved = _normalize(os.path.abspath(candidate_file))
# Only count as internal if it's inside the project root
is_internal = resolved.startswith(_normalize(os.path.abspath(project_root)))
# A sibling such as /repo/app2 must not match /repo/app.
project_prefix = _normalize(os.path.abspath(project_root)).rstrip("/") + "/"
is_internal = resolved.startswith(project_prefix)
return resolved, not is_internal

# Try as a package
candidate_init = os.path.join(candidate_base, "__init__.py")
if os.path.isfile(candidate_init):
resolved = _normalize(os.path.abspath(candidate_init))
is_internal = resolved.startswith(_normalize(os.path.abspath(project_root)))
project_prefix = _normalize(os.path.abspath(project_root)).rstrip("/") + "/"
is_internal = resolved.startswith(project_prefix)
return resolved, not is_internal

# Not found in project → external
Expand Down
99 changes: 96 additions & 3 deletions tests/common/test_extract_graph.py
Original file line number Diff line number Diff line change
Expand Up @@ -64,6 +64,97 @@ class FutureNode(ast.AST):
assert [type(node) for node in _import_analysis_nodes(tree)] == [ast.Import, ast.Call]


@pytest.mark.parametrize("sibling_kind", ["module", "package"])
def test_sibling_with_shared_path_prefix_remains_external(tmp_path, sibling_kind):
project = tmp_path / "app"
project.mkdir()
source = project / "service.py"
source.write_text("import app2\n", encoding="utf-8")
if sibling_kind == "module":
sibling_target = tmp_path / "app2.py"
else:
sibling = tmp_path / "app2"
sibling.mkdir()
sibling_target = sibling / "__init__.py"
sibling_target.write_text("", encoding="utf-8")

graph = extract_graph(str(project), options=CheckOptions(clear_cache=True))

assert any(
edge.source == _normalize(str(source))
and edge.target == _normalize(str(sibling_target))
and edge.external
for edge in graph
)


def test_repeated_absolute_import_reuses_resolution_until_cache_clear(tmp_path, monkeypatch):
extraction = importlib.import_module("archunitpython.common.extraction.extract_graph")
target = tmp_path / "dependency.py"
target.write_text("", encoding="utf-8")
for source_name in ("a.py", "b.py"):
(tmp_path / source_name).write_text("import dependency\n", encoding="utf-8")

original_isfile = os.path.isfile
target_probes = 0

def counted_isfile(path):
nonlocal target_probes
if _normalize(str(path)) == _normalize(str(target)):
target_probes += 1
return original_isfile(path)

monkeypatch.setattr(extraction.os.path, "isfile", counted_isfile)
graph = extract_graph(str(tmp_path), options=CheckOptions(clear_cache=True))
dependency_edges = [
edge
for edge in graph
if edge.target == _normalize(str(target)) and edge.source != edge.target
]
assert len(dependency_edges) == 2
assert all(not edge.external for edge in dependency_edges)
assert target_probes == 1

extract_graph(str(tmp_path), options=CheckOptions(clear_cache=True))
assert target_probes == 2


def test_absolute_resolution_cache_refreshes_when_target_appears(tmp_path):
source = tmp_path / "service.py"
source.write_text("import dependency\n", encoding="utf-8")

before = extract_graph(str(tmp_path), options=CheckOptions(clear_cache=True))
assert any(edge.target == "dependency" and edge.external for edge in before)

target = tmp_path / "dependency.py"
target.write_text("", encoding="utf-8")
after = extract_graph(str(tmp_path), options=CheckOptions(clear_cache=True))
assert any(
edge.target == _normalize(str(target)) and not edge.external for edge in after
)


def test_relative_resolution_is_source_specific_with_session_cache(tmp_path):
expected_targets = set()
for package_name in ("first", "second"):
package = tmp_path / package_name
package.mkdir()
(package / "service.py").write_text("from . import dependency\n", encoding="utf-8")
target = package / "dependency.py"
target.write_text("", encoding="utf-8")
expected_targets.add(_normalize(str(target)))

graph = extract_graph(str(tmp_path), options=CheckOptions(clear_cache=True))
actual_targets = {
edge.target
for edge in graph
if edge.source.endswith("/service.py")
and edge.source != edge.target
and not edge.external
}
assert actual_targets == expected_targets


def test_parallel_extraction_preserves_graph_and_calling_thread_resolution(monkeypatch):
extraction = importlib.import_module("archunitpython.common.extraction.extract_graph")
clear_graph_cache()
Expand All @@ -81,9 +172,9 @@ def extract(path):
worker_threads.add(get_ident())
return original_extract(path)

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

monkeypatch.setattr(extraction, "_extract_located_imports", extract)
monkeypatch.setattr(extraction, "_resolve_import_targets", resolve)
Expand Down Expand Up @@ -147,7 +238,9 @@ def extract(path):

monkeypatch.setattr(extraction, "_extract_located_imports", extract)
monkeypatch.setattr(
extraction, "_resolve_import_targets", lambda import_, *_: [(import_.module_name, True)]
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] == [
Expand Down
Loading