|
2 | 2 |
|
3 | 3 | from pathlib import Path |
4 | 4 |
|
| 5 | +from packaging.markers import default_environment |
| 6 | +from packaging.requirements import Requirement |
| 7 | + |
5 | 8 | try: |
6 | 9 | import tomllib |
7 | 10 | except ModuleNotFoundError: # pragma: no cover - Python 3.10 fallback |
8 | 11 | import tomli as tomllib # type: ignore[no-redef] |
9 | 12 |
|
10 | 13 |
|
11 | 14 | ROOT = Path(__file__).resolve().parents[1] |
| 15 | +ALL_EXTRA = "all" |
| 16 | +HEADROOM_PACKAGE_NAME = "headroom-ai" |
12 | 17 | MACOS_X86_64_TORCH_GUARD = "sys_platform != 'darwin' or platform_machine != 'x86_64'" |
| 18 | +MACOS_X86_64_SYS_PLATFORM = "darwin" |
| 19 | +MACOS_X86_64_PLATFORM_MACHINE = "x86_64" |
| 20 | +PYPROJECT_FILE = "pyproject.toml" |
| 21 | +SYS_PLATFORM_MARKER = "sys_platform" |
| 22 | +PLATFORM_MACHINE_MARKER = "platform_machine" |
| 23 | +TORCH_PACKAGE_NAME = "torch" |
| 24 | +TORCH_TRANSITIVE_PACKAGE_NAMES = frozenset({"sentence-transformers"}) |
| 25 | +UV_LOCK_FILE = "uv.lock" |
| 26 | + |
| 27 | + |
| 28 | +def _selected_dependency_names_for_extra( |
| 29 | + optional_deps: dict[str, list[str]], |
| 30 | + extra_name: str, |
| 31 | + environment: dict[str, str], |
| 32 | + visited: set[str] | None = None, |
| 33 | +) -> set[str]: |
| 34 | + selected: set[str] = set() |
| 35 | + visited = visited or set() |
| 36 | + if extra_name in visited: |
| 37 | + return selected |
| 38 | + visited.add(extra_name) |
| 39 | + |
| 40 | + for dependency in optional_deps[extra_name]: |
| 41 | + requirement = Requirement(dependency) |
| 42 | + if requirement.marker is not None and not requirement.marker.evaluate(environment): |
| 43 | + continue |
| 44 | + if requirement.name == HEADROOM_PACKAGE_NAME: |
| 45 | + for nested_extra in requirement.extras: |
| 46 | + selected.update( |
| 47 | + _selected_dependency_names_for_extra( |
| 48 | + optional_deps, |
| 49 | + nested_extra, |
| 50 | + environment, |
| 51 | + visited, |
| 52 | + ) |
| 53 | + ) |
| 54 | + else: |
| 55 | + selected.add(requirement.name) |
| 56 | + |
| 57 | + return selected |
| 58 | + |
| 59 | + |
| 60 | +def _locked_dependency_names(package_name: str) -> set[str]: |
| 61 | + lock = tomllib.loads((ROOT / UV_LOCK_FILE).read_text(encoding="utf-8")) |
| 62 | + for package in lock["package"]: |
| 63 | + if package["name"] == package_name: |
| 64 | + return {dependency["name"] for dependency in package.get("dependencies", [])} |
| 65 | + raise AssertionError(f"{package_name} not found in {UV_LOCK_FILE}") |
13 | 66 |
|
14 | 67 |
|
15 | 68 | def test_all_extra_does_not_require_torch_on_macos_x86_64() -> None: |
16 | 69 | """Keep `headroom-ai[all]` resolvable where PyTorch publishes no wheel.""" |
17 | 70 |
|
18 | | - pyproject = tomllib.loads((ROOT / "pyproject.toml").read_text(encoding="utf-8")) |
| 71 | + pyproject = tomllib.loads((ROOT / PYPROJECT_FILE).read_text(encoding="utf-8")) |
19 | 72 | optional_deps = pyproject["project"]["optional-dependencies"] |
| 73 | + macos_x86_64_environment = default_environment() |
| 74 | + macos_x86_64_environment.update( |
| 75 | + { |
| 76 | + SYS_PLATFORM_MARKER: MACOS_X86_64_SYS_PLATFORM, |
| 77 | + PLATFORM_MACHINE_MARKER: MACOS_X86_64_PLATFORM_MACHINE, |
| 78 | + } |
| 79 | + ) |
20 | 80 |
|
21 | | - assert "ml" in optional_deps["all"][0] |
22 | | - assert "voice" in optional_deps["all"][0] |
| 81 | + assert "ml" in optional_deps[ALL_EXTRA][0] |
| 82 | + assert "voice" in optional_deps[ALL_EXTRA][0] |
23 | 83 |
|
24 | 84 | torch_deps = [ |
25 | 85 | dep |
26 | 86 | for extra_name in ("ml", "voice") |
27 | 87 | for dep in optional_deps[extra_name] |
28 | | - if dep.startswith("torch") |
| 88 | + if dep.startswith(TORCH_PACKAGE_NAME) |
29 | 89 | ] |
| 90 | + selected_all_dependency_names = _selected_dependency_names_for_extra( |
| 91 | + optional_deps, |
| 92 | + ALL_EXTRA, |
| 93 | + macos_x86_64_environment, |
| 94 | + ) |
| 95 | + locked_torch_transitive_dependency_names = { |
| 96 | + package_name |
| 97 | + for package_name in TORCH_TRANSITIVE_PACKAGE_NAMES |
| 98 | + if TORCH_PACKAGE_NAME in _locked_dependency_names(package_name) |
| 99 | + } |
30 | 100 |
|
31 | 101 | assert torch_deps |
32 | 102 | assert all(MACOS_X86_64_TORCH_GUARD in dep for dep in torch_deps) |
| 103 | + assert locked_torch_transitive_dependency_names |
| 104 | + assert TORCH_PACKAGE_NAME not in selected_all_dependency_names |
| 105 | + assert selected_all_dependency_names.isdisjoint(locked_torch_transitive_dependency_names) |
0 commit comments