diff --git a/kernels/src/kernels/importer.py b/kernels/src/kernels/importer.py index f53c13ea..ccafea8a 100644 --- a/kernels/src/kernels/importer.py +++ b/kernels/src/kernels/importer.py @@ -1,6 +1,5 @@ import importlib import sys -import warnings from dataclasses import dataclass from pathlib import Path from types import ModuleType @@ -66,33 +65,6 @@ def get_loaded_kernels() -> list[LoadedKernel]: return list(_loaded_kernels.values()) -def _warn_if_dirty(metadata: Metadata, variant_str: str) -> None: - """Warn when a kernel variant was built from a dirty git tree. - - A dirty build was produced from a working tree with uncommitted changes, - so its git SHA does not fully identify the sources it was built from and - the build may not be reproducible. - """ - provenance = metadata.provenance - if provenance is None or not provenance.dirty: - return - - dirty_sources = [] - if provenance.kernel is not None and provenance.kernel.dirty: - dirty_sources.append("kernel source") - builder_git = provenance.kernel_builder.git - if builder_git is not None and builder_git.dirty: - dirty_sources.append("kernel-builder") - - warnings.warn( - f"Kernel '{metadata.name}' variant '{variant_str}' was built from a dirty " - f"git tree ({', '.join(dirty_sources)} had uncommitted changes). Its " - "recorded git revision does not fully identify the sources it was built " - "from, so the build may not be reproducible.", - stacklevel=3, - ) - - def _import_from_path( variant_path: Path, deps: dict[str, ModuleType], @@ -102,7 +74,6 @@ def _import_from_path( return loaded_kernel.module metadata = Metadata.read_from_file(variant_path / "metadata.json") - _warn_if_dirty(metadata, variant_path.name) module_name = metadata.name.python_name file_path = variant_path / "__init__.py" diff --git a/kernels/src/kernels/validate.py b/kernels/src/kernels/validate.py index 4d2f4f05..9df7eba9 100644 --- a/kernels/src/kernels/validate.py +++ b/kernels/src/kernels/validate.py @@ -1,3 +1,4 @@ +import warnings from dataclasses import dataclass from typing import Protocol @@ -33,6 +34,30 @@ def validate_metadata(self, *, metadata: Metadata, variant: str) -> None: _check_arch_incompatibility(metadata, variant) +class DirtyValidator: + """Warn when a kernel variant was built from a dirty git tree.""" + + def validate_metadata(self, *, metadata: Metadata, variant: str) -> None: + provenance = metadata.provenance + if provenance is None or not provenance.dirty: + return + + dirty_sources = [] + if provenance.kernel is not None and provenance.kernel.dirty: + dirty_sources.append("kernel source") + builder_git = provenance.kernel_builder.git + if builder_git is not None and builder_git.dirty: + dirty_sources.append("kernel-builder") + + warnings.warn( + f"Kernel '{metadata.name}' variant '{variant}' was built from a dirty " + f"git tree ({', '.join(dirty_sources)} had uncommitted changes). Its " + "recorded git revision does not fully identify the sources it was built " + "from, so the build may not be reproducible.", + stacklevel=3, + ) + + def _installed_version() -> Version | None: """The installed `kernels` version as a numeric version. @@ -96,4 +121,4 @@ def validate_metadata(self, *, metadata: Metadata, variant: str) -> None: def default_metadata_validators() -> list[MetadataValidator]: """The metadata validators that are applied to every kernel dependency tree.""" - return [DependencyValidator(), MinverValidator()] + return [DependencyValidator(), MinverValidator(), DirtyValidator()] diff --git a/kernels/tests/test_dirty_provenance.py b/kernels/tests/test_dirty_provenance.py deleted file mode 100644 index a5bf17aa..00000000 --- a/kernels/tests/test_dirty_provenance.py +++ /dev/null @@ -1,71 +0,0 @@ -import json - -import pytest -from kernels_data import Metadata - -from kernels.importer import _import_from_path, _loaded_kernels, _warn_if_dirty - - -def _write_variant(tmp_path, provenance): - variant_dir = tmp_path / "build" / "torch28-cxx11-cu128-x86_64-linux" - variant_dir.mkdir(parents=True) - metadata = { - "id": "activation_1_cuda", - "name": "activation", - "version": 1, - "license": "Apache-2.0", - "python-depends": ["torch"], - "backend": {"type": "cuda"}, - } - if provenance is not None: - metadata["provenance"] = provenance - (variant_dir / "metadata.json").write_text(json.dumps(metadata)) - return variant_dir - - -CLEAN_PROVENANCE = { - "kernel-builder": {"version": "0.1.0", "commit": "a" * 40, "dirty": False}, - "kernel": {"commit": "b" * 40, "dirty": False}, -} -DIRTY_KERNEL = { - "kernel-builder": {"version": "0.1.0", "commit": "a" * 40, "dirty": False}, - "kernel": {"commit": "b" * 40, "dirty": True}, -} -DIRTY_BUILDER = { - "kernel-builder": {"version": "0.1.0", "commit": "a" * 40, "dirty": True}, - "kernel": {"commit": "b" * 40, "dirty": False}, -} - - -@pytest.mark.parametrize("provenance", [None, CLEAN_PROVENANCE]) -def test_no_warning_when_clean(tmp_path, recwarn, provenance): - variant_dir = _write_variant(tmp_path, provenance) - metadata = Metadata.read_from_file(variant_dir / "metadata.json") - _warn_if_dirty(metadata, variant_dir.name) - assert len(recwarn) == 0 - - -def test_warns_on_dirty_kernel_source(tmp_path): - variant_dir = _write_variant(tmp_path, DIRTY_KERNEL) - metadata = Metadata.read_from_file(variant_dir / "metadata.json") - with pytest.warns(UserWarning, match="dirty git tree"): - _warn_if_dirty(metadata, variant_dir.name) - - -def test_warns_names_dirty_sources(tmp_path): - variant_dir = _write_variant(tmp_path, DIRTY_BUILDER) - metadata = Metadata.read_from_file(variant_dir / "metadata.json") - with pytest.warns(UserWarning, match="kernel-builder"): - _warn_if_dirty(metadata, variant_dir.name) - - -def test_import_from_path_warns_on_dirty(tmp_path): - variant_dir = _write_variant(tmp_path, DIRTY_KERNEL) - (variant_dir / "__init__.py").write_text("value = 42\n") - _loaded_kernels.pop(variant_dir, None) - try: - with pytest.warns(UserWarning, match="dirty git tree"): - module = _import_from_path(variant_dir, deps={}) - assert module.value == 42 - finally: - _loaded_kernels.pop(variant_dir, None) diff --git a/kernels/tests/test_validate.py b/kernels/tests/test_validate.py index 1f258386..238ecd85 100644 --- a/kernels/tests/test_validate.py +++ b/kernels/tests/test_validate.py @@ -1,15 +1,69 @@ +import json import re from pathlib import Path import pytest import torch -from kernels_data import Version +from kernels_data import Metadata, Version import kernels import kernels.validate as validate_module from kernels.deps import DepTreeNode from kernels.resolver import LocalKernel -from kernels.validate import ArchValidator, MinverValidator, _installed_version +from kernels.validate import ( + ArchValidator, + DirtyValidator, + MinverValidator, + _installed_version, + default_metadata_validators, +) + +CLEAN_PROVENANCE = { + "kernel-builder": {"version": "0.1.0", "commit": "a" * 40, "dirty": False}, + "kernel": {"commit": "b" * 40, "dirty": False}, +} +DIRTY_KERNEL = { + "kernel-builder": {"version": "0.1.0", "commit": "a" * 40, "dirty": False}, + "kernel": {"commit": "b" * 40, "dirty": True}, +} +DIRTY_BUILDER = { + "kernel-builder": {"version": "0.1.0", "commit": "a" * 40, "dirty": True}, + "kernel": {"commit": "b" * 40, "dirty": False}, +} + + +def _metadata_with_provenance(provenance): + metadata = { + "id": "activation_1_cuda", + "name": "activation", + "version": 1, + "license": "Apache-2.0", + "python-depends": [], + "backend": {"type": "cuda"}, + } + if provenance is not None: + metadata["provenance"] = provenance + return Metadata.from_bytes(json.dumps(metadata).encode()) + + +@pytest.mark.parametrize("provenance", [None, CLEAN_PROVENANCE]) +def test_dirty_validator_does_not_warn_when_clean(recwarn, provenance): + DirtyValidator().validate_metadata(metadata=_metadata_with_provenance(provenance), variant="test-variant") + assert len(recwarn) == 0 + + +def test_dirty_validator_warns_on_dirty_kernel_source(): + with pytest.warns(UserWarning, match="dirty git tree"): + DirtyValidator().validate_metadata(metadata=_metadata_with_provenance(DIRTY_KERNEL), variant="test-variant") + + +def test_dirty_validator_names_dirty_sources(): + with pytest.warns(UserWarning, match="kernel-builder"): + DirtyValidator().validate_metadata(metadata=_metadata_with_provenance(DIRTY_BUILDER), variant="test-variant") + + +def test_dirty_validator_is_enabled_by_default(): + assert any(isinstance(validator, DirtyValidator) for validator in default_metadata_validators()) @pytest.mark.parametrize("minver", [None, "0.0.1"])