diff --git a/Cargo.lock b/Cargo.lock index 4d1319f6..20423798 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -392,7 +392,7 @@ version = "3.1.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "faf9468729b8cbcea668e36183cb69d317348c2e08e994829fb56ebfdfbaac34" dependencies = [ - "windows-sys 0.59.0", + "windows-sys 0.61.2", ] [[package]] @@ -640,7 +640,7 @@ dependencies = [ "libc", "option-ext", "redox_users", - "windows-sys 0.59.0", + "windows-sys 0.61.2", ] [[package]] @@ -694,7 +694,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "39cab71617ae0d63f51a36d69f866391735b51691dbda63cf6f96d042b63efeb" dependencies = [ "libc", - "windows-sys 0.59.0", + "windows-sys 0.61.2", ] [[package]] @@ -1897,7 +1897,7 @@ version = "0.50.3" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "7957b9740744892f114936ab4a57b3f487491bbeafaf8083688b16841a4240e5" dependencies = [ - "windows-sys 0.59.0", + "windows-sys 0.61.2", ] [[package]] @@ -2540,7 +2540,7 @@ dependencies = [ "errno", "libc", "linux-raw-sys", - "windows-sys 0.59.0", + "windows-sys 0.61.2", ] [[package]] @@ -2597,7 +2597,7 @@ dependencies = [ "security-framework", "security-framework-sys", "webpki-root-certs", - "windows-sys 0.59.0", + "windows-sys 0.61.2", ] [[package]] @@ -3043,7 +3043,7 @@ dependencies = [ "getrandom 0.4.2", "once_cell", "rustix", - "windows-sys 0.59.0", + "windows-sys 0.61.2", ] [[package]] @@ -3848,7 +3848,7 @@ version = "0.1.11" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "c2a7b1c03c876122aa43f3020e6c3c3ee5c05081c9a00739faf7503aeba10d22" dependencies = [ - "windows-sys 0.59.0", + "windows-sys 0.61.2", ] [[package]] diff --git a/kernels-common/src/lib.rs b/kernels-common/src/lib.rs index 30149065..42d32063 100644 --- a/kernels-common/src/lib.rs +++ b/kernels-common/src/lib.rs @@ -4,4 +4,5 @@ pub mod git; pub mod hf; pub mod lock; pub mod metadata; +pub mod signing; pub mod version; diff --git a/kernels-common/src/signing/mod.rs b/kernels-common/src/signing/mod.rs new file mode 100644 index 00000000..fadaadf4 --- /dev/null +++ b/kernels-common/src/signing/mod.rs @@ -0,0 +1 @@ +pub mod receipt; diff --git a/kernels-common/src/signing/receipt.rs b/kernels-common/src/signing/receipt.rs new file mode 100644 index 00000000..86875caf --- /dev/null +++ b/kernels-common/src/signing/receipt.rs @@ -0,0 +1,30 @@ +use serde::{Deserialize, Serialize}; + +use crate::git::Oid; + +/// Kernel location. +/// +/// Every variant must change when the kernel changes, e.g. through +/// an update. For a remote kernel, this is determined by the revision, +/// for local kernels this could e.g. be based on the file +/// names/sizes/mtimes. +#[derive(Clone, Debug, Deserialize, Eq, Hash, PartialEq, Serialize)] +#[serde(tag = "type", rename_all = "snake_case")] +pub enum KernelLocation { + RemoteKernel { + repo_id: String, + revision: Oid, + variant: String, + }, +} + +impl KernelLocation { + /// A Hub kernel. + pub fn remote(repo_id: impl Into, revision: Oid, variant: impl Into) -> Self { + KernelLocation::RemoteKernel { + repo_id: repo_id.into(), + revision, + variant: variant.into(), + } + } +} diff --git a/kernels/rust/git.rs b/kernels/rust/git.rs new file mode 100644 index 00000000..1d995bee --- /dev/null +++ b/kernels/rust/git.rs @@ -0,0 +1,50 @@ +use std::str::FromStr; + +use kernels_common::git::Oid; +use pyo3::exceptions::PyValueError; +use pyo3::prelude::*; + +/// A git object identifier. +#[pyclass(name = "Oid", frozen, eq, hash, ord)] +#[derive(Clone, Debug, Eq, Hash, Ord, PartialEq, PartialOrd)] +pub(crate) struct PyOid { + inner: Oid, +} + +impl From for PyOid { + fn from(inner: Oid) -> Self { + Self { inner } + } +} + +impl PyOid { + pub(crate) fn into_inner(self) -> Oid { + self.inner + } +} + +/// Parse a git object id, mapping a parse failure to a Python `ValueError`. +pub(crate) fn parse_oid(s: &str) -> PyResult { + Oid::from_str(s).map_err(|err| PyValueError::new_err(err.to_string())) +} + +#[pymethods] +impl PyOid { + /// Parse a full SHA-1 or SHA-256 object id. + /// + /// Abbreviated identifiers are rejected: an object id must identify the + /// object unambiguously and permanently. + #[staticmethod] + #[pyo3(name = "from_str")] + fn py_from_str(s: &str) -> PyResult { + parse_oid(s).map(Into::into) + } + + fn __str__(&self) -> &str { + self.inner.as_str() + } + + fn __repr__(&self) -> String { + format!("Oid({:?})", self.inner.as_str()) + } +} diff --git a/kernels/rust/lib.rs b/kernels/rust/lib.rs index fd336079..4f1e2e16 100644 --- a/kernels/rust/lib.rs +++ b/kernels/rust/lib.rs @@ -13,11 +13,15 @@ use pyo3::exceptions::{PyException, PyOSError, PyRuntimeError, PyValueError}; use pyo3::prelude::*; mod config; +mod git; mod lock; +mod signing; mod version; use config::{PyBuild, PyGeneral}; +use git::PyOid; use lock::{PyKernelLock, PyKernelLocks, PyKernelPaths, PyNixKernelLock, PyNixKernelLocks}; +use signing::PyKernelLocation; use version::PyVersion; /// A validated kernel name matching `^[a-z][-a-z0-9]*[a-z0-9]$`. @@ -192,8 +196,8 @@ impl From for PyGitStatus { #[pymethods] impl PyGitStatus { #[getter] - fn commit(&self) -> &str { - self.commit.as_str() + fn commit(&self) -> PyOid { + self.commit.clone().into() } #[getter] @@ -758,6 +762,7 @@ fn data_py(m: &PyBound<'_, PyModule>) -> PyResult<()> { m.add_class::()?; m.add_class::()?; m.add_class::()?; + m.add_class::()?; m.add_class::()?; m.add_class::()?; m.add_class::()?; @@ -772,6 +777,7 @@ fn data_py(m: &PyBound<'_, PyModule>) -> PyResult<()> { m.add_class::()?; m.add_class::()?; m.add_class::()?; + m.add_class::()?; m.add( "DigestValidationError", m.py().get_type::(), diff --git a/kernels/rust/lock.rs b/kernels/rust/lock.rs index 2a7a803d..601b5d3f 100644 --- a/kernels/rust/lock.rs +++ b/kernels/rust/lock.rs @@ -1,6 +1,5 @@ use std::collections::BTreeMap; use std::path::PathBuf; -use std::str::FromStr; use kernels_common::git::Oid; use kernels_common::lock::{KernelLock, KernelLocks, KernelPaths, NixKernelLock, NixKernelLocks}; @@ -9,11 +8,7 @@ use pyo3::exceptions::{PyKeyError, PyValueError}; use pyo3::prelude::*; use crate::PyKernelDependency; - -/// Parse a git object id, mapping a parse failure to a Python `ValueError`. -fn parse_oid(s: &str) -> PyResult { - Oid::from_str(s).map_err(|err| PyValueError::new_err(err.to_string())) -} +use crate::git::{PyOid, parse_oid}; /// A locked kernel revision. #[pyclass(name = "KernelLock", frozen, eq, hash)] @@ -48,8 +43,8 @@ impl PyKernelLock { } #[getter] - fn commit(&self) -> &str { - self.commit.as_str() + fn commit(&self) -> PyOid { + self.commit.clone().into() } /// Parse a `KernelLock` from a JSON string. @@ -247,8 +242,8 @@ impl PyNixKernelLock { } #[getter] - fn commit(&self) -> &str { - self.commit.as_str() + fn commit(&self) -> PyOid { + self.commit.clone().into() } #[getter] diff --git a/kernels/rust/signing.rs b/kernels/rust/signing.rs new file mode 100644 index 00000000..4bf7dafc --- /dev/null +++ b/kernels/rust/signing.rs @@ -0,0 +1,41 @@ +use kernels_common::signing::receipt::KernelLocation; +use pyo3::prelude::*; + +use crate::git::PyOid; + +/// The location of a kernel that a verification applies to. +#[pyclass(name = "KernelLocation", frozen, eq, hash)] +#[derive(Clone, Debug, Eq, Hash, PartialEq)] +pub(crate) struct PyKernelLocation { + inner: KernelLocation, +} + +impl From for PyKernelLocation { + fn from(inner: KernelLocation) -> Self { + Self { inner } + } +} + +#[pymethods] +impl PyKernelLocation { + /// The location of a kernel variant in a Hub repository. + #[staticmethod] + fn remote(repo_id: String, revision: PyOid, variant: String) -> Self { + KernelLocation::remote(repo_id, revision.into_inner(), variant).into() + } + + fn __repr__(&self) -> String { + match &self.inner { + KernelLocation::RemoteKernel { + repo_id, + revision, + variant, + } => format!( + "KernelLocation.remote(repo_id={:?}, revision={:?}, variant={:?})", + repo_id, + revision.as_str(), + variant + ), + } + } +} diff --git a/kernels/src/kernels/_rust.pyi b/kernels/src/kernels/_rust.pyi index bd63a703..626d4206 100644 --- a/kernels/src/kernels/_rust.pyi +++ b/kernels/src/kernels/_rust.pyi @@ -21,12 +21,17 @@ __all__ = [ "KernelLocks", "KernelName", "KernelVersion", + "Oid", "Metadata", "NixKernelLock", "NixKernelLocks", "Digest", "DigestViolation", "DigestValidationError", + "KernelLocation", + "VerificationReceipt", + "ReceiptStore", + "ReceiptError", "Version", "__version__", ] @@ -83,8 +88,8 @@ class GitStatus: """The state of a git working tree.""" @property - def commit(self) -> str: - """Identifier of the `HEAD` commit, as lowercase hexadecimal digits.""" + def commit(self) -> "Oid": + """Identifier of the `HEAD` commit.""" ... @property @@ -285,6 +290,101 @@ class DigestValidationError(Exception): """The individual digest violations.""" ... +@final +class Oid: + """A git object identifier.""" + + @staticmethod + def from_str(s: str) -> "Oid": + """Parse a full SHA-1 or SHA-256 object id. + + Abbreviated identifiers are rejected: an object id must identify the + object unambiguously and permanently. + + Raises: + ValueError: If `s` is not 40 or 64 hexadecimal digits. + """ + ... + + def __str__(self) -> str: ... + def __repr__(self) -> str: ... + +@final +class KernelLocation: + """The location of a kernel that a verification applies to. + + The location changes whenever the kernel's contents change, so that a + receipt is never reused for a kernel that was updated or re-signed. + """ + + @staticmethod + def remote(repo_id: str, revision: Oid, variant: str) -> "KernelLocation": + """The location of a kernel variant in a Hub repository. + + Args: + repo_id: Repository the kernel was downloaded from. + revision: Resolved commit (as a Git SHA). + variant: Build variant of the kernel. + """ + ... + + def __repr__(self) -> str: ... + +@final +class VerificationReceipt: + """Receipt of a successful kernel verification. + + A receipt states that a kernel has already been verified. If the kernel + location changed, the receipt's hash will not match anymore.""" + + def __new__(cls, location: KernelLocation) -> "VerificationReceipt": ... + @property + def location(self) -> KernelLocation: + """The kernel location the verification applies to.""" + ... + + def __repr__(self) -> str: ... + +@final +class ReceiptStore: + """Store of kernel verification receipts.""" + + @staticmethod + def in_kernels_cache() -> "ReceiptStore": + """The receipt store inside the kernels cache. + + The cache location is resolved from the environment, falling back to + the Hub cache and then the user's home directory. + + Raises: + ReceiptError: If the cache directory cannot be determined. + """ + ... + + @staticmethod + def from_path(path: os.PathLike[str] | str) -> "ReceiptStore": + """A receipt store in the given directory.""" + ... + + def load(self, location: KernelLocation) -> Optional[VerificationReceipt]: + """The receipt for `location`, or `None` when the kernel has not been verified yet. + + Raises: + ReceiptError: If a receipt exists but cannot be used. + """ + ... + + def store(self, receipt: VerificationReceipt) -> None: + """Store `receipt`, replacing any existing receipt for its location. + + Raises: + ReceiptError: If the receipt cannot be written. + """ + ... + +class ReceiptError(Exception): + """Raised by `ReceiptStore` when a receipt cannot be read, written, or interpreted.""" + class KernelVersion: """A kernel version: either a numeric version or a git revision string.""" @@ -340,8 +440,8 @@ class KernelLock: ... @property - def commit(self) -> str: - """Locked commit of the kernel, as lowercase hexadecimal digits.""" + def commit(self) -> "Oid": + """Locked commit of the kernel.""" ... @staticmethod @@ -434,8 +534,8 @@ class NixKernelLock: ... @property - def commit(self) -> str: - """Locked commit of the kernel, as lowercase hexadecimal digits.""" + def commit(self) -> "Oid": + """Locked commit of the kernel.""" ... @property diff --git a/kernels/src/kernels/_versions.py b/kernels/src/kernels/_versions.py index 0e997c85..5ea778a9 100644 --- a/kernels/src/kernels/_versions.py +++ b/kernels/src/kernels/_versions.py @@ -1,16 +1,24 @@ import logging -import os from pathlib import Path from huggingface_hub import constants from huggingface_hub.file_download import repo_folder_name from huggingface_hub.hf_api import GitRefInfo -from kernels._rust import KernelDependency, KernelVersion +from kernels._rust import KernelVersion, Oid logger = logging.getLogger(__name__) +def _cached_refs_dir(repo_id: str) -> Path: + """The cache directory that holds the refs of a kernel repository.""" + # Lazy import so that we can mock it in tests. + from kernels.hf_hub import CACHE_DIR + + cache_dir = CACHE_DIR or constants.HF_HUB_CACHE + return Path(cache_dir) / repo_folder_name(repo_id=repo_id, repo_type="kernel") / "refs" + + def _get_available_versions(repo_id: str, *, local_files_only: bool) -> dict[int, GitRefInfo]: """Get kernel versions that are available in the repository.""" from kernels.hf_hub import _get_hf_api @@ -34,30 +42,27 @@ def _get_available_versions(repo_id: str, *, local_files_only: bool) -> dict[int def _get_available_versions_from_cache(repo_id: str) -> dict[int, GitRefInfo]: """Get kernel versions from the local Hugging Face cache.""" - cache_dir = os.environ.get("KERNELS_CACHE") or constants.HF_HUB_CACHE - versions: dict[int, GitRefInfo] = {} - # Tolerate both layouts: the "kernel" repo type used by newer - # huggingface_hub, and the legacy "model" prefix that older caches use. - for repo_type in ("kernel", "model"): - refs_dir = Path(cache_dir) / repo_folder_name(repo_id=repo_id, repo_type=repo_type) / "refs" - if not refs_dir.is_dir(): + + refs_dir = _cached_refs_dir(repo_id) + if not refs_dir.is_dir(): + return versions + + for ref_path in refs_dir.iterdir(): + if not ref_path.is_file(): + continue + ref_name = ref_path.name + if not ref_name.startswith("v"): + continue + try: + version = int(ref_name[1:]) + except ValueError: continue - for ref_path in refs_dir.iterdir(): - if not ref_path.is_file(): - continue - ref_name = ref_path.name - if not ref_name.startswith("v"): - continue - try: - version = int(ref_name[1:]) - except ValueError: - continue - try: - commit = ref_path.read_text().strip() - except OSError: - continue - versions[version] = GitRefInfo(name=ref_name, ref=ref_name, target_commit=commit) + try: + commit = ref_path.read_text().strip() + except OSError: + continue + versions[version] = GitRefInfo(name=ref_name, ref=ref_name, target_commit=commit) return versions @@ -92,7 +97,6 @@ def resolve_version_spec_as_ref(repo_id: str, version_spec: int, local_files_onl return ref -# TODO: maybe we should make this a KernelVersion factory method? def revision_or_version(*, revision: str | None, version: int | None) -> KernelVersion: if revision is not None and version is not None: raise ValueError("Only one of `revision` or `version` must be specified.") @@ -109,33 +113,83 @@ def revision_or_version(*, revision: str | None, version: int | None) -> KernelV ) -def resolve_kernel_version(kernel: KernelDependency, local_files_only: bool) -> str: - if isinstance(kernel.version, KernelVersion.Version): - return resolve_version_spec_as_ref( - kernel.repo_id, kernel.version.version, local_files_only=local_files_only - ).target_commit - elif isinstance(kernel.version, KernelVersion.Revision): - return kernel.version.revision +def _resolve_ref_from_cache(repo_id: str, ref: str) -> str | None: + """Resolve a ref to a commit using the local Hugging Face cache.""" + refs_dir = _cached_refs_dir(repo_id) + ref_path = refs_dir / ref + + # A ref is used as a path segment here, so make sure that a ref like + # `../../elsewhere` cannot read outside the refs directory. + try: + ref_path.resolve().relative_to(refs_dir.resolve()) + except (OSError, ValueError): + return None + + try: + return ref_path.read_text().strip() + except OSError: + return None + + +def _resolve_ref(repo_id: str, ref: str, *, local_files_only: bool) -> Oid: + """Resolve a branch, tag, or commit to the commit it points at. + + If the ref is already a full Git SHA, it is returned as-is and does not + require a Hub request. + """ + try: + return Oid.from_str(ref) + except ValueError: + pass + + if local_files_only: + commit = _resolve_ref_from_cache(repo_id, ref) + if commit is None: + raise ValueError( + f"Cannot resolve revision '{ref}' of '{repo_id}' to a commit: the ref is not in " + "the local cache and Hugging Face Hub is in offline mode. Download the kernel " + "while online first, or pass an explicit `revision=`." + ) + else: + from kernels.hf_hub import _get_hf_api + + commit = _get_hf_api().repo_info(repo_id=repo_id, repo_type="kernel", revision=ref).sha + if commit is None: + raise ValueError(f"Cannot resolve revision '{ref}' of '{repo_id}' to a commit.") + + return Oid.from_str(commit) + + +def resolve_kernel_version(repo_id: str, version: KernelVersion, *, local_files_only: bool) -> Oid: + """Resolve a kernel version to the commit it refers to. + + A `KernelVersion` can either be a version number or a revision (branch, + tag, or commit). This function resolves the version or revision into a + full Git commit SHA. + """ + if isinstance(version, KernelVersion.Version): + ref = resolve_version_spec_as_ref(repo_id, version.version, local_files_only=local_files_only) + return Oid.from_str(ref.target_commit) + elif isinstance(version, KernelVersion.Revision): + return _resolve_ref(repo_id, version.revision, local_files_only=local_files_only) else: - raise ValueError(f"Invalid version type: {kernel.version}") + raise ValueError(f"Invalid version type: {version}") -def select_revision_or_version( +def resolve_revision_or_version( repo_id: str, *, revision: str | None, version: int | None, local_files_only: bool, -) -> str: - if revision is not None and version is not None: - raise ValueError("Only one of `revision` or `version` must be specified.") - elif revision is not None: - return revision - elif version is not None: - return resolve_version_spec_as_ref(repo_id, version, local_files_only=local_files_only).target_commit - else: - raise ValueError( - "A kernel version or revision must be specified. " - "Use `version=` for a stable kernel API version or `revision=` " - "for an explicit Hub revision. See: https://huggingface.co/docs/kernels/migration" - ) +) -> Oid: + """Resolve a `revision` or `version` to a commit. + + The caller must provide either `revision` or `version`, but not both. The + revision can be a commit, tag, or branch. The full Git SHA of the revision + or version is returned.""" + return resolve_kernel_version( + repo_id, + revision_or_version(revision=revision, version=version), + local_files_only=local_files_only, + ) diff --git a/kernels/src/kernels/cli/download.py b/kernels/src/kernels/cli/download.py index 16fc0f08..785f7e55 100644 --- a/kernels/src/kernels/cli/download.py +++ b/kernels/src/kernels/cli/download.py @@ -32,7 +32,7 @@ def download_kernels(args): allow_patterns="build/*", ignore_patterns=_BYTECODE_IGNORE_PATTERNS, cache_dir=CACHE_DIR, - revision=lock.commit, + revision=str(lock.commit), ) else: location = resolve_hub_kernel(dep.repo_id, api=api, backend=None, revision=lock.commit) diff --git a/kernels/src/kernels/cli/verify_signature.py b/kernels/src/kernels/cli/verify_signature.py index 4f2bec8b..5f442fd3 100644 --- a/kernels/src/kernels/cli/verify_signature.py +++ b/kernels/src/kernels/cli/verify_signature.py @@ -8,14 +8,15 @@ else: from typing_extensions import assert_never -from kernels._versions import select_revision_or_version +from kernels._rust import KernelLocation +from kernels._versions import resolve_revision_or_version from kernels.install import install_kernel, install_kernel_all_variants from kernels.variants import get_variants_local from kernels.verify import VerificationResult, verify_variant def verify_signature(args: argparse.Namespace) -> None: - revision = select_revision_or_version( + revision = resolve_revision_or_version( args.repo_id, revision=None, version=args.version, @@ -23,45 +24,34 @@ def verify_signature(args: argparse.Namespace) -> None: ) if args.all_variants: - repo_path = install_kernel_all_variants(args.repo_id, revision=revision) + repo_path = install_kernel_all_variants(args.repo_id, revision=str(revision)) variants = get_variants_local(repo_path) kernel_paths = [repo_path / variant.variant_str for variant in variants] else: - kernel_paths = [install_kernel(args.repo_id, revision=revision)] + kernel_paths = [install_kernel(args.repo_id, revision=str(revision))] failed = False for kernel_path in kernel_paths: - result = verify_variant(kernel_path) variant_str = kernel_path.name + result = verify_variant( + kernel_path, + location=KernelLocation.remote(args.repo_id, revision, variant_str), + # Always fully verify the kernel in this subcommand. + cache=False, + ) + match result: - case VerificationResult.SignatureBundleMissing(): - if not args.filter_unsigned: - print(f"❌ {variant_str}: cannot verify kernel integrity, signature not found") - failed = True - continue - case VerificationResult.SignatureBundleInvalid(reason=reason): - print(f"❌ {variant_str}: cannot verify kernel integrity, invalid signature bundle:\n{reason}") - failed = True - case VerificationResult.MetadataInvalid(reason=reason): - print(f"❌ {variant_str}: cannot verify kernel integrity, invalid metadata:\n{reason}") - failed = True - case VerificationResult.MetadataMissing() | VerificationResult.DigestMissing(): - if not args.filter_no_digest: - print(f"❌ {variant_str}: cannot verify kernel integrity, metadata does not have a digest") - failed = True - continue - case VerificationResult.DigestVerificationFailure(violations=violations): - print(f"❌ {variant_str}: kernel integrity check failed") - for violation in violations: - print(violation) - failed = True - case VerificationResult.SignatureVerificationFailure(reason=reason): - print(f"❌ {variant_str}: metadata signature verification failed:\n{reason}") - failed = True + case VerificationResult.SignatureBundleMissing() if args.filter_unsigned: + pass + case VerificationResult.MetadataMissing() | VerificationResult.DigestMissing() if args.filter_no_digest: + pass case VerificationResult.Success(): - print(f"✅ {variant_str}: kernel metadata is correctly signed") + print(f"✅ {variant_str}: {result}") + case VerificationResult.Failure(): + print(f"❌ {variant_str}: {result}") + failed = True case _ as unreachable: assert_never(unreachable) diff --git a/kernels/src/kernels/compat.py b/kernels/src/kernels/compat.py index abb19dee..d5341c5a 100644 --- a/kernels/src/kernels/compat.py +++ b/kernels/src/kernels/compat.py @@ -7,6 +7,7 @@ import tomli as tomllib +has_sigstore = importlib.util.find_spec("sigstore") is not None has_torch = importlib.util.find_spec("torch") is not None has_tvm_ffi = importlib.util.find_spec("tvm_ffi") is not None diff --git a/kernels/src/kernels/deps.py b/kernels/src/kernels/deps.py index b0e008df..ce2141b9 100644 --- a/kernels/src/kernels/deps.py +++ b/kernels/src/kernels/deps.py @@ -1,7 +1,7 @@ from contextvars import ContextVar from dataclasses import dataclass from types import ModuleType -from typing import TYPE_CHECKING, Generic, TypeVar +from typing import Generic, TypeVar from huggingface_hub.hf_api import HfApi @@ -13,9 +13,7 @@ RemoteKernel, Resolver, ) - -if TYPE_CHECKING: - from kernels.validate import MetadataValidator +from kernels.validate import MetadataValidator # Default state is `None`, to signal that we are not in a kernel # loading context. @@ -72,7 +70,7 @@ def load(self) -> ModuleType: def validate_metadata( self: "DepTreeNode[LocalKernel | RemoteKernel]", - validator: "MetadataValidator", + validator: MetadataValidator, ) -> None: """Validate this kernel and its dependencies with the given validator.""" diff --git a/kernels/src/kernels/hf_hub.py b/kernels/src/kernels/hf_hub.py index ddb30951..634ea58c 100644 --- a/kernels/src/kernels/hf_hub.py +++ b/kernels/src/kernels/hf_hub.py @@ -5,6 +5,7 @@ from huggingface_hub import HfApi, constants +from kernels._rust import Oid from kernels._system import glibc_version from kernels.backends import _select_backend from kernels.compat import has_torch, has_tvm_ffi @@ -87,11 +88,11 @@ class RepoInfo: The following fields are available: - `repo_id` (`str`): the Hub repository containing the kernel. - - `revision` (`str`): the specific revision of the kernel. + - `revision` (`Oid`): the commit of the kernel. """ repo_id: str - revision: str + revision: Oid def _check_trust_remote_code(repo_id: str, local_files_only: bool, trust_remote_code: bool | list[str]) -> None: diff --git a/kernels/src/kernels/importer.py b/kernels/src/kernels/importer.py index 02c15c4e..c18d930a 100644 --- a/kernels/src/kernels/importer.py +++ b/kernels/src/kernels/importer.py @@ -15,10 +15,8 @@ class LoadedKernel: - `metadata` (`Metadata`): kernel metadata. - `module` (`ModuleType`): the imported kernel module. - - `repo_info` (`kernels.hf_hub.RepoInfo | None`): populated only for - kernels loaded via `get_kernel`. Loaders that work from a local path - (`get_local_kernel`) or a lockfile (`get_locked_kernel`, `load_kernel`) - leave this as `None`. + - `repo_info` (`kernels.hf_hub.RepoInfo | None`): populated whenever the + Hub repository the kernel came from is known. The metadata includes the following properties that describe a kernel: diff --git a/kernels/src/kernels/install.py b/kernels/src/kernels/install.py index dd656cfc..c95500d5 100644 --- a/kernels/src/kernels/install.py +++ b/kernels/src/kernels/install.py @@ -92,7 +92,7 @@ def install_kernel_all_variants( allow_patterns="build/*", ignore_patterns=_BYTECODE_IGNORE_PATTERNS, cache_dir=CACHE_DIR, - revision=lock.commit, + revision=str(lock.commit), ) ) ) diff --git a/kernels/src/kernels/layer/func.py b/kernels/src/kernels/layer/func.py index 0e555d5b..e550d610 100644 --- a/kernels/src/kernels/layer/func.py +++ b/kernels/src/kernels/layer/func.py @@ -1,4 +1,3 @@ -import functools import logging from pathlib import Path from types import ModuleType @@ -9,7 +8,6 @@ from kernels._rust import KernelDependency, KernelLocks from kernels.resolver import LockedHubCacheResolver, LockedHubResolver -from .._versions import select_revision_or_version from ..hf_hub import _get_hf_api from ..load import ( get_kernel, @@ -104,19 +102,11 @@ def __init__( self._revision = revision self._version = version - @functools.lru_cache() - def _resolve_revision(self) -> str: - return select_revision_or_version( - repo_id=self._repo_id, - revision=self._revision, - version=self._version, - local_files_only=constants.HF_HUB_OFFLINE, - ) - def load(self) -> Type["nn.Module"]: kernel = get_kernel( self._repo_id, - revision=self._resolve_revision(), + revision=self._revision, + version=self._version, trust_remote_code=self._trust_remote_code, ) return _get_kernel_func(self, kernel) @@ -144,8 +134,11 @@ def __hash__(self): ) ) + def _revision_str(self) -> str: + return self._revision if self._revision is not None else f"version {self._version}" + def __str__(self) -> str: - return f"`{self._repo_id}` (revision: {self._resolve_revision()}), function `{self.func_name}`" + return f"`{self._repo_id}` (revision: {self._revision_str()}), function `{self.func_name}`" class LocalFuncRepository: diff --git a/kernels/src/kernels/layer/layer.py b/kernels/src/kernels/layer/layer.py index 58cafc14..a01fcc7a 100644 --- a/kernels/src/kernels/layer/layer.py +++ b/kernels/src/kernels/layer/layer.py @@ -1,6 +1,5 @@ from __future__ import annotations -import functools import inspect import logging from inspect import Parameter, Signature @@ -13,7 +12,6 @@ from kernels._rust import KernelDependency, KernelLocks from kernels.resolver import LockedHubCacheResolver, LockedHubResolver -from .._versions import select_revision_or_version from ..hf_hub import _get_hf_api from ..load import ( get_kernel, @@ -103,19 +101,11 @@ def __init__( self._revision = revision self._version = version - @functools.lru_cache() - def _resolve_revision(self) -> str: - return select_revision_or_version( - repo_id=self._repo_id, - revision=self._revision, - version=self._version, - local_files_only=constants.HF_HUB_OFFLINE, - ) - def load(self) -> Type["nn.Module"]: kernel = get_kernel( self._repo_id, - revision=self._resolve_revision(), + revision=self._revision, + version=self._version, trust_remote_code=self._trust_remote_code, user_agent=self._user_agent, ) @@ -144,8 +134,11 @@ def __hash__(self): ) ) + def _revision_str(self) -> str: + return self._revision if self._revision is not None else f"version {self._version}" + def __str__(self) -> str: - return f"`{self._repo_id}` (revision: {self._resolve_revision()}), layer `{self.layer_name}`" + return f"`{self._repo_id}` (revision: {self._revision_str()}), layer `{self.layer_name}`" class LocalLayerRepository: diff --git a/kernels/src/kernels/locking.py b/kernels/src/kernels/locking.py index 91da77ba..9f6b3c7f 100644 --- a/kernels/src/kernels/locking.py +++ b/kernels/src/kernels/locking.py @@ -55,16 +55,16 @@ def lock_kernel_tree( trust_remote_code=False, ) - revision = resolve_kernel_version(kernel, local_files_only=False) + revision = resolve_kernel_version(kernel.repo_id, kernel.version, local_files_only=False) - for variant in get_variants(api, repo_id=kernel.repo_id, revision=revision): + for variant in get_variants(api, repo_id=kernel.repo_id, revision=str(revision)): metadata_path = Path( api.hf_hub_download( kernel.repo_id, repo_type="kernel", filename=f"build/{variant.variant_str}/metadata.json", cache_dir=CACHE_DIR, - revision=revision, + revision=str(revision), local_files_only=False, ) ) @@ -77,7 +77,7 @@ def lock_kernel_tree( seen.remove(kernel) kernel_locks[kernel] = KernelLock( - commit=revision, + commit=str(revision), ) return kernel_locks diff --git a/kernels/src/kernels/resolver.py b/kernels/src/kernels/resolver.py index 0a21c8e7..ad1d8748 100644 --- a/kernels/src/kernels/resolver.py +++ b/kernels/src/kernels/resolver.py @@ -5,7 +5,7 @@ from huggingface_hub.errors import LocalEntryNotFoundError from huggingface_hub.hf_api import HfApi -from kernels._rust import KernelDependency, KernelLocks, KernelPaths, Metadata +from kernels._rust import KernelDependency, KernelLocks, KernelPaths, Metadata, Oid from kernels._versions import _get_available_versions, resolve_kernel_version from kernels.hf_hub import CACHE_DIR, _check_trust_remote_code from kernels.variants import ( @@ -50,7 +50,7 @@ class RemoteKernel: """A kernel that can be loaded from a remote path.""" repo_id: str - revision: str + revision: Oid metadata: Metadata variant: Variant @@ -69,7 +69,7 @@ def install(self, *, api: HfApi) -> LocalKernel: allow_patterns=allow_patterns, ignore_patterns=_BYTECODE_IGNORE_PATTERNS, cache_dir=CACHE_DIR, - revision=self.revision, + revision=str(self.revision), local_files_only=False, ) ) @@ -114,12 +114,12 @@ def resolve_hub_kernel( *, api: HfApi, backend: str | None, - revision: str, + revision: Oid, ) -> RemoteKernel: variants = get_variants( api, repo_id=repo_id, - revision=revision, + revision=str(revision), ) variant, trace = resolve_variant(variants, backend) if variant is None: @@ -139,7 +139,7 @@ def resolve_hub_kernel( repo_type="kernel", filename=f"build/{variant.variant_str}/metadata.json", cache_dir=CACHE_DIR, - revision=revision, + revision=str(revision), local_files_only=False, ) ) @@ -186,7 +186,7 @@ def resolve_hub_cache_kernel( api: HfApi, repo_id: str, *, - revision: str, + revision: Oid, backend: str | None, ) -> LocalKernel: """Resolve a kernel variant path from the local Hugging Face cache only. @@ -202,7 +202,7 @@ def resolve_hub_cache_kernel( repo_type="kernel", ignore_patterns=_BYTECODE_IGNORE_PATTERNS, cache_dir=CACHE_DIR, - revision=revision, + revision=str(revision), local_files_only=True, ) ) @@ -226,7 +226,16 @@ def resolve_hub_cache_kernel( raise FileNotFoundError(f"Variant path does not exist: `{variant_path}`") metadata = Metadata.read_from_file(variant_path / "metadata.json") - location = LocalKernel(variant_path=variant_path, metadata=metadata) + location = LocalKernel( + variant_path=variant_path, + metadata=metadata, + origin=RemoteKernel( + repo_id=repo_id, + revision=revision, + metadata=metadata, + variant=variant, + ), + ) return location @@ -278,7 +287,7 @@ def resolve(self, *, api: HfApi, backend: str | None, kernel: KernelDependency) ) # Resolve the revision that we need. - revision = resolve_kernel_version(kernel, local_files_only=False) + revision = resolve_kernel_version(kernel.repo_id, kernel.version, local_files_only=False) # Get the kernel metadata for the revision. return resolve_hub_kernel(kernel.repo_id, api=api, revision=revision, backend=backend) @@ -301,7 +310,7 @@ def resolve(self, *, api: HfApi, backend: str | None, kernel: KernelDependency) ) # Resolve the revision that we need. - revision = resolve_kernel_version(kernel, local_files_only=True) + revision = resolve_kernel_version(kernel.repo_id, kernel.version, local_files_only=True) # Get the kernel metadata for the revision. return resolve_hub_cache_kernel( @@ -312,7 +321,7 @@ def resolve(self, *, api: HfApi, backend: str | None, kernel: KernelDependency) ) -def _locked_revision(kernel_locks: KernelLocks, kernel: KernelDependency) -> str: +def _locked_revision(kernel_locks: KernelLocks, kernel: KernelDependency) -> Oid: kernel_lock = kernel_locks.get(kernel, None) if kernel_lock is None: raise ValueError( diff --git a/kernels/src/kernels/variants.py b/kernels/src/kernels/variants.py index f5ec9168..abb00de1 100644 --- a/kernels/src/kernels/variants.py +++ b/kernels/src/kernels/variants.py @@ -12,7 +12,7 @@ from huggingface_hub.hf_api import RepoFolder from packaging.version import Version, parse -from kernels._versions import select_revision_or_version +from kernels._versions import resolve_revision_or_version from kernels.backends import ( CANN, CUDA, @@ -599,7 +599,7 @@ def get_kernel_variants( """ from kernels.hf_hub import _get_hf_api - revision = select_revision_or_version( + commit = resolve_revision_or_version( repo_id, revision=revision, version=version, @@ -607,7 +607,7 @@ def get_kernel_variants( ) api = _get_hf_api() - variants = get_variants(api, repo_id=repo_id, revision=revision) + variants = get_variants(api, repo_id=repo_id, revision=str(commit)) _, trace = resolve_variants(variants, backend) return trace diff --git a/kernels/src/kernels/verify.py b/kernels/src/kernels/verify.py index a0fff9b3..ff51e7ae 100644 --- a/kernels/src/kernels/verify.py +++ b/kernels/src/kernels/verify.py @@ -1,3 +1,5 @@ +import abc +import logging from dataclasses import dataclass from pathlib import Path from typing import TypeAlias, final @@ -8,7 +10,15 @@ from sigstore.verify import Verifier, policy from sigstore.verify.policy import VerificationPolicy -from kernels._rust import Digest, DigestValidationError, DigestViolation, Metadata +from kernels._rust import ( + Digest, + DigestValidationError, + DigestViolation, + KernelLocation, + Metadata, +) + +logger = logging.getLogger(__name__) class GitHubWorkflowPolicy(VerificationPolicy): @@ -47,30 +57,38 @@ def verify(self, cert: Certificate) -> None: return policy.AllOf(policies).verify(cert) +DEFAULT_POLICY: VerificationPolicy = policy.AnyOf( + [ + GitHubWorkflowPolicy( + repo_id="huggingface/kernels-community", + signer_uris=[ + "https://github.com/huggingface/kernels-community/.github/workflows/build.yaml@refs/heads/main", + "https://github.com/huggingface/kernels-community/.github/workflows/build-mac.yaml@refs/heads/main", + "https://github.com/huggingface/kernels-community/.github/workflows/build-windows.yaml@refs/heads/main", + # This workflow was used to sign existing builds, around Torch 2.10-2.12. Can be removed once these + # Torch versions are ancient. + "https://github.com/huggingface/kernels-community/.github/workflows/sign-old-builds.yaml@refs/heads/main", + ], + ), + ] +) """ -Default policies for the kernels package. +Default verification policy for the kernels package. -This is a curated set of trusted kernel developers. +Accepts kernels signed by a curated set of trusted kernel developers. """ -DEFAULT_POLICIES: list[VerificationPolicy] = [ - GitHubWorkflowPolicy( - repo_id="huggingface/kernels-community", - signer_uris=[ - "https://github.com/huggingface/kernels-community/.github/workflows/build.yaml@refs/heads/main", - "https://github.com/huggingface/kernels-community/.github/workflows/build-mac.yaml@refs/heads/main", - "https://github.com/huggingface/kernels-community/.github/workflows/build-windows.yaml@refs/heads/main", - # This workflow was used to sign existing builds, around Torch 2.10-2.12. Can be removed once these - # Torch versions are ancient. - "https://github.com/huggingface/kernels-community/.github/workflows/sign-old-builds.yaml@refs/heads/main", - ], - ), -] class VerificationResult: + class Failure(abc.ABC): + """A kernel build variant that could not be verified.""" + + @abc.abstractmethod + def __str__(self) -> str: ... + @final @dataclass - class DigestVerificationFailure: + class DigestVerificationFailure(Failure): """ Verification failed because there were digest violations. @@ -79,59 +97,78 @@ class DigestVerificationFailure: violations: list[DigestViolation] + def __str__(self) -> str: + violations = "\n".join(str(violation) for violation in self.violations) + return ( + "the files do not match the digest they were signed with, so they " + f"may have been modified:\n{violations}" + ) + @final @dataclass - class MetadataInvalid: + class MetadataInvalid(Failure): """ The kernel metadata could not be parsed. """ reason: str + def __str__(self) -> str: + return f"the metadata is invalid, so its integrity cannot be verified:\n{self.reason}" + @final @dataclass - class SignatureBundleInvalid: + class SignatureBundleInvalid(Failure): """ The signature bundle could not be parsed. """ reason: str + def __str__(self) -> str: + return f"the signature bundle is invalid, so its integrity cannot be verified:\n{self.reason}" + @final @dataclass - class SignatureVerificationFailure: + class SignatureVerificationFailure(Failure): """ Verification failed because the signature was not valid. """ reason: str + def __str__(self) -> str: + return f"the metadata could not be verified against its signature:\n{self.reason}" + @final @dataclass - class DigestMissing: + class DigestMissing(Failure): """ Verification failed because the metadata did not have a digest. """ - pass + def __str__(self) -> str: + return "the metadata does not record a digest, so its integrity cannot be verified" @final @dataclass - class MetadataMissing: + class MetadataMissing(Failure): """ Verification failed because the kernel did not have metadata. """ - pass + def __str__(self) -> str: + return "the metadata is missing, so its integrity cannot be verified" @final @dataclass - class SignatureBundleMissing: + class SignatureBundleMissing(Failure): """ Verification failed because the kernel metadata was not signed. """ - pass + def __str__(self) -> str: + return "not signed, so its integrity cannot be verified" @final @dataclass @@ -140,7 +177,8 @@ class Success: Verification was successful. """ - pass + def __str__(self) -> str: + return "the metadata is correctly signed" Any: TypeAlias = ( DigestMissing @@ -154,26 +192,37 @@ class Success: ) -def verify_variant(variant_path: Path, policies: list[VerificationPolicy] | None = None) -> VerificationResult.Any: +def verify_variant( + variant_path: Path, + *, + location: KernelLocation, + policy: VerificationPolicy | None = None, + cache: bool = True, +) -> VerificationResult.Any: """ Verify a kernel variant. - The kernel variant at the given path is verified using a set of policies. - This validates that the metadata was signed using a key that is compliant - with the given policies and that the kernel hashes match the digest in the - kernel metadata. + The kernel variant at the given path is verified using a policy. This + validates that the metadata was signed using a key that is compliant with + the given policy and that the kernel hashes match the digest in the kernel + metadata. Args: variant_path (`Path`): Kernel variant path. - policies (`list[VerificationPolicy]`, *optional*): - List of verification policies that should be used while verifying - the kernel. A default set of policies that accepts kernels signed - by a curated set of trusted kernel developers is used if this - argument is set to `None`. + location (`KernelLocation`): + Identity of the kernel, used to cache the verification. + policy (`VerificationPolicy`, *optional*): + Verification policy that should be used while verifying the + kernel. A default policy that accepts kernels signed by a curated + set of trusted kernel developers is used if this argument is set + to `None`. + cache (`bool`): + Whether to use the receipt cache to lookup or store kernel + verifications. Disabling cache use can be useful to validate + the integrity of a kernel downloaded from the hub. """ - if policies is None: - policies = DEFAULT_POLICIES + verify_policy = DEFAULT_POLICY if policy is None else policy bundle_path = variant_path / "metadata.json.sigstore" if not bundle_path.is_file(): @@ -190,7 +239,6 @@ def verify_variant(variant_path: Path, policies: list[VerificationPolicy] | None return VerificationResult.MetadataMissing() verifier = Verifier.production() - verify_policy = policy.AnyOf(policies) metadata_bytes = metadata_path.read_bytes() diff --git a/kernels/tests/test_data_kernel_lock.py b/kernels/tests/test_data_kernel_lock.py index b3db0df4..5c1590f1 100644 --- a/kernels/tests/test_data_kernel_lock.py +++ b/kernels/tests/test_data_kernel_lock.py @@ -49,7 +49,7 @@ def _sample_nix_locks(): def test_kernel_lock_commit(): lock = KernelLock(COMMIT_RELU) - assert lock.commit == COMMIT_RELU + assert str(lock.commit) == COMMIT_RELU def test_kernel_lock_invalid_commit(): @@ -141,7 +141,7 @@ def test_kernel_locks_from_json_rejects_invalid(): def test_nix_kernel_lock_construction(): lock = NixKernelLock(COMMIT_RELU, SRI_HASH) - assert lock.commit == COMMIT_RELU + assert str(lock.commit) == COMMIT_RELU assert lock.hash == SRI_HASH with pytest.raises(ValueError): diff --git a/kernels/tests/test_layer.py b/kernels/tests/test_layer.py index bac5dccd..41186af6 100644 --- a/kernels/tests/test_layer.py +++ b/kernels/tests/test_layer.py @@ -696,6 +696,7 @@ def get_kernel(repo_id, **kwargs): "kernels-test/silu-and-mul", { "revision": "main", + "version": None, "trust_remote_code": False, "user_agent": user_agent, }, diff --git a/kernels/tests/test_loaded_kernels.py b/kernels/tests/test_loaded_kernels.py index e821b24c..75b5b01b 100644 --- a/kernels/tests/test_loaded_kernels.py +++ b/kernels/tests/test_loaded_kernels.py @@ -58,7 +58,7 @@ def test_get_kernel_registers_loaded_kernel(fresh_registry): assert entry.metadata.name.python_name == _PACKAGE_NAME assert entry.repo_info is not None assert entry.repo_info.repo_id == _REPO_ID - assert isinstance(entry.repo_info.revision, str) and entry.repo_info.revision + assert str(entry.repo_info.revision) def test_repeated_get_kernel_is_cached(fresh_registry): diff --git a/kernels/tests/test_resolver.py b/kernels/tests/test_resolver.py index 16b818ac..c739640b 100644 --- a/kernels/tests/test_resolver.py +++ b/kernels/tests/test_resolver.py @@ -15,6 +15,7 @@ KernelPaths, KernelVersion, Metadata, + Oid, ) from kernels._versions import resolve_version_spec_as_ref from kernels.hf_hub import _get_hf_api @@ -241,18 +242,18 @@ def test_hub_resolver_resolves_remote_kernel(api): assert isinstance(location, RemoteKernel) assert location.repo_id == "kernels-community/relu" - assert re.fullmatch(r"[0-9a-f]{40}", location.revision) + assert re.fullmatch(r"[0-9a-f]{40}", str(location.revision)) assert location.metadata.name == KernelName("relu") -def test_hub_resolver_revision_passthrough(api): +def test_hub_resolver_resolves_revision_to_commit(api): location = HubResolver(trust_remote_code=False).resolve( api=api, backend="cpu", kernel=KernelDependency(repo_id="kernels-community/relu", version=KernelVersion.Revision("v1")), ) - assert location.revision == "v1" + assert location.revision == Oid.from_str(str(location.revision)) def test_hub_resolver_blocks_untrusted_org(api): @@ -426,7 +427,7 @@ def test_locked_hub_resolver_resolves_locked_revision(api, relu_locks): ) assert isinstance(location, RemoteKernel) - assert location.revision == commit + assert str(location.revision) == commit def test_locked_hub_resolver_requires_lock(api, relu_locks): @@ -462,7 +463,7 @@ def test_locked_revision_returns_commit(): dep = _dep("test/kernel", version=1) locks = KernelLocks({dep: KernelLock(commit="a" * 40)}) - assert _locked_revision(locks, dep) == "a" * 40 + assert _locked_revision(locks, dep) == Oid.from_str("a" * 40) def test_locked_revision_requires_lock(): diff --git a/kernels/tests/test_validate.py b/kernels/tests/test_validate.py index 3adb7ac6..95f470b9 100644 --- a/kernels/tests/test_validate.py +++ b/kernels/tests/test_validate.py @@ -8,9 +8,10 @@ import kernels import kernels.validate as validate_module -from kernels._rust import Metadata, Version +import kernels.verify as verify_module +from kernels._rust import Metadata, Oid, Version from kernels.deps import DepTreeNode -from kernels.resolver import LocalKernel +from kernels.resolver import LocalKernel, RemoteKernel from kernels.validate import ( ArchValidator, DirtyValidator, @@ -18,6 +19,8 @@ _installed_version, default_metadata_validators, ) +from kernels.variants import parse_variant +from kernels.verify import VerificationResult CLEAN_PROVENANCE = { "kernel-builder": {"version": "0.1.0", "commit": "a" * 40, "dirty": False}, @@ -207,3 +210,35 @@ def test_issue_707_fa3_on_b200(fake_cuda_device, make_metadata): metadata = make_metadata("cuda", ["8.0", "9.0a"]) with pytest.raises(RuntimeError, match="does not support the current device"): ArchValidator().validate_metadata(metadata=metadata, variant="test-variant") + + +_SIGNED_REPO_ID = "kernels-test/signatures" +_SIGNED_REVISION = Oid.from_str("a" * 40) + + +def _hub_kernel(tmp_path, metadata) -> LocalKernel: + variant_path = tmp_path / "torch-cuda" + return LocalKernel( + variant_path=variant_path, + metadata=metadata, + origin=RemoteKernel( + repo_id=_SIGNED_REPO_ID, + revision=_SIGNED_REVISION, + metadata=metadata, + variant=parse_variant("torch-cuda"), + ), + ) + + +@pytest.fixture +def recorded_verifications(monkeypatch): + """Record `verify_variant` calls and control the result it returns.""" + calls = [] + results = [] + + def fake_verify_variant(variant_path, *, location, policy=None, cache=True): + calls.append({"variant_path": variant_path, "policy": policy, "location": location, "cache": cache}) + return results.pop(0) if results else VerificationResult.Success() + + monkeypatch.setattr(verify_module, "verify_variant", fake_verify_variant) + return calls, results diff --git a/kernels/tests/test_verify.py b/kernels/tests/test_verify.py index 18bb6a91..65f9c3a8 100644 --- a/kernels/tests/test_verify.py +++ b/kernels/tests/test_verify.py @@ -1,44 +1,79 @@ +from dataclasses import is_dataclass from pathlib import Path +import pytest from sigstore.verify import policy +import kernels.verify as verify_module from kernels import install_kernel -from kernels._rust import DigestViolation -from kernels._versions import select_revision_or_version +from kernels._rust import DigestViolation, KernelLocation, Oid +from kernels._versions import resolve_revision_or_version from kernels.hf_hub import CACHE_DIR, _get_hf_api from kernels.resolver import _BYTECODE_IGNORE_PATTERNS from kernels.verify import VerificationResult, verify_variant -TEST_POLICIES: list[policy.VerificationPolicy] = [ - policy.Identity(identity="me@danieldk.eu", issuer="https://github.com/login/oauth") -] +TEST_POLICY: policy.VerificationPolicy = policy.Identity( + identity="me@danieldk.eu", issuer="https://github.com/login/oauth" +) + +OTHER_POLICY: policy.VerificationPolicy = policy.Identity( + identity="nobody@example.com", issuer="https://github.com/login/oauth" +) + + +@pytest.fixture +def signed_kernel(): + """A correctly signed kernel, with the location that identifies it.""" + repo_id = "kernels-test/signatures" + revision = resolve_revision_or_version(repo_id, revision=None, version=1, local_files_only=False) + variant_path = install_kernel(repo_id, revision=str(revision)) + return variant_path, KernelLocation.remote(repo_id, revision, variant_path.name) + + +def _verify_uncached(variant_path: Path, **kwargs) -> VerificationResult.Any: + """Verify a variant without reading or writing the receipt cache. + + Used by the tests that exercise verification itself rather than caching, + both to keep them away from the real receipt store and because a location + is required but unused when caching is off. + """ + return verify_variant( + variant_path, + location=KernelLocation.remote("kernels-test/signatures", Oid.from_str("0" * 40), variant_path.name), + cache=False, + **kwargs, + ) + + +def _no_hashing(monkeypatch): + """Make rehashing the variant fail, so that only cache hits can succeed.""" + + class ExplodingDigest: + @staticmethod + def hash_variant(*args, **kwargs): + raise AssertionError("the variant was rehashed, so this was not a cache hit") + + # Patch the name in `kernels.verify`: `Digest` is an extension type, whose + # attributes cannot be set. + monkeypatch.setattr(verify_module, "Digest", ExplodingDigest) def test_correctly_signed_kernel_passes_with_default_policy(): - revision = select_revision_or_version("kernels-community/relu", revision=None, version=1, local_files_only=False) - variant_path = install_kernel("kernels-community/relu", revision=revision) - assert verify_variant(variant_path) == VerificationResult.Success() + revision = resolve_revision_or_version("kernels-community/relu", revision=None, version=1, local_files_only=False) + variant_path = install_kernel("kernels-community/relu", revision=str(revision)) + assert _verify_uncached(variant_path) == VerificationResult.Success() def test_correctly_signed_kernel_passes(): - revision = select_revision_or_version("kernels-test/signatures", revision=None, version=1, local_files_only=False) - variant_path = install_kernel("kernels-test/signatures", revision=revision) - assert ( - verify_variant( - variant_path, - policies=TEST_POLICIES, - ) - == VerificationResult.Success() - ) + revision = resolve_revision_or_version("kernels-test/signatures", revision=None, version=1, local_files_only=False) + variant_path = install_kernel("kernels-test/signatures", revision=str(revision)) + assert _verify_uncached(variant_path, policy=TEST_POLICY) == VerificationResult.Success() def test_invalid_digest_fails(): variant_path = install_kernel("kernels-test/signatures", revision="invalid-digest") - match verify_variant( - variant_path, - policies=TEST_POLICIES, - ): + match _verify_uncached(variant_path, policy=TEST_POLICY): case VerificationResult.DigestVerificationFailure(violations=violations): assert len(violations) == 1 assert isinstance(violations[0], DigestViolation.HashMismatch) @@ -48,7 +83,7 @@ def test_invalid_digest_fails(): def test_invalid_metadata_fails(): # We cannot use regular code paths, because they require valid metadata. - revision = select_revision_or_version( + revision = resolve_revision_or_version( "kernels-test/signatures", revision="invalid-metadata", version=None, @@ -65,17 +100,17 @@ def test_invalid_metadata_fails(): allow_patterns="build/*", ignore_patterns=_BYTECODE_IGNORE_PATTERNS, cache_dir=CACHE_DIR, - revision=revision, + revision=str(revision), ) ) ) / "build" ) - match verify_variant( + match _verify_uncached( # No CUDA dependency, we are only checking metadata. variant_paths / "torch-cuda", - policies=TEST_POLICIES, + policy=TEST_POLICY, ): case VerificationResult.MetadataInvalid(reason=reason): assert "Cannot parse metadata" in reason @@ -85,18 +120,12 @@ def test_invalid_metadata_fails(): def test_missing_digest_fails(): variant_path = install_kernel("kernels-test/signatures", revision="missing-digest") - assert ( - verify_variant( - variant_path, - policies=TEST_POLICIES, - ) - == VerificationResult.DigestMissing() - ) + assert _verify_uncached(variant_path, policy=TEST_POLICY) == VerificationResult.DigestMissing() def test_missing_metadata_fails(): # We cannot use regular code paths, because they require valid metadata. - revision = select_revision_or_version( + revision = resolve_revision_or_version( "kernels-test/signatures", revision="missing-metadata", version=None, @@ -113,7 +142,7 @@ def test_missing_metadata_fails(): allow_patterns="build/*", ignore_patterns=_BYTECODE_IGNORE_PATTERNS, cache_dir=CACHE_DIR, - revision=revision, + revision=str(revision), ) ) ) @@ -121,10 +150,10 @@ def test_missing_metadata_fails(): ) assert ( - verify_variant( + _verify_uncached( # No CUDA dependency, we are only checking metadata. variant_paths / "torch-cuda", - policies=TEST_POLICIES, + policy=TEST_POLICY, ) == VerificationResult.MetadataMissing() ) @@ -132,21 +161,12 @@ def test_missing_metadata_fails(): def test_unsigned_kernel_fails(): variant_path = install_kernel("kernels-test/signatures", revision="signature-missing") - assert ( - verify_variant( - variant_path, - policies=TEST_POLICIES, - ) - == VerificationResult.SignatureBundleMissing() - ) + assert _verify_uncached(variant_path, policy=TEST_POLICY) == VerificationResult.SignatureBundleMissing() def test_broken_signature_bundle_fails(): variant_path = install_kernel("kernels-test/signatures", revision="signature-broken") - match verify_variant( - variant_path, - policies=TEST_POLICIES, - ): + match _verify_uncached(variant_path, policy=TEST_POLICY): case VerificationResult.SignatureBundleInvalid(reason=_): pass case other: @@ -155,11 +175,68 @@ def test_broken_signature_bundle_fails(): def test_invalid_signature_fails(): variant_path = install_kernel("kernels-test/signatures", revision="signature-invalid") - match verify_variant( - variant_path, - policies=TEST_POLICIES, - ): + match _verify_uncached(variant_path, policy=TEST_POLICY): case VerificationResult.SignatureVerificationFailure(reason=_): pass case other: raise RuntimeError(f"Expected SignatureVerificationFailure, was: {other}") + + +ALL_RESULTS = [ + VerificationResult.Success(), + VerificationResult.SignatureBundleMissing(), + VerificationResult.SignatureBundleInvalid(reason="bad bundle"), + VerificationResult.SignatureVerificationFailure(reason="bad signature"), + VerificationResult.MetadataMissing(), + VerificationResult.MetadataInvalid(reason="bad metadata"), + VerificationResult.DigestMissing(), + VerificationResult.DigestVerificationFailure(violations=[DigestViolation.MissingFile("kernel.py")]), +] + + +def test_all_results_are_covered(): + """`ALL_RESULTS` must cover every variant. + + Without this, adding a variant would silently skip the tests below, which + is exactly when they are needed. + """ + variants = { + name for name, member in vars(VerificationResult).items() if isinstance(member, type) and is_dataclass(member) + } + assert {type(result).__name__ for result in ALL_RESULTS} == variants + + +@pytest.mark.parametrize("result", ALL_RESULTS, ids=lambda result: type(result).__name__) +def test_every_result_describes_itself(result): + message = str(result) + assert message + # Prose, rather than the dataclass repr that `str` falls back to. + assert message != repr(result) + assert not message.startswith(type(result).__name__) + + +@pytest.mark.parametrize("result", ALL_RESULTS, ids=lambda result: type(result).__name__) +def test_only_success_is_not_a_failure(result): + is_success = isinstance(result, VerificationResult.Success) + assert isinstance(result, VerificationResult.Failure) != is_success + + +def test_result_messages_include_their_detail(): + assert "bang" in str(VerificationResult.SignatureBundleInvalid(reason="bang")) + assert "bang" in str(VerificationResult.MetadataInvalid(reason="bang")) + assert "bang" in str(VerificationResult.SignatureVerificationFailure(reason="bang")) + + violations = [DigestViolation.MissingFile("kernel.py"), DigestViolation.UnknownFile("extra.so")] + message = str(VerificationResult.DigestVerificationFailure(violations=violations)) + for violation in violations: + assert str(violation) in message + + +def test_failure_must_describe_itself(): + """The base class makes a message mandatory for new failures.""" + + class Undescribed(VerificationResult.Failure): + pass + + with pytest.raises(TypeError, match="abstract"): + Undescribed() # type: ignore[abstract] diff --git a/kernels/tests/test_versions.py b/kernels/tests/test_versions.py new file mode 100644 index 00000000..8a516a6a --- /dev/null +++ b/kernels/tests/test_versions.py @@ -0,0 +1,95 @@ +import pytest +from huggingface_hub.file_download import repo_folder_name + +import kernels.hf_hub as hf_hub +from kernels._rust import KernelVersion, Oid +from kernels._versions import _resolve_ref + +REPO_ID = "kernels-test/signatures" +COMMIT = "d649efb56fb249ac8f7a57fa1866728ad0c60e52" + + +@pytest.fixture +def cached_refs(tmp_path, monkeypatch): + """A cache containing a single ref, so offline resolution is hermetic.""" + monkeypatch.setattr(hf_hub, "CACHE_DIR", str(tmp_path)) + refs = tmp_path / repo_folder_name(repo_id=REPO_ID, repo_type="kernel") / "refs" + refs.mkdir(parents=True) + (refs / "main").write_text(COMMIT) + return refs + + +def test_commit_resolves_without_contacting_the_hub(monkeypatch): + """A revision that is already a commit must not cost a request. + + Lock files and `version=` both produce commits, so this is the common + path and it has to stay free. + """ + + def fail(): + raise AssertionError("the Hub was contacted to resolve a commit") + + monkeypatch.setattr(hf_hub, "_get_hf_api", fail) + + assert _resolve_ref(REPO_ID, COMMIT, local_files_only=False) == Oid.from_str(COMMIT) + + +def test_commit_is_canonicalized(monkeypatch): + def fail(): + raise AssertionError("the Hub was contacted to resolve a commit") + + monkeypatch.setattr(hf_hub, "_get_hf_api", fail) + + assert _resolve_ref(REPO_ID, COMMIT.upper(), local_files_only=False) == Oid.from_str(COMMIT) + + +def test_offline_resolution_uses_cached_refs(cached_refs): + assert _resolve_ref(REPO_ID, "main", local_files_only=True) == Oid.from_str(COMMIT) + + +def test_offline_resolution_reports_an_uncached_ref(cached_refs): + with pytest.raises(ValueError, match="not in the local cache"): + _resolve_ref(REPO_ID, "some-branch", local_files_only=True) + + +def test_offline_resolution_stays_inside_the_refs_directory(cached_refs, tmp_path): + """A ref is used as a path segment, so it must not be able to escape.""" + (tmp_path / "elsewhere").write_text(COMMIT) + + with pytest.raises(ValueError, match="not in the local cache"): + _resolve_ref(REPO_ID, "../../../elsewhere", local_files_only=True) + + +def test_a_cached_ref_must_still_be_a_commit(cached_refs): + (cached_refs / "main").write_text("not-a-commit") + + with pytest.raises(ValueError, match="Invalid git object id"): + _resolve_ref(REPO_ID, "main", local_files_only=True) + + +@pytest.mark.parametrize("invalid", ["d649efb", "", "main", "z" * 40, COMMIT + "0"]) +def test_oid_rejects_anything_but_a_full_object_id(invalid): + with pytest.raises(ValueError, match="Invalid git object id"): + Oid.from_str(invalid) + + +def test_a_version_needs_no_ref_resolution(monkeypatch): + """Looking up a version already yields its commit. + + Guards against reintroducing a resolution step that would suggest a + version can name something other than a commit. + """ + import kernels._versions as versions + + def fail(*args, **kwargs): + raise AssertionError("a version was resolved as if it were a ref") + + monkeypatch.setattr(versions, "_resolve_ref", fail) + + revision = versions.resolve_kernel_version( + "kernels-community/relu", + KernelVersion.Version(1), + local_files_only=False, + ) + + assert revision == Oid.from_str(str(revision))