Skip to content
Merged
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
29 changes: 0 additions & 29 deletions kernels/src/kernels/importer.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,5 @@
import importlib
import sys
import warnings
from dataclasses import dataclass
from pathlib import Path
from types import ModuleType
Expand Down Expand Up @@ -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],
Expand All @@ -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"
Expand Down
27 changes: 26 additions & 1 deletion kernels/src/kernels/validate.py
Original file line number Diff line number Diff line change
@@ -1,3 +1,4 @@
import warnings
from dataclasses import dataclass
from typing import Protocol

Expand Down Expand Up @@ -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.

Expand Down Expand Up @@ -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()]
71 changes: 0 additions & 71 deletions kernels/tests/test_dirty_provenance.py

This file was deleted.

58 changes: 56 additions & 2 deletions kernels/tests/test_validate.py
Original file line number Diff line number Diff line change
@@ -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"])
Expand Down
Loading